diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 7d86106b..97f8855c 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -146,11 +146,11 @@ class GraphDatabase: with self.driver.session() as session: return session.execute_read(query) - def use_database(self, kgdb_name): + def use_database(self, kgdb_name="neo4j"): """切换到指定数据库""" - if kgdb_name != self.kgdb_name: - raise ValueError(f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致") - self.start() + assert kgdb_name == self.kgdb_name, f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致" + if self.status == "closed": + self.start() def txt_add_entity(self, triples, kgdb_name='neo4j'): """添加实体三元组""" @@ -234,38 +234,6 @@ class GraphDatabase: self.txt_add_vector_entity(triples, kgdb_name) - # def batch_create(tx, triples): - # query = """ - # UNWIND $triples AS triple - # MERGE (a:Entity {name: triple.h}) - # MERGE (b:Entity {name: triple.t}) - # MERGE (a)-[r:RELATION {type: triple.r}]->(b) - # """ - # tx.run(query, triples=triples) - # - # def batch_add_embeddings(tx, embeddings): - # query = """ - # UNWIND $embeddings AS embedding - # MATCH (e:Entity {name: embedding.name}) - # SET e.embedding = embedding.vector - # """ - # tx.run(query, embeddings=embeddings) - # - # with self.driver.session() as session: - # session.execute_write(batch_create, triples) - # - # # 获取embedding并批量添加 - # embeddings = [] - # for triple in triples: - # h = triple['h'] - # t = triple['t'] - # embedding_h = self.get_embedding(h) - # embedding_t = self.get_embedding(t) - # embeddings.append({"name": h, "vector": embedding_h}) - # embeddings.append({"name": t, "vector": embedding_t}) - # - # session.execute_write(batch_add_embeddings, embeddings) - self.status = "open" return kgdb_name @@ -292,18 +260,41 @@ class GraphDatabase: """ tx.run(query) - def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2): - """查询图数据库中所有三元组信息""" + def query_node(self, entity_name, hops=2, **kwargs): + # TODO 添加判断节点数量为 0 停止检索 + 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): + result = self.query_by_vector_tep(entity_name=entity_name) + querys = [] + for i in range(num_of_res): + if result[i][1] > threshold: + querys.append(result[i][0]) + else: + break + ans = [] + for query in querys: + tep = self.query_specific_entity(entity_name=query, hops=hops) # 这里是只获取第一个 TODO: 优化 + ans.extend(tep) + return ans + + def query_by_vector_tep(self, entity_name, kgdb_name='neo4j'): + """向量查询""" self.use_database(kgdb_name) - def query(tx, hops): - result = tx.run(f""" - MATCH (n)-[r*1..{hops}]->(m) - RETURN n, r, m - """) + def query(tx, text): + embedding = self.get_embedding(text) + result = tx.run(""" + CALL db.index.vector.queryNodes('entityEmbeddings', 10, $embedding) + YIELD node AS similarEntity, score + RETURN similarEntity.name AS name, score + """, embedding=embedding) return result.values() with self.driver.session() as session: - return session.execute_read(query, hops) + return session.execute_read(query, entity_name) def query_specific_entity(self, entity_name, kgdb_name='neo4j', hops=2): """查询指定实体三元组信息""" @@ -318,6 +309,19 @@ class GraphDatabase: with self.driver.session() as session: return session.execute_read(query, entity_name, hops) + def query_all_nodes_and_relationships(self, kgdb_name='neo4j', hops = 2): + """查询图数据库中所有三元组信息""" + self.use_database(kgdb_name) + def query(tx, hops): + result = tx.run(f""" + MATCH (n)-[r*1..{hops}]->(m) + RETURN n, r, m + """) + return result.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): """查询指定关系三元组信息""" self.use_database(kgdb_name) @@ -346,48 +350,6 @@ class GraphDatabase: with self.driver.session() as session: return session.execute_read(query, keyword, hops) - def query_node(self, entity_name, args): - # TODO 添加判断节点数量为 0 停止检索 - - if args.get("exact_match"): - raise NotImplemented("not implement for `exact_match`") - else: - return self.query_by_vector(entity_name, kgdb_name=args.get("kgdb_name"), hops=args.get("hops")) - - def query_by_vector_tep(self, keyword, kgdb_name='neo4j'): - """向量查询""" - self.use_database(kgdb_name) - def query(tx, text): - embedding = self.get_embedding(text) - result = tx.run(""" - CALL db.index.vector.queryNodes('entityEmbeddings', 10, $embedding) - YIELD node AS similarEntity, score - RETURN similarEntity.name AS name, score - """, embedding=embedding) - # result = result.values() - # query = result[0][0] - return result.values() - - with self.driver.session() as session: - return session.execute_read(query, keyword) - - def query_by_vector(self, entity_name, threshold=0.9, kgdb_name='neo4j', hops=2, num_of_res=2): - self.use_database(kgdb_name) - result = self.query_by_vector_tep(entity_name) - querys = [] - threshold = 0.9 if threshold is None else threshold - num_of_res = 2 if num_of_res is None else num_of_res - for i in range(num_of_res): - if result[i][1] > threshold: - querys.append(result[i][0]) - else: - break - ans = [] - for query in querys: - tep = self.query_specific_entity(query, hops) # 这里是只获取第一个 TODO: 优化 - ans.extend(tep) - return ans - def query_node_info(self, node_name, kgdb_name='neo4j', hops = 2): """查询指定节点的详细信息返回信息""" self.use_database(kgdb_name) # 切换到指定数据库 diff --git a/src/views/database_view.py b/src/views/database_view.py index b461e0a4..ebaa255b 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -1,6 +1,7 @@ import os import json import threading +from functools import wraps from flask import Blueprint, jsonify, request, Response from src.utils import setup_logger, hashstr @@ -12,6 +13,19 @@ logger = setup_logger("server-database") progress = {} # 只针对单个用户的进度 +def handle_exceptions(f): + @wraps(f) + def decorated_function(*args, **kwargs): + try: + logger.debug(f"Entering {f.__name__}") + result = f(*args, **kwargs) + logger.debug(f"Exiting {f.__name__}") + return result + except Exception as e: + logger.error(f"Error in {f.__name__}: {str(e)}") + return jsonify({"message": str(e), "error": "处理请求时发生错误"}), 500 + return decorated_function + @db.route('/', methods=['GET']) def get_databases(): try: @@ -115,41 +129,35 @@ def get_graph_info(): return jsonify(graph_info) @db.route('/graph/node', methods=['GET']) +@handle_exceptions def get_graph_node(): - entity_name = request.args.get('entity_name') - kgdb_name = request.args.get('kgdb_name') - hops = request.args.get('hops') - if not entity_name: - return jsonify({'message': 'entity_name and kgdb_name are required'}), 400 - - logger.debug(f"Get graph node {entity_name} in {kgdb_name} with {hops} hops") - result = startup.dbm.graph_base.query_node(entity_name, request.args) + assert request.args.get("entity_name"), "entity_name is required" + logger.debug(f"Get graph node {request.args.get('entity_name')} with {request.args}") + result = startup.dbm.graph_base.query_node(**request.args) return jsonify({'result': startup.retriever.format_query_results(result), 'message': 'success'}), 200 @db.route('/graph/nodes', methods=['GET']) +@handle_exceptions def get_graph_nodes(): kgdb_name = request.args.get('kgdb_name') num = request.args.get('num') - if not kgdb_name: - return jsonify({'message': 'kgdb_name is required'}), 400 - - if not startup.config.enable_knowledge_graph: - return jsonify({'message': 'Knowledge graph is not enabled'}), 400 + assert kgdb_name, "kgdb_name is required" + assert startup.config.enable_knowledge_graph, "Knowledge graph is not enabled" logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes") result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num) return jsonify({'result': startup.retriever.format_general_results(result), 'message': 'success'}), 200 @db.route('/graph/add', methods=['POST']) +@handle_exceptions def add_graph_entity(): data = json.loads(request.data) kgdb_name = data.get('kgdb_name') file_path = data.get('file_path') + assert file_path.endswith('.jsonl'), "file_path must be a jsonl file" + assert startup.config.enable_knowledge_graph, "Knowledge graph is not enabled" - if file_path.endswith('.jsonl'): - startup.dbm.graph_base.jsonl_file_add_entity(file_path, kgdb_name) - else: - return jsonify({'message': 'Unsupported file type'}), 400 + startup.dbm.graph_base.jsonl_file_add_entity(file_path, kgdb_name) return jsonify({'message': 'Entity successfully added'}), 200