From b60a208c2fbe63e52b244c191fceb9365d35116c Mon Sep 17 00:00:00 2001 From: Parshva Daftari <89991302+parshvadaftari@users.noreply.github.com> Date: Mon, 11 Aug 2025 23:27:29 +0530 Subject: [PATCH] Fixes n_embeddings use and error for memgraph (#3296) --- mem0/memory/memgraph_memory.py | 44 ++++++++++++++++-------------- mem0/vector_stores/qdrant.py | 2 +- tests/vector_stores/test_qdrant.py | 9 ++++-- 3 files changed, 31 insertions(+), 24 deletions(-) diff --git a/mem0/memory/memgraph_memory.py b/mem0/memory/memgraph_memory.py index 63f0dc7da..9aeaddc80 100644 --- a/mem0/memory/memgraph_memory.py +++ b/mem0/memory/memgraph_memory.py @@ -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; """ diff --git a/mem0/vector_stores/qdrant.py b/mem0/vector_stores/qdrant.py index 273b61428..59ee9a92c 100644 --- a/mem0/vector_stores/qdrant.py +++ b/mem0/vector_stores/qdrant.py @@ -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.""" diff --git a/tests/vector_stores/test_qdrant.py b/tests/vector_stores/test_qdrant.py index dfcf145dc..3b2f6be19 100644 --- a/tests/vector_stores/test_qdrant.py +++ b/tests/vector_stores/test_qdrant.py @@ -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):