diff --git a/.gitignore b/.gitignore index b026948c..409b49b2 100644 --- a/.gitignore +++ b/.gitignore @@ -25,13 +25,19 @@ cache ### IDE .vscode +.idea *.nogit.* *.pdf +*.yaml src/data neo4j* */package-lock.json web/package-lock.json saves notebooks -*.yaml \ No newline at end of file +local_neo4j/data +local_neo4j/logs +local_neo4j/import +local_neo4j/plugins +local_neo4j/conf \ No newline at end of file diff --git a/README.md b/README.md index a6bb9d9a..f4c8443e 100644 --- a/README.md +++ b/README.md @@ -22,7 +22,8 @@ zhipuai ### 【可选】配置图数据库 neo4j -使用 docker 部署 neo4j 服务,配置文件见 [local_neo4j/docker-compose.yml](local_neo4j/docker-compose.yml). 默认账号密码见最后一行,可以使用 `http://localhost:7474/` 在浏览器可视化访问。 +使用 docker 部署 neo4j 服务,配置文件见 [local_neo4j/docker-compose.yml](local_neo4j/docker-compose.yml). +默认账号密码见最后一行,可以使用 `http://localhost:7474/` 在浏览器可视化访问。 ```bash cd local_neo4j @@ -30,12 +31,14 @@ docker compose up -d ``` 可以使用 `python test_neo4j.py` 来测试是否正常启动。使用 `docker compose down` 可停止服务。 +如果想要管理 neo4j,也可以使用 `docker ps` 查看容器 id,然后使用 `docker exec -it /bin/bash` 进入容器。 +如果想要删除数据库中的文件,可以进入容器并停止 neo4j 后,执行 `rm -rf /data/databases`。 ## 启动 ```bash -python -m src.api +python -m src.api cd web npm install diff --git a/local_neo4j/docker-compose.yml b/local_neo4j/docker-compose.yml new file mode 100644 index 00000000..aff1e9b2 --- /dev/null +++ b/local_neo4j/docker-compose.yml @@ -0,0 +1,17 @@ +version: '3.9' +services: + + neo4j: + image: neo4j:latest + volumes: + - ./conf:/var/lib/neo4j/conf + - ./import:/var/lib/neo4j/import + - ./plugins:/plugins + - ./data:/data + - ./logs:/var/lib/neo4j/logs + restart: always + ports: + - 7474:7474 + - 7687:7687 + environment: + - NEO4J_AUTH=neo4j/0123456789 diff --git a/local_neo4j/test_neo4j.py b/local_neo4j/test_neo4j.py new file mode 100644 index 00000000..a293e4fc --- /dev/null +++ b/local_neo4j/test_neo4j.py @@ -0,0 +1,34 @@ +from neo4j import GraphDatabase +from neo4j.exceptions import ServiceUnavailable, AuthError + +def check_neo4j_status(uri="bolt://localhost:7687", username="neo4j", password="0123456789"): + """ + 检查 Neo4j 数据库是否可以连接并正常工作。 + + 参数: + uri (str): Neo4j 的 URI,默认为 "bolt://localhost:7687" + username (str): 数据库用户名,默认为 "neo4j" + password (str): 数据库密码,默认为 "0123456789" + + 返回: + str: "OK" 表示连接成功,"UNAVAILABLE" 表示服务不可用,"AUTH_FAILED" 表示认证失败。 + """ + try: + driver = GraphDatabase.driver(uri, auth=(username, password)) + with driver.session() as session: + # 简单的查询来测试连接 + result = session.run("RETURN 1") + if result.single()[0] == 1: + return "OK" + except ServiceUnavailable: + return "UNAVAILABLE" + except AuthError: + return "AUTH_FAILED" + finally: + # 确保关闭驱动 + driver.close() + +# 测试函数 +status = check_neo4j_status() +print(f"Neo4j status: {status}") + diff --git a/src/core/database.py b/src/core/database.py index 443896de..3966b6dc 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -71,13 +71,16 @@ class DataBaseManager: return {"databases": [db.to_dict() for db in self.data["databases"]]} def get_graph(self): - if self.config.enable_graph_base: + if self.config.enable_knowledge_graph: self.data["graph"].update(self.graph_base.get_database_info("neo4j")) return {"graph": self.data["graph"]} else: return {"message": "Graph base not enabled", "graph": {}} def create_database(self, database_name, description, db_type, dimension): + from src.config import EMBED_MODEL_INFO + dimension = dimension or EMBED_MODEL_INFO[self.config.embed_model]["dimension"] + new_database = DataBaseLite(database_name, description, db_type, @@ -94,7 +97,7 @@ class DataBaseManager: if db.embed_model != self.config.embed_model: logger.error(f"Embed model not match, {db.embed_model} != {self.config.embed_model}") - return {"message": "Embed model not match", "status": "failed"} + return {"message": f"Embed model not match, cur: {self.config.embed_model}", "status": "failed"} new_files = [] for file in files: @@ -168,7 +171,6 @@ class DataBaseManager: logger.error(f"File format not supported, only support {support_format}") raise Exception(f"File format not supported, only support {support_format}") - def delete_file(self, db_id, file_id): db = self.get_kb_by_id(db_id) file_idx_to_delete = [idx for idx, f in enumerate(db.files) if f["file_id"] == file_id][0] diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 70d97255..664b4729 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -9,11 +9,12 @@ import warnings from src.plugins import pdf2txt from src.plugins.oneke import OneKE +from src.utils import setup_logger + +logger = setup_logger("server-graphbase") warnings.filterwarnings("ignore", category=UserWarning) - - UIE_MODEL = None class GraphDatabase: @@ -36,6 +37,16 @@ class GraphDatabase: """关闭数据库连接""" self.driver.close() + def get_sample_nodes(self, kgdb_name='neo4j', num=50): + """获取指定数据库的前 num 个节点信息""" + self.use_database(kgdb_name) + def query(tx, num): + result = tx.run("MATCH (n)-[r]->(m) RETURN n, r, m LIMIT $num", num=int(num)) + return result.values() + + with self.driver.session() as session: + return session.execute_read(query, num) + def create_graph_database(self, kgdb_name): """创建新的数据库,如果已存在则返回已有数据库的名称""" with self.driver.session() as session: @@ -116,21 +127,25 @@ class GraphDatabase: MERGE (t:Entity {name: $t}) MERGE (h)-[r:RELATION {type: $r}]->(t) """, h=entry['h'], t=entry['t'], r=entry['r']) - def _create_vector_index(tx): - index_name = "entity-embeddings" + def _create_vector_index(tx, dim): + index_name = "entityEmbeddings" if not _index_exists(tx, index_name): tx.run(f""" CREATE VECTOR INDEX {index_name} FOR (n: Entity) ON (n.embedding) OPTIONS {{indexConfig: {{ - `vector.dimensions`: 1024, + `vector.dimensions`: {dim}, `vector.similarity_function`: 'cosine' }} }}; """) + + from src.config import EMBED_MODEL_INFO + embed_info = EMBED_MODEL_INFO[self.config.embed_model] with self.driver.session() as session: session.execute_write(_create_graph, triples) - session.execute_write(_create_vector_index) - for entry in triples: + session.execute_write(_create_vector_index, embed_info.dimension) + for i, entry in enumerate(triples): + logger.info(f"Adding entity {i+1}/{len(triples)}") embedding_h = self.get_embedding(entry['h']) session.execute_write(self.set_embedding, entry['h'], embedding_h) @@ -148,37 +163,39 @@ class GraphDatabase: triples = list(read_triples(file_path)) - 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) + self.txt_add_vector_entity(triples, kgdb_name) - 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) + # 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 @@ -260,13 +277,21 @@ 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('entity-embeddings', 10, $embedding) + CALL db.index.vector.queryNodes('entityEmbeddings', 10, $embedding) YIELD node AS similarEntity, score RETURN similarEntity.name AS name, score """, embedding=embedding) @@ -277,7 +302,7 @@ class GraphDatabase: 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): + 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 = [] diff --git a/src/core/retriever.py b/src/core/retriever.py index 2025f3a2..2071f3cb 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -83,7 +83,7 @@ class Retriever: r["file"] = kb.id2file(r["entity"]["file_id"]) if self.config.enable_reranker: - RERANK_THRESHOLD = 0.1 + RERANK_THRESHOLD = 0.001 for r in kb_res: r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) @@ -124,7 +124,46 @@ class Retriever: return entities + def foramt_general_results(self, results): + logger.debug(f"Formatting general results: {results}") + formatted_results = {"nodes": [], "edges": []} + + for item in results: + relationship = item[1] + rel_id = relationship.element_id + nodes = relationship.nodes + if len(nodes) != 2: + continue + + source, target = nodes + + source_id = source.element_id + target_id = target.element_id + source_name = source._properties.get('name', 'unknown') + target_name = target._properties.get('name', 'unknown') + + if source_id not in formatted_results["nodes"]: + formatted_results["nodes"].append({"id": source_id, "name": source_name}) + if target_id not in formatted_results["nodes"]: + formatted_results["nodes"].append({"id": target_id, "name": target_name}) + + relationship_type = relationship._properties.get('type', 'unknown') + if relationship_type == 'unknown': + relationship_type = relationship.type + + formatted_results["edges"].append({ + "id": rel_id, + "type": relationship_type, + "source_id": source_id, + "target_id": target_id, + "source_name": source_name, + "target_name": target_name + }) + + return formatted_results + def format_query_results(self, results): + logger.debug(f"Formatting query results: {results}") formatted_results = {"nodes": [], "edges": []} node_dict = {} diff --git a/src/models/embedding.py b/src/models/embedding.py index e7c482d9..9eb81ee2 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -45,10 +45,11 @@ class ZhipuEmbedding: self.query_instruction_for_retrieval = "为这个句子生成表示以用于检索相关文章:" def predict(self, message): - data = [] for i in range(0, len(message), 10): + if len(message) > 10: + logger.info(f"Encoding {i} to {i+10} with {len(message)} messages") group_msg = message[i:i+10] response = self.client.embeddings.create( model=self.model_info.default_path, diff --git a/src/views/database_view.py b/src/views/database_view.py index 43bf1b75..98adcfae 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -123,9 +123,20 @@ def get_graph_node(): 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_by_vector(entity_name, kgdb_name=kgdb_name, hops=hops) + result = startup.dbm.graph_base.query_node(entity_name, request.args) return jsonify({'result': startup.retriever.format_query_results(result), 'message': 'success'}), 200 +@db.route('/graph/nodes', methods=['GET']) +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 + + 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.foramt_general_results(result), 'message': 'success'}), 200 + @db.route('/graph/add', methods=['POST']) def add_graph_entity(): data = json.loads(request.data) diff --git a/web/src/assets/main.css b/web/src/assets/main.css index 238be927..2deb3a7b 100644 --- a/web/src/assets/main.css +++ b/web/src/assets/main.css @@ -8,7 +8,7 @@ .layout-container { width: 100%; - padding: 16px 30px; + padding: 0px 30px; background-color: #FCFEFF; h2 { diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 37374d1a..3b72d028 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -107,11 +107,12 @@ :class="message.role" >

{{ message.text }}

-
+
+
请求错误,请重试

{ } return readChunk() }) + .catch((error) => { + console.error(error) + updateStatus(cur_res_id, "error") + isStreaming.value = false + }) } else { console.log('请输入消息') } @@ -592,6 +599,15 @@ watch( color: black; /* box-shadow: 0px 0.3px 0.9px rgba(0, 0, 0, 0.12), 0px 1.6px 3.6px rgba(0, 0, 0, 0.16); */ /* animation: slideInUp 0.1s ease-in; */ + + .err-msg { + color: red; + border: 1px solid red; + padding: 0.2rem 1rem; + border-radius: 8px; + text-align: center; + background: #FFEBEE; + } } .message-box.sent { @@ -624,8 +640,6 @@ watch( word-wrap: break-word; margin-bottom: 0; } - - } diff --git a/web/src/components/RefsComponent.vue b/web/src/components/RefsComponent.vue index 40249865..895411f9 100644 --- a/web/src/components/RefsComponent.vue +++ b/web/src/components/RefsComponent.vue @@ -1,8 +1,8 @@