From 7563c116fc2da1129967ed9bd1b53d464ad509e8 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Sun, 23 Mar 2025 23:10:32 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E6=96=87=E6=A1=A3=E6=B7=BB?= =?UTF-8?q?=E5=8A=A0=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/core/indexing.py | 3 +- src/core/knowledgebase.py | 165 ++++++++-- src/core/retriever.py | 65 ++-- src/routers/data_router.py | 20 ++ web/src/views/DataBaseInfoView.vue | 502 +++++++++++++++++++++++------ web/src/views/DataBaseView.vue | 2 +- 6 files changed, 601 insertions(+), 156 deletions(-) diff --git a/src/core/indexing.py b/src/core/indexing.py index 12853bdf..15a71986 100644 --- a/src/core/indexing.py +++ b/src/core/indexing.py @@ -29,7 +29,8 @@ def chunk(text_or_path, params=None): chunk_overlap=chunk_overlap, ) - if os.path.isfile(text_or_path) and "uploads" in text_or_path: + # 如果文件存在,并且是当前目录下的文件,则使用文件解析器 + if os.path.isfile(text_or_path) and os.path.exists(text_or_path) and os.path.abspath(text_or_path).startswith(os.getcwd()): parser = SimpleFileNodeParser() file_type = Path(text_or_path).suffix.lower() if file_type in [".txt", ".json", ".md"]: diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 503a0828..4223ae92 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -18,6 +18,12 @@ class KnowledgeBase: self.client = None self.work_dir = os.path.join(config.save_dir, "data") self.database_path = os.path.join(self.work_dir, "database.json") + + # Configuration + self.default_distance_threshold = 0.5 + self.default_rerank_threshold = 0.1 + self.default_max_query_count = 20 + self._load_models() self._load_databases() @@ -29,6 +35,10 @@ class KnowledgeBase: from src.models.embedding import get_embedding_model self.embed_model = get_embedding_model(config) + if config.enable_reranker: + from src.models.rerank_model import get_reranker + self.reranker = get_reranker(config) + if not self.connect_to_milvus(): raise ConnectionError("Failed to connect to Milvus") @@ -126,29 +136,87 @@ class KnowledgeBase: return {"lines": lines} def get_kb_by_id(self, db_id): + if not config.enable_knowledge_base: + return None + return next((db for db in self.data if db.db_id == db_id), None) + def file_to_chunk(self, files, params=None): + """将文件转换为分块 + + 这里主要是将文件转换为分块,但并不保存到数据库,仅仅返回分块后的信息,返回的信息里面也包含文件的id,文件名,文件类型,文件路径,文件状态,文件创建时间等。 + files: list of file path + params: params for chunking + + return: list of chunk info + """ + file_infos = {} + for file in files: + file_id = "file_" + hashstr(file + str(time.time())) + + file_type = file.split(".")[-1].lower() + + if file_type == "pdf": + texts = read_text(file) + nodes = chunk(texts, params=params) + else: + nodes = chunk(file, params=params) + + file_infos[file_id] = { + "file_id": file_id, + "filename": os.path.basename(file), + "path": file, + "type": file_type, + "status": "waiting", + "created_at": time.time(), + "nodes": [node.dict() for node in nodes] + } + + return file_infos + + def url_to_chunk(self, url, params=None): + """将url转换为分块,读取url的内容,并转换为分块""" + raise NotImplementedError("Not implemented") + + def add_chunks(self, db_id, file_chunks): + """添加分块""" + 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}, req: {db.embed_model}", "status": "failed"} + + db.files.update(file_chunks) + self._save_databases() + + for file_id, chunk in file_chunks.items(): + db.files[file_id]["status"] = "processing" + self._save_databases() + + try: + self.add_documents( + file_id=file_id, + collection_name=db.db_id, + docs=[node["text"] for node in chunk["nodes"]], + chunk_infos=chunk["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 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"} + return {"message": f"Embed model not match, cur: {config.embed_model}, req: {db.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 - + new_files = self.file_to_chunk(files, params=params) db.files.update(new_files) # 更新数据库状态 # 先保存一次数据库状态,确保waiting状态被记录 @@ -160,17 +228,11 @@ class KnowledgeBase: 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]) + docs=[node["text"] for node in new_file["nodes"]], + chunk_infos=new_file["nodes"]) db.files[file_id]["status"] = "done" @@ -204,8 +266,63 @@ class KnowledgeBase: self._load_models() self._load_databases() + ################################### + #* Below is the code for retriever # + ################################### + + def get_retriever(self, db_id): + db = self.get_kb_by_id(db_id) + if db is None: + raise Exception(f"database not found, {db_id}") + + return db.retriever + + def query(self, query, db_id, **kwargs): + db = self.get_kb_by_id(db_id) + + distance_threshold = kwargs.get("distance_threshold", self.default_distance_threshold) + rerank_threshold = kwargs.get("rerank_threshold", self.default_rerank_threshold) + max_query_count = kwargs.get("max_query_count", self.default_max_query_count) + + all_db_result = self.search(query, db_id, limit=max_query_count) + for res in all_db_result: + res["file"] = db.files[res["entity"]["file_id"]] + + db_result = [r for r in all_db_result if r["distance"] > distance_threshold] + + if config.enable_reranker and len(db_result) > 0 and self.reranker: + texts = [r["entity"]["text"] for r in db_result] + rerank_scores = self.reranker.compute_score([query, texts], normalize=False) + for i, r in enumerate(db_result): + r["rerank_score"] = rerank_scores[i] + db_result.sort(key=lambda x: x["rerank_score"], reverse=True) + db_result = [_res for _res in db_result if _res["rerank_score"] > rerank_threshold] + + if kwargs.get("top_k", None): + db_result = db_result[:kwargs["top_k"]] + + return { + "results": db_result, + "all_results": all_db_result, + } + + def get_retriever(self, db_id): + retriever_params = { + "distance_threshold": self.default_distance_threshold, + "rerank_threshold": self.default_rerank_threshold, + "max_query_count": self.default_max_query_count, + "top_k": 10, + } + + def retriever(query): + response = self.query(query, db_id, **retriever_params) + return response["results"] + + return retriever + + ################################ - # Below is the code for milvus # + #* Below is the code for milvus # ################################ def connect_to_milvus(self): """ @@ -276,7 +393,7 @@ class KnowledgeBase: return res def search(self, query, collection_name, limit=3): - + """搜索数据库""" query_vectors = self.embed_model.batch_encode([query]) return self.search_by_vector(query_vectors[0], collection_name, limit) diff --git a/src/core/retriever.py b/src/core/retriever.py index 4e010f6a..c5c2118d 100644 --- a/src/core/retriever.py +++ b/src/core/retriever.py @@ -85,48 +85,35 @@ class Retriever: def query_knowledgebase(self, query, history, refs): """查询知识库""" - kb_res = [] - final_res = [] + response = { + "results": [], + "all_results": [], + "rw_query": query, + "message": "", + } - db_id = refs["meta"].get("db_id") + meta = refs["meta"] + + db_id = meta.get("db_id") if not db_id or not config.enable_knowledge_base: - return { - "results": final_res, - "all_results": kb_res, - "rw_query": query, - "message": "Knowledge base is disabled", - } + response["message"] = "知识库未启用、或未指定知识库、或知识库不存在" + return response rw_query = self.rewrite_query(query, history, refs) - kb = knowledge_base.id2db[db_id] - logger.debug(f"{refs['meta']=}") + logger.debug(f"{meta=}") + query_result = knowledge_base.query(query=rw_query, + db_id=db_id, + distance_threshold=meta.get("distanceThreshold", 0.5), + rerank_threshold=meta.get("rerankThreshold", 0.1), + max_query_count=meta.get("maxQueryCount", 20), + top_k=meta.get("topK", 10)) - meta = refs["meta"] - max_query_count = meta.get("maxQueryCount", 10) - rerank_threshold = meta.get("rerankThreshold", 0.1) - distance_threshold = meta.get("distanceThreshold", 0) - top_k = meta.get("topK", 5) + response["results"] = query_result["results"] + response["all_results"] = query_result["all_results"] + response["rw_query"] = rw_query - # 检索 - all_kb_res = knowledge_base.search(rw_query, db_id, limit=max_query_count) - for r in all_kb_res: - r["file"] = kb.files[r["entity"]["file_id"]] - - kb_res = [r for r in all_kb_res if r["distance"] > distance_threshold] - - # 重排序 - if config.enable_reranker and len(kb_res) > 0: - texts = [r["entity"]["text"] for r in kb_res] - rerank_scores = self.reranker.compute_score([rw_query, texts], normalize=True) - for i, r in enumerate(kb_res): - r["rerank_score"] = rerank_scores[i] - kb_res.sort(key=lambda x: x["rerank_score"], reverse=True) - kb_res = [_res for _res in kb_res if _res["rerank_score"] > rerank_threshold] - - kb_res = kb_res[:top_k] - - return {"results": kb_res, "all_results": all_kb_res, "rw_query": rw_query} + return response def query_web(self, query, history, refs): """查询网络""" @@ -144,7 +131,9 @@ class Retriever: def rewrite_query(self, query, history, refs): """重写查询""" - model = select_model(config) + model_provider = config.model_provider_lite + model_name = config.model_name_lite + model = select_model(config, model_provider=model_provider, model_name=model_name) if refs["meta"].get("mode") == "search": # 如果是搜索模式,就使用 meta 的配置,否则就使用全局的配置 rewrite_query_span = refs["meta"].get("use_rewrite_query", "off") else: @@ -168,7 +157,9 @@ class Retriever: def reco_entities(self, query, history, refs): """识别句子中的实体""" query = refs.get("rewritten_query", query) - model = select_model(config) + model_provider = config.model_provider_lite + model_name = config.model_name_lite + model = select_model(config, model_provider=model_provider, model_name=model_name) entities = [] if refs["meta"].get("use_graph"): diff --git a/src/routers/data_router.py b/src/routers/data_router.py index 6b51e531..fd60290e 100644 --- a/src/routers/data_router.py +++ b/src/routers/data_router.py @@ -49,6 +49,12 @@ async def query_test(query: str = Body(...), meta: dict = Body(...)): result = retriever.query_knowledgebase(query, history=None, refs={"meta": meta}) return result +@data.post("/file-to-chunk") +async def file_to_chunk(files: List[str] = Body(...), params: dict = Body(...)): + logger.debug(f"File to chunk: {files}") + result = knowledge_base.file_to_chunk(files, params=params) + return result + @data.post("/add-by-file") async def create_document_by_file(db_id: str = Body(...), files: List[str] = Body(...)): logger.debug(f"Add document in {db_id} by file: {files}") @@ -64,6 +70,20 @@ async def create_document_by_file(db_id: str = Body(...), files: List[str] = Bod logger.error(f"添加文件失败: {e}, {traceback.format_exc()}") return {"message": f"添加文件失败: {e}", "status": "failed"} +@data.post("/add-by-chunks") +async def add_by_chunks(db_id: str = Body(...), file_chunks: dict = Body(...)): + logger.debug(f"Add chunks in {db_id}: {file_chunks}") + try: + loop = asyncio.get_event_loop() + await loop.run_in_executor( + executor, # 使用与chat_router相同的线程池 + lambda: knowledge_base.add_chunks(db_id, file_chunks) + ) + return {"message": "分块添加完成", "status": "success"} + except Exception as 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") diff --git a/web/src/views/DataBaseInfoView.vue b/web/src/views/DataBaseInfoView.vue index 908a2490..0d6f0131 100644 --- a/web/src/views/DataBaseInfoView.vue +++ b/web/src/views/DataBaseInfoView.vue @@ -7,7 +7,7 @@
{{ database.embed_model }} {{ database.dimension }} - {{ database.metadata?.row_count }} 行 · {{ database.files?.length || 0 }} 文件 + {{ database.metadata?.row_count }} 行 · {{ database.files ? Object.keys(database.files).length : 0 }} 文件