From d26fda40537435c0709bd856df1a5af8a8646a5e Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Thu, 20 Mar 2025 19:51:46 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9B=B4=E6=96=B0=E4=BA=86=E5=BE=88=E5=A4=9A?= =?UTF-8?q?=20-=20=E6=B7=BB=E5=8A=A0=E4=BA=86=E7=B3=BB=E7=BB=9F=E6=8F=90?= =?UTF-8?q?=E7=A4=BA=E8=AF=8D=20-=20=E6=B7=BB=E5=8A=A0=E4=BA=86=E7=9F=A5?= =?UTF-8?q?=E8=AF=86=E5=BA=93=E7=9A=84=E5=8A=A0=E8=BD=BD=E6=96=B9=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/__init__.py | 7 +- src/core/__init__.py | 3 +- src/core/database.py | 305 --------------------------- src/core/filereader.py | 24 --- src/core/graphbase.py | 117 +++++----- src/core/history.py | 6 +- src/core/indexing.py | 49 ++++- src/core/knowledgebase.py | 230 +++++++++++++++++++- src/core/retriever.py | 14 +- src/models/embedding.py | 19 +- src/models/rerank_model.py | 2 +- src/routers/base_router.py | 5 +- src/routers/chat_router.py | 2 +- src/routers/data_router.py | 60 +++--- src/utils/prompts.py | 8 +- web/src/components/ChatComponent.vue | 6 +- web/src/views/DataBaseInfoView.vue | 21 +- web/src/views/DataBaseView.vue | 3 +- web/src/views/GraphView.vue | 2 +- 19 files changed, 423 insertions(+), 460 deletions(-) delete mode 100644 src/core/database.py delete mode 100644 src/core/filereader.py diff --git a/src/__init__.py b/src/__init__.py index acf8ed3e..c90b868d 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -8,8 +8,11 @@ executor = ThreadPoolExecutor() from src.config import Config config = Config() -from src.core import DataBaseManager -dbm = DataBaseManager() +from src.core import KnowledgeBase +knowledge_base = KnowledgeBase() + +from src.core import GraphDatabase +graph_base = GraphDatabase() from src.core.retriever import Retriever retriever = Retriever() \ No newline at end of file diff --git a/src/core/__init__.py b/src/core/__init__.py index d59f3c6e..3a868409 100644 --- a/src/core/__init__.py +++ b/src/core/__init__.py @@ -1,2 +1,3 @@ from .history import * -from .database import * \ No newline at end of file +from .knowledgebase import KnowledgeBase +from .graphbase import GraphDatabase diff --git a/src/core/database.py b/src/core/database.py deleted file mode 100644 index ac08b111..00000000 --- a/src/core/database.py +++ /dev/null @@ -1,305 +0,0 @@ -import os -import json -import time -import traceback - -from src import config -from src.utils import hashstr, logger -from src.core.indexing import chunk -from src.models.embedding import get_embedding_model - - -class DataBaseManager: - - def __init__(self) -> None: - self.database_path = os.path.join(config.save_dir, "data", "database.json") - self._load_models() - - def _load_models(self): - """所有需要重启的模型""" - self.embed_model = get_embedding_model(config) - if config.enable_knowledge_base: - from src.core.knowledgebase import KnowledgeBase - self.knowledge_base = KnowledgeBase(config, self.embed_model) - if config.enable_knowledge_graph: - from src.core.graphbase import GraphDatabase - self.graph_base = GraphDatabase(config, self.embed_model) - else: - self.graph_base = None - - self.data = {"databases": [], "graph": {}} - self._load_databases() - self._update_database() - - def _load_databases(self): - """将数据库的信息保存到本地的文件里面""" - if not os.path.exists(self.database_path): - return - - with open(self.database_path, "r") as f: - data = json.load(f) - self.data = { - "databases": [DataBaseLite(**db) for db in data["databases"]], - "graph": data["graph"] - } - - # 检查所有文件,如果出现状态是 processing 的,那么设置为 failed - for db in self.data["databases"]: - for file in db.files: - if file["status"] == "processing" or file["status"] == "waiting": - file["status"] = "failed" - - def _save_databases(self): - """将数据库的信息保存到本地的文件里面""" - self._update_database() - os.makedirs(os.path.dirname(self.database_path), exist_ok=True) - with open(self.database_path, "w+") as f: - json.dump({ - "databases": [db.to_dict() for db in self.data["databases"]], - "graph": self.data["graph"] - }, f, ensure_ascii=False, indent=4) - - def _update_database(self): - self.id2db = {db.db_id: db for db in self.data["databases"]} - self.name2db = {db.name: db for db in self.data["databases"]} - self.metaname2db = {db.metaname: db for db in self.data["databases"]} - - def get_databases(self): - self._update_database() - assert config.enable_knowledge_base, "知识库未启用" - knowledge_base_collections = self.knowledge_base.get_collection_names() - if len(self.data["databases"]) != len(knowledge_base_collections): - logger.warning( - f"Database number not match, {knowledge_base_collections}, " - f"self.data['databases']: {self.get_db_metanames()}, ") - - # 更新每个数据库的状态信息 - for db in self.data["databases"]: - # 获取最新的集合信息 - db.update(self.knowledge_base.get_collection_info(db.metaname)) - - # 检查文件处理状态 - processing_files = [f for f in db.files if f["status"] in ["processing", "waiting"]] - if processing_files: - logger.info(f"数据库 {db.name} 有 {len(processing_files)} 个文件正在处理中") - - return {"databases": [db.to_dict() for db in self.data["databases"]]} - - def get_graph(self): - if 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 is_graph_running(self): - """检查图数据库是否正在运行 - - Returns: - bool: 图数据库是否正在运行 - """ - # 检查是否启用了图数据库 - if not config.enable_knowledge_graph or not hasattr(self, 'graph_base') or self.graph_base is None: - return False - - # 获取图数据库信息,检查状态 - graph_info = self.graph_base.get_database_info("neo4j") - return graph_info.get("status") == "open" - - def create_database(self, database_name, description, db_type, dimension): - from src.config import EMBED_MODEL_INFO - dimension = dimension or EMBED_MODEL_INFO[config.embed_model]["dimension"] - - new_database = DataBaseLite(database_name, - description, - db_type, - embed_model=config.embed_model, - dimension=dimension) - - self.knowledge_base.add_collection(new_database.metaname, dimension) - self.data["databases"].append(new_database) - self._save_databases() - return self.get_databases() - - def add_files(self, db_id, files, params=None): - db = self.get_kb_by_id(db_id) - - if db.embed_model != config.embed_model: - logger.error(f"Embed model not match, {db.embed_model} != {config.embed_model}") - return {"message": f"Embed model not match, cur: {config.embed_model}", "status": "failed"} - - # Preprocessing the files to the queue - new_files = [] - for file in files: - new_file = { - "file_id": "file_" + hashstr(file + str(time.time())), - "filename": os.path.basename(file), - "path": file, - "type": file.split(".")[-1].lower(), - "status": "waiting", - "created_at": time.time() - } - db.files.append(new_file) - new_files.append(new_file) - - # 先保存一次数据库状态,确保waiting状态被记录 - self._save_databases() - - for new_file in new_files: - file_id = new_file["file_id"] - idx = self.get_idx_by_fileid(db, file_id) - db.files[idx]["status"] = "processing" - # 更新处理状态 - self._save_databases() - - try: - if new_file["type"] == "pdf": - texts = self.read_text(new_file["path"]) - nodes = chunk(texts, params=params) - else: - nodes = chunk(new_file["path"], params=params) - - self.knowledge_base.add_documents( - file_id=file_id, - collection_name=db.metaname, - docs=[node.text for node in nodes], - chunk_infos=[node.dict() for node in nodes]) - - idx = self.get_idx_by_fileid(db, file_id) - db.files[idx]["status"] = "done" - - except Exception as e: - logger.error(f"Failed to add documents to collection {db.metaname}, {e}, {traceback.format_exc()}") - idx = self.get_idx_by_fileid(db, file_id) - db.files[idx]["status"] = "failed" - - # 每个文件处理完成后立即保存数据库状态 - self._save_databases() - - - def get_database_info(self, db_id): - db = self.get_kb_by_id(db_id) - if db is None: - return None - else: - db.update(self.knowledge_base.get_collection_info(db.metaname)) - return db.to_dict() - - def read_text(self, file, params=None): - support_format = [".pdf", ".txt", ".md"] - assert os.path.exists(file), "File not found" - logger.info(f"Try to read file {file}") - - if not os.path.isfile(file): - logger.error(f"Directory not supported now!") - raise NotImplementedError("Directory not supported now!") - - if file.endswith(".pdf"): - from src.plugins import ocr - return ocr.process_pdf(file) - - elif file.endswith(".txt") or file.endswith(".md"): - from src.core.filereader import plainreader - return plainreader(file) - - else: - 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 = self.get_idx_by_fileid(db, file_id) - - self.knowledge_base.client.delete( - collection_name=db.metaname, - filter=f"file_id == '{file_id}'"), - - del db.files[file_idx_to_delete] - self._save_databases() - - def get_file_info(self, db_id, file_id): - db = self.get_kb_by_id(db_id) - if db is None: - return {"message": "database not found"}, 404 - lines = self.knowledge_base.client.query( - collection_name=db.metaname, - filter=f"file_id == '{file_id}'", - output_fields=None - ) - # 删除 vector 字段 - for line in lines: - line.pop("vector") - - lines.sort(key=lambda x: x.get("start_char_idx") or 0) - # logger.debug(f"lines[0]: {lines[0]}") - return {"lines": lines} - - def get_db_metanames(self): - return [db.metaname for db in self.data["databases"]] - - def delete_database(self, db_id): - db = self.get_kb_by_id(db_id) - if db is None: - return {"message": "database not found"}, 404 - - self.knowledge_base.client.drop_collection(db.metaname) - self.data["databases"] = [d for d in self.data["databases"] if d.db_id != db_id] - self._save_databases() - return {"message": "删除成功"} - - def get_kb_by_id(self, db_id): - for db in self.data["databases"]: - if db.db_id == db_id: - return db - return None - - def get_idx_by_fileid(self, db, file_id): - for idx, f in enumerate(db.files): - if f["file_id"] == file_id: - return idx - - def restart(self): - self.embed_model = get_embedding_model(config) - self._load_databases() - self._update_database() - - -class DataBaseLite: - def __init__(self, name, description, db_type, dimension=None, **kwargs) -> None: - self.name = name - self.description = description - self.db_type = db_type - self.dimension = dimension - self.db_id = kwargs.get("db_id", hashstr(name)) - self.metaname = kwargs.get("metaname", f"{db_type[:1]}{hashstr(name)}") - self.metadata = kwargs.get("metadata", {}) - self.files = kwargs.get("files", []) - self.embed_model = kwargs.get("embed_model", None) - - def id2file(self, file_id): - for f in self.files: - if f["file_id"] == file_id: - return f - return None - - def update(self, metadata): - self.metadata = metadata - - def to_dict(self): - return { - "name": self.name, - "description": self.description, - "db_type": self.db_type, - "db_id": self.db_id, - "embed_model": self.embed_model, - "metaname": self.metaname, - "metadata": self.metadata, - "files": self.files, - "dimension": self.dimension - } - - def to_json(self): - return json.dumps(self.to_dict(), ensure_ascii=False) - - def __str__(self): - return self.to_json() \ No newline at end of file diff --git a/src/core/filereader.py b/src/core/filereader.py deleted file mode 100644 index 62fd11a6..00000000 --- a/src/core/filereader.py +++ /dev/null @@ -1,24 +0,0 @@ -import os - -from pathlib import Path - -def pdfreader(file_path): - """读取PDF文件并返回text文本""" - assert os.path.exists(file_path), "File not found" - assert file_path.endswith(".pdf"), "File format not supported" - - from llama_index.readers.file import PDFReader - doc = PDFReader().load_data(file=Path(file_path)) - - # 简单的拼接起来之后返回纯文本 - text = "\n\n".join([d.get_content() for d in doc]) - return text - -def plainreader(file_path): - """读取普通文本文件并返回text文本""" - assert os.path.exists(file_path), "File not found" - - with open(file_path, "r") as f: - text = f.read() - return text - diff --git a/src/core/graphbase.py b/src/core/graphbase.py index 15a162eb..178d6ef5 100644 --- a/src/core/graphbase.py +++ b/src/core/graphbase.py @@ -6,6 +6,7 @@ import traceback import torch from neo4j import GraphDatabase as GD +from src import config from src.utils import logger warnings.filterwarnings("ignore", category=UserWarning) @@ -14,16 +15,13 @@ warnings.filterwarnings("ignore", category=UserWarning) UIE_MODEL = None class GraphDatabase: - def __init__(self, config, embed_model=None, kgdb_name="neo4j"): - self.config = config + def __init__(self): self.driver = None self.files = [] self.status = "closed" - self.kgdb_name = kgdb_name - assert embed_model, "embed_model=None" - self.embed_model = embed_model + self.kgdb_name = "neo4j" self.embed_model_name = None - self.work_dir = os.path.join(config.save_dir, "knowledge_graph", kgdb_name) + self.work_dir = os.path.join(config.save_dir, "knowledge_graph", self.kgdb_name) os.makedirs(self.work_dir, exist_ok=True) # 尝试加载已保存的图数据库信息 @@ -33,6 +31,8 @@ class GraphDatabase: self.start() def start(self): + if not config.enable_knowledge_graph or not config.enable_knowledge_base: + return uri = os.environ.get("NEO4J_URI", "bolt://localhost:7687") username = os.environ.get("NEO4J_USERNAME", "neo4j") password = os.environ.get("NEO4J_PASSWORD", "0123456789") @@ -40,9 +40,9 @@ class GraphDatabase: try: self.driver = GD.driver(f"{uri}/{self.kgdb_name}", auth=(username, password)) self.status = "open" - logger.info(f"Connected to Neo4j at {uri}/{self.kgdb_name}, {self.get_database_info()}") + logger.info(f"Connected to Neo4j at {uri}/{self.kgdb_name}, {self.get_graph_info(self.kgdb_name)}") # 连接成功后保存图数据库信息 - self.save_graph_info() + self.save_graph_info(self.kgdb_name) except Exception as e: logger.error(f"Failed to connect to Neo4j: {e}, {uri}, {self.kgdb_name}, {username}, {password}") self.config.enable_knowledge_graph = False @@ -51,6 +51,13 @@ class GraphDatabase: """关闭数据库连接""" self.driver.close() + def is_running(self): + """检查图数据库是否正在运行""" + if not config.enable_knowledge_graph or not config.enable_knowledge_base: + return False + else: + return self.status == "open" + def get_sample_nodes(self, kgdb_name='neo4j', num=50): """获取指定数据库的 num 个节点信息""" self.use_database(kgdb_name) @@ -75,30 +82,6 @@ class GraphDatabase: print(f"数据库 '{kgdb_name}' 创建成功.") return kgdb_name # 返回创建的数据库名称 - def get_database_info(self, db_name="neo4j"): - """获取指定数据库的信息""" - self.use_database(db_name) - def query(tx): - entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"] - relationship_count = tx.run("MATCH ()-[r]->() RETURN count(r) AS count").single()["count"] - triples_count = tx.run("MATCH (n)-[r]->(m) RETURN count(n) AS count").single()["count"] - - # 获取所有标签 - labels = tx.run("CALL db.labels() YIELD label RETURN collect(label) AS labels").single()["labels"] - - return { - "database_name": db_name, - "entity_count": entity_count, - "relationship_count": relationship_count, - "triples_count": triples_count, - "labels": labels, - "status": self.status, - "embed_model_name": self.embed_model_name - } - - with self.driver.session() as session: - return session.execute_read(query) - def use_database(self, kgdb_name="neo4j"): """切换到指定数据库""" assert kgdb_name == self.kgdb_name, f"传入的数据库名称 '{kgdb_name}' 与当前实例的数据库名称 '{self.kgdb_name}' 不一致" @@ -159,7 +142,7 @@ class GraphDatabase: # 判断模型名称是否匹配 from src.config import EMBED_MODEL_INFO - cur_embed_info = EMBED_MODEL_INFO[self.config.embed_model] + cur_embed_info = EMBED_MODEL_INFO[config.embed_model] self.embed_model_name = self.embed_model_name or cur_embed_info.get('name') assert self.embed_model_name == cur_embed_info.get('name') or self.embed_model_name is None, \ f"embed_model_name={self.embed_model_name}, {cur_embed_info.get('name')=}" @@ -167,7 +150,7 @@ class GraphDatabase: with self.driver.session() as session: logger.info(f"Adding entity to {kgdb_name}") session.execute_write(_create_graph, triples) - logger.info(f"Creating vector index for {kgdb_name} with {self.config.embed_model}") + logger.info(f"Creating vector index for {kgdb_name} with {config.embed_model}") session.execute_write(_create_vector_index, cur_embed_info['dimension']) # NOTE 这里需要异步处理 for i, entry in enumerate(triples): @@ -329,7 +312,8 @@ class GraphDatabase: def get_embedding(self, text): with torch.no_grad(): - outputs = self.embed_model.encode([text])[0] + from src import knowledge_base + outputs = knowledge_base.embed_model.encode([text])[0] return outputs def set_embedding(self, tx, entity_name, embedding): @@ -338,34 +322,53 @@ class GraphDatabase: CALL db.create.setNodeVectorProperty(e, 'embedding', $embedding) """, name=entity_name, embedding=embedding) - def save_graph_info(self): + def get_graph_info(self, graph_name="neo4j"): + self.use_database(graph_name) + def query(tx): + entity_count = tx.run("MATCH (n) RETURN count(n) AS count").single()["count"] + relationship_count = tx.run("MATCH ()-[r]->() RETURN count(r) AS count").single()["count"] + triples_count = tx.run("MATCH (n)-[r]->(m) RETURN count(n) AS count").single()["count"] + + # 获取所有标签 + labels = tx.run("CALL db.labels() YIELD label RETURN collect(label) AS labels").single()["labels"] + + return { + "graph_name": graph_name, + "entity_count": entity_count, + "relationship_count": relationship_count, + "triples_count": triples_count, + "labels": labels, + "status": self.status, + "embed_model_name": self.embed_model_name, + "unindexed_node_count": self.query_nodes_without_embedding(graph_name) + } + + try: + if self.status == "open" and self.driver and self.is_running(): + # 获取数据库信息 + with self.driver.session() as session: + graph_info = session.execute_read(query) + + # 添加时间戳 + from datetime import datetime + graph_info["last_updated"] = datetime.now().isoformat() + return graph_info + + except Exception as e: + logger.error(f"获取图数据库信息失败:{e}, {traceback.format_exc()}") + return None + + def save_graph_info(self, graph_name="neo4j"): """ 将图数据库的基本信息保存到工作目录中的JSON文件 保存的信息包括:数据库名称、状态、嵌入模型名称等 """ try: - # 获取数据库信息 - db_info = None - if self.status == "open" and self.driver: - try: - db_info = self.get_database_info(self.kgdb_name) - except Exception as e: - logger.warning(f"无法获取数据库信息:{e}") + graph_info = self.get_graph_info(graph_name) + if graph_info is None: + logger.error(f"图数据库信息为空,无法保存") + return False - # 构建要保存的信息字典 - graph_info = { - "kgdb_name": self.kgdb_name, - "status": self.status, - "embed_model_name": self.embed_model_name, - "last_updated": None, # 这里可以添加时间戳 - "database_info": db_info - } - - # 添加时间戳 - from datetime import datetime - graph_info["last_updated"] = datetime.now().isoformat() - - # 保存到JSON文件 info_file_path = os.path.join(self.work_dir, "graph_info.json") with open(info_file_path, 'w', encoding='utf-8') as f: json.dump(graph_info, f, ensure_ascii=False, indent=2) diff --git a/src/core/history.py b/src/core/history.py index dae392e6..277a8d38 100644 --- a/src/core/history.py +++ b/src/core/history.py @@ -1,9 +1,13 @@ +from src.utils.prompts import get_system_prompt from src.utils.logging_config import logger class HistoryManager(): - def __init__(self, history=None): + def __init__(self, history=None, system_prompt=None): self.messages = history or [] + system_prompt = system_prompt or get_system_prompt() + self.add_system(system_prompt) + def add(self, role, content): self.messages.append({"role": role, "content": content}) return self.messages diff --git a/src/core/indexing.py b/src/core/indexing.py index 5dd74f7e..12853bdf 100644 --- a/src/core/indexing.py +++ b/src/core/indexing.py @@ -5,12 +5,13 @@ from llama_index.core.node_parser import SimpleFileNodeParser from llama_index.core.node_parser import SentenceSplitter from llama_index.readers.file import FlatReader, DocxReader -from src.utils import hashstr +from src.utils import hashstr, logger + def chunk(text_or_path, params=None): """ 将文本或文件切分成固定大小的块 - + Args: text_or_path: 文本或文件路径 params: 参数 @@ -48,3 +49,47 @@ def chunk(text_or_path, params=None): nodes = splitter.get_nodes_from_documents(docs) return nodes + + + +def pdfreader(file_path): + """读取PDF文件并返回text文本""" + assert os.path.exists(file_path), "File not found" + assert file_path.endswith(".pdf"), "File format not supported" + + from llama_index.readers.file import PDFReader + doc = PDFReader().load_data(file=Path(file_path)) + + # 简单的拼接起来之后返回纯文本 + text = "\n\n".join([d.get_content() for d in doc]) + return text + +def plainreader(file_path): + """读取普通文本文件并返回text文本""" + assert os.path.exists(file_path), "File not found" + + with open(file_path, "r") as f: + text = f.read() + return text + +def read_text(file, params=None): + support_format = [".pdf", ".txt", ".md"] + assert os.path.exists(file), "File not found" + logger.info(f"Try to read file {file}") + + if not os.path.isfile(file): + logger.error(f"Directory not supported now!") + raise NotImplementedError("Directory not supported now!") + + if file.endswith(".pdf"): + from src.plugins import ocr + return ocr.process_pdf(file) + + elif file.endswith(".txt") or file.endswith(".md"): + return plainreader(file) + + else: + logger.error(f"File format not supported, only support {support_format}") + raise Exception(f"File format not supported, only support {support_format}") + + diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index de64eb65..dbce831b 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -1,27 +1,202 @@ import os +import json +import time +import traceback from pymilvus import MilvusClient, MilvusException + +from src import config from src.utils import logger, hashstr +from src.core.indexing import chunk, read_text + class KnowledgeBase: - def __init__(self, config=None, embed_model=None) -> None: - self.config = config or {} - assert embed_model, "embed_model=None" - self.embed_model = embed_model - + def __init__(self) -> None: + self.data = [] self.client = None + self.database_path = os.path.join(config.save_dir, "data", "database.json") + self._load_models() + self._load_databases() + + def _load_models(self): + """所有需要重启的模型""" + if not config.enable_knowledge_base: + return + + from src.models.embedding import get_embedding_model + self.embed_model = get_embedding_model(config) + if not self.connect_to_milvus(): raise ConnectionError("Failed to connect to Milvus") + def _load_databases(self): + """将数据库的信息保存到本地的文件里面""" + if not os.path.exists(self.database_path): + return + + with open(self.database_path, "r") as f: + data = json.load(f) + self.data = [DataBaseLite(**db) for db in data["databases"]] + + self._update_database() + + def _save_databases(self): + """将数据库的信息保存到本地的文件里面""" + self._update_database() + os.makedirs(os.path.dirname(self.database_path), exist_ok=True) + with open(self.database_path, "w") as f: + json.dump({ + "databases": [db.to_dict() for db in self.data], + }, f, ensure_ascii=False, indent=4) + + def _update_database(self): + self.id2db = {db.db_id: db for db in self.data} + self.name2db = {db.name: db for db in self.data} + + def create_database(self, database_name, description, dimension=None): + """创建一个数据库""" + dimension = dimension or self.embed_model.get_dimension() + db = DataBaseLite(database_name, + description, + embed_model=self.embed_model.embed_model_fullname, + dimension=dimension) + + self.add_collection(db.db_id, dimension) + self.data.append(db) + self._save_databases() + + def get_databases(self): + assert config.enable_knowledge_base, "知识库未启用" + + for db in self.data: + db.update(self.get_collection_info(db.db_id)) + processing_files = [f for fid, f in db.files.items() if f["status"] in ["processing", "waiting"]] + if processing_files: + logger.info(f"数据库 {db.name} 有 {len(processing_files)} 个文件正在处理中") + + self._save_databases() + return {"databases": [db.to_dict() for db in self.data]} + + def get_database_info(self, db_id): + db = self.get_kb_by_id(db_id) + if db is None: + return None + else: + db.update(self.get_collection_info(db.db_id)) + return db.to_dict() + + def get_database_id(self): + return [db.db_id for db in self.data] + + def get_file_info(self, db_id, file_id): + db = self.get_kb_by_id(db_id) + if db is None: + raise Exception(f"database not found, {db_id}") + + lines = self.client.query( + collection_name=db.db_id, + filter=f"file_id == '{file_id}'", + output_fields=None + ) + # 删除 vector 字段 + for line in lines: + line.pop("vector") + + lines.sort(key=lambda x: x.get("start_char_idx") or 0) + # logger.debug(f"lines[0]: {lines[0]}") + return {"lines": lines} + + def get_kb_by_id(self, db_id): + return next((db for db in self.data if db.db_id == db_id), None) + + def add_files(self, db_id, files, params=None): + db = self.get_kb_by_id(db_id) + + if db.embed_model != config.embed_model: + logger.error(f"Embed model not match, {db.embed_model} != {config.embed_model}") + return {"message": f"Embed model not match, cur: {config.embed_model}", "status": "failed"} + + # Preprocessing the files to the queue + new_files = {} + for file in files: + file_id = "file_" + hashstr(file + str(time.time())) + new_file = { + "file_id": file_id, + "filename": os.path.basename(file), + "path": file, + "type": file.split(".")[-1].lower(), + "status": "waiting", + "created_at": time.time() + } + new_files[file_id] = new_file + + db.files.update(new_files) # 更新数据库状态 + + # 先保存一次数据库状态,确保waiting状态被记录 + self._save_databases() + + for file_id, new_file in new_files.items(): + db.files[file_id]["status"] = "processing" + # 更新处理状态 + self._save_databases() + + try: + if new_file["type"] == "pdf": + texts = read_text(new_file["path"]) + nodes = chunk(texts, params=params) + else: + nodes = chunk(new_file["path"], params=params) + + self.add_documents( + file_id=file_id, + collection_name=db.db_id, + docs=[node.text for node in nodes], + chunk_infos=[node.dict() for node in nodes]) + + db.files[file_id]["status"] = "done" + + except Exception as e: + logger.error(f"Failed to add documents to collection {db.db_id}, {e}, {traceback.format_exc()}") + db.files[file_id]["status"] = "failed" + + # 每个文件处理完成后立即保存数据库状态 + self._save_databases() + + def delete_file(self, db_id, file_id): + db = self.get_kb_by_id(db_id) + if db is None: + raise Exception(f"database not found, {db_id}") + + self.client.delete(collection_name=db.db_id, filter=f"file_id == '{file_id}'") + del db.files[file_id] + self._save_databases() + + def delete_database(self, db_id): + db = self.get_kb_by_id(db_id) + if db is None: + raise Exception(f"database not found, {db_id}") + + self.client.drop_collection(collection_name=db.db_id) + self.data.remove(db) + self._save_databases() + return {"message": "删除成功"} + + def restart(self): + self.embed_model = get_embedding_model(config) + self._load_databases() + + ################################ + # Below is the code for milvus # + ################################ def connect_to_milvus(self): """ 连接到 Milvus 服务。 使用配置中的 URI,如果没有配置,则使用默认值。 """ try: - uri = os.getenv('MILVUS_URI', self.config.get('milvus_uri', "http://milvus:19530")) + uri = os.getenv('MILVUS_URI', config.get('milvus_uri', "http://milvus:19530")) self.client = MilvusClient(uri=uri) # 可以添加一个简单的测试来确保连接成功 self.client.list_collections() @@ -109,3 +284,46 @@ class KnowledgeBase: def search_by_id(self, collection_name, id, output_fields=["id", "text"]): res = self.client.get(collection_name, id, output_fields=output_fields) return res + + +class DataBaseLite: + def __init__(self, name, description, dimension=None, **kwargs) -> None: + self.name = name + self.description = description + self.dimension = dimension + self.metadata = kwargs.get("metadata", {}) + # logger.debug(f"DataBaseLite init: {self.metadata}") + self.db_id = self.metadata.get("collection_name", kwargs.get("db_id")) # metaname 的历史遗留问题 + self.db_id = self.db_id or f"kb_{hashstr(name, with_salt=True)}" + self.files = kwargs.get("files", []) + + if isinstance(self.files, list): + self.files = {f["file_id"]: f for f in self.files} + + self.embed_model = kwargs.get("embed_model", None) + + def id2file(self, file_id): + for f in self.files: + if f["file_id"] == file_id: + return f + return None + + def update(self, metadata): + self.metadata = metadata + + def to_dict(self): + return { + "name": self.name, + "description": self.description, + "db_id": self.db_id, + "embed_model": self.embed_model, + "metadata": self.metadata, + "files": self.files, + "dimension": self.dimension + } + + def to_json(self): + return json.dumps(self.to_dict(), ensure_ascii=False) + + def __str__(self): + return self.to_json() \ No newline at end of file diff --git a/src/core/retriever.py b/src/core/retriever.py index 92444f1d..4e010f6a 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -1,4 +1,4 @@ -from src import config, dbm +from src import config, knowledge_base, graph_base from src.models.rerank_model import get_reranker from src.utils.logging_config import logger from src.models import select_model @@ -76,7 +76,7 @@ class Retriever: results = [] if refs["meta"].get("use_graph") and config.enable_knowledge_base: for entity in refs["entities"]: - result = dbm.graph_base.query_by_vector(entity) + result = graph_base.query_by_vector(entity) if result != []: results.extend(result) return {"results": self.format_query_results(results)} @@ -88,8 +88,8 @@ class Retriever: kb_res = [] final_res = [] - db_name = refs["meta"].get("db_name") - if not db_name or not config.enable_knowledge_base: + db_id = refs["meta"].get("db_id") + if not db_id or not config.enable_knowledge_base: return { "results": final_res, "all_results": kb_res, @@ -99,7 +99,7 @@ class Retriever: rw_query = self.rewrite_query(query, history, refs) - kb = dbm.metaname2db[db_name] + kb = knowledge_base.id2db[db_id] logger.debug(f"{refs['meta']=}") meta = refs["meta"] @@ -109,9 +109,9 @@ class Retriever: top_k = meta.get("topK", 5) # 检索 - all_kb_res = dbm.knowledge_base.search(rw_query, db_name, limit=max_query_count) + all_kb_res = knowledge_base.search(rw_query, db_id, limit=max_query_count) for r in all_kb_res: - r["file"] = kb.id2file(r["entity"]["file_id"]) + r["file"] = kb.files[r["entity"]["file_id"]] kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold] diff --git a/src/models/embedding.py b/src/models/embedding.py index ea10bde9..558de9be 100644 --- a/src/models/embedding.py +++ b/src/models/embedding.py @@ -4,13 +4,22 @@ import requests from FlagEmbedding import FlagModel from zhipuai import ZhipuAI +from src import config from src.config import EMBED_MODEL_INFO from src.utils import hashstr, logger, get_docker_safe_url class BaseEmbeddingModel: embed_state = {} - EMBED_MODEL_INFO = EMBED_MODEL_INFO + + def get_dimension(self): + if hasattr(self, "dimension"): + return self.dimension + + if hasattr("embed_model_fullname"): + return EMBED_MODEL_INFO[self.embed_model_fullname].get("dimension", None) + + return EMBED_MODEL_INFO[self.model].get("dimension", None) def encode(self, message): return self.predict(message) @@ -49,6 +58,8 @@ class LocalEmbeddingModel(FlagModel, BaseEmbeddingModel): self.model = config.model_local_paths.get(info["name"], info.get("local_path")) self.model = self.model or info["name"] + self.dimension = info["dimension"] + self.embed_model_fullname = config.embed_model if os.path.exists(_path := os.path.join(os.getenv("MODEL_DIR"), self.model)): self.model = _path @@ -69,7 +80,9 @@ class ZhipuEmbedding(BaseEmbeddingModel): def __init__(self, config) -> None: self.config = config self.model = EMBED_MODEL_INFO[config.embed_model]["name"] + self.dimension = EMBED_MODEL_INFO[config.embed_model]["dimension"] self.client = ZhipuAI(api_key=os.getenv("ZHIPUAI_API_KEY")) + self.embed_model_fullname = config.embed_model def predict(self, message): response = self.client.embeddings.create( @@ -86,6 +99,8 @@ class OllamaEmbedding(BaseEmbeddingModel): self.model = self.info["name"] self.url = self.info.get("url", "http://localhost:11434/api/embed") self.url = get_docker_safe_url(self.url) + self.dimension = self.info.get("dimension", None) + self.embed_model_fullname = config.embed_model def predict(self, message: list[str] | str): if isinstance(message, str): @@ -105,6 +120,8 @@ class OtherEmbedding(BaseEmbeddingModel): def __init__(self, config) -> None: self.info = EMBED_MODEL_INFO[config.embed_model] + self.embed_model_fullname = config.embed_model + self.dimension = self.info.get("dimension", None) self.model = self.info["name"] self.api_key = os.getenv(self.info["api_key"], None) self.url = get_docker_safe_url(self.info["url"]) diff --git a/src/models/rerank_model.py b/src/models/rerank_model.py index 2bd1c2fc..afe3e124 100644 --- a/src/models/rerank_model.py +++ b/src/models/rerank_model.py @@ -40,7 +40,7 @@ class SilconFlowReranker(): payload = self.build_payload(query, sentences, max_length) response = requests.request("POST", self.url, json=payload, headers=self.headers) response = json.loads(response.text) - logger.debug(f"SiliconFlow Reranker response: {response}") + # logger.debug(f"SiliconFlow Reranker response: {response}") results = sorted(response["results"], key=lambda x: x["index"]) all_scores = [result["relevance_score"] for result in results] diff --git a/src/routers/base_router.py b/src/routers/base_router.py index c33fb9a1..cbf8d1fe 100644 --- a/src/routers/base_router.py +++ b/src/routers/base_router.py @@ -4,7 +4,7 @@ from fastapi import Request, Body base = APIRouter() -from src import config, dbm, retriever +from src import config, retriever, knowledge_base, graph_base from src.utils import logger @@ -27,7 +27,8 @@ async def update_config(key = Body(...), value = Body(...)): @base.post("/restart") async def restart(): - dbm.restart() + knowledge_base.restart() + graph_base.restart() retriever.restart() return {"message": "Restarted!"} diff --git a/src/routers/chat_router.py b/src/routers/chat_router.py index e25776ab..b898a714 100644 --- a/src/routers/chat_router.py +++ b/src/routers/chat_router.py @@ -37,7 +37,7 @@ def chat_post( }, ensure_ascii=False).encode('utf-8') + b"\n" def need_retrieve(meta): - return meta.get("use_web") or meta.get("use_graph") or meta.get("db_name") + return meta.get("use_web") or meta.get("use_graph") or meta.get("db_id") def generate_response(): modified_query = query diff --git a/src/routers/data_router.py b/src/routers/data_router.py index f719225a..71cdb31d 100644 --- a/src/routers/data_router.py +++ b/src/routers/data_router.py @@ -5,7 +5,7 @@ from typing import List, Optional from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body from src.utils import logger, hashstr -from src import executor, dbm, retriever, config +from src import executor, retriever, config, knowledge_base, graph_base data = APIRouter(prefix="/data") @@ -13,8 +13,9 @@ data = APIRouter(prefix="/data") @data.get("/") async def get_databases(): try: - database = dbm.get_databases() + database = knowledge_base.get_databases() except Exception as e: + logger.error(f"获取数据库列表失败 {e}, {traceback.format_exc()}") return {"message": f"获取数据库列表失败 {e}", "databases": []} return database @@ -22,22 +23,24 @@ async def get_databases(): async def create_database( database_name: str = Body(...), description: str = Body(...), - db_type: str = Body(...), dimension: Optional[int] = Body(None) ): logger.debug(f"Create database {database_name}") - database_info = dbm.create_database( - database_name, - description, - db_type, - dimension=dimension - ) + try: + database_info = knowledge_base.create_database( + database_name, + description, + dimension=dimension + ) + except Exception as e: + logger.error(f"创建数据库失败 {e}, {traceback.format_exc()}") + return {"message": f"创建数据库失败 {e}", "status": "failed"} return database_info @data.delete("/") async def delete_database(db_id): logger.debug(f"Delete database {db_id}") - dbm.delete_database(db_id) + knowledge_base.delete_database(db_id) return {"message": "删除成功"} @data.post("/query-test") @@ -54,17 +57,17 @@ async def create_document_by_file(db_id: str = Body(...), files: List[str] = Bod loop = asyncio.get_event_loop() await loop.run_in_executor( executor, # 使用与chat_router相同的线程池 - lambda: dbm.add_files(db_id, files) + lambda: knowledge_base.add_files(db_id, files) ) return {"message": "文件添加完成", "status": "success"} except Exception as e: - logger.error(f"添加文件失败: {e}") + logger.error(f"添加文件失败: {e}, {traceback.format_exc()}") return {"message": f"添加文件失败: {e}", "status": "failed"} @data.get("/info") async def get_database_info(db_id: str): logger.debug(f"Get database {db_id} info") - database = dbm.get_database_info(db_id) + database = knowledge_base.get_database_info(db_id) if database is None: raise HTTPException(status_code=404, detail="Database not found") return database @@ -72,7 +75,7 @@ async def get_database_info(db_id: str): @data.delete("/document") async def delete_document(db_id: str = Body(...), file_id: str = Body(...)): logger.debug(f"DELETE document {file_id} info in {db_id}") - dbm.delete_file(db_id, file_id) + knowledge_base.delete_file(db_id, file_id) return {"message": "删除成功"} @data.get("/document") @@ -80,10 +83,10 @@ async def get_document_info(db_id: str, file_id: str): logger.debug(f"GET document {file_id} info in {db_id}") try: - info = dbm.get_file_info(db_id, file_id) + info = knowledge_base.get_file_info(db_id, file_id) except Exception as e: logger.error(f"Failed to get file info, {e}, {db_id=}, {file_id=}, {traceback.format_exc()}") - info = {"message": "Failed to get file info", "status": "failed"}, 500 + info = {"message": "Failed to get file info", "status": "failed"} return info @@ -105,36 +108,27 @@ async def upload_file(file: UploadFile = File(...)): @data.get("/graph") async def get_graph_info(): - graph_info = dbm.get_graph() - - # 获取未索引节点数量 - unindexed_count = 0 - if dbm.is_graph_running(): - # 调用GraphDatabase的query_nodes_without_embedding方法 - unindexed_nodes = dbm.graph_base.query_nodes_without_embedding() - unindexed_count = len(unindexed_nodes) if unindexed_nodes else 0 - - # 将未索引节点数量添加到返回结果中 - graph_info["graph"]["unindexed_node_count"] = unindexed_count - + graph_info = graph_base.get_graph_info() + if graph_info is None: + raise HTTPException(status_code=400, detail="图数据库获取出错") return graph_info @data.post("/graph/index-nodes") async def index_nodes(data: dict = Body(default={})): - if not dbm.is_graph_running(): + if not graph_base.is_running(): raise HTTPException(status_code=400, detail="图数据库未启动") # 获取参数或使用默认值 kgdb_name = data.get('kgdb_name', 'neo4j') # 调用GraphDatabase的add_embedding_to_nodes方法 - count = dbm.graph_base.add_embedding_to_nodes(kgdb_name=kgdb_name) + count = graph_base.add_embedding_to_nodes(kgdb_name=kgdb_name) return {"status": "success", "message": f"已成功为{count}个节点添加嵌入向量", "indexed_count": count} @data.get("/graph/node") async def get_graph_node(entity_name: str): - result = dbm.graph_base.query_node(entity_name=entity_name) + result = graph_base.query_node(entity_name=entity_name) return {"result": retriever.format_query_results(result), "message": "success"} @data.get("/graph/nodes") @@ -143,7 +137,7 @@ async def get_graph_nodes(kgdb_name: str, num: int): raise HTTPException(status_code=400, detail="Knowledge graph is not enabled") logger.debug(f"Get graph nodes in {kgdb_name} with {num} nodes") - result = dbm.graph_base.get_sample_nodes(kgdb_name, num) + result = graph_base.get_sample_nodes(kgdb_name, num) return {"result": retriever.format_general_results(result), "message": "success"} @data.post("/graph/add-by-jsonl") @@ -154,6 +148,6 @@ async def add_graph_entity(file_path: str = Body(...), kgdb_name: Optional[str] if not file_path.endswith('.jsonl'): raise HTTPException(status_code=400, detail="file_path must be a jsonl file") - dbm.graph_base.jsonl_file_add_entity(file_path, kgdb_name) + graph_base.jsonl_file_add_entity(file_path, kgdb_name) return {"message": "Entity successfully added"} diff --git a/src/utils/prompts.py b/src/utils/prompts.py index f105ce4c..489fc671 100644 --- a/src/utils/prompts.py +++ b/src/utils/prompts.py @@ -1,9 +1,11 @@ -system_prompt = """ -""" +from datetime import datetime + +def get_system_prompt(): + return (f"当前时间:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n") knowbase_qa_template = """ -请利用查询到的资料回答问题,回答问题时,不要过度的分点作答。如果非要分点作答,可以使用 一、二、等: +请利用查询到的资料回答问题,回答问题时,不要过度的分点作答。 <参考资料>: {external} diff --git a/web/src/components/ChatComponent.vue b/web/src/components/ChatComponent.vue index cc541399..0c09a1fa 100644 --- a/web/src/components/ChatComponent.vue +++ b/web/src/components/ChatComponent.vue @@ -226,7 +226,7 @@ const meta = reactive(JSON.parse(localStorage.getItem('meta')) || { stream: true, summary_title: false, history_round: 5, - db_name: null, + db_id: null, }) const marked = new Marked( @@ -524,7 +524,7 @@ const fetchChatResponse = (user_input, cur_res_id) => { // 更新后的 sendMessage 函数 const sendMessage = () => { const user_input = conv.value.inputText.trim(); - const dbName = opts.databases.length > 0 ? opts.databases[meta.selectedKB]?.metaname : null; + const dbID = opts.databases.length > 0 ? opts.databases[meta.selectedKB]?.db_id : null; if (isStreaming.value) { message.error('请等待上一条消息处理完成'); return @@ -537,7 +537,7 @@ const sendMessage = () => { const cur_res_id = conv.value.messages[conv.value.messages.length - 1].id; conv.value.inputText = ''; - meta.db_name = dbName; + meta.db_id = dbID; fetchChatResponse(user_input, cur_res_id) } else { console.log('请输入消息'); diff --git a/web/src/views/DataBaseInfoView.vue b/web/src/views/DataBaseInfoView.vue index ad676216..5680bebd 100644 --- a/web/src/views/DataBaseInfoView.vue +++ b/web/src/views/DataBaseInfoView.vue @@ -54,7 +54,7 @@ 刷新状态 - +