优化上传逻辑
This commit is contained in:
parent
f1f2459d60
commit
f797d66128
@ -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, "知识库未启用"
|
||||
|
||||
|
||||
@ -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():
|
||||
|
||||
@ -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"
|
||||
>
|
||||
|
||||
Loading…
Reference in New Issue
Block a user