优化部分并行性能

This commit is contained in:
Wenjie Zhang 2025-03-06 23:55:20 +08:00
parent 9e46215379
commit 952841c68d
5 changed files with 49 additions and 18 deletions

View File

@ -69,9 +69,16 @@ class DataBaseManager:
f"Database number not match, {knowledge_base_collections}, " f"Database number not match, {knowledge_base_collections}, "
f"self.data['databases']: {self.data['databases']}, ") f"self.data['databases']: {self.data['databases']}, ")
# 更新每个数据库的状态信息
for db in self.data["databases"]: for db in self.data["databases"]:
# 获取最新的集合信息
db.update(self.knowledge_base.get_collection_info(db.metaname)) 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"]]} return {"databases": [db.to_dict() for db in self.data["databases"]]}
def get_graph(self): def get_graph(self):
@ -117,11 +124,16 @@ class DataBaseManager:
db.files.append(new_file) db.files.append(new_file)
new_files.append(new_file) new_files.append(new_file)
# 先保存一次数据库状态确保waiting状态被记录
self._save_databases()
from src.core.indexing import chunk from src.core.indexing import chunk
for new_file in new_files: for new_file in new_files:
file_id = new_file["file_id"] file_id = new_file["file_id"]
idx = self.get_idx_by_fileid(db, file_id) idx = self.get_idx_by_fileid(db, file_id)
db.files[idx]["status"] = "processing" db.files[idx]["status"] = "processing"
# 更新处理状态
self._save_databases()
try: try:
if new_file["type"] == "pdf": if new_file["type"] == "pdf":
@ -143,9 +155,9 @@ class DataBaseManager:
idx = self.get_idx_by_fileid(db, file_id) idx = self.get_idx_by_fileid(db, file_id)
db.files[idx]["status"] = "failed" db.files[idx]["status"] = "failed"
# 每个文件处理完成后立即保存数据库状态
self._save_databases() self._save_databases()
return {"message": "全部解析完成", "status": "success"}
def get_database_info(self, db_id): def get_database_info(self, db_id):
db = self.get_kb_by_id(db_id) db = self.get_kb_by_id(db_id)

View File

@ -1,10 +1,15 @@
import os import os
from concurrent.futures import ThreadPoolExecutor
from src.core import DataBaseManager from src.core import DataBaseManager
from src.core.retriever import Retriever from src.core.retriever import Retriever
from src.models import select_model from src.models import select_model
from src.config import Config from src.config import Config
from src.utils import logger from src.utils import logger
# 创建线程池
executor = ThreadPoolExecutor()
class Startup: class Startup:
def __init__(self): def __init__(self):

View File

@ -2,14 +2,12 @@ import json
import asyncio import asyncio
from fastapi import APIRouter, Body from fastapi import APIRouter, Body
from fastapi.responses import StreamingResponse, Response from fastapi.responses import StreamingResponse, Response
from concurrent.futures import ThreadPoolExecutor
from src.core import HistoryManager 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 from src.utils.logging_config import logger
chat = APIRouter(prefix="/chat") chat = APIRouter(prefix="/chat")
# 创建线程池
executor = ThreadPoolExecutor()
@chat.get("/") @chat.get("/")

View File

@ -1,15 +1,16 @@
import os import os
import asyncio
from typing import List, Optional from typing import List, Optional
from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body
from src.utils import logger, hashstr from src.utils import logger, hashstr
from src.core.startup import startup from src.core.startup import startup, executor
data = APIRouter(prefix="/data") data = APIRouter(prefix="/data")
@data.get("/") @data.get("/")
def get_databases(): async def get_databases():
try: try:
database = startup.dbm.get_databases() database = startup.dbm.get_databases()
except Exception as e: except Exception as e:
@ -47,8 +48,17 @@ async def query_test(query: str = Body(...), meta: dict = Body(...)):
@data.post("/add-by-file") @data.post("/add-by-file")
async def create_document_by_file(db_id: str = Body(...), files: List[str] = Body(...)): 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}") logger.debug(f"Add document in {db_id} by file: {files}")
msg = startup.dbm.add_files(db_id, files) try:
return msg # 使用线程池执行耗时操作
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") @data.get("/info")
async def get_database_info(db_id: str): 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") upload_dir = os.path.join(startup.config.save_dir, "data/uploads")
os.makedirs(upload_dir, exist_ok=True) os.makedirs(upload_dir, exist_ok=True)
basename, ext = os.path.splitext(file.filename) 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) file_path = os.path.join(upload_dir, filename)
with open(file_path, "wb") as buffer: with open(file_path, "wb") as buffer:

View File

@ -206,7 +206,7 @@
</template> </template>
<script setup> <script setup>
import { onMounted, reactive, ref, watch, toRaw } from 'vue'; import { onMounted, reactive, ref, watch, toRaw, onUnmounted } from 'vue';
import { message, Modal } from 'ant-design-vue'; import { message, Modal } from 'ant-design-vue';
import { useRoute, useRouter } from 'vue-router'; import { useRoute, useRouter } from 'vue-router';
import { useConfigStore } from '@/stores/config' import { useConfigStore } from '@/stores/config'
@ -490,9 +490,7 @@ const addDocumentByFile = () => {
state.loading = true state.loading = true
state.lock = true state.lock = true
state.refreshInterval = setInterval(() => {
getDatabaseInfo();
}, 1000);
fetch('/api/data/add-by-file', { fetch('/api/data/add-by-file', {
method: "POST", method: "POST",
headers: { headers: {
@ -519,7 +517,7 @@ const addDocumentByFile = () => {
}) })
.finally(() => { .finally(() => {
getDatabaseInfo() getDatabaseInfo()
clearInterval(state.refreshInterval) //
state.loading = false state.loading = false
}) })
} }
@ -559,9 +557,17 @@ const useQueryExample = (example) => {
onMounted(() => { onMounted(() => {
getDatabaseInfo(); getDatabaseInfo();
// const refreshInterval = setInterval(() => { state.refreshInterval = setInterval(() => {
// getDatabaseInfo(); getDatabaseInfo();
// }, 10000); }, 10000);
})
// onUnmounted
onUnmounted(() => {
if (state.refreshInterval) {
clearInterval(state.refreshInterval);
state.refreshInterval = null;
}
}) })