This commit is contained in:
Wenjie Zhang 2024-07-17 20:01:01 +08:00
commit 1a21158f56
2 changed files with 16 additions and 26 deletions

View File

@ -123,9 +123,7 @@ class GraphDatabase:
return result.values() return result.values()
with self.driver.session() as session: with self.driver.session() as session:
original_results = session.execute_read(query, hops) return session.execute_read(query, hops)
formatted_results = self.format_query_results(original_results)
return formatted_results, original_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):
"""查询指定实体三元组信息""" """查询指定实体三元组信息"""
@ -138,9 +136,7 @@ class GraphDatabase:
return result.values() return result.values()
with self.driver.session() as session: with self.driver.session() as session:
original_results = session.execute_read(query, entity_name, hops) return session.execute_read(query, entity_name, hops)
formatted_results = self.format_query_results(original_results)
return formatted_results, original_results
def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2): def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2):
"""查询指定关系三元组信息""" """查询指定关系三元组信息"""
@ -153,9 +149,7 @@ class GraphDatabase:
return result.values() return result.values()
with self.driver.session() as session: with self.driver.session() as session:
original_results = session.execute_read(query, relationship_type, hops) return session.execute_read(query, relationship_type, hops)
formatted_results = self.format_query_results(original_results)
return formatted_results, original_results
def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2): def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2):
"""模糊查询""" """模糊查询"""
@ -170,9 +164,7 @@ class GraphDatabase:
return result.values() return result.values()
with self.driver.session() as session: with self.driver.session() as session:
original_results = session.execute_read(query, keyword, hops) return session.execute_read(query, keyword, hops)
formatted_results = self.format_query_results(original_results)
return formatted_results, original_results
def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2):
"""查询指定节点的详细信息返回信息""" """查询指定节点的详细信息返回信息"""
@ -186,20 +178,18 @@ class GraphDatabase:
return result.values() return result.values()
with self.driver.session() as session: with self.driver.session() as session:
original_results = session.execute_read(query, node_name, hops) return session.execute_read(query, node_name, hops)
formatted_results = self.format_query_results(original_results)
return formatted_results, original_results
def format_query_results(self, results): # def format_query_results(self, results):
formatted_results = [] # formatted_results = []
for row in results: # for row in results:
n, rs, m = row # n, rs, m = row
entity_a = n['name'] # entity_a = n['name']
entity_b = m['name'] # entity_b = m['name']
for rel in rs: # for rel in rs:
relationship = rel.type # relationship = rel.type
formatted_results.append(f"实体 {entity_a} 和 实体 {entity_b} 的关系是 {relationship}") # formatted_results.append(f"实体 {entity_a} 和 实体 {entity_b} 的关系是 {relationship}")
return formatted_results # return formatted_results

View File

@ -52,7 +52,7 @@ class Retriever:
results = [] results = []
_, entities = self.rewrite_query(query, history) _, entities = self.rewrite_query(query, history)
for entitie in entities: for entitie in entities:
result, _ = dbm.graph_base.query_entity_like(entitie) result = dbm.graph_base.query_entity_like(entitie)
results.append(result) if result else None results.append(result) if result else None
return results return results