From 1a2edf739d91fe72061034cdc5eb89d901d740e0 Mon Sep 17 00:00:00 2001 From: Soumil Rathi Date: Mon, 13 Apr 2026 12:17:35 -0700 Subject: [PATCH] fix: delete graph test files that import removed graph modules These test files import mem0.memory.graph_memory, kuzu_memory, memgraph_memory, apache_age_memory, and mem0.graphs.neptune which were deleted in the graph store removal commit. Co-Authored-By: Claude Opus 4.6 (1M context) --- tests/memory/test_apache_age_e2e.py | 1105 ----------------- tests/memory/test_apache_age_memory.py | 226 ---- tests/memory/test_graph_memory_soft_delete.py | 315 ----- tests/memory/test_kuzu.py | 253 ---- tests/memory/test_memgraph_memory.py | 107 -- tests/memory/test_neptune_analytics_memory.py | 338 ----- tests/memory/test_neptune_memory.py | 411 ------ 7 files changed, 2755 deletions(-) delete mode 100644 tests/memory/test_apache_age_e2e.py delete mode 100644 tests/memory/test_apache_age_memory.py delete mode 100644 tests/memory/test_graph_memory_soft_delete.py delete mode 100644 tests/memory/test_kuzu.py delete mode 100644 tests/memory/test_memgraph_memory.py delete mode 100644 tests/memory/test_neptune_analytics_memory.py delete mode 100644 tests/memory/test_neptune_memory.py diff --git a/tests/memory/test_apache_age_e2e.py b/tests/memory/test_apache_age_e2e.py deleted file mode 100644 index c55ea74a0..000000000 --- a/tests/memory/test_apache_age_e2e.py +++ /dev/null @@ -1,1105 +0,0 @@ -"""End-to-end integration tests for Apache AGE graph memory. - -These tests run against a real Apache AGE instance (via Docker) and exercise -every layer of the MemoryGraph class: connection, node MERGE, relationship -MERGE, embedding storage/retrieval, similarity search, deletion, and the -full add/search/get_all/delete_all/reset public API. - -Requirements: - docker run --name age-test \ - -e POSTGRES_DB=mem0_test -e POSTGRES_USER=mem0_user \ - -e POSTGRES_PASSWORD=mem0_pass -p 15432:5432 -d apache/age - -Run: - pytest tests/memory/test_apache_age_e2e.py -v -s -""" - -import json -import os -from unittest.mock import MagicMock, patch - -import age -import pytest - -from mem0.memory.apache_age_memory import MemoryGraph # noqa: E402 - -# -- E2E test configuration --------------------------------------------------- - -AGE_HOST = os.environ.get("AGE_HOST", "localhost") -AGE_PORT = int(os.environ.get("AGE_PORT", "15432")) -AGE_DB = os.environ.get("AGE_DB", "mem0_test") -AGE_USER = os.environ.get("AGE_USER", "mem0_user") -AGE_PASS = os.environ.get("AGE_PASS", "mem0_pass") -GRAPH_NAME = "e2e_test_graph" - - -def _age_available(): - """Check if the AGE database is reachable.""" - try: - ag = age.connect( - graph=GRAPH_NAME, - host=AGE_HOST, port=AGE_PORT, - dbname=AGE_DB, user=AGE_USER, password=AGE_PASS, - ) - ag.close() - return True - except Exception: - return False - - -skip_no_age = pytest.mark.skipif( - not _age_available(), - reason="Apache AGE not available (start Docker container first)", -) - - -# -- Helpers ------------------------------------------------------------------- - -def _make_e2e_instance(graph_name=GRAPH_NAME): - """Create a MemoryGraph instance wired to the real AGE database, - but with LLM and embedding model mocked out.""" - with patch.object(MemoryGraph, "__init__", return_value=None): - mg = MemoryGraph.__new__(MemoryGraph) - - mg.ag = age.connect( - graph=graph_name, - host=AGE_HOST, port=AGE_PORT, - dbname=AGE_DB, user=AGE_USER, password=AGE_PASS, - ) - mg.graph_name = graph_name - mg.threshold = 0.7 - - # Mock LLM — not needed for DB-layer tests - mg.llm_provider = "openai" - mg.llm = MagicMock() - mg.config = MagicMock() - mg.config.graph_store.custom_prompt = None - - # Deterministic embedding model: return a fixed vector derived from the - # entity name so similarity searches are predictable. Uses a longer - # vector (16-dim) with multiple hash seeds to minimize collisions. - def _fake_embed(text): - """Map text to a deterministic 16-dim vector for testing.""" - import hashlib - digest = hashlib.sha256(text.encode()).digest() - return [b / 255.0 for b in digest[:16]] - - mg.embedding_model = MagicMock() - mg.embedding_model.embed = _fake_embed - mg.user_id = None - - return mg - - -def _cleanup(mg): - """Remove all nodes and close the connection.""" - try: - mg._exec_cypher("MATCH (n) DETACH DELETE n") - mg.ag.commit() - except Exception: - pass - try: - mg.ag.close() - except Exception: - pass - - -# ============================================================================== -# Test: Low-level _exec_cypher -# ============================================================================== - -@skip_no_age -class TestExecCypher: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_create_and_return_vertex(self): - results = self.mg._exec_cypher( - "CREATE (n {name: %s, val: %s}) RETURN n", - params=("test_node", 42), - ) - self.mg.ag.commit() - assert len(results) == 1 - props = results[0] - assert props["name"] == "test_node" - assert props["val"] == 42 - - def test_return_scalars_with_cols(self): - self.mg._exec_cypher( - "CREATE (a {name: %s, user_id: %s})", params=("x", "u1") - ) - self.mg._exec_cypher( - "CREATE (b {name: %s, user_id: %s})", params=("y", "u1") - ) - self.mg._exec_cypher( - "MATCH (a {name: %s}), (b {name: %s}) CREATE (a)-[:LINK]->(b)", - params=("x", "y"), - ) - self.mg.ag.commit() - - results = self.mg._exec_cypher( - "MATCH (n)-[r]->(m) RETURN n.name, type(r), m.name", - cols=["source", "rel", "target"], - ) - assert len(results) == 1 - assert results[0] == {"source": "x", "rel": "LINK", "target": "y"} - - def test_empty_result(self): - results = self.mg._exec_cypher( - "MATCH (n {name: %s}) RETURN n", params=("nonexistent",) - ) - assert results == [] - - -# ============================================================================== -# Test: _merge_node -# ============================================================================== - -@skip_no_age -class TestMergeNode: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_creates_node_on_first_merge(self): - self.mg._merge_node("u1", "alice", [0.1, 0.2, 0.3]) - self.mg.ag.commit() - - results = self.mg._exec_cypher( - "MATCH (n {user_id: %s, name: %s}) RETURN n", - params=("u1", "alice"), - ) - assert len(results) == 1 - props = results[0] - assert props["name"] == "alice" - assert props["mentions"] == 1 - assert props["created"] is not None - assert json.loads(props["embedding"]) == [0.1, 0.2, 0.3] - - def test_merge_is_idempotent_increments_mentions(self): - self.mg._merge_node("u1", "bob", [0.4, 0.5]) - self.mg.ag.commit() - self.mg._merge_node("u1", "bob", [0.4, 0.5]) - self.mg.ag.commit() - - results = self.mg._exec_cypher( - "MATCH (n {user_id: %s, name: %s}) RETURN n", - params=("u1", "bob"), - ) - assert len(results) == 1 - assert results[0]["mentions"] == 2 - - def test_merge_with_agent_id(self): - self.mg._merge_node("u1", "carol", [0.6], agent_id="agent_1") - self.mg.ag.commit() - - results = self.mg._exec_cypher( - "MATCH (n {user_id: %s, name: %s}) RETURN n", - params=("u1", "carol"), - ) - assert results[0]["agent_id"] == "agent_1" - - -# ============================================================================== -# Test: Relationship creation and retrieval -# ============================================================================== - -@skip_no_age -class TestRelationships: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_create_and_query_relationship(self): - self.mg._merge_node("u1", "alice", [0.0]*16) - self.mg._merge_node("u1", "bob", [0.0]*16) - self.mg.ag.commit() - - self.mg._exec_cypher( - "MATCH (s {user_id: %s, name: %s}), (d {user_id: %s, name: %s}) " - "MERGE (s)-[r:KNOWS]->(d)", - params=("u1", "alice", "u1", "bob"), - ) - self.mg.ag.commit() - - results = self.mg._exec_cypher( - "MATCH (n {user_id: %s})-[r]->(m) RETURN n.name, type(r), m.name", - cols=["source", "rel", "target"], - params=("u1",), - ) - assert len(results) == 1 - assert results[0]["source"] == "alice" - assert results[0]["rel"] == "KNOWS" - assert results[0]["target"] == "bob" - - def test_multiple_relationships(self): - for name in ["alice", "bob", "carol"]: - self.mg._merge_node("u1", name, [0.0]) - self.mg.ag.commit() - - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)", - params=("alice", "bob"), - ) - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:LIKES]->(d)", - params=("alice", "carol"), - ) - self.mg.ag.commit() - - results = self.mg._exec_cypher( - "MATCH (n {name: %s})-[r]->(m) RETURN n.name, type(r), m.name", - cols=["source", "rel", "target"], - params=("alice",), - ) - assert len(results) == 2 - rels = {r["rel"] for r in results} - assert rels == {"KNOWS", "LIKES"} - - -# ============================================================================== -# Test: Embedding storage + similarity search -# ============================================================================== - -@skip_no_age -class TestEmbeddingSimilarity: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_embedding_roundtrip(self): - emb = [0.1, 0.2, 0.3, 0.4] + [0.0]*12 - self.mg._merge_node("u1", "node_a", emb) - self.mg.ag.commit() - - results = self.mg._exec_cypher( - "MATCH (n {name: %s}) RETURN n", params=("node_a",) - ) - stored = json.loads(results[0]["embedding"]) - assert stored == emb - - def test_find_similar_node_exact_match(self): - emb = [1.0, 0.0, 0.0, 0.0] + [0.0]*12 - self.mg._merge_node("u1", "target", emb) - self.mg.ag.commit() - - match = self.mg._find_similar_node(emb, {"user_id": "u1"}, threshold=0.99) - assert match is not None - assert match["name"] == "target" - - def test_find_similar_node_no_match_below_threshold(self): - self.mg._merge_node("u1", "far_away", [1.0, 0.0, 0.0, 0.0] + [0.0]*12) - self.mg.ag.commit() - - orthogonal = [0.0, 1.0, 0.0, 0.0] + [0.0]*12 - match = self.mg._find_similar_node(orthogonal, {"user_id": "u1"}, threshold=0.5) - assert match is None - - def test_find_similar_node_picks_closest(self): - # Use vectors where "close" is clearly more similar to the query - self.mg._merge_node("u1", "close", [0.9, 0.1, 0.0, 0.0] + [0.0]*12) - self.mg._merge_node("u1", "far", [0.0, 0.0, 1.0, 0.0] + [0.0]*12) - self.mg.ag.commit() - - query = [1.0, 0.0, 0.0, 0.0] + [0.0]*12 - match = self.mg._find_similar_node(query, {"user_id": "u1"}, threshold=0.5) - assert match is not None - assert match["name"] == "close" - - def test_find_similar_node_respects_user_id(self): - vec = [1.0, 0.0, 0.0, 0.0] + [0.0]*12 - self.mg._merge_node("u1", "mine", vec) - self.mg._merge_node("u2", "theirs", vec) - self.mg.ag.commit() - - match = self.mg._find_similar_node( - vec, {"user_id": "u2"}, threshold=0.9 - ) - assert match is not None - assert match["name"] == "theirs" - - -# ============================================================================== -# Test: Public API — get_all, delete_all, reset -# ============================================================================== - -@skip_no_age -class TestPublicAPICRUD: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_get_all_returns_relationships(self): - self.mg._merge_node("u1", "alice", [0.0]) - self.mg._merge_node("u1", "bob", [0.0]) - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FRIEND]->(d)", - params=("alice", "bob"), - ) - self.mg.ag.commit() - - results = self.mg.get_all({"user_id": "u1"}) - assert len(results) == 1 - assert results[0]["source"] == "alice" - assert results[0]["relationship"] == "FRIEND" - assert results[0]["target"] == "bob" - - def test_get_all_empty_for_different_user(self): - self.mg._merge_node("u1", "alice", [0.0]) - self.mg._merge_node("u1", "bob", [0.0]) - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FRIEND]->(d)", - params=("alice", "bob"), - ) - self.mg.ag.commit() - - results = self.mg.get_all({"user_id": "u999"}) - assert results == [] - - def test_get_all_respects_limit(self): - for i in range(5): - self.mg._merge_node("u1", f"src_{i}", [0.0]) - self.mg._merge_node("u1", f"dst_{i}", [0.0]) - self.mg.ag.commit() - for i in range(5): - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:REL]->(d)", - params=(f"src_{i}", f"dst_{i}"), - ) - self.mg.ag.commit() - - results = self.mg.get_all({"user_id": "u1"}, top_k=3) - assert len(results) == 3 - - def test_delete_all_removes_user_data(self): - self.mg._merge_node("u1", "alice", [0.0]) - self.mg._merge_node("u1", "bob", [0.0]) - self.mg._merge_node("u2", "carol", [0.0]) - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)", - params=("alice", "bob"), - ) - self.mg.ag.commit() - - self.mg.delete_all({"user_id": "u1"}) - - # u1's data should be gone - results = self.mg._exec_cypher( - "MATCH (n {user_id: %s}) RETURN n", params=("u1",) - ) - assert results == [] - - # u2's data should still exist - results = self.mg._exec_cypher( - "MATCH (n {user_id: %s}) RETURN n", params=("u2",) - ) - assert len(results) == 1 - - def test_reset_clears_everything(self): - self.mg._merge_node("u1", "alice", [0.0]) - self.mg._merge_node("u2", "bob", [0.0]) - self.mg.ag.commit() - - self.mg.reset() - - results = self.mg._exec_cypher("MATCH (n) RETURN n") - assert results == [] - - -# ============================================================================== -# Test: _delete_entities -# ============================================================================== - -@skip_no_age -class TestDeleteEntities: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_deletes_specific_relationship(self): - self.mg._merge_node("u1", "alice", [0.0]) - self.mg._merge_node("u1", "bob", [0.0]) - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)", - params=("alice", "bob"), - ) - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:LIKES]->(d)", - params=("alice", "bob"), - ) - self.mg.ag.commit() - - # Delete only KNOWS - self.mg._delete_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "u1"}, - ) - - # LIKES should remain - results = self.mg._exec_cypher( - "MATCH (n {name: %s})-[r]->(m) RETURN n.name, type(r), m.name", - cols=["source", "rel", "target"], - params=("alice",), - ) - assert len(results) == 1 - assert results[0]["rel"] == "LIKES" - - -# ============================================================================== -# Test: _add_entities (full flow with merge + similarity) -# ============================================================================== - -@skip_no_age -class TestAddEntities: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_creates_new_nodes_and_relationship(self): - result = self.mg._add_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "u1"}, - entity_type_map={"alice": "person", "bob": "person"}, - ) - assert len(result) == 1 - assert result[0][0]["source"] == "alice" - assert result[0][0]["relationship"] == "KNOWS" - assert result[0][0]["target"] == "bob" - - # Verify nodes exist in DB - nodes = self.mg._exec_cypher( - "MATCH (n {user_id: %s}) RETURN n", params=("u1",) - ) - names = {n["name"] for n in nodes} - assert names == {"alice", "bob"} - - def test_merges_to_existing_similar_node(self): - # Pre-create "alice" with a known embedding - emb = self.mg.embedding_model.embed("alice") - self.mg._merge_node("u1", "alice", emb) - self.mg.ag.commit() - - # Now add an entity where source="alice" — should merge to existing - self.mg.threshold = 0.99 # high threshold, but same embedding = exact match - self.mg._add_entities( - [{"source": "alice", "destination": "carol", "relationship": "LIKES"}], - {"user_id": "u1"}, - entity_type_map={"alice": "person", "carol": "person"}, - ) - - # Should still have exactly one "alice" node (not a duplicate) - nodes = self.mg._exec_cypher( - "MATCH (n {user_id: %s, name: %s}) RETURN n", - params=("u1", "alice"), - ) - assert len(nodes) == 1 - # mentions should be > 1 from merge - assert nodes[0]["mentions"] >= 2 - - def test_add_multiple_relationships(self): - entities = [ - {"source": "alice", "destination": "bob", "relationship": "KNOWS"}, - {"source": "alice", "destination": "carol", "relationship": "LIKES"}, - {"source": "bob", "destination": "carol", "relationship": "WORKS_WITH"}, - ] - self.mg._add_entities(entities, {"user_id": "u1"}, entity_type_map={}) - - results = self.mg.get_all({"user_id": "u1"}) - assert len(results) == 3 - rels = {(r["source"], r["relationship"], r["target"]) for r in results} - assert ("alice", "KNOWS", "bob") in rels - assert ("alice", "LIKES", "carol") in rels - assert ("bob", "WORKS_WITH", "carol") in rels - - -# ============================================================================== -# Test: _search_graph_db -# ============================================================================== - -@skip_no_age -class TestSearchGraphDB: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_finds_related_entities(self): - # Create a small graph - emb_alice = self.mg.embedding_model.embed("alice") - emb_bob = self.mg.embedding_model.embed("bob") - self.mg._merge_node("u1", "alice", emb_alice) - self.mg._merge_node("u1", "bob", emb_bob) - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)", - params=("alice", "bob"), - ) - self.mg.ag.commit() - - # Search with "alice" embedding — should find the KNOWS relationship - self.mg.threshold = 0.99 # exact match only - results = self.mg._search_graph_db(["alice"], {"user_id": "u1"}) - assert len(results) >= 1 - found_knows = any( - r["source"] == "alice" and r["relationship"] == "KNOWS" and r["destination"] == "bob" - for r in results - ) - assert found_knows, f"Expected KNOWS relationship in {results}" - - def test_search_returns_empty_for_no_matches(self): - self.mg.threshold = 0.99 - results = self.mg._search_graph_db(["nonexistent"], {"user_id": "u1"}) - assert results == [] - - -# ============================================================================== -# Test: Full add() + search() integration (mocking LLM, real DB) -# ============================================================================== - -@skip_no_age -class TestAddSearchIntegration: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_full_add_and_search_cycle(self): - """Simulates the full add() → search() cycle with mocked LLM responses.""" - filters = {"user_id": "test_user_1"} - - # Mock LLM: _retrieve_nodes_from_data - self.mg.llm.generate_response.side_effect = [ - # 1st call: extract entities for add() - {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Alice", "entity_type": "person"}, - {"entity": "Bob", "entity_type": "person"}, - ]}}]}, - # 2nd call: establish relations for add() - {"tool_calls": [{"name": "add_entities", "arguments": {"entities": [ - {"source": "Alice", "relationship": "knows", "destination": "Bob"}, - ]}}]}, - # 3rd call: get_delete_entities (nothing to delete) - {"tool_calls": []}, - # 4th call: extract entities for search() - {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Alice", "entity_type": "person"}, - ]}}]}, - ] - - # Add - add_result = self.mg.add("Alice knows Bob", filters) - assert "added_entities" in add_result - assert "deleted_entities" in add_result - - # Verify in DB - all_rels = self.mg.get_all(filters) - assert len(all_rels) == 1 - assert all_rels[0]["source"] == "alice" - # Relationship labels are lowercased by _remove_spaces_from_entities - assert all_rels[0]["relationship"] == "knows" - assert all_rels[0]["target"] == "bob" - - # Search - search_results = self.mg.search("Who does Alice know?", filters) - assert len(search_results) >= 1 - assert any(r["source"] == "alice" and r["destination"] == "bob" for r in search_results) - - # Delete all - self.mg.delete_all(filters) - remaining = self.mg.get_all(filters) - assert remaining == [] - - def test_add_then_update_relationship(self): - """Tests that adding conflicting data removes old relationships.""" - filters = {"user_id": "test_user_2"} - - # First add: Alice likes cats - self.mg.llm.generate_response.side_effect = [ - {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Alice", "entity_type": "person"}, - {"entity": "cats", "entity_type": "animal"}, - ]}}]}, - {"tool_calls": [{"name": "add_entities", "arguments": {"entities": [ - {"source": "Alice", "relationship": "likes", "destination": "cats"}, - ]}}]}, - {"tool_calls": []}, # nothing to delete - ] - self.mg.add("Alice likes cats", filters) - - all_rels = self.mg.get_all(filters) - assert len(all_rels) == 1 - assert all_rels[0]["relationship"] == "likes" - - # Second add: Alice now dislikes cats (delete old, add new) - self.mg.llm.generate_response.side_effect = [ - {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Alice", "entity_type": "person"}, - {"entity": "cats", "entity_type": "animal"}, - ]}}]}, - {"tool_calls": [{"name": "add_entities", "arguments": {"entities": [ - {"source": "Alice", "relationship": "dislikes", "destination": "cats"}, - ]}}]}, - # LLM says to delete the old likes relationship - {"tool_calls": [{"name": "delete_graph_memory", "arguments": { - "source": "alice", "relationship": "likes", "destination": "cats", - }}]}, - ] - self.mg.add("Alice dislikes cats", filters) - - all_rels = self.mg.get_all(filters) - rels = {r["relationship"] for r in all_rels} - assert "likes" not in rels - assert "dislikes" in rels - - self.mg.delete_all(filters) - - -# ============================================================================== -# Test: Multi-tenant isolation -# ============================================================================== - -@skip_no_age -class TestMultiTenantIsolation: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_users_cant_see_each_others_data(self): - # User 1 - self.mg._merge_node("user_1", "alice", [0.0]*16) - self.mg._merge_node("user_1", "bob", [0.0]*16) - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {user_id: %s, name: %s}), (d {user_id: %s, name: %s}) " - "MERGE (s)-[:KNOWS]->(d)", - params=("user_1", "alice", "user_1", "bob"), - ) - self.mg.ag.commit() - - # User 2 - self.mg._merge_node("user_2", "carol", [0.0]*16) - self.mg._merge_node("user_2", "dave", [0.0]*16) - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {user_id: %s, name: %s}), (d {user_id: %s, name: %s}) " - "MERGE (s)-[:WORKS_WITH]->(d)", - params=("user_2", "carol", "user_2", "dave"), - ) - self.mg.ag.commit() - - # User 1 sees only their data - u1_results = self.mg.get_all({"user_id": "user_1"}) - assert len(u1_results) == 1 - assert u1_results[0]["source"] == "alice" - - # User 2 sees only their data - u2_results = self.mg.get_all({"user_id": "user_2"}) - assert len(u2_results) == 1 - assert u2_results[0]["source"] == "carol" - - # Delete user 1 doesn't affect user 2 - self.mg.delete_all({"user_id": "user_1"}) - u2_after = self.mg.get_all({"user_id": "user_2"}) - assert len(u2_after) == 1 - - -# ============================================================================== -# Test: agent_id / run_id filtering in delete_all and get_all -# ============================================================================== - -@skip_no_age -class TestAgentRunIdFiltering: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_get_all_filters_by_agent_id(self): - # Create nodes for two different agents under the same user - self.mg._merge_node("u1", "alice", [0.0]*16, agent_id="agent_a") - self.mg._merge_node("u1", "bob", [0.0]*16, agent_id="agent_a") - self.mg._merge_node("u1", "carol", [0.0]*16, agent_id="agent_b") - self.mg._merge_node("u1", "dave", [0.0]*16, agent_id="agent_b") - self.mg.ag.commit() - - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS]->(d)", - params=("alice", "bob"), - ) - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:LIKES]->(d)", - params=("carol", "dave"), - ) - self.mg.ag.commit() - - # get_all with agent_a should only return alice->bob - results_a = self.mg.get_all({"user_id": "u1", "agent_id": "agent_a"}) - assert len(results_a) == 1 - assert results_a[0]["source"] == "alice" - - # get_all with agent_b should only return carol->dave - results_b = self.mg.get_all({"user_id": "u1", "agent_id": "agent_b"}) - assert len(results_b) == 1 - assert results_b[0]["source"] == "carol" - - def test_delete_all_filters_by_agent_id(self): - self.mg._merge_node("u1", "alice", [0.0]*16, agent_id="agent_a") - self.mg._merge_node("u1", "bob", [0.0]*16, agent_id="agent_a") - self.mg._merge_node("u1", "carol", [0.0]*16, agent_id="agent_b") - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:REL]->(d)", - params=("alice", "bob"), - ) - self.mg.ag.commit() - - # Delete only agent_a's data - self.mg.delete_all({"user_id": "u1", "agent_id": "agent_a"}) - - # agent_a nodes should be gone - nodes_a = self.mg._exec_cypher( - "MATCH (n) WHERE n.user_id = %s AND n.agent_id = %s RETURN n", - params=("u1", "agent_a"), - ) - assert nodes_a == [] - - # agent_b's data should still exist - nodes_b = self.mg._exec_cypher( - "MATCH (n) WHERE n.user_id = %s AND n.agent_id = %s RETURN n", - params=("u1", "agent_b"), - ) - assert len(nodes_b) == 1 - - def test_get_all_filters_by_run_id(self): - self.mg._merge_node("u1", "alice", [0.0]*16, run_id="run_1") - self.mg._merge_node("u1", "bob", [0.0]*16, run_id="run_1") - self.mg._merge_node("u1", "carol", [0.0]*16, run_id="run_2") - self.mg._merge_node("u1", "dave", [0.0]*16, run_id="run_2") - self.mg.ag.commit() - - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:R1]->(d)", - params=("alice", "bob"), - ) - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:R2]->(d)", - params=("carol", "dave"), - ) - self.mg.ag.commit() - - results = self.mg.get_all({"user_id": "u1", "run_id": "run_1"}) - assert len(results) == 1 - assert results[0]["source"] == "alice" - - def test_delete_all_filters_by_run_id(self): - self.mg._merge_node("u1", "alice", [0.0]*16, run_id="run_1") - self.mg._merge_node("u1", "bob", [0.0]*16, run_id="run_2") - self.mg.ag.commit() - - self.mg.delete_all({"user_id": "u1", "run_id": "run_1"}) - - # run_1 node should be gone - nodes_1 = self.mg._exec_cypher( - "MATCH (n) WHERE n.user_id = %s AND n.run_id = %s RETURN n", - params=("u1", "run_1"), - ) - assert nodes_1 == [] - - # run_2 node should remain - nodes_2 = self.mg._exec_cypher( - "MATCH (n) WHERE n.user_id = %s AND n.run_id = %s RETURN n", - params=("u1", "run_2"), - ) - assert len(nodes_2) == 1 - - -# ============================================================================== -# Test: Special characters and edge cases -# ============================================================================== - -@skip_no_age -class TestEdgeCases: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_node_name_with_underscores(self): - """Entity names go through _remove_spaces_from_entities which lowercases - and replaces spaces with underscores.""" - self.mg._merge_node("u1", "new_york_city", [0.0]*16) - self.mg._merge_node("u1", "united_states", [0.0]*16) - self.mg.ag.commit() - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:LOCATED_IN]->(d)", - params=("new_york_city", "united_states"), - ) - self.mg.ag.commit() - - results = self.mg.get_all({"user_id": "u1"}) - assert len(results) == 1 - assert results[0]["source"] == "new_york_city" - assert results[0]["target"] == "united_states" - - def test_node_with_apostrophe_in_name(self): - """The AGE Python driver has a known limitation where single quotes in - parameterized values cause a syntax error due to double-quoting in - buildCypher(). In practice this is not hit because entity names go - through _remove_spaces_from_entities which sanitizes them.""" - import psycopg2 - with pytest.raises(psycopg2.errors.SyntaxError): - self.mg._merge_node("u1", "o'brien", [0.0]*16) - - def test_empty_graph_get_all(self): - results = self.mg.get_all({"user_id": "u1"}) - assert results == [] - - def test_empty_graph_delete_all_no_error(self): - # Should not raise even on empty graph - self.mg.delete_all({"user_id": "u1"}) - - def test_empty_graph_reset_no_error(self): - self.mg.reset() - - def test_duplicate_relationship_merge_is_idempotent(self): - """MERGE on the same relationship twice should not create duplicates.""" - self.mg._merge_node("u1", "alice", [0.0]*16) - self.mg._merge_node("u1", "bob", [0.0]*16) - self.mg.ag.commit() - - for _ in range(3): - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FRIENDS]->(d)", - params=("alice", "bob"), - ) - self.mg.ag.commit() - - results = self.mg.get_all({"user_id": "u1"}) - assert len(results) == 1 # Only one relationship, not 3 - - def test_bidirectional_relationships(self): - """Two nodes can have relationships in both directions.""" - self.mg._merge_node("u1", "alice", [0.0]*16) - self.mg._merge_node("u1", "bob", [0.0]*16) - self.mg.ag.commit() - - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FOLLOWS]->(d)", - params=("alice", "bob"), - ) - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:FOLLOWS]->(d)", - params=("bob", "alice"), - ) - self.mg.ag.commit() - - results = self.mg.get_all({"user_id": "u1"}) - assert len(results) == 2 - pairs = {(r["source"], r["target"]) for r in results} - assert ("alice", "bob") in pairs - assert ("bob", "alice") in pairs - - def test_self_referencing_relationship(self): - """A node can have a relationship to itself.""" - self.mg._merge_node("u1", "alice", [0.0]*16) - self.mg.ag.commit() - - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:KNOWS_SELF]->(d)", - params=("alice", "alice"), - ) - self.mg.ag.commit() - - results = self.mg.get_all({"user_id": "u1"}) - assert len(results) == 1 - assert results[0]["source"] == "alice" - assert results[0]["target"] == "alice" - - -# ============================================================================== -# Test: _search_graph_db with agent_id/run_id filtering -# ============================================================================== - -@skip_no_age -class TestSearchGraphDBFiltering: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_search_filters_by_agent_id(self): - emb = self.mg.embedding_model.embed("alice") - self.mg._merge_node("u1", "alice", emb, agent_id="agent_a") - self.mg._merge_node("u1", "bob", [0.0]*16, agent_id="agent_a") - self.mg._merge_node("u1", "alice_clone", emb, agent_id="agent_b") - self.mg._merge_node("u1", "carol", [0.0]*16, agent_id="agent_b") - self.mg.ag.commit() - - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:R1]->(d)", - params=("alice", "bob"), - ) - self.mg._exec_cypher( - "MATCH (s {name: %s}), (d {name: %s}) MERGE (s)-[:R2]->(d)", - params=("alice_clone", "carol"), - ) - self.mg.ag.commit() - - self.mg.threshold = 0.99 - results = self.mg._search_graph_db( - ["alice"], {"user_id": "u1", "agent_id": "agent_a"} - ) - # Should only find relationships for agent_a - for r in results: - assert r["source"] != "alice_clone", f"Leaked agent_b data: {r}" - - def test_find_similar_node_filters_by_run_id(self): - emb = [1.0, 0.0, 0.0, 0.0] + [0.0]*12 - self.mg._merge_node("u1", "target_run1", emb, run_id="run_1") - self.mg._merge_node("u1", "target_run2", emb, run_id="run_2") - self.mg.ag.commit() - - match = self.mg._find_similar_node( - emb, {"user_id": "u1", "run_id": "run_1"}, threshold=0.99 - ) - assert match is not None - assert match["name"] == "target_run1" - - -# ============================================================================== -# Test: Full lifecycle with agent_id -# ============================================================================== - -@skip_no_age -class TestFullLifecycleWithAgentId: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_add_search_delete_with_agent_id(self): - """Full cycle: add → get_all → search → delete_all, all scoped by agent_id.""" - filters = {"user_id": "u1", "agent_id": "agent_x"} - - self.mg.llm.generate_response.side_effect = [ - # extract entities - {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Python", "entity_type": "language"}, - {"entity": "Alice", "entity_type": "person"}, - ]}}]}, - # establish relations - {"tool_calls": [{"name": "add_entities", "arguments": {"entities": [ - {"source": "Alice", "relationship": "uses", "destination": "Python"}, - ]}}]}, - # nothing to delete - {"tool_calls": []}, - ] - - self.mg.add("Alice uses Python", filters) - - # Verify nodes have agent_id - nodes = self.mg._exec_cypher( - "MATCH (n) WHERE n.user_id = %s AND n.agent_id = %s RETURN n", - params=("u1", "agent_x"), - ) - assert len(nodes) == 2 - names = {n["name"] for n in nodes} - assert names == {"alice", "python"} - - # get_all with agent_id filter - all_rels = self.mg.get_all(filters) - assert len(all_rels) == 1 - assert all_rels[0]["relationship"] == "uses" - - # search - self.mg.llm.generate_response.side_effect = [ - {"tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Alice", "entity_type": "person"}, - ]}}]}, - ] - search_results = self.mg.search("What does Alice use?", filters) - assert len(search_results) >= 1 - - # delete only agent_x - self.mg.delete_all(filters) - remaining = self.mg.get_all(filters) - assert remaining == [] - - -# ============================================================================== -# Test: _merge_node preserves created timestamp -# ============================================================================== - -@skip_no_age -class TestMergeNodeTimestamp: - - def setup_method(self): - self.mg = _make_e2e_instance() - - def teardown_method(self): - _cleanup(self.mg) - - def test_created_preserved_across_merges(self): - """The created timestamp should be set on first merge and preserved on subsequent merges.""" - self.mg._merge_node("u1", "alice", [0.0]*16) - self.mg.ag.commit() - - nodes1 = self.mg._exec_cypher( - "MATCH (n {name: %s}) RETURN n", params=("alice",) - ) - created1 = nodes1[0]["created"] - assert created1 is not None - - # Second merge — created should not change - import time - time.sleep(0.05) # Ensure clock moves - self.mg._merge_node("u1", "alice", [0.0]*16) - self.mg.ag.commit() - - nodes2 = self.mg._exec_cypher( - "MATCH (n {name: %s}) RETURN n", params=("alice",) - ) - created2 = nodes2[0]["created"] - assert created2 == created1, f"created changed from {created1} to {created2}" - assert nodes2[0]["mentions"] == 2 diff --git a/tests/memory/test_apache_age_memory.py b/tests/memory/test_apache_age_memory.py deleted file mode 100644 index c0cc5e086..000000000 --- a/tests/memory/test_apache_age_memory.py +++ /dev/null @@ -1,226 +0,0 @@ -from unittest.mock import MagicMock, Mock, patch - -# age and rank_bm25 are optional deps — mock them so tests run without install -_age_mock = Mock() -patch.dict("sys.modules", { - "age": _age_mock, - "age.models": Mock(), - "rank_bm25": Mock(), -}).start() - -from mem0.memory.apache_age_memory import MemoryGraph, _cosine_similarity # noqa: E402 - - -def _make_instance(): - with patch.object(MemoryGraph, "__init__", return_value=None): - instance = MemoryGraph.__new__(MemoryGraph) - instance.llm_provider = "openai" - instance.llm = MagicMock() - instance.embedding_model = MagicMock() - instance.config = MagicMock() - instance.config.graph_store.custom_prompt = None - instance.ag = MagicMock() - instance.graph_name = "test_graph" - instance.threshold = 0.7 - return instance - - -class TestCosineSimilarity: - """Tests for the _cosine_similarity helper.""" - - def test_identical_vectors(self): - assert abs(_cosine_similarity([1, 0, 0], [1, 0, 0]) - 1.0) < 1e-6 - - def test_orthogonal_vectors(self): - assert abs(_cosine_similarity([1, 0, 0], [0, 1, 0])) < 1e-6 - - def test_zero_vector(self): - assert _cosine_similarity([0, 0, 0], [1, 2, 3]) == 0.0 - - -class TestRetrieveNodesFromData: - """Tests for _retrieve_nodes_from_data in Apache AGE MemoryGraph.""" - - def test_normal_entities_extracted(self): - instance = _make_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Alice", "entity_type": "person"}, - {"entity": "hiking", "entity_type": "activity"}, - ]}}] - } - result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"}) - assert result == {"alice": "person", "hiking": "activity"} - - def test_malformed_entity_missing_entity_type_is_skipped(self): - instance = _make_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "matrix multiplication", "entity_type": "task"}, - {"entity": "task"}, - {"entity": "ReLU", "entity_type": "task"}, - ]}}] - } - result = instance._retrieve_nodes_from_data("some text", {"user_id": "u1"}) - assert "matrix_multiplication" in result - assert "relu" in result - assert "task" not in result - - def test_missing_entities_key_returns_empty(self): - instance = _make_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}] - } - result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"}) - assert result == {} - - def test_none_tool_calls_returns_empty(self): - instance = _make_instance() - instance.llm.generate_response.return_value = {"tool_calls": None} - result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"}) - assert result == {} - - -class TestEstablishNodesRelationsFromData: - """Tests for _establish_nodes_relations_from_data in Apache AGE MemoryGraph.""" - - def test_none_response_does_not_crash(self): - instance = _make_instance() - instance.llm.generate_response.return_value = None - result = instance._establish_nodes_relations_from_data( - "Hello world", {"user_id": "u1"}, {} - ) - assert result == [] - - def test_empty_tool_calls_returns_empty(self): - instance = _make_instance() - instance.llm.generate_response.return_value = {"tool_calls": []} - result = instance._establish_nodes_relations_from_data( - "Hello world", {"user_id": "u1"}, {} - ) - assert result == [] - - def test_valid_entities_returned(self): - instance = _make_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "add_entities", "arguments": {"entities": [ - {"source": "alice", "relationship": "loves", "destination": "hiking"} - ]}}] - } - result = instance._establish_nodes_relations_from_data( - "Alice loves hiking", {"user_id": "u1"}, {"alice": "person"} - ) - assert len(result) == 1 - assert result[0]["source"] == "alice" - - -class TestRemoveSpacesFromEntities: - """Tests for _remove_spaces_from_entities.""" - - def test_spaces_and_case(self): - instance = _make_instance() - entities = [{"source": "Alice Smith", "relationship": "Works At", "destination": "Big Corp"}] - result = instance._remove_spaces_from_entities(entities) - assert result[0]["source"] == "alice_smith" - assert result[0]["relationship"] == "works_at" - assert result[0]["destination"] == "big_corp" - - -class TestFindSimilarNode: - """Tests for _find_similar_node.""" - - def test_returns_none_when_no_nodes(self): - instance = _make_instance() - instance._exec_cypher = MagicMock(return_value=[]) - result = instance._find_similar_node([1.0, 0.0], {"user_id": "u1"}, threshold=0.9) - assert result is None - - def test_returns_best_match_above_threshold(self): - instance = _make_instance() - instance._exec_cypher = MagicMock(return_value=[ - {"name": "alice", "embedding": [1.0, 0.0], "user_id": "u1"}, - {"name": "bob", "embedding": [0.0, 1.0], "user_id": "u1"}, - ]) - result = instance._find_similar_node([1.0, 0.0], {"user_id": "u1"}, threshold=0.9) - assert result["name"] == "alice" - - def test_filters_by_agent_id(self): - instance = _make_instance() - instance._exec_cypher = MagicMock(return_value=[ - {"name": "alice", "embedding": [1.0, 0.0], "user_id": "u1", "agent_id": "a2"}, - ]) - result = instance._find_similar_node( - [1.0, 0.0], {"user_id": "u1", "agent_id": "a1"}, threshold=0.9 - ) - assert result is None - - -class TestDeleteAll: - """Tests for delete_all.""" - - def test_calls_exec_cypher_and_commits(self): - instance = _make_instance() - instance._exec_cypher = MagicMock(return_value=[]) - instance.delete_all({"user_id": "u1"}) - instance._exec_cypher.assert_called_once() - instance.ag.commit.assert_called_once() - - -class TestGetAll: - """Tests for get_all.""" - - def test_returns_formatted_results(self): - instance = _make_instance() - instance._exec_cypher = MagicMock(return_value=[ - {"source": "alice", "relationship": "KNOWS", "target": "bob"}, - {"source": "alice", "relationship": "LIKES", "target": "hiking"}, - ]) - results = instance.get_all({"user_id": "u1"}, top_k=10) - assert len(results) == 2 - assert results[0]["source"] == "alice" - assert results[0]["relationship"] == "KNOWS" - assert results[0]["target"] == "bob" - - def test_passes_limit_to_cypher(self): - """Limit is enforced via LIMIT in the Cypher query, not Python slicing.""" - instance = _make_instance() - instance._exec_cypher = MagicMock(return_value=[ - {"source": "n0", "relationship": "R", "target": "m0"}, - ]) - instance.get_all({"user_id": "u1"}, top_k=3) - # Verify limit was passed as a parameter to the query - cypher_stmt = instance._exec_cypher.call_args[0][0] - assert "LIMIT %s" in cypher_stmt - params = instance._exec_cypher.call_args[1].get("params") or instance._exec_cypher.call_args[0][2] - assert 3 in params - - -class TestAdd: - """Tests for the add orchestration method.""" - - def test_add_returns_added_and_deleted(self): - instance = _make_instance() - instance._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"}) - instance._establish_nodes_relations_from_data = MagicMock(return_value=[ - {"source": "alice", "relationship": "knows", "destination": "bob"} - ]) - instance._search_graph_db = MagicMock(return_value=[]) - instance._get_delete_entities_from_search_output = MagicMock(return_value=[]) - instance._delete_entities = MagicMock(return_value=[]) - instance._add_entities = MagicMock(return_value=["added"]) - - result = instance.add("Alice knows Bob", {"user_id": "u1"}) - assert "deleted_entities" in result - assert "added_entities" in result - assert result["added_entities"] == ["added"] - - -class TestSearch: - """Tests for the search method.""" - - def test_returns_empty_when_no_search_output(self): - instance = _make_instance() - instance._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"}) - instance._search_graph_db = MagicMock(return_value=[]) - result = instance.search("Who is Alice?", {"user_id": "u1"}) - assert result == [] diff --git a/tests/memory/test_graph_memory_soft_delete.py b/tests/memory/test_graph_memory_soft_delete.py deleted file mode 100644 index 7b17f5c5d..000000000 --- a/tests/memory/test_graph_memory_soft_delete.py +++ /dev/null @@ -1,315 +0,0 @@ -"""Tests for graph memory soft-delete behavior. - -Verifies that _delete_entities marks relationships as invalid (soft-delete) -rather than permanently removing them, and that search/retrieval queries -exclude soft-deleted relationships by default. - -See: https://github.com/mem0ai/mem0/issues/4187 -""" - -from unittest.mock import Mock, patch - -# Mock optional deps at module level so the import works across all Python -# versions without triggering transitive C-extension reloads (numpy via -# qdrant_client). This matches the pattern in test_memgraph_memory.py. -_neo4j_mock = Mock() -patch.dict("sys.modules", { - "langchain_neo4j": _neo4j_mock, - "rank_bm25": Mock(), -}).start() - -from mem0.memory.graph_memory import MemoryGraph # noqa: E402 - - -def _create_graph_memory(): - """Create a MemoryGraph instance with mocked dependencies.""" - with patch.object(MemoryGraph, "__init__", lambda self, *a, **kw: None): - mg = MemoryGraph.__new__(MemoryGraph) - mg.graph = Mock() - mg.graph.query = Mock(return_value=[]) - mg.embedding_model = Mock() - mg.embedding_model.embed = Mock(return_value=[0.1] * 128) - mg.llm = Mock() - mg.node_label = ":Entity" - mg.threshold = 0.7 - mg.llm_provider = "openai" - return mg - - -class TestSoftDelete: - """Verify _delete_entities uses SET r.valid = false, not DELETE r.""" - - def test_delete_entities_sends_soft_delete_cypher(self): - mg = _create_graph_memory() - mg.graph.query.return_value = [ - {"source": "Alice", "target": "Bob", "relationship": "KNOWS"} - ] - - mg._delete_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "user1"}, - ) - - cypher = mg.graph.query.call_args[0][0] - assert "SET r.valid = false" in cypher - assert "r.invalidated_at = datetime()" in cypher - assert "DELETE r" not in cypher - - def test_delete_entities_only_targets_valid_edges(self): - mg = _create_graph_memory() - mg._delete_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "user1"}, - ) - - cypher = mg.graph.query.call_args[0][0] - assert "r.valid IS NULL OR r.valid = true" in cypher - - def test_delete_entities_is_idempotent(self): - mg = _create_graph_memory() - item = [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}] - filters = {"user_id": "user1"} - - mg.graph.query.return_value = [ - {"source": "Alice", "target": "Bob", "relationship": "KNOWS"} - ] - mg._delete_entities(item, filters) - - mg.graph.query.return_value = [] - mg._delete_entities(item, filters) - - # Both calls should have the same WHERE filter - for c in mg.graph.query.call_args_list: - assert "r.valid IS NULL OR r.valid = true" in c[0][0] - - -class TestSearchExcludesSoftDeleted: - """Verify search and get_all filter out soft-deleted relationships.""" - - def test_get_all_filters_soft_deleted(self): - mg = _create_graph_memory() - mg.get_all(filters={"user_id": "user1"}, top_k=10) - - cypher = mg.graph.query.call_args[0][0] - assert "r.valid IS NULL OR r.valid = true" in cypher - - def test_search_graph_db_filters_both_directions(self): - """_search_graph_db must filter soft-deleted edges in both outgoing and incoming queries.""" - mg = _create_graph_memory() - mg.graph.query.return_value = [] - - mg._search_graph_db(node_list=["alice"], filters={"user_id": "user1"}) - - cypher = mg.graph.query.call_args[0][0] - # The UNION query has two MATCH branches — both must filter - occurrences = cypher.count("r.valid IS NULL OR r.valid = true") - assert occurrences >= 2, ( - f"_search_graph_db has {occurrences} valid-filter(s) but needs >= 2 " - "(one for outgoing, one for incoming relationships)" - ) - - def test_delete_all_still_hard_deletes(self): - mg = _create_graph_memory() - mg.delete_all(filters={"user_id": "user1"}) - - cypher = mg.graph.query.call_args[0][0] - assert "DETACH DELETE" in cypher - - -class TestMergeResetsValidFlag: - """Verify MERGE in _add_entities sets r.valid = true. - - Critical: after soft-delete, a MERGE that matches the existing - (invalidated) edge must reset valid=true, or the edge becomes - a zombie -- exists but invisible to queries. - """ - - def _run_add_entities(self, source_found, dest_found): - """Helper: call _add_entities with configurable node search results.""" - mg = _create_graph_memory() - - source_result = ( - [{"elementId(source_candidate)": "src_id_1"}] if source_found else [] - ) - dest_result = ( - [{"elementId(destination_candidate)": "dst_id_1"}] if dest_found else [] - ) - - mg._search_source_node = Mock(return_value=source_result) - mg._search_destination_node = Mock(return_value=dest_result) - - mg._add_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "user1"}, - entity_type_map={}, - ) - - cypher = mg.graph.query.call_args[0][0] - return cypher - - def test_merge_sets_valid_true_when_source_found(self): - cypher = self._run_add_entities(source_found=True, dest_found=False) - assert "r.valid = true" in cypher - - def test_merge_sets_valid_true_when_dest_found(self): - cypher = self._run_add_entities(source_found=False, dest_found=True) - assert "r.valid = true" in cypher - - def test_merge_sets_valid_true_when_both_found(self): - cypher = self._run_add_entities(source_found=True, dest_found=True) - assert "r.valid = true" in cypher - - def test_merge_sets_valid_true_when_neither_found(self): - cypher = self._run_add_entities(source_found=False, dest_found=False) - assert "r.valid = true" in cypher - - def test_merge_clears_invalidated_at_on_resurrection(self): - """When a soft-deleted edge is resurrected via MERGE, invalidated_at must be cleared. - - Without this, a resurrected edge (valid=true) still carries stale - invalidated_at metadata, which corrupts temporal reasoning queries. - """ - for label, src, dst in [ - ("source found", True, False), - ("dest found", False, True), - ("both found", True, True), - ("neither found", False, False), - ]: - cypher = self._run_add_entities(source_found=src, dest_found=dst) - assert "r.invalidated_at = null" in cypher, ( - f"MERGE path '{label}': ON MATCH SET does not clear r.invalidated_at. " - "Resurrected edges will have stale invalidation timestamps." - ) - - -class TestCypherConsistency: - """Verify all MERGE blocks use consistent property names and variable aliases.""" - - def _get_merge_cypher(self, source_found, dest_found): - mg = _create_graph_memory() - mg._search_source_node = Mock( - return_value=[{"elementId(source_candidate)": "id1"}] if source_found else [] - ) - mg._search_destination_node = Mock( - return_value=[{"elementId(destination_candidate)": "id2"}] if dest_found else [] - ) - mg._add_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "user1"}, - entity_type_map={}, - ) - return mg.graph.query.call_args[0][0] - - def test_all_blocks_use_created_at_not_created(self): - """All MERGE blocks must use r.created_at, not r.created.""" - for label, src, dst in [ - ("source found", True, False), - ("dest found", False, True), - ("both found", True, True), - ("neither found", False, False), - ]: - cypher = self._get_merge_cypher(src, dst) - assert "r.created_at" in cypher, ( - f"MERGE path '{label}': uses r.created instead of r.created_at" - ) - - def test_all_blocks_use_r_not_rel(self): - """All MERGE blocks must use 'r' as the relationship variable, not 'rel'.""" - for label, src, dst in [ - ("source found", True, False), - ("dest found", False, True), - ("both found", True, True), - ("neither found", False, False), - ]: - cypher = self._get_merge_cypher(src, dst) - assert "rel." not in cypher, ( - f"MERGE path '{label}': uses 'rel' variable instead of 'r'" - ) - - def test_all_blocks_set_updated_at_on_create(self): - """All MERGE blocks must set r.updated_at on CREATE for consistent timestamps.""" - for label, src, dst in [ - ("source found", True, False), - ("dest found", False, True), - ("both found", True, True), - ("neither found", False, False), - ]: - cypher = self._get_merge_cypher(src, dst) - assert "r.updated_at = timestamp()" in cypher, ( - f"MERGE path '{label}': missing r.updated_at on CREATE SET" - ) - - -class TestSoftDeleteWithFilters: - """Verify soft-delete works correctly with agent_id and run_id filters.""" - - def test_delete_entities_with_agent_id(self): - mg = _create_graph_memory() - mg._delete_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "user1", "agent_id": "agent1"}, - ) - - cypher = mg.graph.query.call_args[0][0] - params = mg.graph.query.call_args[1]["params"] - assert "SET r.valid = false" in cypher - assert "agent_id: $agent_id" in cypher - assert params["agent_id"] == "agent1" - - def test_delete_entities_with_run_id(self): - mg = _create_graph_memory() - mg._delete_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "user1", "run_id": "run1"}, - ) - - cypher = mg.graph.query.call_args[0][0] - params = mg.graph.query.call_args[1]["params"] - assert "SET r.valid = false" in cypher - assert "run_id: $run_id" in cypher - assert params["run_id"] == "run1" - - def test_get_all_with_agent_id_filters_soft_deleted(self): - mg = _create_graph_memory() - mg.get_all(filters={"user_id": "user1", "agent_id": "agent1"}, top_k=10) - - cypher = mg.graph.query.call_args[0][0] - assert "r.valid IS NULL OR r.valid = true" in cypher - assert "agent_id: $agent_id" in cypher - - def test_merge_with_agent_id_sets_valid_true(self): - mg = _create_graph_memory() - mg._search_source_node = Mock( - return_value=[{"elementId(source_candidate)": "id1"}] - ) - mg._search_destination_node = Mock(return_value=[]) - - mg._add_entities( - [{"source": "alice", "destination": "bob", "relationship": "KNOWS"}], - {"user_id": "user1", "agent_id": "agent1"}, - entity_type_map={}, - ) - - cypher = mg.graph.query.call_args[0][0] - assert "r.valid = true" in cypher - assert "agent_id: $agent_id" in cypher - - -class TestResetAndCleanup: - """Verify reset and delete_all use hard-delete (DETACH DELETE).""" - - def test_reset_uses_detach_delete(self): - mg = _create_graph_memory() - mg.reset() - - cypher = mg.graph.query.call_args[0][0] - assert "DETACH DELETE" in cypher - assert "valid" not in cypher.lower() - - def test_delete_all_does_not_soft_delete(self): - mg = _create_graph_memory() - mg.delete_all(filters={"user_id": "user1"}) - - cypher = mg.graph.query.call_args[0][0] - assert "DETACH DELETE" in cypher - assert "r.valid = false" not in cypher diff --git a/tests/memory/test_kuzu.py b/tests/memory/test_kuzu.py deleted file mode 100644 index 14ed1d175..000000000 --- a/tests/memory/test_kuzu.py +++ /dev/null @@ -1,253 +0,0 @@ -from unittest.mock import MagicMock, Mock, patch - -import numpy as np -import pytest - -from mem0.memory.kuzu_memory import MemoryGraph - - -class TestKuzu: - """Test that Kuzu memory works correctly""" - - # Create distinct embeddings that won't match with threshold=0.7 - # Each embedding is mostly zeros with ones in different positions to ensure low similarity - alice_emb = np.zeros(384) - alice_emb[0:96] = 1.0 - - bob_emb = np.zeros(384) - bob_emb[96:192] = 1.0 - - charlie_emb = np.zeros(384) - charlie_emb[192:288] = 1.0 - - dave_emb = np.zeros(384) - dave_emb[288:384] = 1.0 - - embeddings = { - "alice": alice_emb.tolist(), - "bob": bob_emb.tolist(), - "charlie": charlie_emb.tolist(), - "dave": dave_emb.tolist(), - } - - @pytest.fixture - def mock_config(self): - """Create a mock configuration for testing""" - config = Mock() - - # Mock embedder config - config.embedder.provider = "mock_embedder" - config.embedder.config = {"model": "mock_model"} - config.vector_store.config = {"dimensions": 384} - - # Mock graph store config - config.graph_store.config.db = ":memory:" - config.graph_store.threshold = 0.7 - - # Mock LLM config - config.llm.provider = "mock_llm" - config.llm.config = {"api_key": "test_key"} - - return config - - @pytest.fixture - def mock_embedding_model(self): - """Create a mock embedding model""" - mock_model = Mock() - mock_model.config.embedding_dims = 384 - - def mock_embed(text): - return self.embeddings[text] - - mock_model.embed.side_effect = mock_embed - return mock_model - - @pytest.fixture - def mock_llm(self): - """Create a mock LLM""" - mock_llm = Mock() - mock_llm.generate_response.return_value = { - "tool_calls": [ - { - "name": "extract_entities", - "arguments": {"entities": [{"entity": "test_entity", "entity_type": "test_type"}]}, - } - ] - } - return mock_llm - - @patch("mem0.memory.kuzu_memory.EmbedderFactory") - @patch("mem0.memory.kuzu_memory.LlmFactory") - def test_kuzu_memory_initialization( - self, mock_llm_factory, mock_embedder_factory, mock_config, mock_embedding_model, mock_llm - ): - """Test that Kuzu memory initializes correctly""" - # Setup mocks - mock_embedder_factory.create.return_value = mock_embedding_model - mock_llm_factory.create.return_value = mock_llm - - # Create instance - kuzu_memory = MemoryGraph(mock_config) - - # Verify initialization - assert kuzu_memory.config == mock_config - assert kuzu_memory.embedding_model == mock_embedding_model - assert kuzu_memory.embedding_dims == 384 - assert kuzu_memory.llm == mock_llm - assert kuzu_memory.threshold == 0.7 - - @pytest.mark.parametrize( - "embedding_dims", - [None, 0, -1], - ) - @patch("mem0.memory.kuzu_memory.EmbedderFactory") - def test_kuzu_memory_initialization_invalid_embedding_dims( - self, mock_embedder_factory, embedding_dims, mock_config - ): - """Test that Kuzu memory raises ValuError when initialized with invalid embedding_dims""" - # Setup mocks - mock_embedding_model = Mock() - mock_embedding_model.config.embedding_dims = embedding_dims - mock_embedder_factory.create.return_value = mock_embedding_model - - with pytest.raises(ValueError, match="must be a positive"): - MemoryGraph(mock_config) - - @patch("mem0.memory.kuzu_memory.EmbedderFactory") - @patch("mem0.memory.kuzu_memory.LlmFactory") - def test_kuzu(self, mock_llm_factory, mock_embedder_factory, mock_config, mock_embedding_model, mock_llm): - """Test adding memory to the graph""" - mock_embedder_factory.create.return_value = mock_embedding_model - mock_llm_factory.create.return_value = mock_llm - - kuzu_memory = MemoryGraph(mock_config) - - filters = {"user_id": "test_user", "agent_id": "test_agent", "run_id": "test_run"} - data1 = [ - {"source": "alice", "destination": "bob", "relationship": "knows"}, - {"source": "bob", "destination": "charlie", "relationship": "knows"}, - {"source": "charlie", "destination": "alice", "relationship": "knows"}, - ] - data2 = [ - {"source": "charlie", "destination": "alice", "relationship": "likes"}, - ] - - result = kuzu_memory._add_entities(data1, filters, {}) - assert result[0] == [{"source": "alice", "relationship": "knows", "target": "bob"}] - assert result[1] == [{"source": "bob", "relationship": "knows", "target": "charlie"}] - assert result[2] == [{"source": "charlie", "relationship": "knows", "target": "alice"}] - assert get_node_count(kuzu_memory) == 3 - assert get_edge_count(kuzu_memory) == 3 - - result = kuzu_memory._add_entities(data2, filters, {}) - assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}] - assert get_node_count(kuzu_memory) == 3 - assert get_edge_count(kuzu_memory) == 4 - - data3 = [ - {"source": "dave", "destination": "alice", "relationship": "admires"} - ] - result = kuzu_memory._add_entities(data3, filters, {}) - assert result[0] == [{"source": "dave", "relationship": "admires", "target": "alice"}] - assert get_node_count(kuzu_memory) == 4 # dave is new - assert get_edge_count(kuzu_memory) == 5 - - results = kuzu_memory.get_all(filters) - assert set([f"{result['source']}_{result['relationship']}_{result['target']}" for result in results]) == set([ - "alice_knows_bob", - "bob_knows_charlie", - "charlie_likes_alice", - "charlie_knows_alice", - "dave_admires_alice" - ]) - - results = kuzu_memory._search_graph_db(["bob"], filters, threshold=0.8) - assert set([f"{result['source']}_{result['relationship']}_{result['destination']}" for result in results]) == set([ - "alice_knows_bob", - "bob_knows_charlie", - ]) - - result = kuzu_memory._delete_entities(data2, filters) - assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}] - assert get_node_count(kuzu_memory) == 4 - assert get_edge_count(kuzu_memory) == 4 - - result = kuzu_memory._delete_entities(data1, filters) - assert result[0] == [{"source": "alice", "relationship": "knows", "target": "bob"}] - assert result[1] == [{"source": "bob", "relationship": "knows", "target": "charlie"}] - assert result[2] == [{"source": "charlie", "relationship": "knows", "target": "alice"}] - assert get_node_count(kuzu_memory) == 4 - assert get_edge_count(kuzu_memory) == 1 - - result = kuzu_memory.delete_all(filters) - assert get_node_count(kuzu_memory) == 0 - assert get_edge_count(kuzu_memory) == 0 - - result = kuzu_memory._add_entities(data2, filters, {}) - assert result[0] == [{"source": "charlie", "relationship": "likes", "target": "alice"}] - assert get_node_count(kuzu_memory) == 2 - assert get_edge_count(kuzu_memory) == 1 - - result = kuzu_memory.reset() - assert get_node_count(kuzu_memory) == 0 - assert get_edge_count(kuzu_memory) == 0 - -def _make_kuzu_instance(): - with patch.object(MemoryGraph, "__init__", return_value=None): - instance = MemoryGraph.__new__(MemoryGraph) - instance.llm_provider = "openai" - instance.llm = MagicMock() - instance.embedding_model = MagicMock() - instance.config = MagicMock() - instance.config.graph_store.custom_prompt = None - return instance - - -class TestRetrieveNodesFromData: - """Tests for _retrieve_nodes_from_data in KuzuMemoryGraph.""" - - def test_missing_entities_key_returns_empty(self): - """LLM returns extract_entities tool call without 'entities' key — should not crash. - Reproduces the exact scenario from issue #4238.""" - instance = _make_kuzu_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}] - } - result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"}) - assert result == {} - - def test_normal_entities_extracted(self): - instance = _make_kuzu_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Alice", "entity_type": "person"}, - {"entity": "hiking", "entity_type": "activity"}, - ]}}] - } - result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"}) - assert result == {"alice": "person", "hiking": "activity"} - - def test_none_tool_calls_returns_empty(self): - instance = _make_kuzu_instance() - instance.llm.generate_response.return_value = {"tool_calls": None} - result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"}) - assert result == {} - - -def get_node_count(kuzu_memory): - results = kuzu_memory.kuzu_execute( - """ - MATCH (n) - RETURN COUNT(n) as count - """ - ) - return int(results[0]['count']) - -def get_edge_count(kuzu_memory): - results = kuzu_memory.kuzu_execute( - """ - MATCH (n)-[e]->(m) - RETURN COUNT(e) as count - """ - ) - return int(results[0]['count']) diff --git a/tests/memory/test_memgraph_memory.py b/tests/memory/test_memgraph_memory.py deleted file mode 100644 index 58d6acce4..000000000 --- a/tests/memory/test_memgraph_memory.py +++ /dev/null @@ -1,107 +0,0 @@ -from unittest.mock import MagicMock, Mock, patch - -# langchain_memgraph and rank_bm25 are optional deps — mock them so tests run without install -_memgraph_mock = Mock() -patch.dict("sys.modules", { - "langchain_memgraph": _memgraph_mock, - "langchain_memgraph.graphs": _memgraph_mock, - "langchain_memgraph.graphs.memgraph": _memgraph_mock, - "rank_bm25": Mock(), -}).start() - -from mem0.memory.memgraph_memory import MemoryGraph as MemgraphMemoryGraph # noqa: E402 - -MemoryGraph = MemgraphMemoryGraph - - -def _make_instance(): - with patch.object(MemoryGraph, "__init__", return_value=None): - instance = MemoryGraph.__new__(MemoryGraph) - instance.llm_provider = "openai" - instance.llm = MagicMock() - instance.embedding_model = MagicMock() - instance.config = MagicMock() - instance.config.graph_store.custom_prompt = None - return instance - - -class TestRetrieveNodesFromData: - """Tests for _retrieve_nodes_from_data in MemoryGraph.""" - - def test_normal_entities_extracted(self): - instance = _make_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "Alice", "entity_type": "person"}, - {"entity": "hiking", "entity_type": "activity"}, - ]}}] - } - result = instance._retrieve_nodes_from_data("Alice loves hiking", {"user_id": "u1"}) - assert result == {"alice": "person", "hiking": "activity"} - - def test_malformed_entity_missing_entity_type_is_skipped(self): - """LLM returns entity dict without entity_type — should skip it, keep valid ones. - Reproduces the exact data from issue #4055.""" - instance = _make_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "extract_entities", "arguments": {"entities": [ - {"entity": "matrix multiplication", "entity_type": "task"}, - {"entity": "task"}, - {"entity": "ReLU", "entity_type": "task"}, - ]}}] - } - result = instance._retrieve_nodes_from_data("some text", {"user_id": "u1"}) - assert "matrix_multiplication" in result - assert "relu" in result - assert "task" not in result - - def test_missing_entities_key_returns_empty(self): - """LLM returns extract_entities tool call without 'entities' key — should not crash. - Reproduces the exact scenario from issue #4238.""" - instance = _make_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "extract_entities", "arguments": {"text": "Hello."}}] - } - result = instance._retrieve_nodes_from_data("Hello.", {"user_id": "u1"}) - assert result == {} - - def test_none_tool_calls_returns_empty(self): - instance = _make_instance() - instance.llm.generate_response.return_value = {"tool_calls": None} - result = instance._retrieve_nodes_from_data("hello world", {"user_id": "u1"}) - assert result == {} - - -class TestEstablishNodesRelationsFromData: - """Tests for _establish_nodes_relations_from_data in MemoryGraph.""" - - def test_none_response_does_not_crash(self): - """openai_structured returns None when no relations found — must not crash. - Exact crash from issue #4055: TypeError: 'NoneType' object is not subscriptable.""" - instance = _make_instance() - instance.llm.generate_response.return_value = None - result = instance._establish_nodes_relations_from_data( - "Hello world", {"user_id": "u1"}, {} - ) - assert result == [] - - def test_empty_tool_calls_returns_empty(self): - instance = _make_instance() - instance.llm.generate_response.return_value = {"tool_calls": []} - result = instance._establish_nodes_relations_from_data( - "Hello world", {"user_id": "u1"}, {} - ) - assert result == [] - - def test_valid_entities_returned(self): - instance = _make_instance() - instance.llm.generate_response.return_value = { - "tool_calls": [{"name": "add_entities", "arguments": {"entities": [ - {"source": "alice", "relationship": "loves", "destination": "hiking"} - ]}}] - } - result = instance._establish_nodes_relations_from_data( - "Alice loves hiking", {"user_id": "u1"}, {"alice": "person"} - ) - assert len(result) == 1 - assert result[0]["source"] == "alice" diff --git a/tests/memory/test_neptune_analytics_memory.py b/tests/memory/test_neptune_analytics_memory.py deleted file mode 100644 index dbec73fd8..000000000 --- a/tests/memory/test_neptune_analytics_memory.py +++ /dev/null @@ -1,338 +0,0 @@ -import unittest -from unittest.mock import MagicMock, patch - -import pytest - -from mem0.graphs.neptune.base import NeptuneBase -from mem0.graphs.neptune.neptunegraph import MemoryGraph - - -class TestNeptuneMemory(unittest.TestCase): - """Test suite for the Neptune Memory implementation.""" - - def setUp(self): - """Set up test fixtures before each test method.""" - - # Create a mock config - self.config = MagicMock() - self.config.graph_store.config.endpoint = "neptune-graph://test-graph" - self.config.graph_store.config.base_label = True - self.config.graph_store.threshold = 0.7 - self.config.llm.provider = "openai_structured" - self.config.graph_store.llm = None - self.config.graph_store.custom_prompt = None - - # Create mock for NeptuneAnalyticsGraph - self.mock_graph = MagicMock() - self.mock_graph.client.get_graph.return_value = {"status": "AVAILABLE"} - - # Create mocks for static methods - self.mock_embedding_model = MagicMock() - self.mock_llm = MagicMock() - - # Patch the necessary components - self.neptune_analytics_graph_patcher = patch("mem0.graphs.neptune.neptunegraph.NeptuneAnalyticsGraph") - self.mock_neptune_analytics_graph = self.neptune_analytics_graph_patcher.start() - self.mock_neptune_analytics_graph.return_value = self.mock_graph - - # Patch the static methods - self.create_embedding_model_patcher = patch.object(NeptuneBase, "_create_embedding_model") - self.mock_create_embedding_model = self.create_embedding_model_patcher.start() - self.mock_create_embedding_model.return_value = self.mock_embedding_model - - self.create_llm_patcher = patch.object(NeptuneBase, "_create_llm") - self.mock_create_llm = self.create_llm_patcher.start() - self.mock_create_llm.return_value = self.mock_llm - - # Create the MemoryGraph instance - self.memory_graph = MemoryGraph(self.config) - - # Set up common test data - self.user_id = "test_user" - self.test_filters = {"user_id": self.user_id} - - def tearDown(self): - """Tear down test fixtures after each test method.""" - self.neptune_analytics_graph_patcher.stop() - self.create_embedding_model_patcher.stop() - self.create_llm_patcher.stop() - - def test_initialization(self): - """Test that the MemoryGraph is initialized correctly.""" - self.assertEqual(self.memory_graph.graph, self.mock_graph) - self.assertEqual(self.memory_graph.embedding_model, self.mock_embedding_model) - self.assertEqual(self.memory_graph.llm, self.mock_llm) - self.assertEqual(self.memory_graph.llm_provider, "openai_structured") - self.assertEqual(self.memory_graph.node_label, ":`__Entity__`") - self.assertEqual(self.memory_graph.threshold, 0.7) - - def test_init(self): - """Test the class init functions""" - - # Create a mock config with bad endpoint - config_no_endpoint = MagicMock() - config_no_endpoint.graph_store.config.endpoint = None - - # Create the MemoryGraph instance - with pytest.raises(ValueError): - MemoryGraph(config_no_endpoint) - - # Create a mock config with bad endpoint - config_ndb_endpoint = MagicMock() - config_ndb_endpoint.graph_store.config.endpoint = "neptune-db://test-graph" - - with pytest.raises(ValueError): - MemoryGraph(config_ndb_endpoint) - - def test_add_method(self): - """Test the add method with mocked components.""" - - # Mock the necessary methods that add() calls - self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person", "bob": "person"}) - self.memory_graph._establish_nodes_relations_from_data = MagicMock( - return_value=[{"source": "alice", "relationship": "knows", "destination": "bob"}] - ) - self.memory_graph._search_graph_db = MagicMock(return_value=[]) - self.memory_graph._get_delete_entities_from_search_output = MagicMock(return_value=[]) - self.memory_graph._delete_entities = MagicMock(return_value=[]) - self.memory_graph._add_entities = MagicMock( - return_value=[{"source": "alice", "relationship": "knows", "target": "bob"}] - ) - - # Call the add method - result = self.memory_graph.add("Alice knows Bob", self.test_filters) - - # Verify the method calls - self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Alice knows Bob", self.test_filters) - self.memory_graph._establish_nodes_relations_from_data.assert_called_once() - self.memory_graph._search_graph_db.assert_called_once() - self.memory_graph._get_delete_entities_from_search_output.assert_called_once() - self.memory_graph._delete_entities.assert_called_once_with([], self.user_id) - self.memory_graph._add_entities.assert_called_once() - - # Check the result structure - self.assertIn("deleted_entities", result) - self.assertIn("added_entities", result) - - def test_search_method(self): - """Test the search method with mocked components.""" - # Mock the necessary methods that search() calls - self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"}) - - # Mock search results - mock_search_results = [ - {"source": "alice", "relationship": "knows", "destination": "bob"}, - {"source": "alice", "relationship": "works_with", "destination": "charlie"}, - ] - self.memory_graph._search_graph_db = MagicMock(return_value=mock_search_results) - - # Mock BM25Okapi - with patch("mem0.graphs.neptune.base.BM25Okapi") as mock_bm25: - mock_bm25_instance = MagicMock() - mock_bm25.return_value = mock_bm25_instance - - # Mock get_top_n to return reranked results - reranked_results = [["alice", "knows", "bob"], ["alice", "works_with", "charlie"]] - mock_bm25_instance.get_top_n.return_value = reranked_results - - # Call the search method - result = self.memory_graph.search("Find Alice", self.test_filters, top_k=5) - - # Verify the method calls - self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Find Alice", self.test_filters) - self.memory_graph._search_graph_db.assert_called_once_with(node_list=["alice"], filters=self.test_filters) - - # Check the result structure - self.assertEqual(len(result), 2) - self.assertEqual(result[0]["source"], "alice") - self.assertEqual(result[0]["relationship"], "knows") - self.assertEqual(result[0]["destination"], "bob") - - def test_get_all_method(self): - """Test the get_all method.""" - - # Mock the _get_all_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"user_id": self.user_id, "limit": 10} - self.memory_graph._get_all_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [ - {"source": "alice", "relationship": "knows", "target": "bob"}, - {"source": "bob", "relationship": "works_with", "target": "charlie"}, - ] - self.mock_graph.query.return_value = mock_query_result - - # Call the get_all method - result = self.memory_graph.get_all(self.test_filters, top_k=10) - - # Verify the method calls - self.memory_graph._get_all_cypher.assert_called_once_with(self.test_filters, 10) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result structure - self.assertEqual(len(result), 2) - self.assertEqual(result[0]["source"], "alice") - self.assertEqual(result[0]["relationship"], "knows") - self.assertEqual(result[0]["target"], "bob") - - def test_delete_all_method(self): - """Test the delete_all method.""" - # Mock the _delete_all_cypher method - mock_cypher = "MATCH (n) DETACH DELETE n" - mock_params = {"user_id": self.user_id} - self.memory_graph._delete_all_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Call the delete_all method - self.memory_graph.delete_all(self.test_filters) - - # Verify the method calls - self.memory_graph._delete_all_cypher.assert_called_once_with(self.test_filters) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - def test_search_source_node(self): - """Test the _search_source_node method.""" - # Mock embedding - mock_embedding = [0.1, 0.2, 0.3] - - # Mock the _search_source_node_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"source_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9} - self.memory_graph._search_source_node_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}] - self.mock_graph.query.return_value = mock_query_result - - # Call the _search_source_node method - result = self.memory_graph._search_source_node(mock_embedding, self.user_id, threshold=0.9) - - # Verify the method calls - self.memory_graph._search_source_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result - self.assertEqual(result, mock_query_result) - - def test_search_destination_node(self): - """Test the _search_destination_node method.""" - # Mock embedding - mock_embedding = [0.1, 0.2, 0.3] - - # Mock the _search_destination_node_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"destination_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9} - self.memory_graph._search_destination_node_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}] - self.mock_graph.query.return_value = mock_query_result - - # Call the _search_destination_node method - result = self.memory_graph._search_destination_node(mock_embedding, self.user_id, threshold=0.9) - - # Verify the method calls - self.memory_graph._search_destination_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result - self.assertEqual(result, mock_query_result) - - def test_search_graph_db(self): - """Test the _search_graph_db method.""" - # Mock node list - node_list = ["alice", "bob"] - - # Mock embedding - mock_embedding = [0.1, 0.2, 0.3] - self.mock_embedding_model.embed.return_value = mock_embedding - - # Mock the _search_graph_db_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"n_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.7, "limit": 10} - self.memory_graph._search_graph_db_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query results - mock_query_result1 = [{"source": "alice", "relationship": "knows", "destination": "bob"}] - mock_query_result2 = [{"source": "bob", "relationship": "works_with", "destination": "charlie"}] - self.mock_graph.query.side_effect = [mock_query_result1, mock_query_result2] - - # Call the _search_graph_db method - result = self.memory_graph._search_graph_db(node_list, self.test_filters, top_k=10) - - # Verify the method calls - self.assertEqual(self.mock_embedding_model.embed.call_count, 2) - self.assertEqual(self.memory_graph._search_graph_db_cypher.call_count, 2) - self.assertEqual(self.mock_graph.query.call_count, 2) - - # Check the result - expected_result = mock_query_result1 + mock_query_result2 - self.assertEqual(result, expected_result) - - def test_add_entities(self): - """Test the _add_entities method.""" - # Mock data - to_be_added = [{"source": "alice", "relationship": "knows", "destination": "bob"}] - entity_type_map = {"alice": "person", "bob": "person"} - - # Mock embeddings - mock_embedding = [0.1, 0.2, 0.3] - self.mock_embedding_model.embed.return_value = mock_embedding - - # Mock search results - mock_source_search = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}] - mock_dest_search = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}] - - # Mock the search methods - self.memory_graph._search_source_node = MagicMock(return_value=mock_source_search) - self.memory_graph._search_destination_node = MagicMock(return_value=mock_dest_search) - - # Mock the _add_entities_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"source_id": 123, "destination_id": 456} - self.memory_graph._add_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}] - self.mock_graph.query.return_value = mock_query_result - - # Call the _add_entities method - result = self.memory_graph._add_entities(to_be_added, self.user_id, entity_type_map) - - # Verify the method calls - self.assertEqual(self.mock_embedding_model.embed.call_count, 2) - self.memory_graph._search_source_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.7) - self.memory_graph._search_destination_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.7) - self.memory_graph._add_entities_cypher.assert_called_once() - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result - self.assertEqual(result, [mock_query_result]) - - def test_delete_entities(self): - """Test the _delete_entities method.""" - # Mock data - to_be_deleted = [{"source": "alice", "relationship": "knows", "destination": "bob"}] - - # Mock the _delete_entities_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"source_name": "alice", "dest_name": "bob", "user_id": self.user_id} - self.memory_graph._delete_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}] - self.mock_graph.query.return_value = mock_query_result - - # Call the _delete_entities method - result = self.memory_graph._delete_entities(to_be_deleted, self.user_id) - - # Verify the method calls - self.memory_graph._delete_entities_cypher.assert_called_once_with("alice", "bob", "knows", self.user_id) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result - self.assertEqual(result, [mock_query_result]) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/memory/test_neptune_memory.py b/tests/memory/test_neptune_memory.py deleted file mode 100644 index ad28b6e34..000000000 --- a/tests/memory/test_neptune_memory.py +++ /dev/null @@ -1,411 +0,0 @@ -import unittest -from datetime import datetime, timezone -from unittest.mock import MagicMock, patch - -import pytest - -from mem0.graphs.neptune.base import NeptuneBase -from mem0.graphs.neptune.neptunedb import MemoryGraph - - -class TestNeptuneMemory(unittest.TestCase): - """Test suite for the Neptune Memory implementation.""" - - def setUp(self): - """Set up test fixtures before each test method.""" - - # Create a mock config - self.config = MagicMock() - self.config.graph_store.config.endpoint = "neptune-db://test-graph" - self.config.graph_store.config.base_label = True - self.config.graph_store.threshold = 0.7 - self.config.llm.provider = "openai_structured" - self.config.graph_store.llm = None - self.config.graph_store.custom_prompt = None - self.config.vector_store.provider = "qdrant" - self.config.vector_store.config = MagicMock() - - # Create mock for NeptuneGraph - self.mock_graph = MagicMock() - - # Create mocks for static methods - self.mock_embedding_model = MagicMock() - self.mock_llm = MagicMock() - self.mock_vector_store = MagicMock() - - # Patch the necessary components - self.neptune_graph_patcher = patch("mem0.graphs.neptune.neptunedb.NeptuneGraph") - self.mock_neptune_graph = self.neptune_graph_patcher.start() - self.mock_neptune_graph.return_value = self.mock_graph - - # Patch the static methods - self.create_embedding_model_patcher = patch.object(NeptuneBase, "_create_embedding_model") - self.mock_create_embedding_model = self.create_embedding_model_patcher.start() - self.mock_create_embedding_model.return_value = self.mock_embedding_model - - self.create_llm_patcher = patch.object(NeptuneBase, "_create_llm") - self.mock_create_llm = self.create_llm_patcher.start() - self.mock_create_llm.return_value = self.mock_llm - - self.create_vector_store_patcher = patch.object(NeptuneBase, "_create_vector_store") - self.mock_create_vector_store = self.create_vector_store_patcher.start() - self.mock_create_vector_store.return_value = self.mock_vector_store - - # Create the MemoryGraph instance - self.memory_graph = MemoryGraph(self.config) - - # Set up common test data - self.user_id = "test_user" - self.test_filters = {"user_id": self.user_id} - - def tearDown(self): - """Tear down test fixtures after each test method.""" - self.neptune_graph_patcher.stop() - self.create_embedding_model_patcher.stop() - self.create_llm_patcher.stop() - self.create_vector_store_patcher.stop() - - def test_initialization(self): - """Test that the MemoryGraph is initialized correctly.""" - self.assertEqual(self.memory_graph.graph, self.mock_graph) - self.assertEqual(self.memory_graph.embedding_model, self.mock_embedding_model) - self.assertEqual(self.memory_graph.llm, self.mock_llm) - self.assertEqual(self.memory_graph.vector_store, self.mock_vector_store) - self.assertEqual(self.memory_graph.llm_provider, "openai_structured") - self.assertEqual(self.memory_graph.node_label, ":`__Entity__`") - self.assertEqual(self.memory_graph.threshold, 0.7) - self.assertEqual(self.memory_graph.vector_store_limit, 5) - - def test_collection_name_variants(self): - """Test all collection_name configuration variants.""" - - # Test 1: graph_store.config.collection_name is set - config1 = MagicMock() - config1.graph_store.config.endpoint = "neptune-db://test-graph" - config1.graph_store.config.base_label = True - config1.graph_store.config.collection_name = "custom_collection" - config1.llm.provider = "openai" - config1.graph_store.llm = None - config1.vector_store.provider = "qdrant" - config1.vector_store.config = MagicMock() - - MemoryGraph(config1) - self.assertEqual(config1.vector_store.config.collection_name, "custom_collection") - - # Test 2: vector_store.config.collection_name exists, graph_store.config.collection_name is None - config2 = MagicMock() - config2.graph_store.config.endpoint = "neptune-db://test-graph" - config2.graph_store.config.base_label = True - config2.graph_store.config.collection_name = None - config2.llm.provider = "openai" - config2.graph_store.llm = None - config2.vector_store.provider = "qdrant" - config2.vector_store.config = MagicMock() - config2.vector_store.config.collection_name = "existing_collection" - - MemoryGraph(config2) - self.assertEqual(config2.vector_store.config.collection_name, "existing_collection_neptune_vector_store") - - # Test 3: Neither collection_name is set (default case) - config3 = MagicMock() - config3.graph_store.config.endpoint = "neptune-db://test-graph" - config3.graph_store.config.base_label = True - config3.graph_store.config.collection_name = None - config3.llm.provider = "openai" - config3.graph_store.llm = None - config3.vector_store.provider = "qdrant" - config3.vector_store.config = MagicMock() - config3.vector_store.config.collection_name = None - - MemoryGraph(config3) - self.assertEqual(config3.vector_store.config.collection_name, "mem0_neptune_vector_store") - - def test_init(self): - """Test the class init functions""" - - # Create a mock config with bad endpoint - config_no_endpoint = MagicMock() - config_no_endpoint.graph_store.config.endpoint = None - - # Create the MemoryGraph instance - with pytest.raises(ValueError): - MemoryGraph(config_no_endpoint) - - # Create a mock config with wrong endpoint type - config_wrong_endpoint = MagicMock() - config_wrong_endpoint.graph_store.config.endpoint = "neptune-graph://test-graph" - - with pytest.raises(ValueError): - MemoryGraph(config_wrong_endpoint) - - def test_add_method(self): - """Test the add method with mocked components.""" - - # Mock the necessary methods that add() calls - self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person", "bob": "person"}) - self.memory_graph._establish_nodes_relations_from_data = MagicMock( - return_value=[{"source": "alice", "relationship": "knows", "destination": "bob"}] - ) - self.memory_graph._search_graph_db = MagicMock(return_value=[]) - self.memory_graph._get_delete_entities_from_search_output = MagicMock(return_value=[]) - self.memory_graph._delete_entities = MagicMock(return_value=[]) - self.memory_graph._add_entities = MagicMock( - return_value=[{"source": "alice", "relationship": "knows", "target": "bob"}] - ) - - # Call the add method - result = self.memory_graph.add("Alice knows Bob", self.test_filters) - - # Verify the method calls - self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Alice knows Bob", self.test_filters) - self.memory_graph._establish_nodes_relations_from_data.assert_called_once() - self.memory_graph._search_graph_db.assert_called_once() - self.memory_graph._get_delete_entities_from_search_output.assert_called_once() - self.memory_graph._delete_entities.assert_called_once_with([], self.user_id) - self.memory_graph._add_entities.assert_called_once() - - # Check the result structure - self.assertIn("deleted_entities", result) - self.assertIn("added_entities", result) - - def test_search_method(self): - """Test the search method with mocked components.""" - # Mock the necessary methods that search() calls - self.memory_graph._retrieve_nodes_from_data = MagicMock(return_value={"alice": "person"}) - - # Mock search results - mock_search_results = [ - {"source": "alice", "relationship": "knows", "destination": "bob"}, - {"source": "alice", "relationship": "works_with", "destination": "charlie"}, - ] - self.memory_graph._search_graph_db = MagicMock(return_value=mock_search_results) - - # Mock BM25Okapi - with patch("mem0.graphs.neptune.base.BM25Okapi") as mock_bm25: - mock_bm25_instance = MagicMock() - mock_bm25.return_value = mock_bm25_instance - - # Mock get_top_n to return reranked results - reranked_results = [["alice", "knows", "bob"], ["alice", "works_with", "charlie"]] - mock_bm25_instance.get_top_n.return_value = reranked_results - - # Call the search method - result = self.memory_graph.search("Find Alice", self.test_filters, top_k=5) - - # Verify the method calls - self.memory_graph._retrieve_nodes_from_data.assert_called_once_with("Find Alice", self.test_filters) - self.memory_graph._search_graph_db.assert_called_once_with(node_list=["alice"], filters=self.test_filters) - - # Check the result structure - self.assertEqual(len(result), 2) - self.assertEqual(result[0]["source"], "alice") - self.assertEqual(result[0]["relationship"], "knows") - self.assertEqual(result[0]["destination"], "bob") - - def test_get_all_method(self): - """Test the get_all method.""" - - # Mock the _get_all_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"user_id": self.user_id, "limit": 10} - self.memory_graph._get_all_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [ - {"source": "alice", "relationship": "knows", "target": "bob"}, - {"source": "bob", "relationship": "works_with", "target": "charlie"}, - ] - self.mock_graph.query.return_value = mock_query_result - - # Call the get_all method - result = self.memory_graph.get_all(self.test_filters, top_k=10) - - # Verify the method calls - self.memory_graph._get_all_cypher.assert_called_once_with(self.test_filters, 10) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result structure - self.assertEqual(len(result), 2) - self.assertEqual(result[0]["source"], "alice") - self.assertEqual(result[0]["relationship"], "knows") - self.assertEqual(result[0]["target"], "bob") - - def test_delete_all_method(self): - """Test the delete_all method.""" - # Mock the _delete_all_cypher method - mock_cypher = "MATCH (n) DETACH DELETE n" - mock_params = {"user_id": self.user_id} - self.memory_graph._delete_all_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Call the delete_all method - self.memory_graph.delete_all(self.test_filters) - - # Verify the method calls - self.memory_graph._delete_all_cypher.assert_called_once_with(self.test_filters) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - def test_search_source_node(self): - """Test the _search_source_node method.""" - # Mock embedding - mock_embedding = [0.1, 0.2, 0.3] - - # Mock the _search_source_node_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"source_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9} - self.memory_graph._search_source_node_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}] - self.mock_graph.query.return_value = mock_query_result - - # Call the _search_source_node method - result = self.memory_graph._search_source_node(mock_embedding, self.user_id, threshold=0.9) - - # Verify the method calls - self.memory_graph._search_source_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result - self.assertEqual(result, mock_query_result) - - def test_search_destination_node(self): - """Test the _search_destination_node method.""" - # Mock embedding - mock_embedding = [0.1, 0.2, 0.3] - - # Mock the _search_destination_node_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"destination_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.9} - self.memory_graph._search_destination_node_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}] - self.mock_graph.query.return_value = mock_query_result - - # Call the _search_destination_node method - result = self.memory_graph._search_destination_node(mock_embedding, self.user_id, threshold=0.9) - - # Verify the method calls - self.memory_graph._search_destination_node_cypher.assert_called_once_with(mock_embedding, self.user_id, 0.9) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result - self.assertEqual(result, mock_query_result) - - def test_add_new_entities_payloads_use_utc_timestamps(self): - """Test that Neptune vector-store payloads use UTC timestamps.""" - self.memory_graph._add_new_entities_cypher( - source="alice", - source_embedding=[0.1, 0.2], - source_type="person", - destination="bob", - dest_embedding=[0.3, 0.4], - destination_type="person", - relationship="KNOWS", - user_id=self.user_id, - ) - - _, kwargs = self.mock_vector_store.insert.call_args - for payload in kwargs["payloads"]: - parsed = datetime.fromisoformat(payload["created_at"]) - self.assertEqual(parsed.tzinfo, timezone.utc) - self.assertEqual(parsed.utcoffset().total_seconds(), 0) - - def test_search_graph_db(self): - """Test the _search_graph_db method.""" - # Mock node list - node_list = ["alice", "bob"] - - # Mock embedding - mock_embedding = [0.1, 0.2, 0.3] - self.mock_embedding_model.embed.return_value = mock_embedding - - # Mock the _search_graph_db_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"n_embedding": mock_embedding, "user_id": self.user_id, "threshold": 0.7, "limit": 10} - self.memory_graph._search_graph_db_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query results - mock_query_result1 = [{"source": "alice", "relationship": "knows", "destination": "bob"}] - mock_query_result2 = [{"source": "bob", "relationship": "works_with", "destination": "charlie"}] - self.mock_graph.query.side_effect = [mock_query_result1, mock_query_result2] - - # Call the _search_graph_db method - result = self.memory_graph._search_graph_db(node_list, self.test_filters, top_k=10) - - # Verify the method calls - self.assertEqual(self.mock_embedding_model.embed.call_count, 2) - self.assertEqual(self.memory_graph._search_graph_db_cypher.call_count, 2) - self.assertEqual(self.mock_graph.query.call_count, 2) - - # Check the result - expected_result = mock_query_result1 + mock_query_result2 - self.assertEqual(result, expected_result) - - def test_add_entities(self): - """Test the _add_entities method.""" - # Mock data - to_be_added = [{"source": "alice", "relationship": "knows", "destination": "bob"}] - entity_type_map = {"alice": "person", "bob": "person"} - - # Mock embeddings - mock_embedding = [0.1, 0.2, 0.3] - self.mock_embedding_model.embed.return_value = mock_embedding - - # Mock search results - mock_source_search = [{"id(source_candidate)": 123, "cosine_similarity": 0.95}] - mock_dest_search = [{"id(destination_candidate)": 456, "cosine_similarity": 0.92}] - - # Mock the search methods - self.memory_graph._search_source_node = MagicMock(return_value=mock_source_search) - self.memory_graph._search_destination_node = MagicMock(return_value=mock_dest_search) - - # Mock the _add_entities_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"source_id": 123, "destination_id": 456} - self.memory_graph._add_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}] - self.mock_graph.query.return_value = mock_query_result - - # Call the _add_entities method - result = self.memory_graph._add_entities(to_be_added, self.user_id, entity_type_map) - - # Verify the method calls - self.assertEqual(self.mock_embedding_model.embed.call_count, 2) - self.memory_graph._search_source_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.7) - self.memory_graph._search_destination_node.assert_called_once_with(mock_embedding, self.user_id, threshold=0.7) - self.memory_graph._add_entities_cypher.assert_called_once() - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result - self.assertEqual(result, [mock_query_result]) - - def test_delete_entities(self): - """Test the _delete_entities method.""" - # Mock data - to_be_deleted = [{"source": "alice", "relationship": "knows", "destination": "bob"}] - - # Mock the _delete_entities_cypher method - mock_cypher = "MATCH (n) RETURN n" - mock_params = {"source_name": "alice", "dest_name": "bob", "user_id": self.user_id} - self.memory_graph._delete_entities_cypher = MagicMock(return_value=(mock_cypher, mock_params)) - - # Mock the graph.query result - mock_query_result = [{"source": "alice", "relationship": "knows", "target": "bob"}] - self.mock_graph.query.return_value = mock_query_result - - # Call the _delete_entities method - result = self.memory_graph._delete_entities(to_be_deleted, self.user_id) - - # Verify the method calls - self.memory_graph._delete_entities_cypher.assert_called_once_with("alice", "bob", "knows", self.user_id) - self.mock_graph.query.assert_called_once_with(mock_cypher, params=mock_params) - - # Check the result - self.assertEqual(result, [mock_query_result]) - - -if __name__ == "__main__": - unittest.main()