From e421ab93bca13a1988b29b0a79faeaf6e2a8d6b2 Mon Sep 17 00:00:00 2001 From: lan Date: Tue, 18 Jun 2024 23:04:52 +0800 Subject: [PATCH] =?UTF-8?q?update:=E5=AD=98=E5=82=A8=E6=94=B9=E6=88=90?= =?UTF-8?q?=E5=8D=95=E4=BE=8B=E6=A8=A1=E5=BC=8F=EF=BC=8C=E4=BF=AE=E5=A4=8D?= =?UTF-8?q?=E8=AF=BB=E5=8F=96=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- apps/admin/views.py | 4 +++- apps/base/views.py | 6 +++++- core/storage.py | 12 +++++++++--- core/tasks.py | 4 +++- 4 files changed, 20 insertions(+), 6 deletions(-) diff --git a/apps/admin/views.py b/apps/admin/views.py index 93baf46..c5a5525 100644 --- a/apps/admin/views.py +++ b/apps/admin/views.py @@ -11,7 +11,7 @@ from apps.admin.pydantics import IDData from apps.base.models import FileCodes, KeyValue from core.response import APIResponse from core.settings import settings -from core.storage import file_storage +from core.storage import FileStorageInterface, storages admin_api = APIRouter( prefix='/admin', @@ -26,6 +26,7 @@ async def login(): @admin_api.delete('/file/delete', dependencies=[Depends(admin_required)]) async def file_delete(data: IDData): + file_storage: FileStorageInterface = storages[settings.file_storage]() file_code = await FileCodes.get(id=data.id) await file_storage.delete_file(file_code) await file_code.delete() @@ -79,6 +80,7 @@ async def get_file_by_id(id): @admin_api.get('/file/download', dependencies=[Depends(admin_required)]) async def file_download(id: int): + file_storage: FileStorageInterface = storages[settings.file_storage]() has, file_code = await get_file_by_id(id) # 检查文件是否存在 if not has: diff --git a/apps/base/views.py b/apps/base/views.py index af48d92..318a391 100644 --- a/apps/base/views.py +++ b/apps/base/views.py @@ -10,7 +10,7 @@ from apps.base.pydantics import SelectFileModel from apps.base.utils import get_expire_info, get_file_path_name, error_ip_limit, upload_ip_limit from core.response import APIResponse from core.settings import settings -from core.storage import file_storage +from core.storage import storages, FileStorageInterface from core.utils import get_select_token # 创建一个API路由 @@ -56,6 +56,7 @@ async def share_file(expire_value: int = Form(default=1, gt=0), expire_style: st # 获取文件路径和名称 path, suffix, prefix, uuid_file_name, save_path = await get_file_path_name(file) # 保存文件 + file_storage: FileStorageInterface = storages[settings.file_storage]() await file_storage.save_file(file, save_path) # 创建一个新的FileCodes实例 await FileCodes.create( @@ -94,6 +95,7 @@ async def get_code_file_by_code(code, check=True): # 获取文件的API @share_api.get('/select/') async def get_code_file(code: str, ip: str = Depends(error_ip_limit)): + file_storage: FileStorageInterface = storages[settings.file_storage]() # 获取文件 has, file_code = await get_code_file_by_code(code) # 检查文件是否存在 @@ -115,6 +117,7 @@ async def get_code_file(code: str, ip: str = Depends(error_ip_limit)): # 选择文件的API @share_api.post('/select/') async def select_file(data: SelectFileModel, ip: str = Depends(error_ip_limit)): + file_storage: FileStorageInterface = storages[settings.file_storage]() # 获取文件 has, file_code = await get_code_file_by_code(data.code) # 检查文件是否存在 @@ -141,6 +144,7 @@ async def select_file(data: SelectFileModel, ip: str = Depends(error_ip_limit)): # 下载文件的API @share_api.get('/download') async def download_file(key: str, code: str, ip: str = Depends(error_ip_limit)): + file_storage: FileStorageInterface = storages[settings.file_storage]() # 检查token是否有效 is_valid = await get_select_token(code) == key if not is_valid: diff --git a/core/storage.py b/core/storage.py index af43114..c3fc3a9 100644 --- a/core/storage.py +++ b/core/storage.py @@ -2,6 +2,8 @@ # @Author : Lan # @File : storage.py # @Software: PyCharm +from typing import Optional + import aiohttp import asyncio from pathlib import Path @@ -20,6 +22,12 @@ from fastapi.responses import FileResponse class FileStorageInterface: + _instance: Optional['FileStorageInterface'] = None + + def __new__(cls, *args, **kwargs): + if cls._instance is None: + cls._instance = super(FileStorageInterface, cls).__new__(cls, *args, **kwargs) + return cls._instance async def save_file(self, file: UploadFile, save_path: str): """ @@ -79,7 +87,7 @@ class SystemFileStorage(FileStorageInterface): return await get_file_url(file_code.code) async def get_file_response(self, file_code: FileCodes): - file_path = file_storage.root_path / await file_code.get_file_path() + file_path = self.root_path / await file_code.get_file_path() if not file_path.exists(): return APIResponse(code=404, detail='文件已过期删除') return FileResponse(file_path, filename=file_code.prefix + file_code.suffix) @@ -288,5 +296,3 @@ storages = { 'onedrive': OneDriveFileStorage, 'opendal': OpenDALFileStorage, } - -file_storage: FileStorageInterface = storages[settings.file_storage]() diff --git a/core/tasks.py b/core/tasks.py index bcd07e2..fc6ff07 100644 --- a/core/tasks.py +++ b/core/tasks.py @@ -8,11 +8,13 @@ from tortoise.expressions import Q from apps.base.models import FileCodes from apps.base.utils import error_ip_limit, upload_ip_limit -from core.storage import file_storage +from core.settings import settings +from core.storage import FileStorageInterface, storages from core.utils import get_now async def delete_expire_files(): + file_storage: FileStorageInterface = storages[settings.file_storage]() while True: try: await error_ip_limit.remove_expired_ip()