优化上传逻辑

This commit is contained in:
Wenjie Zhang 2025-03-20 20:31:14 +08:00
parent f1f2459d60
commit f797d66128
3 changed files with 31 additions and 7 deletions

View File

@ -16,7 +16,8 @@ class KnowledgeBase:
def __init__(self) -> None:
self.data = []
self.client = None
self.database_path = os.path.join(config.save_dir, "data", "database.json")
self.work_dir = os.path.join(config.save_dir, "data")
self.database_path = os.path.join(self.work_dir, "database.json")
self._load_models()
self._load_databases()
@ -63,10 +64,26 @@ class KnowledgeBase:
embed_model=self.embed_model.embed_model_fullname,
dimension=dimension)
# 创建数据库对应的文件夹
self._ensure_db_folders(db.db_id)
self.add_collection(db.db_id, dimension)
self.data.append(db)
self._save_databases()
def _ensure_db_folders(self, db_id):
"""确保数据库文件夹存在"""
db_folder = os.path.join(self.work_dir, db_id)
uploads_folder = os.path.join(db_folder, "uploads")
os.makedirs(db_folder, exist_ok=True)
os.makedirs(uploads_folder, exist_ok=True)
return db_folder, uploads_folder
def get_db_upload_path(self, db_id=None):
"""获取上传文件夹路径如果没有指定db_id则使用默认路径"""
_, uploads_folder = self._ensure_db_folders(db_id)
return uploads_folder
def get_databases(self):
assert config.enable_knowledge_base, "知识库未启用"

View File

@ -2,7 +2,7 @@ import os
import asyncio
import traceback
from typing import List, Optional
from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body
from fastapi import APIRouter, File, UploadFile, HTTPException, Depends, Body, Form, Query
from src.utils import logger, hashstr
from src import executor, retriever, config, knowledge_base, graph_base
@ -91,12 +91,19 @@ async def get_document_info(db_id: str, file_id: str):
return info
@data.post("/upload")
async def upload_file(file: UploadFile = File(...)):
async def upload_file(
file: UploadFile = File(...),
db_id: Optional[str] = Query(None)
):
if not file.filename:
raise HTTPException(status_code=400, detail="No selected file")
upload_dir = os.path.join(config.save_dir, "data/uploads")
os.makedirs(upload_dir, exist_ok=True)
# 根据db_id获取上传路径如果db_id为None则使用默认路径
if db_id:
upload_dir = knowledge_base.get_db_upload_path(db_id)
else:
upload_dir = os.path.join(config.save_dir, "data", "uploads")
basename, ext = os.path.splitext(file.filename)
filename = f"{basename}_{hashstr(basename, 4, with_salt=True)}{ext}".lower()
file_path = os.path.join(upload_dir, filename)
@ -104,7 +111,7 @@ async def upload_file(file: UploadFile = File(...)):
with open(file_path, "wb") as buffer:
buffer.write(await file.read())
return {"message": "File successfully uploaded", "file_path": file_path}
return {"message": "File successfully uploaded", "file_path": file_path, "db_id": db_id}
@data.get("/graph")
async def get_graph_info():

View File

@ -32,7 +32,7 @@
name="file"
:multiple="true"
:disabled="state.loading"
action="/api/data/upload"
:action="'/api/data/upload?db_id=' + databaseId"
@change="handleFileUpload"
@drop="handleDrop"
>