fix: 图查询出错的问题
This commit is contained in:
parent
9f191e739e
commit
5925e359e9
@ -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) # 切换到指定数据库
|
||||
|
||||
@ -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
|
||||
|
||||
|
||||
Loading…
Reference in New Issue
Block a user