From 612bd9f7b0fad8b676a1d10ed1afc010ee216de9 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 30 Jul 2025 10:50:06 +0800 Subject: [PATCH] =?UTF-8?q?refactor(graph):=20=E9=87=8D=E6=9E=84=E7=9F=A5?= =?UTF-8?q?=E8=AF=86=E5=9B=BE=E8=B0=B1=E6=9F=A5=E8=AF=A2=E5=92=8C=E6=A0=BC?= =?UTF-8?q?=E5=BC=8F=E5=8C=96=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 修改查询方法以支持不同返回格式(graph/triples) - 移除不再使用的格式化方法和废弃查询方法 - 优化节点和关系的查询性能及结果处理 - 更新相关API端点以直接返回原始查询结果 --- server/routers/graph_router.py | 8 +- src/agents/chatbot/graph.py | 2 +- src/agents/tools_factory.py | 4 +- src/knowledge/graphbase.py | 233 +++++++++++++++------------------ 4 files changed, 109 insertions(+), 138 deletions(-) diff --git a/server/routers/graph_router.py b/server/routers/graph_router.py index 8cc402b6..0b348e38 100644 --- a/server/routers/graph_router.py +++ b/server/routers/graph_router.py @@ -189,16 +189,15 @@ async def get_neo4j_nodes( raise HTTPException(status_code=400, detail="图数据库未启动") result = graph_base.get_sample_nodes(kgdb_name, num) - formatted_result = graph_base.format_general_results(result) return { "success": True, - "result": formatted_result, + "result": result, "message": "success" } except Exception as e: - logger.error(f"获取图节点数据失败: {e}") + logger.error(f"获取图节点数据失败: {e}\n{traceback.format_exc()}") raise HTTPException(status_code=500, detail=f"获取图节点数据失败: {str(e)}") @graph.get("/neo4j/node") @@ -214,11 +213,10 @@ async def get_neo4j_node( raise HTTPException(status_code=400, detail="图数据库未启动") result = graph_base.query_node(entity_name=entity_name) - formatted_result = graph_base.format_query_result_to_graph(result) return { "success": True, - "result": formatted_result, + "result": result, "message": "success" } diff --git a/src/agents/chatbot/graph.py b/src/agents/chatbot/graph.py index 0fe36864..510e3a69 100644 --- a/src/agents/chatbot/graph.py +++ b/src/agents/chatbot/graph.py @@ -18,7 +18,7 @@ from src.agents.chatbot.configuration import ChatbotConfiguration from src.agents.tools_factory import get_runnable_tools class ChatbotAgent(BaseAgent): - name = "对话机器人(Chatbot)" + name = "问答助手" description = "基础的对话机器人,可以回答问题,默认不使用任何工具,可在配置中启用需要的工具。" config_schema = ChatbotConfiguration diff --git a/src/agents/tools_factory.py b/src/agents/tools_factory.py index 4ebc2d4b..3d1c2e08 100644 --- a/src/agents/tools_factory.py +++ b/src/agents/tools_factory.py @@ -139,8 +139,8 @@ def calculator(a: float, b: float, operation: str) -> float: @tool def query_knowledge_graph(query: Annotated[str, "The keyword to query knowledge graph."]): - """Use this to query knowledge graph.""" - return graph_base.query_node(query, hops=2) + """Use this to query knowledge graph, which include some food domain knowledge.""" + return graph_base.query_node(query, hops=2, return_format='triples') # 更新工具注册表 _TOOLS_REGISTRY.update({ diff --git a/src/knowledge/graphbase.py b/src/knowledge/graphbase.py index d0094142..7f9191d4 100644 --- a/src/knowledge/graphbase.py +++ b/src/knowledge/graphbase.py @@ -61,11 +61,26 @@ class GraphDatabase: self.use_database(kgdb_name) def query(tx, num): """Note: 这里只查询带有 Entity 标签的节点""" - result = tx.run("MATCH (n:Entity)-[r]->(m:Entity) RETURN n, r, m LIMIT $num", num=int(num)) - return result.values() + query_str = """ + MATCH (n:Entity)-[r]->(m:Entity) + RETURN + {id: elementId(n), name: n.name} AS h, + {type: r.type, source_id: elementId(n), target_id: elementId(m)} AS r, + {id: elementId(m), name: m.name} AS t + LIMIT $num + """ + results = tx.run(query_str, num=int(num)) + formatted_results = {'nodes': [], 'edges': []} + + for item in results: + formatted_results['nodes'].extend([item['h'], item['t']]) + formatted_results['edges'].append(item['r']) + + return formatted_results with self.driver.session() as session: - return session.execute_read(query, num) + results = session.execute_read(query, num) + return results def create_graph_database(self, kgdb_name): """创建新的数据库,如果已存在则返回已有数据库的名称""" @@ -266,7 +281,7 @@ class GraphDatabase: """ tx.run(query) - def query_node(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, max_entities=5, **kwargs): + def query_node(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, max_entities=5, return_format='graph', **kwargs): """知识图谱查询节点的入口:""" assert self.driver is not None, "Database is not connected" assert self.is_running(), "图数据库未启动" @@ -274,26 +289,32 @@ class GraphDatabase: self.use_database(kgdb_name) # 使用向量索引进行查询 - results_sim = self._query_with_vector_sim(entity_name, kgdb_name, hops, threshold) - results_fuzzy = self._query_with_fuzzy_match(entity_name, kgdb_name, hops) + results_sim = self._query_with_vector_sim(entity_name, kgdb_name, threshold) + results_fuzzy = self._query_with_fuzzy_match(entity_name, kgdb_name) results = results_sim + results_fuzzy qualified_entities = [result[0] for result in results][:max_entities] logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}") # 对每个合格的实体进行查询 - all_query_results = [] + all_query_results = {'nodes': [], 'edges': [], 'triples': []} for entity in qualified_entities: - query_result = self._query_specific_entity(entity_name=entity, hops=hops, kgdb_name=kgdb_name) - all_query_results.extend(query_result) + query_result = self._query_specific_entity(entity_name=entity, kgdb_name=kgdb_name, hops=hops) + if return_format == 'graph': + all_query_results['nodes'].extend(query_result['nodes']) + all_query_results['edges'].extend(query_result['edges']) + elif return_format == 'triples': + all_query_results['triples'].extend(query_result['triples']) + else: + raise ValueError(f"Invalid return_format: {return_format}") return all_query_results - def _query_with_fuzzy_match(self, keyword, kgdb_name='neo4j', hops = 2): + def _query_with_fuzzy_match(self, keyword, kgdb_name='neo4j'): """模糊查询""" assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) - def query_fuzzy_match(tx, keyword, hops): + def query_fuzzy_match(tx, keyword): result = tx.run(""" MATCH (n:Entity) WHERE n.name CONTAINS $keyword @@ -304,9 +325,9 @@ class GraphDatabase: return values with self.driver.session() as session: - return session.execute_read(query_fuzzy_match, keyword, hops) + return session.execute_read(query_fuzzy_match, keyword) - def _query_with_vector_sim(self, keyword, kgdb_name='neo4j', hops = 2, threshold=0.9): + def _query_with_vector_sim(self, keyword, kgdb_name='neo4j', threshold=0.9): """向量查询""" assert self.driver is not None, "Database is not connected" self.use_database(kgdb_name) @@ -334,7 +355,6 @@ class GraphDatabase: with self.driver.session() as session: results = session.execute_read(query_by_vector, keyword, threshold=threshold) - results = clean_triples_embedding(results) return results @@ -349,21 +369,35 @@ class GraphDatabase: def query(tx, entity_name, hops, limit): try: - query_str = f""" - MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m) - RETURN n AS n, r, m AS m + query_str = """ + MATCH (n {name: $entity_name})-[r1]-(m1) + RETURN + {id: elementId(n), name: n.name} AS h, + {type: r1.type, source_id: elementId(n), target_id: elementId(m1)} AS r, + {id: elementId(m1), name: m1.name} AS t + UNION + MATCH (n {name: $entity_name})-[r1]-(m1)-[r2]-(m2) + RETURN + {id: elementId(m1), name: m1.name} AS h, + {type: r2.type, source_id: elementId(m1), target_id: elementId(m2)} AS r, + {id: elementId(m2), name: m2.name} AS t LIMIT $limit """ - result = tx.run(query_str, entity_name=entity_name, limit=limit) + results = tx.run(query_str, entity_name=entity_name, limit=limit) - if not result: + if not results: logger.info(f"未找到实体 {entity_name} 的相关信息") - return [] + return {} - values = result.values() - # 安全地处理embedding属性 - values = clean_triples_embedding(values) - return values + formatted_results = {'nodes': [], 'edges': [], 'triples': []} + + for item in results: + formatted_results['nodes'].extend([item['h'], item['t']]) + formatted_results['edges'].append(item['r']) + formatted_results['triples'].append((item['h']['name'], item['r']['type'], item['t']['name'])) + + logger.debug(f"Query Results: {results}") + return formatted_results except Exception as e: logger.error(f"查询实体 {entity_name} 失败: {str(e)}") @@ -372,62 +406,11 @@ class GraphDatabase: try: with self.driver.session() as session: return session.execute_read(query, entity_name, hops, limit) + except Exception as e: logger.error(f"数据库会话异常: {str(e)}") return [] - def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2): - """查询图数据库中所有三元组信息 NEVER USE""" - raise Exception("NEVER USE") - assert self.driver is not None, "Database is not connected" - self.use_database(kgdb_name) - def query(tx, hops): - result = tx.run(f""" - MATCH (n)-[r*1..{hops}]->(m) - RETURN n AS n, r, m AS m - """) - values = result.values() - values = clean_triples_embedding(values) - return values - - with self.driver.session() as session: - return session.execute_read(query, hops) - - def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2): - """查询指定关系三元组信息 NEVER USE""" - raise Exception("NEVER USE") - assert self.driver is not None, "Database is not connected" - self.use_database(kgdb_name) - def query(tx, relationship_type, hops): - result = tx.run(f""" - MATCH (n)-[r:`{relationship_type}`*1..{hops}]->(m) - RETURN n AS n, r, m AS m - """) - values = result.values() - values = clean_triples_embedding(values) - return values - - with self.driver.session() as session: - return session.execute_read(query, relationship_type, hops) - - def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): - """查询指定节点的详细信息返回信息 NEVER USE""" - raise Exception("NEVER USE") - assert self.driver is not None, "Database is not connected" - self.use_database(kgdb_name) # 切换到指定数据库 - def query(tx, node_name, hops): - result = tx.run(f""" - MATCH (n {{name: $node_name}}) - OPTIONAL MATCH (n)-[r*1..{hops}]->(m) - RETURN n AS n, r, m AS m - """, node_name=node_name) - values = result.values() - values = clean_triples_embedding(values) - return values - - with self.driver.session() as session: - return session.execute_read(query, node_name, hops) - async def aget_embedding(self, text): if isinstance(text, list): outputs = await self.embed_model.abatch_encode(text, batch_size=40) @@ -587,64 +570,15 @@ 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 或 target_name,则需要 node_dict - if source_name is None or target_name is None: - assert node_dict is not None, "node_dict is required when source_name or target_name is None" - 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": []} + 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) + nodes.extend([item['h'], item['t']]) + edges.append(item['r']) + formatted_results = {"nodes": nodes, "edges": edges} return formatted_results def format_query_result_to_graph(self, query_results): @@ -709,6 +643,45 @@ class GraphDatabase: return formatted_results + 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 或 target_name,则需要 node_dict + if source_name is None or target_name is None: + assert node_dict is not None, "node_dict is required when source_name or target_name is None" + 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 clean_triples_embedding(triples): for item in triples: if hasattr(item[0], '_properties'):