fix graph bugs
This commit is contained in:
parent
7ce9b5bd8a
commit
152bfbf49f
@ -17,10 +17,13 @@ class DataBaseManager:
|
|||||||
|
|
||||||
if self.config.enable_knowledge_base:
|
if self.config.enable_knowledge_base:
|
||||||
from src.core.knowledgebase import KnowledgeBase
|
from src.core.knowledgebase import KnowledgeBase
|
||||||
from src.core.graphbase import GraphDatabase
|
|
||||||
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
||||||
self.graph_base = GraphDatabase(self.config, self.embed_model)
|
if self.config.enable_knowledge_graph:
|
||||||
self.graph_base.start()
|
from src.core.graphbase import GraphDatabase
|
||||||
|
self.graph_base = GraphDatabase(self.config, self.embed_model)
|
||||||
|
self.graph_base.start()
|
||||||
|
else:
|
||||||
|
self.graph_base = None
|
||||||
|
|
||||||
self.data = {"databases": [], "graph": {}}
|
self.data = {"databases": [], "graph": {}}
|
||||||
|
|
||||||
|
|||||||
@ -83,19 +83,21 @@ kg.close()
|
|||||||
UIE_MODEL = None
|
UIE_MODEL = None
|
||||||
|
|
||||||
class GraphDatabase:
|
class GraphDatabase:
|
||||||
def __init__(self, config, embed_model=None):
|
def __init__(self, config, embed_model=None, kgdb_name="neo4j"):
|
||||||
self.config = config
|
self.config = config
|
||||||
self.driver = None
|
self.driver = None
|
||||||
self.files = []
|
self.files = []
|
||||||
self.status = "closed"
|
self.status = "closed"
|
||||||
|
self.kgdb_name = kgdb_name
|
||||||
assert embed_model, "embed_model=None"
|
assert embed_model, "embed_model=None"
|
||||||
self.embed_model = embed_model
|
self.embed_model = embed_model
|
||||||
|
|
||||||
def start(self):
|
def start(self):
|
||||||
uri = os.environ.get("NEO4J_URI")
|
uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687")
|
||||||
username = os.environ.get("NEO4J_USERNAME")
|
username = os.environ.get("NEO4J_USERNAME", "neo4j")
|
||||||
password = os.environ.get("NEO4J_PASSWORD")
|
password = os.environ.get("NEO4J_PASSWORD", "0123456789")
|
||||||
self.driver = GD.driver(uri, auth=(username, password))
|
logger.info(f"Connecting to Neo4j at {uri}/{self.kgdb_name}")
|
||||||
|
self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password))
|
||||||
self.status = "open"
|
self.status = "open"
|
||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
@ -146,7 +148,9 @@ class GraphDatabase:
|
|||||||
|
|
||||||
def use_database(self, kgdb_name):
|
def use_database(self, kgdb_name):
|
||||||
"""切换到指定数据库"""
|
"""切换到指定数据库"""
|
||||||
self.driver = GD.driver(f"{os.environ.get('NEO4J_URI')}/{kgdb_name}", auth=(os.environ.get('NEO4J_USERNAME'), os.environ.get('NEO4J_PASSWORD')))
|
if kgdb_name != self.kgdb_name:
|
||||||
|
raise ValueError(f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致")
|
||||||
|
self.start()
|
||||||
|
|
||||||
def txt_add_entity(self, triples, kgdb_name='neo4j'):
|
def txt_add_entity(self, triples, kgdb_name='neo4j'):
|
||||||
"""添加实体三元组"""
|
"""添加实体三元组"""
|
||||||
|
|||||||
@ -133,6 +133,9 @@ def get_graph_nodes():
|
|||||||
if not kgdb_name:
|
if not kgdb_name:
|
||||||
return jsonify({'message': 'kgdb_name is required'}), 400
|
return jsonify({'message': 'kgdb_name is required'}), 400
|
||||||
|
|
||||||
|
if not startup.config.enable_knowledge_graph:
|
||||||
|
return jsonify({'message': 'Knowledge graph is not enabled'}), 400
|
||||||
|
|
||||||
logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
|
logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes")
|
||||||
result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num)
|
result = startup.dbm.graph_base.get_sample_nodes(kgdb_name, num)
|
||||||
return jsonify({'result': startup.retriever.format_general_results(result), 'message': 'success'}), 200
|
return jsonify({'result': startup.retriever.format_general_results(result), 'message': 'success'}), 200
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user