fix graph bugs
This commit is contained in:
parent
7ce9b5bd8a
commit
152bfbf49f
@ -17,10 +17,13 @@ class DataBaseManager:
|
||||
|
||||
if self.config.enable_knowledge_base:
|
||||
from src.core.knowledgebase import KnowledgeBase
|
||||
from src.core.graphbase import GraphDatabase
|
||||
self.knowledge_base = KnowledgeBase(config, self.embed_model)
|
||||
self.graph_base = GraphDatabase(self.config, self.embed_model)
|
||||
self.graph_base.start()
|
||||
if self.config.enable_knowledge_graph:
|
||||
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": {}}
|
||||
|
||||
|
||||
@ -83,19 +83,21 @@ kg.close()
|
||||
UIE_MODEL = None
|
||||
|
||||
class GraphDatabase:
|
||||
def __init__(self, config, embed_model=None):
|
||||
def __init__(self, config, embed_model=None, kgdb_name="neo4j"):
|
||||
self.config = config
|
||||
self.driver = None
|
||||
self.files = []
|
||||
self.status = "closed"
|
||||
self.kgdb_name = kgdb_name
|
||||
assert embed_model, "embed_model=None"
|
||||
self.embed_model = embed_model
|
||||
|
||||
def start(self):
|
||||
uri = os.environ.get("NEO4J_URI")
|
||||
username = os.environ.get("NEO4J_USERNAME")
|
||||
password = os.environ.get("NEO4J_PASSWORD")
|
||||
self.driver = GD.driver(uri, auth=(username, password))
|
||||
uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687")
|
||||
username = os.environ.get("NEO4J_USERNAME", "neo4j")
|
||||
password = os.environ.get("NEO4J_PASSWORD", "0123456789")
|
||||
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"
|
||||
|
||||
def close(self):
|
||||
@ -146,7 +148,9 @@ class GraphDatabase:
|
||||
|
||||
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'):
|
||||
"""添加实体三元组"""
|
||||
|
||||
@ -133,6 +133,9 @@ def get_graph_nodes():
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user