Fix Neo4j Cypher syntax error with agent_id filtering (#3158)

Co-authored-by: parshvadaftari <daftariparshva@gmail.com>
This commit is contained in:
DrJsPBs
2025-08-08 12:08:13 -04:00
committed by GitHub
parent 148bbf0a5c
commit 26732771eb
3 changed files with 407 additions and 39 deletions
+105 -39
View File
@@ -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
+292
View File
@@ -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
+10
View File
@@ -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",