diff --git a/app/controller/add_knowledge_route.py b/app/controller/add_knowledge_route.py index 33cd544..040f755 100644 --- a/app/controller/add_knowledge_route.py +++ b/app/controller/add_knowledge_route.py @@ -1,41 +1,70 @@ +import asyncio +from http.client import HTTPException from typing import List -from fastapi.params import Param -from llama_index.core import Document +from fastapi import UploadFile, File, Form -from app.utils.milvus_utils import milvus_service -from app.utils.next_id import next_id +from app.utils.db_utils import AsyncSessionDep +from app.utils.knowledge_utils import knowledge_service def add_knowledge_route(app): - @app.get("/knowledge/search") - async def knowledge_search(text: str = Param()): - print("text", text) - query_result = await milvus_service.async_search(text or "hello") - return query_result + # @app.get("/knowledge/search") + # async def knowledge_search(text: str = Param()): + # print("text", text) + # query_result = await milvus_service.async_search(text or "hello") + # return query_result + # + # @app.post("/knowledge/embed_text") + # async def knowledge_embed_text(text_list: List[str]): + # id_list = await next_id(len(text_list)) + # document_list = [ + # Document( + # text=text, + # id=id_list[index], + # metadata={"kb_id": "cde"} + # ) + # for index, text in enumerate(text_list) + # ] + # await milvus_service.async_create_index_from_documents(document_list) + # search_result = await milvus_service.async_search("hello") + # return { + # "result": "嵌入成功:", + # "origin_documents": [document.to_dict() for document in document_list], + # "search_documents": search_result, + # } + # + # @app.post("/knowledge/delete") + # async def knowledge_embed_text(body: dict): + # await milvus_service.async_delete(body.get('id')) + # return { + # "result": "删除成功", + @app.post("/knowledge/upload_files") + async def knowledge_search( + session: AsyncSessionDep, + files: List[UploadFile] = File(...), + kb_code: str = Form(..., description="所属知识库的编码"), + ): + """ + 批量上传文档并嵌入到向量数据库(同步响应版本) + """ + if not files: + raise HTTPException(status_code=400, detail="No files provided") - @app.post("/knowledge/embed_text") - async def knowledge_embed_text(text_list: List[str]): - id_list = await next_id(len(text_list)) - document_list = [ - Document( - text=text, - id=id_list[index], - metadata={"kb_id": "cde"} - ) - for index, text in enumerate(text_list) - ] - await milvus_service.async_create_index_from_documents(document_list) - search_result = await milvus_service.async_search("hello") - return { - "result": "嵌入成功:", - "origin_documents": [document.to_dict() for document in document_list], - "search_documents": search_result, - } + # 限制文件数量 + if len(files) > 20: + raise HTTPException(status_code=400, detail="Maximum 20 files allowed per upload") - @app.post("/knowledge/delete") - async def knowledge_embed_text(body: dict): - await milvus_service.async_delete(body.get('id')) - return { - "result": "删除成功", - } + # 先将所有文件直接保存到本地 + task_list = [asyncio.create_task(knowledge_service.save_file_with_new_session(file=file)) for file in files] + task_result_list = await asyncio.gather(*task_list) + file_dict_list = [item["result"] for item in task_result_list] + + # 插入对应的文档对象记录 + doc_cls_list = await knowledge_service.save_knowledge_doc_list(session=session, file_dict_list=file_dict_list, kb_code=kb_code) + + # 异步处理嵌入文档,不再等待 + [asyncio.create_task(knowledge_service.process_file_dict(file_dict=file_dict, kb_code=kb_code)) for file_dict in file_dict_list] + + # 直接返回文档对象 + return {"message": f"正在处理 {len(doc_cls_list)} 个文档,请刷新列表查看文档状态。", "result": [item.model_dump() for item in doc_cls_list]} diff --git a/app/utils/knowledge_utils.py b/app/utils/knowledge_utils.py new file mode 100644 index 0000000..335ffc0 --- /dev/null +++ b/app/utils/knowledge_utils.py @@ -0,0 +1,41 @@ +import asyncio +from typing import List + +from fastapi import UploadFile + +from app.model.FileModel import FileSaveService +from app.model.KnowledgeDoc import KnowledgeDocModel, KnowledgeDocService +from app.utils.db_utils import AsyncSessionDep, async_session + + +class KnowledgeService: + def __init__(self): + pass + + # 为每个文件创建独立的 session 进行文件保存 + async def save_file_with_new_session(self, file: UploadFile): + # 这里需要根据您的数据库配置创建新的 session + # 假设您有一个 session factory + async with async_session() as session: + return await FileSaveService.saveFile(session=session, file=file, filename=file.filename, file_record={}) + + async def save_knowledge_doc_list(self, session: AsyncSessionDep, file_dict_list: List[dict], kb_code: str): + # 先插入文档对象 + + tobe_insert_kd_list = [ + KnowledgeDocModel(id=file_dict["id"], name=file_dict["name"], path=file_dict["path"], parent_code=kb_code, status="process").model_dump() + for file_dict in file_dict_list + ] + + insert_cls_list: List[KnowledgeDocModel] = await KnowledgeDocService.batch_insert(session=session, row_dict_list=tobe_insert_kd_list) + + return insert_cls_list + + async def process_file_dict(self, file_dict: dict, kb_code: str): + print("process start", file_dict['name']) + await asyncio.sleep(3) + print("process end", file_dict['name']) + return {} + + +knowledge_service = KnowledgeService()