From 26732771eba1c318f7e18482ff8b5d4e434d2a8f Mon Sep 17 00:00:00 2001 From: DrJsPBs Date: Fri, 8 Aug 2025 12:08:13 -0400 Subject: [PATCH] Fix Neo4j Cypher syntax error with agent_id filtering (#3158) Co-authored-by: parshvadaftari --- mem0/memory/graph_memory.py | 144 ++++++++--- tests/memory/test_neo4j_cypher_syntax.py | 292 +++++++++++++++++++++++ tests/test_main.py | 10 + 3 files changed, 407 insertions(+), 39 deletions(-) create mode 100644 tests/memory/test_neo4j_cypher_syntax.py diff --git a/mem0/memory/graph_memory.py b/mem0/memory/graph_memory.py index f68a6ef66..913d9ef0e 100644 --- a/mem0/memory/graph_memory.py +++ b/mem0/memory/graph_memory.py @@ -122,18 +122,23 @@ class MemoryGraph: return search_results def delete_all(self, filters): + # Build node properties for filtering + node_props = ["user_id: $user_id"] if filters.get("agent_id"): - cypher = f""" - MATCH (n {self.node_label} {{user_id: $user_id, agent_id: $agent_id}}) - DETACH DELETE n - """ - params = {"user_id": filters["user_id"], "agent_id": filters["agent_id"]} - else: - cypher = f""" - MATCH (n {self.node_label} {{user_id: $user_id}}) - DETACH DELETE n - """ - params = {"user_id": filters["user_id"]} + node_props.append("agent_id: $agent_id") + if filters.get("run_id"): + node_props.append("run_id: $run_id") + node_props_str = ", ".join(node_props) + + cypher = f""" + MATCH (n {self.node_label} {{{node_props_str}}}) + DETACH DELETE n + """ + params = {"user_id": filters["user_id"]} + if filters.get("agent_id"): + params["agent_id"] = filters["agent_id"] + if filters.get("run_id"): + params["run_id"] = filters["run_id"] self.graph.query(cypher, params=params) def get_all(self, filters, limit=100): @@ -147,15 +152,20 @@ class MemoryGraph: - 'contexts': The base data store response for each memory. - 'entities': A list of strings representing the nodes and relationships """ - agent_filter = "" params = {"user_id": filters["user_id"], "limit": limit} + + # Build node properties based on filters + node_props = ["user_id: $user_id"] if filters.get("agent_id"): - agent_filter = "AND n.agent_id = $agent_id AND m.agent_id = $agent_id" + node_props.append("agent_id: $agent_id") params["agent_id"] = filters["agent_id"] + if filters.get("run_id"): + node_props.append("run_id: $run_id") + params["run_id"] = filters["run_id"] + node_props_str = ", ".join(node_props) query = f""" - MATCH (n {self.node_label} {{user_id: $user_id}})-[r]->(m {self.node_label} {{user_id: $user_id}}) - WHERE 1=1 {agent_filter} + MATCH (n {self.node_label} {{{node_props_str}}})-[r]->(m {self.node_label} {{{node_props_str}}}) RETURN n.name AS source, type(r) AS relationship, m.name AS target LIMIT $limit """ @@ -215,6 +225,8 @@ class MemoryGraph: user_identity = f"user_id: {filters['user_id']}" if filters.get("agent_id"): user_identity += f", agent_id: {filters['agent_id']}" + if filters.get("run_id"): + user_identity += f", run_id: {filters['run_id']}" if self.config.graph_store.custom_prompt: system_content = EXTRACT_RELATIONS_PROMPT.replace("USER_ID", user_identity) @@ -251,26 +263,30 @@ class MemoryGraph: def _search_graph_db(self, node_list, filters, limit=100): """Search similar nodes among and their respective incoming and outgoing relations.""" result_relations = [] - agent_filter = "" + + # Build node properties for filtering + node_props = ["user_id: $user_id"] if filters.get("agent_id"): - agent_filter = "AND n.agent_id = $agent_id AND m.agent_id = $agent_id" + node_props.append("agent_id: $agent_id") + if filters.get("run_id"): + node_props.append("run_id: $run_id") + node_props_str = ", ".join(node_props) for node in node_list: n_embedding = self.embedding_model.embed(node) cypher_query = f""" - MATCH (n {self.node_label}) - WHERE n.embedding IS NOT NULL AND n.user_id = $user_id - {agent_filter} + MATCH (n {self.node_label} {{{node_props_str}}}) + WHERE n.embedding IS NOT NULL WITH n, round(2 * vector.similarity.cosine(n.embedding, $n_embedding) - 1, 4) AS similarity // denormalize for backward compatibility WHERE similarity >= $threshold CALL {{ - MATCH (n)-[r]->(m) - WHERE m.user_id = $user_id {agent_filter.replace("n.", "m.")} + WITH n + MATCH (n)-[r]->(m {self.node_label} {{{node_props_str}}}) RETURN n.name AS source, elementId(n) AS source_id, type(r) AS relationship, elementId(r) AS relation_id, m.name AS destination, elementId(m) AS destination_id UNION - MATCH (m)-[r]->(n) - WHERE m.user_id = $user_id {agent_filter.replace("n.", "m.")} + WITH n + MATCH (n)<-[r]-(m {self.node_label} {{{node_props_str}}}) RETURN m.name AS source, elementId(m) AS source_id, type(r) AS relationship, elementId(r) AS relation_id, n.name AS destination, elementId(n) AS destination_id }} WITH distinct source, source_id, relationship, relation_id, destination, destination_id, similarity @@ -287,6 +303,8 @@ class MemoryGraph: } if filters.get("agent_id"): params["agent_id"] = filters["agent_id"] + if filters.get("run_id"): + params["run_id"] = filters["run_id"] ans = self.graph.query(cypher_query, params=params) result_relations.extend(ans) @@ -301,6 +319,8 @@ class MemoryGraph: user_identity = f"user_id: {filters['user_id']}" if filters.get("agent_id"): user_identity += f", agent_id: {filters['agent_id']}" + if filters.get("run_id"): + user_identity += f", run_id: {filters['run_id']}" system_prompt, user_prompt = get_delete_messages(search_output_string, data, user_identity) @@ -331,6 +351,7 @@ class MemoryGraph: """Delete the entities from the graph.""" user_id = filters["user_id"] agent_id = filters.get("agent_id", None) + run_id = filters.get("run_id", None) results = [] for item in to_be_deleted: @@ -339,7 +360,7 @@ class MemoryGraph: relationship = item["relationship"] # Build the agent filter for the query - agent_filter = "" + params = { "source_name": source, "dest_name": destination, @@ -347,15 +368,28 @@ class MemoryGraph: } if agent_id: - agent_filter = "AND n.agent_id = $agent_id AND m.agent_id = $agent_id" params["agent_id"] = agent_id + if run_id: + params["run_id"] = run_id + # Build node properties for filtering + source_props = ["name: $source_name", "user_id: $user_id"] + dest_props = ["name: $dest_name", "user_id: $user_id"] + if agent_id: + source_props.append("agent_id: $agent_id") + dest_props.append("agent_id: $agent_id") + if run_id: + source_props.append("run_id: $run_id") + dest_props.append("run_id: $run_id") + source_props_str = ", ".join(source_props) + dest_props_str = ", ".join(dest_props) + # Delete the specific relationship between nodes cypher = f""" - MATCH (n {self.node_label} {{name: $source_name, user_id: $user_id}}) + MATCH (n {self.node_label} {{{source_props_str}}}) -[r:{relationship}]-> - (m {self.node_label} {{name: $dest_name, user_id: $user_id}}) - WHERE 1=1 {agent_filter} + (m {self.node_label} {{{dest_props_str}}}) + DELETE r RETURN n.name AS source, @@ -372,6 +406,7 @@ class MemoryGraph: """Add the new entities to the graph. Merge the nodes if they already exist.""" user_id = filters["user_id"] agent_id = filters.get("agent_id", None) + run_id = filters.get("run_id", None) results = [] for item in to_be_added: # entities @@ -401,6 +436,8 @@ class MemoryGraph: merge_props = ["name: $destination_name", "user_id: $user_id"] if agent_id: merge_props.append("agent_id: $agent_id") + if run_id: + merge_props.append("run_id: $run_id") merge_props_str = ", ".join(merge_props) cypher = f""" @@ -435,12 +472,16 @@ class MemoryGraph: } if agent_id: params["agent_id"] = agent_id + if run_id: + params["run_id"] = run_id elif destination_node_search_result and not source_node_search_result: # Build source MERGE properties merge_props = ["name: $source_name", "user_id: $user_id"] if agent_id: merge_props.append("agent_id: $agent_id") + if run_id: + merge_props.append("run_id: $run_id") merge_props_str = ", ".join(merge_props) cypher = f""" @@ -475,6 +516,8 @@ class MemoryGraph: } if agent_id: params["agent_id"] = agent_id + if run_id: + params["run_id"] = run_id elif source_node_search_result and destination_node_search_result: cypher = f""" @@ -501,6 +544,8 @@ class MemoryGraph: } if agent_id: params["agent_id"] = agent_id + if run_id: + params["run_id"] = run_id else: # Build dynamic MERGE props for both source and destination @@ -509,6 +554,9 @@ class MemoryGraph: if agent_id: source_props.append("agent_id: $agent_id") dest_props.append("agent_id: $agent_id") + if run_id: + source_props.append("run_id: $run_id") + dest_props.append("run_id: $run_id") source_props_str = ", ".join(source_props) dest_props_str = ", ".join(dest_props) @@ -544,6 +592,8 @@ class MemoryGraph: } if agent_id: params["agent_id"] = agent_id + if run_id: + params["run_id"] = run_id result = self.graph.query(cypher, params=params) results.append(result) return results @@ -556,15 +606,21 @@ class MemoryGraph: return entity_list def _search_source_node(self, source_embedding, filters, threshold=0.9): - agent_filter = "" + + # Build WHERE conditions + where_conditions = [ + "source_candidate.embedding IS NOT NULL", + "source_candidate.user_id = $user_id" + ] if filters.get("agent_id"): - agent_filter = "AND source_candidate.agent_id = $agent_id" + where_conditions.append("source_candidate.agent_id = $agent_id") + if filters.get("run_id"): + where_conditions.append("source_candidate.run_id = $run_id") + where_clause = " AND ".join(where_conditions) cypher = f""" MATCH (source_candidate {self.node_label}) - WHERE source_candidate.embedding IS NOT NULL - AND source_candidate.user_id = $user_id - {agent_filter} + WHERE {where_clause} WITH source_candidate, round(2 * vector.similarity.cosine(source_candidate.embedding, $source_embedding) - 1, 4) AS source_similarity // denormalize for backward compatibility @@ -584,20 +640,28 @@ class MemoryGraph: } if filters.get("agent_id"): params["agent_id"] = filters["agent_id"] + if filters.get("run_id"): + params["run_id"] = filters["run_id"] result = self.graph.query(cypher, params=params) return result def _search_destination_node(self, destination_embedding, filters, threshold=0.9): - agent_filter = "" + + # Build WHERE conditions + where_conditions = [ + "destination_candidate.embedding IS NOT NULL", + "destination_candidate.user_id = $user_id" + ] if filters.get("agent_id"): - agent_filter = "AND destination_candidate.agent_id = $agent_id" + where_conditions.append("destination_candidate.agent_id = $agent_id") + if filters.get("run_id"): + where_conditions.append("destination_candidate.run_id = $run_id") + where_clause = " AND ".join(where_conditions) cypher = f""" MATCH (destination_candidate {self.node_label}) - WHERE destination_candidate.embedding IS NOT NULL - AND destination_candidate.user_id = $user_id - {agent_filter} + WHERE {where_clause} WITH destination_candidate, round(2 * vector.similarity.cosine(destination_candidate.embedding, $destination_embedding) - 1, 4) AS destination_similarity // denormalize for backward compatibility @@ -618,6 +682,8 @@ class MemoryGraph: } if filters.get("agent_id"): params["agent_id"] = filters["agent_id"] + if filters.get("run_id"): + params["run_id"] = filters["run_id"] result = self.graph.query(cypher, params=params) return result diff --git a/tests/memory/test_neo4j_cypher_syntax.py b/tests/memory/test_neo4j_cypher_syntax.py new file mode 100644 index 000000000..45eab377a --- /dev/null +++ b/tests/memory/test_neo4j_cypher_syntax.py @@ -0,0 +1,292 @@ +import os +from unittest.mock import Mock, patch + + +class TestNeo4jCypherSyntaxFix: + """Test that Neo4j Cypher syntax fixes work correctly""" + + def test_get_all_generates_valid_cypher_with_agent_id(self): + """Test that get_all method generates valid Cypher with agent_id""" + # Mock the langchain_neo4j module to avoid import issues + with patch.dict('sys.modules', {'langchain_neo4j': Mock()}): + from mem0.memory.graph_memory import MemoryGraph + + # Create instance (will fail on actual connection, but that's fine for syntax testing) + try: + _ = MemoryGraph(url="bolt://localhost:7687", username="test", password="test") + except Exception: + # Expected to fail on connection, just test the class exists + assert MemoryGraph is not None + return + + def test_cypher_syntax_validation(self): + """Test that our Cypher fixes don't contain problematic patterns""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Ensure the old buggy pattern is not present + assert "AND n.agent_id = $agent_id AND m.agent_id = $agent_id" not in content + assert "WHERE 1=1 {agent_filter}" not in content + + # Ensure proper node property syntax is present + assert "node_props" in content + assert "agent_id: $agent_id" in content + + # Ensure run_id follows the same pattern + # Check for absence of problematic run_id patterns + assert "AND n.run_id = $run_id AND m.run_id = $run_id" not in content + assert "WHERE 1=1 {run_id_filter}" not in content + + def test_no_undefined_variables_in_cypher(self): + """Test that we don't have undefined variable patterns""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Check for patterns that would cause "Variable 'm' not defined" errors + lines = content.split('\n') + for i, line in enumerate(lines): + # Look for WHERE clauses that reference variables not in MATCH + if 'WHERE' in line and 'm.agent_id' in line: + # Check if there's a MATCH clause before this that defines 'm' + preceding_lines = lines[max(0, i-10):i] + match_found = any('MATCH' in prev_line and ' m ' in prev_line for prev_line in preceding_lines) + assert match_found, f"Line {i+1}: WHERE clause references 'm' without MATCH definition" + + # Also check for run_id patterns that might have similar issues + if 'WHERE' in line and 'm.run_id' in line: + # Check if there's a MATCH clause before this that defines 'm' + preceding_lines = lines[max(0, i-10):i] + match_found = any('MATCH' in prev_line and ' m ' in prev_line for prev_line in preceding_lines) + assert match_found, f"Line {i+1}: WHERE clause references 'm.run_id' without MATCH definition" + + def test_agent_id_integration_syntax(self): + """Test that agent_id is properly integrated into MATCH clauses""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Should have node property building logic + assert 'node_props = [' in content + assert 'node_props.append("agent_id: $agent_id")' in content + assert 'node_props_str = ", ".join(node_props)' in content + + # Should use the node properties in MATCH clauses + assert '{{{node_props_str}}}' in content or '{node_props_str}' in content + + def test_run_id_integration_syntax(self): + """Test that run_id is properly integrated into MATCH clauses""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Should have node property building logic for run_id + assert 'node_props = [' in content + assert 'node_props.append("run_id: $run_id")' in content + assert 'node_props_str = ", ".join(node_props)' in content + + # Should use the node properties in MATCH clauses + assert '{{{node_props_str}}}' in content or '{node_props_str}' in content + + def test_agent_id_filter_patterns(self): + """Test that agent_id filtering follows the correct pattern""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Check that agent_id is handled in filters + assert 'if filters.get("agent_id"):' in content + assert 'params["agent_id"] = filters["agent_id"]' in content + + # Check that agent_id is used in node properties + assert 'node_props.append("agent_id: $agent_id")' in content + + def test_run_id_filter_patterns(self): + """Test that run_id filtering follows the same pattern as agent_id""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Check that run_id is handled in filters + assert 'if filters.get("run_id"):' in content + assert 'params["run_id"] = filters["run_id"]' in content + + # Check that run_id is used in node properties + assert 'node_props.append("run_id: $run_id")' in content + + def test_agent_id_cypher_generation(self): + """Test that agent_id is properly included in Cypher query generation""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Check that the dynamic property building pattern exists + assert 'node_props = [' in content + assert 'node_props_str = ", ".join(node_props)' in content + + # Check that agent_id is handled in the pattern + assert 'if filters.get(' in content + assert 'node_props.append(' in content + + # Verify the pattern is used in MATCH clauses + assert '{{{node_props_str}}}' in content or '{node_props_str}' in content + + def test_run_id_cypher_generation(self): + """Test that run_id is properly included in Cypher query generation""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Check that the dynamic property building pattern exists + assert 'node_props = [' in content + assert 'node_props_str = ", ".join(node_props)' in content + + # Check that run_id is handled in the pattern + assert 'if filters.get(' in content + assert 'node_props.append(' in content + + # Verify the pattern is used in MATCH clauses + assert '{{{node_props_str}}}' in content or '{node_props_str}' in content + + def test_agent_id_implementation_pattern(self): + """Test that the code structure supports agent_id implementation""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Verify that agent_id pattern is used consistently + assert 'node_props = [' in content + assert 'node_props_str = ", ".join(node_props)' in content + assert 'if filters.get("agent_id"):' in content + assert 'node_props.append("agent_id: $agent_id")' in content + + def test_run_id_implementation_pattern(self): + """Test that the code structure supports run_id implementation""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Verify that run_id pattern is used consistently + assert 'node_props = [' in content + assert 'node_props_str = ", ".join(node_props)' in content + assert 'if filters.get("run_id"):' in content + assert 'node_props.append("run_id: $run_id")' in content + + def test_user_identity_integration(self): + """Test that both agent_id and run_id are properly integrated into user identity""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Check that user_identity building includes both agent_id and run_id + assert 'user_identity = f"user_id: {filters[\'user_id\']}"' in content + assert 'user_identity += f", agent_id: {filters[\'agent_id\']}"' in content + assert 'user_identity += f", run_id: {filters[\'run_id\']}"' in content + + def test_search_methods_integration(self): + """Test that both agent_id and run_id are properly integrated into search methods""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Check that search methods handle both agent_id and run_id + assert 'where_conditions.append("source_candidate.agent_id = $agent_id")' in content + assert 'where_conditions.append("source_candidate.run_id = $run_id")' in content + assert 'where_conditions.append("destination_candidate.agent_id = $agent_id")' in content + assert 'where_conditions.append("destination_candidate.run_id = $run_id")' in content + + def test_add_entities_integration(self): + """Test that both agent_id and run_id are properly integrated into add_entities""" + graph_memory_path = 'mem0/memory/graph_memory.py' + + # Check if file exists before reading + if not os.path.exists(graph_memory_path): + # Skip test if file doesn't exist (e.g., in CI environment) + return + + with open(graph_memory_path, 'r') as f: + content = f.read() + + # Check that add_entities handles both agent_id and run_id + assert 'agent_id = filters.get("agent_id", None)' in content + assert 'run_id = filters.get("run_id", None)' in content + + # Check that merge properties include both + assert 'if agent_id:' in content + assert 'if run_id:' in content + assert 'merge_props.append("agent_id: $agent_id")' in content + assert 'merge_props.append("run_id: $run_id")' in content + diff --git a/tests/test_main.py b/tests/test_main.py index c89e7bc34..2f548e315 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -23,11 +23,16 @@ def memory_instance(): patch("mem0.utils.factory.LlmFactory") as mock_llm, patch("mem0.memory.telemetry.capture_event"), patch("mem0.memory.graph_memory.MemoryGraph"), + patch("mem0.memory.main.GraphStoreFactory") as mock_graph_store, ): mock_embedder.create.return_value = Mock() mock_vector_store.create.return_value = Mock() mock_vector_store.create.return_value.search.return_value = [] mock_llm.create.return_value = Mock() + + # Create a mock instance that won't try to access config attributes + mock_graph_instance = Mock() + mock_graph_store.create.return_value = mock_graph_instance config = MemoryConfig(version="v1.1") config.graph_store.config = {"some_config": "value"} @@ -42,11 +47,16 @@ def memory_custom_instance(): patch("mem0.utils.factory.LlmFactory") as mock_llm, patch("mem0.memory.telemetry.capture_event"), patch("mem0.memory.graph_memory.MemoryGraph"), + patch("mem0.memory.main.GraphStoreFactory") as mock_graph_store, ): mock_embedder.create.return_value = Mock() mock_vector_store.create.return_value = Mock() mock_vector_store.create.return_value.search.return_value = [] mock_llm.create.return_value = Mock() + + # Create a mock instance that won't try to access config attributes + mock_graph_instance = Mock() + mock_graph_store.create.return_value = mock_graph_instance config = MemoryConfig( version="v1.1",