Fixes n_embeddings use and error for memgraph (#3296)

This commit is contained in:
Parshva Daftari
2025-08-11 23:27:29 +05:30
committed by GitHub
parent c2792c6558
commit b60a208c2f
3 changed files with 31 additions and 24 deletions
+24 -20
View File
@@ -274,23 +274,25 @@ class MemoryGraph:
# Build query based on whether agent_id is provided
if filters.get("agent_id"):
cypher_query = """
MATCH (n:Entity {user_id: $user_id, agent_id: $agent_id})-[r]->(m:Entity)
MATCH (n:Entity {user_id: $user_id, agent_id: $agent_id})
WHERE n.embedding IS NOT NULL
WITH collect(n) AS nodes1, collect(m) AS nodes2, r
CALL node_similarity.cosine_pairwise("embedding", nodes1, nodes2)
WITH n, $n_embedding as n_embedding
CALL node_similarity.cosine_pairwise("embedding", [n_embedding], [n.embedding])
YIELD node1, node2, similarity
WITH node1, node2, similarity, r
WITH n, similarity
WHERE similarity >= $threshold
RETURN node1.name AS source, id(node1) AS source_id, type(r) AS relationship, id(r) AS relation_id, node2.name AS destination, id(node2) AS destination_id, similarity
MATCH (n)-[r]->(m:Entity)
RETURN n.name AS source, id(n) AS source_id, type(r) AS relationship, id(r) AS relation_id, m.name AS destination, id(m) AS destination_id, similarity
UNION
MATCH (n:Entity {user_id: $user_id, agent_id: $agent_id})<-[r]-(m:Entity)
MATCH (n:Entity {user_id: $user_id, agent_id: $agent_id})
WHERE n.embedding IS NOT NULL
WITH collect(n) AS nodes1, collect(m) AS nodes2, r
CALL node_similarity.cosine_pairwise("embedding", nodes1, nodes2)
WITH n, $n_embedding as n_embedding
CALL node_similarity.cosine_pairwise("embedding", [n_embedding], [n.embedding])
YIELD node1, node2, similarity
WITH node1, node2, similarity, r
WITH n, similarity
WHERE similarity >= $threshold
RETURN node2.name AS source, id(node2) AS source_id, type(r) AS relationship, id(r) AS relation_id, node1.name AS destination, id(node1) AS destination_id, similarity
MATCH (m:Entity)-[r]->(n)
RETURN m.name AS source, id(m) AS source_id, type(r) AS relationship, id(r) AS relation_id, n.name AS destination, id(n) AS destination_id, similarity
ORDER BY similarity DESC
LIMIT $limit;
"""
@@ -303,23 +305,25 @@ class MemoryGraph:
}
else:
cypher_query = """
MATCH (n:Entity {user_id: $user_id})-[r]->(m:Entity)
MATCH (n:Entity {user_id: $user_id})
WHERE n.embedding IS NOT NULL
WITH collect(n) AS nodes1, collect(m) AS nodes2, r
CALL node_similarity.cosine_pairwise("embedding", nodes1, nodes2)
WITH n, $n_embedding as n_embedding
CALL node_similarity.cosine_pairwise("embedding", [n_embedding], [n.embedding])
YIELD node1, node2, similarity
WITH node1, node2, similarity, r
WITH n, similarity
WHERE similarity >= $threshold
RETURN node1.name AS source, id(node1) AS source_id, type(r) AS relationship, id(r) AS relation_id, node2.name AS destination, id(node2) AS destination_id, similarity
MATCH (n)-[r]->(m:Entity)
RETURN n.name AS source, id(n) AS source_id, type(r) AS relationship, id(r) AS relation_id, m.name AS destination, id(m) AS destination_id, similarity
UNION
MATCH (n:Entity {user_id: $user_id})<-[r]-(m:Entity)
MATCH (n:Entity {user_id: $user_id})
WHERE n.embedding IS NOT NULL
WITH collect(n) AS nodes1, collect(m) AS nodes2, r
CALL node_similarity.cosine_pairwise("embedding", nodes1, nodes2)
WITH n, $n_embedding as n_embedding
CALL node_similarity.cosine_pairwise("embedding", [n_embedding], [n.embedding])
YIELD node1, node2, similarity
WITH node1, node2, similarity, r
WITH n, similarity
WHERE similarity >= $threshold
RETURN node2.name AS source, id(node2) AS source_id, type(r) AS relationship, id(r) AS relation_id, node1.name AS destination, id(node1) AS destination_id, similarity
MATCH (m:Entity)-[r]->(n)
RETURN m.name AS source, id(m) AS source_id, type(r) AS relationship, id(r) AS relation_id, n.name AS destination, id(n) AS destination_id, similarity
ORDER BY similarity DESC
LIMIT $limit;
"""
+1 -1
View File
@@ -261,7 +261,7 @@ class Qdrant(VectorStoreBase):
with_payload=True,
with_vectors=False,
)
return result.points
return result
def reset(self):
"""Reset the index by deleting and recreating it."""
+6 -3
View File
@@ -231,7 +231,7 @@ class TestQdrant(unittest.TestCase):
score=0.95,
payload={"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}
)
self.client_mock.scroll.return_value = MagicMock(points=[mock_point])
self.client_mock.scroll.return_value = [mock_point]
filters = {"user_id": "alice", "agent_id": "agent1", "run_id": "run1"}
results = self.qdrant.list(filters=filters, limit=10)
@@ -247,6 +247,7 @@ class TestQdrant(unittest.TestCase):
self.assertIsInstance(scroll_filter, Filter)
self.assertEqual(len(scroll_filter.must), 3) # user_id, agent_id, run_id
# The list method returns the result directly
self.assertEqual(len(results), 1)
self.assertEqual(results[0].payload["user_id"], "alice")
self.assertEqual(results[0].payload["agent_id"], "agent1")
@@ -259,7 +260,7 @@ class TestQdrant(unittest.TestCase):
score=0.95,
payload={"user_id": "alice"}
)
self.client_mock.scroll.return_value = MagicMock(points=[mock_point])
self.client_mock.scroll.return_value = [mock_point]
filters = {"user_id": "alice"}
results = self.qdrant.list(filters=filters, limit=10)
@@ -270,19 +271,21 @@ class TestQdrant(unittest.TestCase):
self.assertIsInstance(scroll_filter, Filter)
self.assertEqual(len(scroll_filter.must), 1) # Only user_id
# The list method returns the result directly
self.assertEqual(len(results), 1)
self.assertEqual(results[0].payload["user_id"], "alice")
def test_list_with_no_filters(self):
"""Test list with no filters."""
mock_point = MagicMock(id=str(uuid.uuid4()), score=0.95, payload={"key": "value"})
self.client_mock.scroll.return_value = MagicMock(points=[mock_point])
self.client_mock.scroll.return_value = [mock_point]
results = self.qdrant.list(filters=None, limit=10)
call_args = self.client_mock.scroll.call_args[1]
self.assertIsNone(call_args["scroll_filter"])
# The list method returns the result directly
self.assertEqual(len(results), 1)
def test_delete_col(self):