From 8323279ff46218aa8041d06978fa16effe6a4d43 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sun, 13 Apr 2025 21:06:23 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E7=9F=A5=E8=AF=86=E5=9B=BE?= =?UTF-8?q?=E8=B0=B1=E6=A3=80=E7=B4=A2bug?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/core/graphbase.py | 95 +++++++++++++++++++++++++++++-------------- src/core/retriever.py | 7 ++-- 2 files changed, 68 insertions(+), 34 deletions(-) diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 9c61063b..6e8237df 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -206,19 +206,12 @@ class GraphDatabase: """ tx.run(query) - def query_node(self, entity_name, hops=2, **kwargs): + def query_node(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, max_entities=5, **kwargs): # TODO 添加判断节点数量为 0 停止检索 # 判断是否启动 if not self.is_running(): raise Exception("图数据库未启动") - logger.debug(f"Query graph node {entity_name} with {hops=}") - if kwargs.get("exact_match"): - raise NotImplemented("not implement for `exact_match`") - else: - return self.query_by_vector(entity_name=entity_name, **kwargs) - - def query_by_vector(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, num_of_res=5): self.use_database(kgdb_name) def _index_exists(tx, index_name): """检查索引是否存在""" @@ -251,7 +244,7 @@ class GraphDatabase: raise e # 筛选出分数高于阈值的实体 - qualified_entities = [result[0] for result in results[:num_of_res] if result[1] > threshold] + qualified_entities = [result[0] for result in results[:max_entities] if result[1] > threshold] logger.debug(f"Graph Query Entities: {entity_name}, {qualified_entities=}") # 对每个合格的实体进行查询 @@ -262,70 +255,102 @@ class GraphDatabase: return all_query_results - def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2): + def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2, limit=100): """查询指定实体三元组信息(无向关系)""" - self.use_database(kgdb_name) - def query(tx, entity_name, hops): - result = tx.run(f""" - MATCH (n {{name: $entity_name}})-[r*1..{hops}]-(m) - RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m - """, entity_name=entity_name) - return result.values() + if not entity_name: + logger.warning("实体名称为空") + return [] - with self.driver.session() as session: - return session.execute_read(query, entity_name, hops) + self.use_database(kgdb_name) + + 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 + LIMIT $limit + """ + result = tx.run(query_str, entity_name=entity_name, limit=limit) + + if not result: + logger.info(f"未找到实体 {entity_name} 的相关信息") + return [] + + values = result.values() + # 安全地处理embedding属性 + values = clean_triples_embedding(values) + return values + + except Exception as e: + logger.error(f"查询实体 {entity_name} 失败: {str(e)}") + return [] + + 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""" self.use_database(kgdb_name) def query(tx, hops): result = tx.run(f""" MATCH (n)-[r*1..{hops}]->(m) - RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m + RETURN n AS n, r, m AS m """) - return result.values() + 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""" 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 {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m + RETURN n AS n, r, m AS m """) - return result.values() + 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_entity_like(self, keyword, kgdb_name='neo4j', hops = 2): - """模糊查询""" + """模糊查询 NEVER USE""" self.use_database(kgdb_name) def query(tx, keyword, hops): result = tx.run(f""" MATCH (n:Entity) WHERE n.name CONTAINS $keyword MATCH (n)-[r*1..{hops}]->(m) - RETURN n {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m + RETURN n AS n, r, m AS m """, keyword=keyword) - return result.values() + values = result.values() + values = clean_triples_embedding(values) + return values with self.driver.session() as session: return session.execute_read(query, keyword, hops) def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): - """查询指定节点的详细信息返回信息""" + """查询指定节点的详细信息返回信息 NEVER USE""" 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 {{.*, embedding: null}} AS n, r, m {{.*, embedding: null}} AS m + RETURN n AS n, r, m AS m """, node_name=node_name) - return result.values() + values = result.values() + values = clean_triples_embedding(values) + return values with self.driver.session() as session: return session.execute_read(query, node_name, hops) @@ -473,6 +498,14 @@ class GraphDatabase: return count +def clean_triples_embedding(triples): + for item in triples: + if hasattr(item[0], '_properties'): + item[0]._properties['embedding'] = None + if hasattr(item[2], '_properties'): + item[2]._properties['embedding'] = None + return triples + if __name__ == "__main__": pass diff --git a/src/core/retriever.py b/src/core/retriever.py index bc053afc..7fa7a249 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,10 +1,11 @@ +import traceback + from src import config, knowledge_base, graph_base from src.models.rerank_model import get_reranker from src.utils.logging_config import logger from src.models import select_model from src.core.operators import HyDEOperator - class Retriever: def __init__(self): @@ -80,7 +81,7 @@ class Retriever: for entity in refs["entities"]: if entity == "": continue - result = graph_base.query_by_vector(entity) + result = graph_base.query_node(entity) if result != []: results.extend(result) return {"results": self.format_query_results(results)} @@ -256,7 +257,7 @@ class Retriever: # 添加边 formatted_results["edges"].append(edge_info) except Exception as e: - logger.error(f"处理关系时出错: {e}, 关系: {relationship}") + logger.error(f"处理关系时出错: {e}, 关系: {relationship}, {traceback.format_exc()}") continue # 将节点字典转换为列表