fix: 图查询出错的问题

This commit is contained in:
Wenjie Zhang 2024-09-14 02:44:53 +08:00
parent 9f191e739e
commit 5925e359e9
2 changed files with 73 additions and 103 deletions

View File

@ -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) # 切换到指定数据库

View File

@ -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