diff --git a/src/core/knowledgebase.py b/src/core/knowledgebase.py index a79380f5..503a0828 100644 --- a/src/core/knowledgebase.py +++ b/src/core/knowledgebase.py @@ -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, "知识库未启用" diff --git a/src/routers/data_router.py b/src/routers/data_router.py index 71cdb31d..a352e0c4 100644 --- a/src/routers/data_router.py +++ b/src/routers/data_router.py @@ -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(): diff --git a/web/src/views/DataBaseInfoView.vue b/web/src/views/DataBaseInfoView.vue index 5680bebd..908a2490 100644 --- a/web/src/views/DataBaseInfoView.vue +++ b/web/src/views/DataBaseInfoView.vue @@ -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" >