From 577b78821f7b6210ce5dbb41b4cad076ea5cc328 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Fri, 28 Feb 2025 02:42:21 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E5=AD=90=E5=9B=BE=E6=A3=80?= =?UTF-8?q?=E7=B4=A2=E7=9A=84bug?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/core/graphbase.py | 2 +- src/core/retriever.py | 31 +++++++++++++++++++++---------- 2 files changed, 22 insertions(+), 11 deletions(-) diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 14844a0c..6294c683 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -280,7 +280,7 @@ class GraphDatabase: def query(tx, entity_name, hops): result = tx.run(f""" MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m) - RETURN n.name AS node_name, r, m.name AS neighbor_name + RETURN n, r, m """, entity_name=entity_name) return result.values() diff --git a/src/core/retriever.py b/src/core/retriever.py index b31f7ef8..55d21a9b 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -175,7 +175,7 @@ class Retriever: return entities - def _extract_relationship_info(self, relationship, source_name, target_name): + def _extract_relationship_info(self, relationship, source_name=None, target_name=None, node_dict=None): """ 提取关系信息并返回格式化的节点和边信息 """ @@ -188,6 +188,9 @@ class Retriever: source_id = source.element_id target_id = target.element_id + source_name = node_dict[source_id]["name"] if source_name is None else source_name + target_name = node_dict[target_id]["name"] if target_name is None else target_name + relationship_type = relationship._properties.get("type", "unknown") if relationship_type == "unknown": relationship_type = relationship.type @@ -234,20 +237,28 @@ class Retriever: node_dict = {} for item in results: - if not isinstance(item[1], list) or not item[1]: + # 检查数据格式 + if len(item) < 2 or not isinstance(item[1], list): continue - relationship = item[1][0] - source_name = item[0] - target_name = item[2] if len(item) > 2 else "unknown" + node_dict[item[0].element_id] = dict(id=item[0].element_id, name=item[0]._properties.get("name", "Unknown")) + node_dict[item[2].element_id] = dict(id=item[2].element_id, name=item[2]._properties.get("name", "Unknown")) - node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name) - if node_info is None or edge_info is None: - continue + # 处理关系列表中的每个关系 + for i, relationship in enumerate(item[1]): + try: + # 提取关系信息 + node_info, edge_info = self._extract_relationship_info(relationship, node_dict=node_dict) + if node_info is None or edge_info is None: + continue - node_dict.update({node["id"]: node for node in node_info}) - formatted_results["edges"].append(edge_info) + # 添加边 + formatted_results["edges"].append(edge_info) + except Exception as e: + logger.error(f"处理关系时出错: {e}, 关系: {relationship}") + continue + # 将节点字典转换为列表 formatted_results["nodes"] = list(node_dict.values()) return formatted_results