From ce6f88b8b24eb6f79a4d9a383c3c95eb2b0ad3e3 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Tue, 15 Apr 2025 10:49:30 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E5=9B=BE=E8=B0=B1=E6=A3=80?= =?UTF-8?q?=E7=B4=A2=E7=BB=93=E6=9E=9C=E9=87=8D=E5=A4=8D=E7=9A=84=E9=97=AE?= =?UTF-8?q?=E9=A2=98=EF=BC=8C=E5=90=8C=E6=97=B6=E5=B0=86=E6=A3=80=E7=B4=A2?= =?UTF-8?q?=E9=83=A8=E5=88=86=E7=9A=84=E9=80=BB=E8=BE=91=E4=BB=A3=E7=A0=81?= =?UTF-8?q?=EF=BC=8C=E8=BF=81=E7=A7=BB=E5=88=B0graphbase=20=E9=87=8C?= =?UTF-8?q?=E9=9D=A2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- server/routers/data_router.py | 4 +- src/core/graphbase.py | 119 ++++++++++++++++++++++++++++++++++ src/core/retriever.py | 90 +------------------------ 3 files changed, 122 insertions(+), 91 deletions(-) diff --git a/server/routers/data_router.py b/server/routers/data_router.py index 27bb04ca..275237f8 100644 --- a/server/routers/data_router.py +++ b/server/routers/data_router.py @@ -157,7 +157,7 @@ async def index_nodes(data: dict = Body(default={})): @data.get("/graph/node") async def get_graph_node(entity_name: str): result = graph_base.query_node(entity_name=entity_name) - return {"result": retriever.format_query_results(result), "message": "success"} + return {"result": graph_base.format_query_result_to_graph(result), "message": "success"} @data.get("/graph/nodes") async def get_graph_nodes(kgdb_name: str, num: int): @@ -166,7 +166,7 @@ async def get_graph_nodes(kgdb_name: str, num: int): logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes") result = graph_base.get_sample_nodes(kgdb_name, num) - return {"result": retriever.format_general_results(result), "message": "success"} + return {"result": graph_base.format_general_results(result), "message": "success"} @data.post("/graph/add-by-jsonl") async def add_graph_entity(file_path: str = Body(...), kgdb_name: Optional[str] = Body(None)): diff --git a/src/core/graphbase.py b/src/core/graphbase.py index f395a335..f8274568 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -498,6 +498,125 @@ class GraphDatabase: return count + + def _extract_relationship_info(self, relationship, source_name=None, target_name=None, node_dict=None): + """ + 提取关系信息并返回格式化的节点和边信息 + """ + rel_id = relationship.element_id + nodes = relationship.nodes + if len(nodes) != 2: + return None, None + + source, target = nodes + 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 + + edge_info = { + "id": rel_id, + "type": relationship_type, + "source_id": source_id, + "target_id": target_id, + "source_name": source_name, + "target_name": target_name, + } + + node_info = [ + {"id": source_id, "name": source_name}, + {"id": target_id, "name": target_name}, + ] + + return node_info, edge_info + + def format_general_results(self, results): + formatted_results = {"nodes": [], "edges": []} + + for item in results: + relationship = item[1] + source_name = item[0]._properties.get("name", "unknown") + target_name = item[2]._properties.get("name", "unknown") if len(item) > 2 else "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 node in node_info: + if node["id"] not in [n["id"] for n in formatted_results["nodes"]]: + formatted_results["nodes"].append(node) + + formatted_results["edges"].append(edge_info) + + return formatted_results + + def format_query_result_to_graph(self, query_results): + """将检索到的结果转换为 {"nodes": [], "edges": []} 的格式 + + 例如: + { + "nodes": [ + { + "id": "4:5efbff88-72ef-44f9-b867-6c0e164a4a13:103", + "name": "张若锦" + }, + { + "id": "4:5efbff88-72ef-44f9-b867-6c0e164a4a13:20", + "name": "贾宝玉" + }, + .... + ], + "edges": [ + { + "id": "5:5efbff88-72ef-44f9-b867-6c0e164a4a13:71", + "type": "奴仆", + "source_id": "4:5efbff88-72ef-44f9-b867-6c0e164a4a13:88", + "target_id": "4:5efbff88-72ef-44f9-b867-6c0e164a4a13:20", + "source_name": "宋嬷嬷", + "target_name": "贾宝玉" + }, + .... + ] + } + """ + formatted_results = {"nodes": [], "edges": []} + node_dict = {} + edge_dict = {} + + for item in query_results: + # 检查数据格式 + if len(item) < 2 or not isinstance(item[1], list): + continue + + 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")) + + # 处理关系列表中的每个关系 + 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 + + # 添加边 + edge_dict[edge_info["id"]] = edge_info + except Exception as e: + logger.error(f"处理关系时出错: {e}, 关系: {relationship}, {traceback.format_exc()}") + continue + + # 将节点字典转换为列表 + formatted_results["nodes"] = list(node_dict.values()) + formatted_results["edges"] = list(edge_dict.values()) + + + return formatted_results + def clean_triples_embedding(triples): for item in triples: if hasattr(item[0], '_properties'): diff --git a/src/core/retriever.py b/src/core/retriever.py index 7fa7a249..89b4b873 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -84,7 +84,7 @@ class Retriever: result = graph_base.query_node(entity) if result != []: results.extend(result) - return {"results": self.format_query_results(results)} + return {"results": graph_base.format_query_result_to_graph(results)} def query_knowledgebase(self, query, history, refs): @@ -177,94 +177,6 @@ class Retriever: return entities - def _extract_relationship_info(self, relationship, source_name=None, target_name=None, node_dict=None): - """ - 提取关系信息并返回格式化的节点和边信息 - """ - rel_id = relationship.element_id - nodes = relationship.nodes - if len(nodes) != 2: - return None, None - - source, target = nodes - 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 - - edge_info = { - "id": rel_id, - "type": relationship_type, - "source_id": source_id, - "target_id": target_id, - "source_name": source_name, - "target_name": target_name, - } - - node_info = [ - {"id": source_id, "name": source_name}, - {"id": target_id, "name": target_name}, - ] - - return node_info, edge_info - - def format_general_results(self, results): - formatted_results = {"nodes": [], "edges": []} - - for item in results: - relationship = item[1] - source_name = item[0]._properties.get("name", "unknown") - target_name = item[2]._properties.get("name", "unknown") if len(item) > 2 else "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 node in node_info: - if node["id"] not in [n["id"] for n in formatted_results["nodes"]]: - formatted_results["nodes"].append(node) - - formatted_results["edges"].append(edge_info) - - return formatted_results - - def format_query_results(self, results): - # logger.debug(f"Graph Query Results: {results}") - formatted_results = {"nodes": [], "edges": []} - node_dict = {} - - for item in results: - # 检查数据格式 - if len(item) < 2 or not isinstance(item[1], list): - continue - - 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")) - - # 处理关系列表中的每个关系 - 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 - - # 添加边 - formatted_results["edges"].append(edge_info) - except Exception as e: - logger.error(f"处理关系时出错: {e}, 关系: {relationship}, {traceback.format_exc()}") - continue - - # 将节点字典转换为列表 - formatted_results["nodes"] = list(node_dict.values()) - - return formatted_results - def __call__(self, query, history, meta): refs = self.retrieval(query, history, meta) query = self.construct_query(query, refs, meta)