From 152bfbf49fa41a2ae4aea4f22a0d6d75aca0f266 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Wed, 11 Sep 2024 01:08:13 +0800 Subject: [PATCH] fix graph bugs --- src/core/database.py | 9 ++++++--- src/core/graphbase.py | 16 ++++++++++------ src/views/database_view.py | 3 +++ 3 files changed, 19 insertions(+), 9 deletions(-) diff --git a/src/core/database.py b/src/core/database.py index 3966b6dc..440e2300 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -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": {}} diff --git a/src/core/graphbase.py b/src/core/graphbase.py index f6db98af..7d86106b 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -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'): """添加实体三元组""" diff --git a/src/views/database_view.py b/src/views/database_view.py index 5f5057c6..b461e0a4 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -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