From 952841c68dea2912f5c7519de25bca896ab470d1 Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Thu, 6 Mar 2025 23:55:20 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BC=98=E5=8C=96=E9=83=A8=E5=88=86=E5=B9=B6?= =?UTF-8?q?=E8=A1=8C=E6=80=A7=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/core/database.py | 14 +++++++++++++- src/core/startup.py | 5 +++++ src/routers/chat_router.py | 6 ++---- src/routers/data_router.py | 20 +++++++++++++++----- web/src/views/DataBaseInfoView.vue | 22 ++++++++++++++-------- 5 files changed, 49 insertions(+), 18 deletions(-) diff --git a/src/core/database.py b/src/core/database.py index 2b80e22c..1e6a0318 100644 --- a/src/core/database.py +++ b/src/core/database.py @@ -69,9 +69,16 @@ class DataBaseManager: f"Database number not match, {knowledge_base_collections}, " f"self.data['databases']: {self.data['databases']}, ") + # 更新每个数据库的状态信息 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): @@ -117,11 +124,16 @@ class DataBaseManager: db.files.append(new_file) new_files.append(new_file) + # 先保存一次数据库状态,确保waiting状态被记录 + self._save_databases() + from src.core.indexing import chunk 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": @@ -143,9 +155,9 @@ class DataBaseManager: idx = self.get_idx_by_fileid(db, file_id) db.files[idx]["status"] = "failed" + # 每个文件处理完成后立即保存数据库状态 self._save_databases() - return {"message": "全部解析完成", "status": "success"} def get_database_info(self, db_id): db = self.get_kb_by_id(db_id) diff --git a/src/core/startup.py b/src/core/startup.py index 11092a2d..69e301a7 100644 --- a/src/core/startup.py +++ b/src/core/startup.py @@ -1,10 +1,15 @@ import os +from concurrent.futures import ThreadPoolExecutor + from src.core import DataBaseManager from src.core.retriever import Retriever from src.models import select_model from src.config import Config from src.utils import logger +# 创建线程池 +executor = ThreadPoolExecutor() + class Startup: def __init__(self): diff --git a/src/routers/chat_router.py b/src/routers/chat_router.py index 78a8ab70..16f97ba3 100644 --- a/src/routers/chat_router.py +++ b/src/routers/chat_router.py @@ -2,14 +2,12 @@ import json import asyncio from fastapi import APIRouter, Body from fastapi.responses import StreamingResponse, Response -from concurrent.futures import ThreadPoolExecutor from src.core import HistoryManager -from src.core.startup import startup +from src.core.startup import startup, executor from src.utils.logging_config import logger chat = APIRouter(prefix="/chat") -# 创建线程池 -executor = ThreadPoolExecutor() + @chat.get("/") diff --git a/src/routers/data_router.py b/src/routers/data_router.py index 6a2bf2ce..06f24403 100644 --- a/src/routers/data_router.py +++ b/src/routers/data_router.py @@ -1,15 +1,16 @@ import os +import asyncio from typing import List, Optional from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body from src.utils import logger, hashstr -from src.core.startup import startup +from src.core.startup import startup, executor data = APIRouter(prefix="/data") @data.get("/") -def get_databases(): +async def get_databases(): try: database = startup.dbm.get_databases() except Exception as e: @@ -47,8 +48,17 @@ async def query_test(query: str = Body(...), meta: dict = Body(...)): @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}") - msg = startup.dbm.add_files(db_id, files) - return msg + try: + # 使用线程池执行耗时操作 + loop = asyncio.get_event_loop() + await loop.run_in_executor( + executor, # 使用与chat_router相同的线程池 + lambda: startup.dbm.add_files(db_id, files) + ) + return {"message": "文件添加完成", "status": "success"} + except Exception as e: + logger.error(f"添加文件失败: {e}") + return {"message": f"添加文件失败: {e}", "status": "failed"} @data.get("/info") async def get_database_info(db_id: str): @@ -84,7 +94,7 @@ async def upload_file(file: UploadFile = File(...)): upload_dir = os.path.join(startup.config.save_dir, "data/uploads") os.makedirs(upload_dir, exist_ok=True) basename, ext = os.path.splitext(file.filename) - filename = f"{basename}_{hashstr(basename, 4, with_salt=True)}_{ext}".lower() + filename = f"{basename}_{hashstr(basename, 4, with_salt=True)}{ext}".lower() file_path = os.path.join(upload_dir, filename) with open(file_path, "wb") as buffer: diff --git a/web/src/views/DataBaseInfoView.vue b/web/src/views/DataBaseInfoView.vue index db9f5207..85e772d7 100644 --- a/web/src/views/DataBaseInfoView.vue +++ b/web/src/views/DataBaseInfoView.vue @@ -206,7 +206,7 @@