From f9cabc54da8b307a6217189adec05102b8a34e24 Mon Sep 17 00:00:00 2001 From: QiyiJiang <3442721524@qq.com> Date: Wed, 17 Jul 2024 11:34:05 +0000 Subject: [PATCH] change output --- src/core/graphbase.py | 40 +++++++++++++++------------------------- src/core/retriever.py | 2 +- 2 files changed, 16 insertions(+), 26 deletions(-) diff --git a/src/core/graphbase.py b/src/core/graphbase.py index b329a28b..12beacf9 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -123,9 +123,7 @@ class GraphDatabase: return result.values() with self.driver.session() as session: - original_results = session.execute_read(query, hops) - formatted_results = self.format_query_results(original_results) - return formatted_results, original_results + return session.execute_read(query, hops) def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops = 2): """查询指定实体三元组信息""" @@ -138,9 +136,7 @@ class GraphDatabase: return result.values() with self.driver.session() as session: - original_results = session.execute_read(query, entity_name, hops) - formatted_results = self.format_query_results(original_results) - return formatted_results, original_results + return session.execute_read(query, entity_name, hops) def query_by_relationship_type(self, relationship_type, kgdb_name='neo4j', hops = 2): """查询指定关系三元组信息""" @@ -153,9 +149,7 @@ class GraphDatabase: return result.values() with self.driver.session() as session: - original_results = session.execute_read(query, relationship_type, hops) - formatted_results = self.format_query_results(original_results) - return formatted_results, original_results + return session.execute_read(query, relationship_type, hops) def query_entity_like(self, keyword, kgdb_name='neo4j', hops = 2): """模糊查询""" @@ -170,9 +164,7 @@ class GraphDatabase: return result.values() with self.driver.session() as session: - original_results = session.execute_read(query, keyword, hops) - formatted_results = self.format_query_results(original_results) - return formatted_results, original_results + return session.execute_read(query, keyword, hops) def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): """查询指定节点的详细信息返回信息""" @@ -186,20 +178,18 @@ class GraphDatabase: return result.values() with self.driver.session() as session: - original_results = session.execute_read(query, node_name, hops) - formatted_results = self.format_query_results(original_results) - return formatted_results, original_results + return session.execute_read(query, node_name, hops) - def format_query_results(self, results): - formatted_results = [] - for row in results: - n, rs, m = row - entity_a = n['name'] - entity_b = m['name'] - for rel in rs: - relationship = rel.type - formatted_results.append(f"实体 {entity_a} 和 实体 {entity_b} 的关系是 {relationship}") - return formatted_results + # def format_query_results(self, results): + # formatted_results = [] + # for row in results: + # n, rs, m = row + # entity_a = n['name'] + # entity_b = m['name'] + # for rel in rs: + # relationship = rel.type + # formatted_results.append(f"实体 {entity_a} 和 实体 {entity_b} 的关系是 {relationship}") + # return formatted_results diff --git a/src/core/retriever.py b/src/core/retriever.py index 94434789..f3408cc8 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -44,7 +44,7 @@ class Retriever: results = [] _, entities = self.rewrite_query(query, history) 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 return results