feat: 知识库批量上传文档
This commit is contained in:
@@ -1,41 +1,70 @@
|
|||||||
|
import asyncio
|
||||||
|
from http.client import HTTPException
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
from fastapi.params import Param
|
from fastapi import UploadFile, File, Form
|
||||||
from llama_index.core import Document
|
|
||||||
|
|
||||||
from app.utils.milvus_utils import milvus_service
|
from app.utils.db_utils import AsyncSessionDep
|
||||||
from app.utils.next_id import next_id
|
from app.utils.knowledge_utils import knowledge_service
|
||||||
|
|
||||||
|
|
||||||
def add_knowledge_route(app):
|
def add_knowledge_route(app):
|
||||||
@app.get("/knowledge/search")
|
# @app.get("/knowledge/search")
|
||||||
async def knowledge_search(text: str = Param()):
|
# async def knowledge_search(text: str = Param()):
|
||||||
print("text", text)
|
# print("text", text)
|
||||||
query_result = await milvus_service.async_search(text or "hello")
|
# query_result = await milvus_service.async_search(text or "hello")
|
||||||
return query_result
|
# 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]):
|
if len(files) > 20:
|
||||||
id_list = await next_id(len(text_list))
|
raise HTTPException(status_code=400, detail="Maximum 20 files allowed per upload")
|
||||||
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):
|
task_list = [asyncio.create_task(knowledge_service.save_file_with_new_session(file=file)) for file in files]
|
||||||
await milvus_service.async_delete(body.get('id'))
|
task_result_list = await asyncio.gather(*task_list)
|
||||||
return {
|
file_dict_list = [item["result"] for item in task_result_list]
|
||||||
"result": "删除成功",
|
|
||||||
}
|
# 插入对应的文档对象记录
|
||||||
|
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]}
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user