This commit is contained in:
lan
2024-07-28 19:03:37 +08:00
parent 99c2bee1a6
commit a28b49e6cd
4 changed files with 25 additions and 18 deletions
+4 -4
View File
@@ -86,7 +86,7 @@ async def get_random_code(style='num'):
return code return code
# 错误IP限制器 ip_limit = {
error_ip_limit = IPRateLimit(count=settings.errorCount, minutes=settings.errorMinute) 'error': IPRateLimit(count=settings.uploadCount, minutes=settings.errorMinute),
# 上传文件限制器 'upload': IPRateLimit(count=settings.errorCount, minutes=settings.errorMinute)
upload_ip_limit = IPRateLimit(count=settings.uploadCount, minutes=settings.errorMinute) }
+12 -11
View File
@@ -7,7 +7,7 @@ from fastapi import APIRouter, Form, UploadFile, File, Depends, HTTPException
from apps.admin.depends import admin_required from apps.admin.depends import admin_required
from apps.base.models import FileCodes from apps.base.models import FileCodes
from apps.base.pydantics import SelectFileModel 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 apps.base.utils import get_expire_info, get_file_path_name, ip_limit
from core.response import APIResponse from core.response import APIResponse
from core.settings import settings from core.settings import settings
from core.storage import storages, FileStorageInterface from core.storage import storages, FileStorageInterface
@@ -22,7 +22,7 @@ share_api = APIRouter(
# 分享文本的API # 分享文本的API
@share_api.post('/text/', dependencies=[Depends(admin_required)]) @share_api.post('/text/', dependencies=[Depends(admin_required)])
async def share_text(text: str = Form(...), expire_value: int = Form(default=1, gt=0), expire_style: str = Form(default='day'), ip: str = Depends(upload_ip_limit)): async def share_text(text: str = Form(...), expire_value: int = Form(default=1, gt=0), expire_style: str = Form(default='day'), ip: str = Depends(ip_limit['upload'])):
# 获取过期信息 # 获取过期信息
expired_at, expired_count, used_count, code = await get_expire_info(expire_value, expire_style) expired_at, expired_count, used_count, code = await get_expire_info(expire_value, expire_style)
# 创建一个新的FileCodes实例 # 创建一个新的FileCodes实例
@@ -36,7 +36,7 @@ async def share_text(text: str = Form(...), expire_value: int = Form(default=1,
prefix='文本分享' prefix='文本分享'
) )
# 添加IP到限制列表 # 添加IP到限制列表
upload_ip_limit.add_ip(ip) ip_limit['upload'].add_ip(ip)
# 返回API响应 # 返回API响应
return APIResponse(detail={ return APIResponse(detail={
'code': code, 'code': code,
@@ -45,7 +45,8 @@ async def share_text(text: str = Form(...), expire_value: int = Form(default=1,
# 分享文件的API # 分享文件的API
@share_api.post('/file/', dependencies=[Depends(admin_required)]) @share_api.post('/file/', dependencies=[Depends(admin_required)])
async def share_file(expire_value: int = Form(default=1, gt=0), expire_style: str = Form(default='day'), file: UploadFile = File(...), ip: str = Depends(upload_ip_limit)): async def share_file(expire_value: int = Form(default=1, gt=0), expire_style: str = Form(default='day'), file: UploadFile = File(...),
ip: str = Depends(ip_limit['upload'])):
# 检查文件大小是否超过限制 # 检查文件大小是否超过限制
if file.size > int(settings.uploadSize): if file.size > int(settings.uploadSize):
raise HTTPException(status_code=403, detail=f'文件大小超过限制,最大为{settings.uploadSize}字节') raise HTTPException(status_code=403, detail=f'文件大小超过限制,最大为{settings.uploadSize}字节')
@@ -71,7 +72,7 @@ async def share_file(expire_value: int = Form(default=1, gt=0), expire_style: st
used_count=used_count, used_count=used_count,
) )
# 添加IP到限制列表 # 添加IP到限制列表
upload_ip_limit.add_ip(ip) ip_limit['upload'].add_ip(ip)
# 返回API响应 # 返回API响应
return APIResponse(detail={ return APIResponse(detail={
'code': code, 'code': code,
@@ -94,14 +95,14 @@ async def get_code_file_by_code(code, check=True):
# 获取文件的API # 获取文件的API
@share_api.get('/select/') @share_api.get('/select/')
async def get_code_file(code: str, ip: str = Depends(error_ip_limit)): async def get_code_file(code: str, ip: str = Depends(ip_limit['error'])):
file_storage: FileStorageInterface = storages[settings.file_storage]() file_storage: FileStorageInterface = storages[settings.file_storage]()
# 获取文件 # 获取文件
has, file_code = await get_code_file_by_code(code) has, file_code = await get_code_file_by_code(code)
# 检查文件是否存在 # 检查文件是否存在
if not has: if not has:
# 添加IP到限制列表 # 添加IP到限制列表
error_ip_limit.add_ip(ip) ip_limit['error'].add_ip(ip)
# 返回API响应 # 返回API响应
return APIResponse(code=404, detail=file_code) return APIResponse(code=404, detail=file_code)
# 更新文件的使用次数和过期次数 # 更新文件的使用次数和过期次数
@@ -116,14 +117,14 @@ async def get_code_file(code: str, ip: str = Depends(error_ip_limit)):
# 选择文件的API # 选择文件的API
@share_api.post('/select/') @share_api.post('/select/')
async def select_file(data: SelectFileModel, ip: str = Depends(error_ip_limit)): async def select_file(data: SelectFileModel, ip: str = Depends(ip_limit['error'])):
file_storage: FileStorageInterface = storages[settings.file_storage]() file_storage: FileStorageInterface = storages[settings.file_storage]()
# 获取文件 # 获取文件
has, file_code = await get_code_file_by_code(data.code) has, file_code = await get_code_file_by_code(data.code)
# 检查文件是否存在 # 检查文件是否存在
if not has: if not has:
# 添加IP到限制列表 # 添加IP到限制列表
error_ip_limit.add_ip(ip) ip_limit['error'].add_ip(ip)
# 返回API响应 # 返回API响应
return APIResponse(code=404, detail=file_code) return APIResponse(code=404, detail=file_code)
# 更新文件的使用次数和过期次数 # 更新文件的使用次数和过期次数
@@ -143,13 +144,13 @@ async def select_file(data: SelectFileModel, ip: str = Depends(error_ip_limit)):
# 下载文件的API # 下载文件的API
@share_api.get('/download') @share_api.get('/download')
async def download_file(key: str, code: str, ip: str = Depends(error_ip_limit)): async def download_file(key: str, code: str, ip: str = Depends(ip_limit['error'])):
file_storage: FileStorageInterface = storages[settings.file_storage]() file_storage: FileStorageInterface = storages[settings.file_storage]()
# 检查token是否有效 # 检查token是否有效
is_valid = await get_select_token(code) == key is_valid = await get_select_token(code) == key
if not is_valid: if not is_valid:
# 添加IP到限制列表 # 添加IP到限制列表
error_ip_limit.add_ip(ip) ip_limit['error'].add_ip(ip)
# 获取文件 # 获取文件
has, file_code = await get_code_file_by_code(code, False) has, file_code = await get_code_file_by_code(code, False)
# 检查文件是否存在 # 检查文件是否存在
+3 -3
View File
@@ -7,7 +7,7 @@ import asyncio
from tortoise.expressions import Q from tortoise.expressions import Q
from apps.base.models import FileCodes from apps.base.models import FileCodes
from apps.base.utils import error_ip_limit, upload_ip_limit from apps.base.utils import ip_limit
from core.settings import settings from core.settings import settings
from core.storage import FileStorageInterface, storages from core.storage import FileStorageInterface, storages
from core.utils import get_now from core.utils import get_now
@@ -17,8 +17,8 @@ async def delete_expire_files():
file_storage: FileStorageInterface = storages[settings.file_storage]() file_storage: FileStorageInterface = storages[settings.file_storage]()
while True: while True:
try: try:
await error_ip_limit.remove_expired_ip() await ip_limit['error'].remove_expired_ip()
await upload_ip_limit.remove_expired_ip() await ip_limit['upload'].remove_expired_ip()
expire_data = await FileCodes.filter(Q(expired_at__lt=await get_now()) | Q(expired_count=0)).all() expire_data = await FileCodes.filter(Q(expired_at__lt=await get_now()) | Q(expired_count=0)).all()
for exp in expire_data: for exp in expire_data:
await file_storage.delete_file(exp) await file_storage.delete_file(exp)
+6
View File
@@ -10,7 +10,9 @@ from fastapi.responses import HTMLResponse
from fastapi.staticfiles import StaticFiles from fastapi.staticfiles import StaticFiles
from tortoise.contrib.fastapi import register_tortoise from tortoise.contrib.fastapi import register_tortoise
from apps.base.depends import IPRateLimit
from apps.base.models import KeyValue from apps.base.models import KeyValue
from apps.base.utils import ip_limit
from apps.base.views import share_api from apps.base.views import share_api
from apps.admin.views import admin_api from apps.admin.views import admin_api
from core.response import APIResponse from core.response import APIResponse
@@ -59,6 +61,10 @@ async def startup_event():
# 读取用户配置 # 读取用户配置
user_config, created = await KeyValue.get_or_create(key='settings', defaults={'value': DEFAULT_CONFIG}) user_config, created = await KeyValue.get_or_create(key='settings', defaults={'value': DEFAULT_CONFIG})
settings.user_config = user_config.value settings.user_config = user_config.value
ip_limit['error'].minutes = settings.errorMinute
ip_limit['error'].count = settings.errorCount
ip_limit['upload'].minutes = settings.uploadMinute
ip_limit['upload'].count = settings.uploadCount
@app.get('/') @app.get('/')