diff --git a/.gitignore b/.gitignore index 409b49b2..f602ffa0 100644 --- a/.gitignore +++ b/.gitignore @@ -40,4 +40,5 @@ local_neo4j/data local_neo4j/logs local_neo4j/import local_neo4j/plugins -local_neo4j/conf \ No newline at end of file +local_neo4j/conf +graphrag \ No newline at end of file diff --git a/local_neo4j/test_neo4j.py b/local_neo4j/test_neo4j.py index a293e4fc..53b967a0 100644 --- a/local_neo4j/test_neo4j.py +++ b/local_neo4j/test_neo4j.py @@ -1,15 +1,17 @@ from neo4j import GraphDatabase from neo4j.exceptions import ServiceUnavailable, AuthError +# sudo ln -s /snap/core22/1586/usr/sbin/iptables /usr/sbin/iptables + 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" 表示认证失败。 """ diff --git a/src/config/__init__.py b/src/config/__init__.py index 50ae6ce7..1828e406 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -49,7 +49,7 @@ class Config(SimpleConfig): # 模型配置 ## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径 ## 如果需要自定义路径,则在 config/base.yaml 中配置 model_local_paths - self.add_item("model_provider", default="zhipu", des="模型提供商", choices=["qianfan", "vllm", "zhipu", "deepseek", "dashscope"]) + self.add_item("model_provider", default="zhipu", des="模型提供商", choices=["openai", "qianfan", "vllm", "zhipu", "deepseek", "dashscope"]) self.add_item("model_name", default=None, des="模型名称") self.add_item("embed_model", default="zhipu-embedding-3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys())) self.add_item("reranker", default="bge-reranker-v2-m3", des="Re-Ranker 模型", choices=["bge-reranker-v2-m3"]) @@ -145,6 +145,13 @@ class Config(SimpleConfig): MODEL_NAMES = { # https://platform.deepseek.com/api-docs/zh-cn/pricing + "openai": [ + "gpt-4", + "gpt-4o", + "gpt-4-0125-preview", + "gpt-4o-mini", + ], + "deepseek": [ "deepseek-chat", "deepseek-coder" diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 664b4729..f6db98af 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -15,6 +15,71 @@ logger = setup_logger("server-graphbase") warnings.filterwarnings("ignore", category=UserWarning) + +""" +from neo4j import GraphDatabase +import random + +class KnowledgeGraph: + def __init__(self, uri, user, password): + self._driver = GraphDatabase.driver(uri, auth=(user, password)) + + def close(self): + self._driver.close() + + def use_database(self, kgdb_name): + with self._driver.session() as session: + session.run(f"USE {kgdb_name}") + + def get_sample_nodes(self, kgdb_name='neo4j', num=50): + self.use_database(kgdb_name) + selected_nodes = set() + nodes_to_expand = set() + result_nodes = [] + + with self._driver.session() as session: + while len(selected_nodes) < num: + # 如果需要扩展的节点为空,随机选择一个新节点 + if not nodes_to_expand: + result = session.run("MATCH (n) RETURN n, rand() as r ORDER BY r LIMIT 1") + for record in result: + nodes_to_expand.add(record['n'].id) + result_nodes.append({'n': record['n'], 'r': None, 'm': None}) + + # 从需要扩展的节点中随机选择一个节点 + current_node_id = random.choice(list(nodes_to_expand)) + nodes_to_expand.remove(current_node_id) + + # 获取当前节点的邻居节点,最多5个 + result = session.run( + f"MATCH (n)-[r]-(m) WHERE id(n) = {current_node_id} RETURN n, r, m LIMIT 5" + ) + + for record in result: + neighbor_node_id = record['m'].id + if neighbor_node_id not in selected_nodes: + selected_nodes.add(neighbor_node_id) + nodes_to_expand.add(neighbor_node_id) + result_nodes.append({'n': record['n'], 'r': record['r'], 'm': record['m']}) + + # 如果已经达到最大值,停止扩展 + if len(selected_nodes) >= num: + break + + return result_nodes[:num] + +# 示例用法 +uri = "bolt://localhost:7687" +user = "neo4j" +password = "password" +kg = KnowledgeGraph(uri, user, password) +sample_nodes = kg.get_sample_nodes(num=50) +for node in sample_nodes: + print(f"Node: {node['n'].id}, Relationship: {node['r']}, Neighbor: {node['m'].id}") +kg.close() + +""" + UIE_MODEL = None class GraphDatabase: @@ -38,7 +103,7 @@ class GraphDatabase: self.driver.close() def get_sample_nodes(self, kgdb_name='neo4j', num=50): - """获取指定数据库的前 num 个节点信息""" + """获取指定数据库的 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)) diff --git a/src/core/history.py b/src/core/history.py index 98faa01c..dae392e6 100644 --- a/src/core/history.py +++ b/src/core/history.py @@ -38,5 +38,6 @@ class HistoryManager(): def __str__(self): history_str = "" for message in self.messages: - history_str += f"{message['role']}: {message['content']}" + msg = message["content"].replace('\n', ' ') + history_str += f"\n{message['role']}: {msg}" return history_str \ No newline at end of file diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 327384f0..873f925a 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -63,7 +63,7 @@ class KnowledgeBase: "id": int(random.random() * 1e12), "vector": vectors[i], "text": docs[i], - "hash": hashstr(docs[i] + str(random.random())), + "hash": hashstr(docs[i], with_salt=True), **kwargs} for i in range(len(vectors))] res = self.client.insert(collection_name=collection_name, data=data) diff --git a/src/core/retriever.py b/src/core/retriever.py index 2071f3cb..d2bcd128 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,5 +1,6 @@ from src.models.embedding import Reranker from src.utils.logging_config import setup_logger + logger = setup_logger("server-common") @@ -37,12 +38,14 @@ class Retriever: # 解析图数据库的结果 db_res = refs.get("graph_base").get("results", []) if len(db_res["nodes"]) > 0: - db_text = '\n'.join([f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res['edges']]) + db_text = "\n".join( + [f"{edge['source_name']}和{edge['target_name']}的关系是{edge['type']}" for edge in db_res["edges"]] + ) external += f"图数据库信息: \n\n{db_text}" # 构造查询 if len(external) > 0: - query = f"以下是参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答:{query}" + query = f"参考资料:\n\n\n{external}\n\n\n请根据前面的知识回答问题。\n\n问题:{query}\n\n回答:" return query @@ -70,42 +73,53 @@ class Retriever: kb_res = [] final_res = [] if not refs["meta"].get("db_name") or not self.config.enable_knowledge_base: - return {"results": final_res, "all_results": kb_res, "rw_query": query, "message": "Knowledge base is disabled"} + return { + "results": final_res, + "all_results": kb_res, + "rw_query": query, + "message": "Knowledge base is disabled", + } rw_query = self.rewrite_query(query, history, refs) db_name = refs["meta"]["db_name"] kb = self.dbm.metaname2db[db_name] - limit = refs["meta"].get("queryCount", 10) - kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=limit) - for r in kb_res: + max_query_count = refs["meta"].get("maxQueryCount", 10) + rerank_threshold = refs["meta"].get("rerankThreshold", 0.1) + distance_threshold = refs["meta"].get("distanceThreshold", 0) + top_k = refs["meta"].get("topK", 5) + + all_kb_res = self.dbm.knowledge_base.search(rw_query, db_name, limit=max_query_count) + for r in all_kb_res: r["file"] = kb.id2file(r["entity"]["file_id"]) + # use distance threshold to filter results + kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold] + if self.config.enable_reranker: - RERANK_THRESHOLD = 0.001 for r in kb_res: - r["rerank_score"] = self.reranker.compute_score([query, r["entity"]["text"]], normalize=True) + r["rerank_score"] = self.reranker.compute_score([rw_query, r["entity"]["text"]], normalize=True) kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) - final_res = [_res for _res in kb_res if _res["rerank_score"] > RERANK_THRESHOLD] + kb_res = [_res for _res in kb_res if _res["rerank_score"] > rerank_threshold] - else: - final_res = kb_res[:5] + kb_res = kb_res[:top_k] - return {"results": final_res, "all_results": kb_res, "rw_query": rw_query} + return {"results": kb_res, "all_results": all_kb_res, "rw_query": rw_query} def rewrite_query(self, query, history, refs): """重写查询""" - rewrite_query_span = refs["meta"].get("rewrite_query", None) - if rewrite_query_span is None or rewrite_query_span == "OFF": + rewrite_query_span = refs["meta"].get("rewriteQuery", "off") + if rewrite_query_span == "off": rewritten_query = query else: from src.utils.prompts import rewritten_query_prompt_template - history_query = [entry['content'] for entry in history if entry['role'] == 'user'] if history else "" + + history_query = [entry["content"] for entry in history if entry["role"] == "user"] if history else "" rewritten_query_prompt = rewritten_query_prompt_template.format(history=history_query, query=query) rewritten_query = self.model.predict(rewritten_query_prompt).content - if rewrite_query_span == "HyDE": + if rewrite_query_span == "hyde": hy_doc = self.model.predict(rewritten_query).content rewritten_query = f"{rewritten_query} {hy_doc}" @@ -118,54 +132,68 @@ class Retriever: entities = [] if refs["meta"].get("use_graph"): from src.utils.prompts import entity_extraction_prompt_template + entity_extraction_prompt = entity_extraction_prompt_template.format(text=query) entities = self.model.predict(entity_extraction_prompt).content.split(",") - entities = [entity for entity in entities if all(char.isalnum() or char in '汉字' for char in entity)] + entities = [entity for entity in entities if all(char.isalnum() or char in "汉字" for char in entity)] return entities - def foramt_general_results(self, results): - logger.debug(f"Formatting general results: {results}") + def _extract_relationship_info(self, relationship, source_name, target_name): + """ + 提取关系信息并返回格式化的节点和边信息 + """ + rel_id = relationship.element_id + nodes = relationship.nodes + if len(nodes) != 2: + return None, None + + source, target = nodes + source_id = source.element_id + target_id = target.element_id + + relationship_type = relationship._properties.get("type", "unknown") + if relationship_type == "unknown": + relationship_type = relationship.type + + edge_info = { + "id": rel_id, + "type": relationship_type, + "source_id": source_id, + "target_id": target_id, + "source_name": source_name, + "target_name": target_name, + } + + node_info = [ + {"id": source_id, "name": source_name}, + {"id": target_id, "name": target_name}, + ] + + return node_info, edge_info + + def format_general_results(self, results): formatted_results = {"nodes": [], "edges": []} for item in results: relationship = item[1] - rel_id = relationship.element_id - nodes = relationship.nodes - if len(nodes) != 2: + source_name = item[0]._properties.get("name", "unknown") + target_name = item[2]._properties.get("name", "unknown") if len(item) > 2 else "unknown" + + node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name) + if node_info is None or edge_info is None: continue - source, target = nodes + for node in node_info: + if node["id"] not in [n["id"] for n in formatted_results["nodes"]]: + formatted_results["nodes"].append(node) - 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 - }) + formatted_results["edges"].append(edge_info) return formatted_results def format_query_results(self, results): - logger.debug(f"Formatting query results: {results}") formatted_results = {"nodes": [], "edges": []} - node_dict = {} for item in results: @@ -173,35 +201,18 @@ class Retriever: continue relationship = item[1][0] - rel_id = relationship.element_id - nodes = relationship.nodes - if len(nodes) != 2: + source_name = item[0] + target_name = item[2] if len(item) > 2 else "unknown" + + node_info, edge_info = self._extract_relationship_info(relationship, source_name, target_name) + if node_info is None or edge_info is None: continue - source, target = nodes + for node in node_info: + if node["id"] not in node_dict: + node_dict[node["id"]] = node - source_id = source.element_id - target_id = target.element_id - source_name = item[0] - target_name = item[2] if len(item) > 2 else 'unknown' - - if source_id not in node_dict: - node_dict[source_id] = {"id": source_id, "name": source_name} - if target_id not in node_dict: - node_dict[target_id] = {"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 - }) + formatted_results["edges"].append(edge_info) formatted_results["nodes"] = list(node_dict.values()) @@ -210,4 +221,4 @@ class Retriever: def __call__(self, query, history, meta): refs = self.retrieval(query, history, meta) query = self.construct_query(query, refs, meta) - return query, refs \ No newline at end of file + return query, refs diff --git a/src/models/README.md b/src/models/README.md index 433e24ce..fe1fc347 100644 --- a/src/models/README.md +++ b/src/models/README.md @@ -38,7 +38,7 @@ python -m vllm.entrypoints.openai.api_server \ |`bge-large-zh-v1.5`|`BAAI/bge-large-zh-v1.5`|`bge-large-zh-v1.5`(*修改为本地路径)| |`zhipu`|`embedding-2`|`ZHIPUAPI` (`.env`)| -### 3. 重排序模型支持 + 例如(`config/base.yaml`): @@ -52,3 +52,7 @@ model_name: null # for default model_local_paths: bge-large-zh-v1.5: /models/bge-large-zh-v1.5 ``` + +### 3. 重排序模型支持 + +目前仅支持 `BAAI/bge-reranker-v2-m3`。 diff --git a/src/models/__init__.py b/src/models/__init__.py index 857e0b80..8e1f6592 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -28,6 +28,10 @@ def select_model(config): from src.models.chat_model import DashScope return DashScope(model_name) + elif model_provider == "openai": + from src.models.chat_model import OpenModel + return OpenModel(model_name) + elif model_provider is None: raise ValueError("Model provider not specified, please modify `model_provider` in `src/config/base.yaml`") else: diff --git a/src/models/chat_model.py b/src/models/chat_model.py index be612215..2ed63c18 100644 --- a/src/models/chat_model.py +++ b/src/models/chat_model.py @@ -39,6 +39,14 @@ class OpenAIBase(): return response.choices[0].message +class OpenModel(OpenAIBase): + def __init__(self, model_name=None): + model_name = model_name or "gpt-4o-mini" + api_key = os.getenv("OPENAI_API_KEY") + base_url = os.getenv("OPENAI_API_BASE") + super().__init__(api_key=api_key, base_url=base_url, model_name=model_name) + + class DeepSeek(OpenAIBase): def __init__(self, model_name=None): model_name = model_name or "deepseek-chat" diff --git a/src/models/dify_fording.py b/src/models/dify_fording.py deleted file mode 100644 index e69de29b..00000000 diff --git a/src/models/embedding.py b/src/models/embedding.py index 9eb81ee2..21d32a03 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -1,12 +1,16 @@ import os +import uuid from FlagEmbedding import FlagModel, FlagReranker from src.config import EMBED_MODEL_INFO, RERANKER_LIST from src.utils.logging_config import setup_logger +from src.utils import hashstr logger = setup_logger("EmbeddingModel") +GLOBAL_EMBED_STATE = {} + class EmbeddingModel(FlagModel): def __init__(self, model_info, config, **kwargs): @@ -47,9 +51,21 @@ class ZhipuEmbedding: def predict(self, message): data = [] + if len(message) > 10: + global GLOBAL_EMBED_STATE + task_id = hashstr(message) + logger.info(f"Creating new state for process {task_id}") + GLOBAL_EMBED_STATE[task_id] = { + 'status': 'in-progress', + 'total': len(message), + 'progress': 0 + } + for i in range(0, len(message), 10): if len(message) > 10: logger.info(f"Encoding {i} to {i+10} with {len(message)} messages") + GLOBAL_EMBED_STATE[task_id]['progress'] = i + group_msg = message[i:i+10] response = self.client.embeddings.create( model=self.model_info.default_path, @@ -58,6 +74,10 @@ class ZhipuEmbedding: data.extend([a.embedding for a in response.data]) + if len(message) > 10: + GLOBAL_EMBED_STATE[task_id]['progress'] = len(message) + GLOBAL_EMBED_STATE[task_id]['status'] = 'completed' + return data def encode(self, message): diff --git a/src/plugins/oneke.py b/src/plugins/oneke.py index 9f150d26..aea4c5c8 100644 --- a/src/plugins/oneke.py +++ b/src/plugins/oneke.py @@ -16,7 +16,7 @@ logger = setup_logger("OneKE") dotenv.load_dotenv() -MODEL_NAME_OR_PATH = os.path.join(os.getenv('MODEL_ROOT_DIR'), 'OneKE') +MODEL_NAME_OR_PATH = os.path.join(os.getenv('MODEL_ROOT_DIR', './'), 'OneKE') instruction_mapper = { 'NERzh': "你是专门进行实体抽取的专家。请从input中抽取出符合schema定义的实体,不存在的实体类型返回空列表。请按照JSON字符串的格式回答。", diff --git a/src/plugins/pdf2txt.py b/src/plugins/pdf2txt.py index a0e4f255..ae64ea39 100644 --- a/src/plugins/pdf2txt.py +++ b/src/plugins/pdf2txt.py @@ -1,25 +1,53 @@ import os +import uuid import fitz # fitz就是pip install PyMuPDF import cv2 from copy import deepcopy from tqdm import tqdm +from argparse import ArgumentParser +from src.utils import logger + +GOLBAL_STATE = {} def pdf2txt(pdf_path, return_text=False): + if not os.path.exists(pdf_path): + raise FileNotFoundError(f"File not found: {pdf_path}") + + # Importing these modules here to avoid unnecessary imports in other files from paddleocr import PPStructure, save_structure_res from paddleocr.ppstructure.recovery.recovery_to_doc import sorted_layout_boxes, convert_info_docx - output_dir = os.path.join('tmp', 'pdf2txt', os.path.basename(pdf_path).split('.')[0]) + filename = os.path.basename(pdf_path).split('.')[0] + output_dir = os.path.join('saves', 'data', 'pdf2txt', filename) os.makedirs(output_dir, exist_ok=True) table_engine = PPStructure(recovery=True, lang='ch') imgs = convert_imgs(pdf_path, output_dir) - respath = os.path.join(output_dir, 'res.txt') + respath = os.path.join(output_dir, f'{filename}.txt') + + global GOLBAL_STATE + task_id = str(uuid.uuid4()) + if task_id in GOLBAL_STATE: + logger.info(f"Reusing previous state for process {task_id}") + return GOLBAL_STATE[task_id] + else: + logger.info(f"Creating new state for process {task_id}") + GOLBAL_STATE[task_id] = { + 'pdf_path': pdf_path, + 'return_text': return_text, + 'output_dir': output_dir, + 'status': 'in-progress', + 'total': len(imgs), + 'progress': 0 + } + text = [] for img_name in tqdm(imgs, desc='to txt', ncols=100): img = cv2.imread(img_name) result = table_engine(img) + GOLBAL_STATE[task_id]['progress'] += 1 save_structure_res(result, output_dir, "structure_result") @@ -43,6 +71,9 @@ def pdf2txt(pdf_path, return_text=False): whole_text = ''.join(text) with open(respath, 'w', encoding='utf-8') as f: f.write(whole_text) + logger.info(f"Extracted text saved to {respath}") + + GOLBAL_STATE[task_id]['status'] = 'completed' if return_text: return whole_text @@ -72,6 +103,13 @@ def convert_imgs(pdf_path, output_dir): return imgs +def get_state(task_id): + return GOLBAL_STATE.get(task_id, {}) + if __name__ == "__main__": - pdf_path = r'saves/data/uploads/2e04d5_保健食品.pdf' - print(pdf2txt(pdf_path)) + parser = ArgumentParser() + parser.add_argument('--pdf-path', type=str, required=True, help='Path to the PDF file') + parser.add_argument('--return-text', action='store_true', help='Return the extracted text') + args = parser.parse_args() + + pdf2txt(args.pdf_path, args.return_text) diff --git a/src/utils/__init__.py b/src/utils/__init__.py index 1df9e7e6..5b613e7a 100644 --- a/src/utils/__init__.py +++ b/src/utils/__init__.py @@ -1,4 +1,5 @@ import time +import random from src.utils.logging_config import setup_logger, logger def is_text_pdf(pdf_path): @@ -15,7 +16,7 @@ def hashstr(input_string, length=8, with_salt=False): import hashlib # 添加时间戳作为干扰 if with_salt: - input_string += str(time.time()) + input_string += str(time.time() + random.random()) hash = hashlib.md5(str(input_string).encode()).hexdigest() return hash[:length] \ No newline at end of file diff --git a/src/utils/logging_config.py b/src/utils/logging_config.py index 2eb64b50..67a4d1f9 100644 --- a/src/utils/logging_config.py +++ b/src/utils/logging_config.py @@ -5,10 +5,10 @@ from datetime import datetime DATETIME = datetime.now().strftime('%Y-%m-%d-%H%M%S') # DATETIME = "debug" # 为了方便,调试的时候输出到 debug.log 文件 -LOG_FILE = f'log/project-{DATETIME}.log' +LOG_FILE = f'saves/log/project-{DATETIME}.log' def setup_logger(name, level=logging.DEBUG, console=False): - os.makedirs("log", exist_ok=True) + os.makedirs("saves/log", exist_ok=True) """Function to setup logger with the given name and log file.""" logger = logging.getLogger(name) diff --git a/src/views/common_view.py b/src/views/common_view.py index 3eec200f..730ee785 100644 --- a/src/views/common_view.py +++ b/src/views/common_view.py @@ -1,3 +1,4 @@ +import itertools import json from flask import Blueprint, jsonify, request, Response @@ -89,6 +90,7 @@ def restart(): def get_log(): from src.utils.logging_config import LOG_FILE with open(LOG_FILE, 'r') as f: - log = f.read() + log_lines = itertools.islice(f, 1000) + log = ''.join(log_lines) return jsonify({"log": log}) \ No newline at end of file diff --git a/src/views/database_view.py b/src/views/database_view.py index 98adcfae..5f5057c6 100644 --- a/src/views/database_view.py +++ b/src/views/database_view.py @@ -135,7 +135,7 @@ def get_graph_nodes(): 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 + return jsonify({'result': startup.retriever.format_general_results(result), 'message': 'success'}), 200 @db.route('/graph/add', methods=['POST']) def add_graph_entity(): diff --git a/web/public/icon-narrow.png b/web/public/icon-narrow.png new file mode 100644 index 00000000..5fa011b2 Binary files /dev/null and b/web/public/icon-narrow.png differ diff --git a/web/public/icon.png b/web/public/icon.png new file mode 100644 index 00000000..8b200a98 Binary files /dev/null and b/web/public/icon.png differ diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index 3b72d028..dd602682 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -30,10 +30,10 @@ @@ -75,21 +76,23 @@ import { Graph } from "@antv/g6"; import { computed, onMounted, reactive, ref } from 'vue'; import { message } from "ant-design-vue"; import { useConfigStore } from '@/stores/config'; +import { UploadOutlined } from '@ant-design/icons-vue'; const configStore = useConfigStore() let graphInstance -const graph = ref(null) +const graphInfo = ref(null) const container = ref(null); const fileList = ref([]); const sampleNodeCount = ref(100); -const subgraph = reactive({ +const graphData = reactive({ nodes: [], edges: [], }); const state = reactive({ - graphloading: false, + fetching: false, + loadingGraphInfo: false, searchInput: '', searchLoading: false, showModal: false, @@ -98,27 +101,27 @@ const state = reactive({ }) -const loadGraph = () => { - state.graphloading = true +const loadGraphInfo = () => { + state.loadingGraphInfo = true fetch('/api/database/graph', { method: "GET", }) .then(response => response.json()) .then(data => { console.log(data) - graph.value = data.graph - state.graphloading = false + graphInfo.value = data.graph + state.loadingGraphInfo = false }) .catch(error => { console.error(error) message.error(error.message) - state.graphloading = false + state.loadingGraphInfo = false }) } -const graphData = computed(() => { +const getGraphData = () => { return { - nodes: subgraph.nodes.map(node => { + nodes: graphData.nodes.map(node => { return { id: node.id, data: { @@ -126,7 +129,7 @@ const graphData = computed(() => { }, } }), - edges: subgraph.edges.map(edge => { + edges: graphData.edges.map(edge => { return { source: edge.source_id, target: edge.target_id, @@ -136,7 +139,7 @@ const graphData = computed(() => { } }), } -}) +} const addDocumentByFile = () => { state.precessing = true @@ -159,6 +162,7 @@ const addDocumentByFile = () => { }; const loadSampleNodes = () => { + state.fetching = true fetch(`/api/database/graph/nodes?kgdb_name=neo4j&num=${sampleNodeCount.value}`) .then((res) => { if (res.ok) { @@ -168,17 +172,15 @@ const loadSampleNodes = () => { } }) .then((data) => { - subgraph.nodes = data.result.nodes - subgraph.edges = data.result.edges - console.log(data) - console.log(subgraph) - setTimeout(() => { - randerGraph() - }, 500) + graphData.nodes = data.result.nodes + graphData.edges = data.result.edges + console.log(graphData) + randerGraph() }) .catch((error) => { message.error(error.message); }) + .finally(() => state.fetching = false) } const onSearch = () => { @@ -187,6 +189,12 @@ const onSearch = () => { return } + const cur_embed_model = configStore.config.embed_model + if (cur_embed_model !== 'zhipu-embedding-3') { + message.error('当前不支持实体检索,请在设置中选择向量模型为 zhipu-embedding-3') + return + } + state.searchLoading = true fetch(`/api/database/graph/node?entity_name=${state.searchInput}`) .then((res) => { @@ -197,10 +205,13 @@ const onSearch = () => { } }) .then((data) => { - subgraph.nodes = data.result.nodes - subgraph.edges = data.result.edges + graphData.nodes = data.result.nodes + graphData.edges = data.result.edges + if (graphData.nodes.length === 0) { + message.info('未找到相关实体') + } console.log(data) - console.log(subgraph) + console.log(graphData) randerGraph() }) .catch((error) => { @@ -210,54 +221,58 @@ const onSearch = () => { }; const randerGraph = () => { - graphInstance.setData(graphData.value); + + if (graphInstance) { + graphInstance.destroy(); + } + + initGraph(); + graphInstance.setData(getGraphData()); graphInstance.render(); } +const initGraph = () => { + graphInstance = new Graph({ + container: container.value, + width: container.value.offsetWidth, + height: container.value.offsetHeight, + autoFit: true, + autoResize: true, + layout: { + type: 'd3-force', + preventOverlap: true, + kr: 20, + collide: { + strength: 1.0, + }, + }, + node: { + type: 'circle', + style: { + labelText: (d) => d.data.label, + size: 70, + }, + palette: { + field: 'label', + color: 'tableau', + }, + }, + edge: { + type: 'line', + style: { + labelText: (d) => d.data.label, + labelBackground: '#fff', + endArrow: true, + }, + }, + behaviors: ['drag-element', 'zoom-canvas', 'drag-canvas'], + }); + window.addEventListener('resize', randerGraph); +} + onMounted(() => { - loadGraph(); + loadGraphInfo(); loadSampleNodes(); - setTimeout(() => { - if (state.showPage) { - graphInstance = new Graph({ - container: container.value, - width: container.value.offsetWidth, - height: container.value.offsetHeight, - autoFit: true, - autoResize: true, - layout: { - type: 'd3-force', - preventOverlap: true, - kr: 100, - collide: { - strength: 0.5, - }, - }, - node: { - type: 'circle', - style: { - labelText: (d) => d.data.label, - size: 40, - }, - palette: { - field: 'label', - color: 'tableau', - }, - }, - edge: { - type: 'line', - style: { - labelText: (d) => d.data.label, - labelBackground: '#fff', - }, - }, - behaviors: ['drag-element', 'zoom-canvas', 'drag-canvas'], - }); - graphInstance.setData(graphData.value); - graphInstance.render(); - window.addEventListener('resize', randerGraph); - } - }, 400) }); @@ -303,7 +318,7 @@ const handleDrop = (event) => { justify-content: space-between; margin-bottom: 20px; - .actions-left { + .actions-left, .actions-right { display: flex; align-items: center; gap: 10px; @@ -311,7 +326,6 @@ const handleDrop = (event) => { input { width: 100px; - margin-right: 10px; border-radius: 8px; padding: 4px 12px; border: 2px solid var(--main-300); @@ -324,6 +338,7 @@ const handleDrop = (event) => { } button { + border-width: 2px; height: 40px; box-shadow: none; } @@ -343,8 +358,9 @@ const handleDrop = (event) => { margin: 20px 0; border-radius: 16px; width: 100%; - height: 800px; + height: calc(100vh - 200px); resize: horizontal; + overflow: hidden; } .database-empty { diff --git a/web/src/views/SettingView.vue b/web/src/views/SettingView.vue index aef7d244..54b150d2 100644 --- a/web/src/views/SettingView.vue +++ b/web/src/views/SettingView.vue @@ -42,7 +42,8 @@
- {{ items?.embed_model.des }}   + + {{ items?.embed_model.des }}   需要重启 @@ -58,7 +59,8 @@
- {{ items?.reranker.des }}   + + {{ items?.reranker.des }}   需要重启 @@ -89,7 +91,11 @@ />
- {{ items?.enable_knowledge_graph.des }} + {{ items?.enable_knowledge_graph.des }} + + 需要重启 + + { .label { margin-right: 20px; + + button { + margin-left: 10px; + height: 24px; + padding: 0 8px; + font-size: smaller; + } } }