Fixes n_embeddings use and error for memgraph (#3296)
This commit is contained in:
@@ -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;
|
||||
"""
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user