From 144632b703df942dc359eba94c22c97ba51ffb34 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Fri, 9 May 2025 23:47:16 +0800 Subject: [PATCH] =?UTF-8?q?=E8=BF=9B=E4=B8=80=E6=AD=A5=E4=BC=98=E5=8C=96?= =?UTF-8?q?=E5=B9=B6=E8=A1=8C=E5=A4=84=E7=90=86=E8=83=BD=E5=8A=9B=EF=BC=9A?= =?UTF-8?q?file-to-chunk=20=E9=83=A8=E5=88=86=E4=BC=98=E5=8C=96=E5=AE=8C?= =?UTF-8?q?=E6=88=90?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- server/routers/data_router.py | 11 +++-------- src/core/indexing.py | 3 +++ src/core/knowledgebase.py | 26 +++++++++++++------------- 3 files changed, 19 insertions(+), 21 deletions(-) diff --git a/server/routers/data_router.py b/server/routers/data_router.py index 55d07121..b7ad34fa 100644 --- a/server/routers/data_router.py +++ b/server/routers/data_router.py @@ -55,25 +55,20 @@ async def query_test(query: str = Body(...), meta: dict = Body(...), current_use @data.post("/file-to-chunk") async def file_to_chunk(files: List[str] = Body(...), params: dict = Body(...), current_user: User = Depends(get_admin_user)): logger.debug(f"File to chunk: {files}") - result = knowledge_base.file_to_chunk(files, params=params) + result = await knowledge_base.file_to_chunk(files, params=params) return result @data.post("/url-to-chunk") async def url_to_chunk(urls: List[str] = Body(...), params: dict = Body(...), current_user: User = Depends(get_admin_user)): logger.debug(f"Url to chunk: {urls}") - result = knowledge_base.url_to_chunk(urls, params=params) + result = await knowledge_base.url_to_chunk(urls, params=params) return result @data.post("/add-by-file") async def create_document_by_file(db_id: str = Body(...), files: List[str] = Body(...), current_user: User = Depends(get_admin_user)): logger.debug(f"Add document in {db_id} by file: {files}") try: - # 使用线程池执行耗时操作 - loop = asyncio.get_event_loop() - await loop.run_in_executor( - executor, # 使用与chat_router相同的线程池 - lambda: knowledge_base.add_files(db_id, files) - ) + await knowledge_base.add_files(db_id, files) return {"message": "文件添加完成", "status": "success"} except Exception as e: logger.error(f"添加文件失败: {e}, {traceback.format_exc()}") diff --git a/src/core/indexing.py b/src/core/indexing.py index 15a71986..aafd61eb 100644 --- a/src/core/indexing.py +++ b/src/core/indexing.py @@ -1,4 +1,5 @@ import os +import asyncio from pathlib import Path from llama_index.core import Document from llama_index.core.node_parser import SimpleFileNodeParser @@ -94,3 +95,5 @@ def read_text(file, params=None): raise Exception(f"File format not supported, only support {support_format}") +async def read_text_async(file): + return await asyncio.to_thread(read_text, file) diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index 78587aae..fcb246fc 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -8,7 +8,7 @@ from pymilvus import MilvusClient, MilvusException from src import config from src.utils import logger, hashstr -from src.core.indexing import chunk, read_text +from src.core.indexing import chunk, read_text_async from src.core.kb_db_manager import kb_db_manager class KnowledgeBase: @@ -175,7 +175,7 @@ class KnowledgeBase: return self.db_manager.get_database_by_id(db_id) - def file_to_chunk(self, files, params=None): + async def file_to_chunk(self, files, params=None): """将文件转换为分块 这里主要是将文件转换为分块,但并不保存到数据库,仅仅返回分块后的信息,返回的信息里面也包含文件的id,文件名,文件类型,文件路径,文件状态,文件创建时间等。 @@ -191,7 +191,7 @@ class KnowledgeBase: file_type = file.split(".")[-1].lower() if file_type == "pdf": - texts = read_text(file) + texts = await read_text_async(file) nodes = chunk(texts, params=params) else: nodes = chunk(file, params=params) @@ -208,7 +208,7 @@ class KnowledgeBase: return file_infos - def url_to_chunk(self, urls, params=None): + async def url_to_chunk(self, urls, params=None): """将url转换为分块,读取url的内容,并转换为分块 Args: @@ -226,7 +226,7 @@ class KnowledgeBase: file_infos = {} # 使用UnstructuredURLLoader加载URL内容 - loader = UnstructuredURLLoader(urls=urls, continue_on_failure=True) + # loader = UnstructuredURLLoader(urls=urls, continue_on_failure=True) for url_idx, url in enumerate(urls): file_id = "url_" + hashstr(url + str(time.time())) @@ -234,7 +234,7 @@ class KnowledgeBase: try: # 加载单个URL内容 single_loader = UnstructuredURLLoader(urls=[url], continue_on_failure=False) - documents = single_loader.load() + documents = await single_loader.aload() # 将文档内容合并 text_content = "\n\n".join([doc.page_content for doc in documents]) @@ -278,7 +278,7 @@ class KnowledgeBase: return file_infos - def add_chunks(self, db_id, file_chunks): + async def add_chunks(self, db_id, file_chunks): """添加分块""" db = self.get_kb_by_id(db_id) @@ -298,7 +298,7 @@ class KnowledgeBase: ) try: - self.add_documents( + await self.add_documents( file_id=file_id, collection_name=db_id, docs=[node["text"] for node in chunk_info["nodes"]], @@ -312,7 +312,7 @@ class KnowledgeBase: # 更新文件状态为失败 self.db_manager.update_file_status(file_id, "failed") - def add_files(self, db_id, files, params=None): + async def add_files(self, db_id, files, params=None): db = self.get_kb_by_id(db_id) if not self.check_embed_model(db_id): @@ -320,7 +320,7 @@ class KnowledgeBase: return {"message": f"Embed model not match, cur: {self.embed_model.embed_model_fullname}, req: {db['embed_model']}", "status": "failed"} # Preprocessing the files to the queue - new_files = self.file_to_chunk(files, params=params) + new_files = await self.file_to_chunk(files, params=params) for file_id, new_file in new_files.items(): # 在数据库中创建文件记录 @@ -334,7 +334,7 @@ class KnowledgeBase: ) try: - self.add_documents( + await self.add_documents( file_id=file_id, collection_name=db_id, docs=[node["text"] for node in new_file["nodes"]], @@ -498,7 +498,7 @@ class KnowledgeBase: dimension= dimension, # The vectors we will use in this demo has 768 dimensions ) - def add_documents(self, docs, collection_name, chunk_infos=None, **kwargs): + async def add_documents(self, docs, collection_name, chunk_infos=None, **kwargs): """添加已经分块之后的文本""" # 检查 collection 是否存在 import random @@ -508,7 +508,7 @@ class KnowledgeBase: chunk_infos = chunk_infos or [{}] * len(docs) - vectors = self.embed_model.batch_encode(docs) + vectors = await self.embed_model.abatch_encode(docs) data = [{ "id": int(random.random() * 1e12),