diff --git a/apps/base/depends.py b/apps/base/depends.py
new file mode 100644
index 0000000..e768266
--- /dev/null
+++ b/apps/base/depends.py
@@ -0,0 +1,49 @@
+# @Time : 2023/8/14 12:20
+# @Author : Lan
+# @File : depends.py
+# @Software: PyCharm
+from typing import Union
+from datetime import datetime, timedelta
+
+from fastapi import Header, HTTPException, Request
+
+from core.response import APIResponse
+
+
+async def admin_required(pwd: Union[str, None] = Header(default=None), request: Request = None):
+ return False
+
+
+class IPRateLimit:
+ def __init__(self, count, minutes):
+ self.ips = {}
+ self.count = count
+ self.minutes = minutes
+
+ def check_ip(self, ip):
+ # 检查ip是否被禁止
+ if ip in self.ips:
+ if self.ips[ip]['count'] >= self.count:
+ if self.ips[ip]['time'] + timedelta(minutes=self.minutes) > datetime.now():
+ return False
+ else:
+ self.ips.pop(ip)
+ return True
+
+ def add_ip(self, ip):
+ ip_info = self.ips.get(ip, {'count': 0, 'time': datetime.now()})
+ ip_info['count'] += 1
+ ip_info['time'] = datetime.now()
+ self.ips[ip] = ip_info
+ return ip_info['count']
+
+ async def remove_expired_ip(self):
+ for ip in list(self.ips.keys()):
+ if self.ips[ip]['time'] + timedelta(minutes=self.minutes) < datetime.now():
+ self.ips.pop(ip)
+
+ def __call__(self, request: Request):
+ ip = request.headers.get('X-Real-IP', request.headers.get('X-Forwarded-For', request.client.host))
+ if not self.check_ip(ip):
+ raise HTTPException(status_code=423, detail=f"请求次数过多,请稍后再试")
+ return ip
diff --git a/apps/base/models.py b/apps/base/models.py
index 58d88bb..a50af1e 100644
--- a/apps/base/models.py
+++ b/apps/base/models.py
@@ -28,7 +28,7 @@ class FileCodes(Model):
created_at: Optional[datetime] = fields.DatetimeField(auto_now_add=True, description='创建时间')
async def is_expired(self):
- if self.expired_at:
+ if self.expired_at and (self.expired_count == -1 or self.used_count < self.expired_count):
return self.expired_at < await get_now()
else:
return self.expired_count != -1 and self.used_count >= self.expired_count
diff --git a/apps/base/pydantics.py b/apps/base/pydantics.py
index e69de29..b5d677d 100644
--- a/apps/base/pydantics.py
+++ b/apps/base/pydantics.py
@@ -0,0 +1,5 @@
+from pydantic import BaseModel
+
+
+class SelectFileModel(BaseModel):
+ code: str
diff --git a/apps/base/utils.py b/apps/base/utils.py
index 5c5ed06..ced0e0b 100644
--- a/apps/base/utils.py
+++ b/apps/base/utils.py
@@ -7,6 +7,7 @@ import uuid
import os
from fastapi import UploadFile
+from apps.base.depends import IPRateLimit
from apps.base.models import FileCodes
from core.utils import get_random_num, get_random_string
@@ -41,7 +42,6 @@ async def get_expire_info(expire_value: int, expire_style: str):
:return: expired_at 过期时间, expired_count 可用次数, used_count 已用次数, code 随机码
"""
expired_count, used_count, now, code = -1, 0, datetime.datetime.now(), None
-
if expire_style == 'day':
expired_at = now + datetime.timedelta(days=expire_value)
elif expire_style == 'hour':
@@ -58,6 +58,7 @@ async def get_expire_info(expire_value: int, expire_style: str):
expired_at = now + datetime.timedelta(days=1)
if not code:
code = await get_random_code()
+ print(expire_style, expire_value, expired_at, expired_count, used_count, code)
return expired_at, expired_count, used_count, code
@@ -70,3 +71,9 @@ async def get_random_code(style='num'):
code = await get_random_num() if style == 'num' else await get_random_string()
if not await FileCodes.filter(code=code).exists():
return code
+
+
+# 错误IP限制器
+error_ip_limit = IPRateLimit(1, 1)
+# 上传文件限制器
+upload_ip_limit = IPRateLimit(10, 1)
diff --git a/apps/base/views.py b/apps/base/views.py
index b1f27fa..55aefe1 100644
--- a/apps/base/views.py
+++ b/apps/base/views.py
@@ -2,11 +2,12 @@
# @Author : Lan
# @File : views.py
# @Software: PyCharm
-from fastapi import APIRouter, Form, UploadFile, File
-from pydantic import BaseModel
+from fastapi import APIRouter, Form, UploadFile, File, Depends
from apps.base.models import FileCodes
-from apps.base.utils import get_expire_info, get_file_path_name
+from apps.base.pydantics import SelectFileModel
+from apps.base.utils import get_expire_info, get_file_path_name, error_ip_limit
+from core.response import APIResponse
from core.storage import file_storage
share_api = APIRouter(
@@ -27,13 +28,9 @@ async def share_text(text: str = Form(...), expire_value: int = Form(default=1,
size=len(text),
prefix='文本分享'
)
- return {
- 'code': 200,
- 'msg': 'success',
- 'data': {
- 'code': code,
- }
- }
+ return APIResponse(detail={
+ 'code': code,
+ })
@share_api.post('/file/')
@@ -52,40 +49,25 @@ async def share_file(expire_value: int = Form(default=1, gt=0), expire_style: st
expired_count=expired_count,
used_count=used_count,
)
- return {
- 'code': 200,
- 'msg': 'success',
- 'data': {
- 'code': code,
- 'name': file.filename,
- }
- }
-
-
-class SelectFileModel(BaseModel):
- code: str
+ return APIResponse(detail={
+ 'code': code,
+ 'name': file.filename,
+ })
@share_api.post('/select/')
-async def select_file(data: SelectFileModel):
+async def select_file(data: SelectFileModel, ip: str = Depends(error_ip_limit)):
file_code = await FileCodes.filter(code=data.code).first()
if not file_code:
- return {
- 'code': 404,
- 'msg': '文件不存在',
- }
+ error_ip_limit.add_ip(ip)
+ return APIResponse(code=404, detail='文件不存在')
if await file_code.is_expired():
- return {
- 'code': 403,
- 'msg': '文件已过期',
- }
- return {
- 'code': 200,
- 'msg': 'success',
- 'data': {
- 'code': file_code.code,
- 'name': file_code.prefix + file_code.suffix,
- 'size': file_code.size,
- 'text': await file_storage.get_file_url(file_code),
- }
- }
+ return APIResponse(code=403, detail='文件已过期')
+ file_code.used_count += 1
+ await file_code.save()
+ return APIResponse(detail={
+ 'code': file_code.code,
+ 'name': file_code.prefix + file_code.suffix,
+ 'size': file_code.size,
+ 'text': await file_storage.get_file_url(file_code),
+ })
diff --git a/core/response.py b/core/response.py
new file mode 100644
index 0000000..f6b8a38
--- /dev/null
+++ b/core/response.py
@@ -0,0 +1,15 @@
+# @Time : 2023/8/14 11:48
+# @Author : Lan
+# @File : response.py
+# @Software: PyCharm
+from typing import Generic, TypeVar
+
+from pydantic.v1.generics import GenericModel
+
+T = TypeVar('T')
+
+
+class APIResponse(GenericModel, Generic[T]):
+ code: int = 200
+ message: str = 'ok'
+ detail: T
diff --git a/core/storage.py b/core/storage.py
index 4c66e36..88027af 100644
--- a/core/storage.py
+++ b/core/storage.py
@@ -63,4 +63,4 @@ class S3FileStorage:
return result
-file_storage = S3FileStorage()
+file_storage = SystemFileStorage()
diff --git a/core/utils.py b/core/utils.py
index f6a96cd..ab856b9 100644
--- a/core/utils.py
+++ b/core/utils.py
@@ -6,6 +6,8 @@ import datetime
import random
import string
+from apps.base.depends import IPRateLimit
+
async def get_random_num():
"""
diff --git a/fcb-fronted/components.d.ts b/fcb-fronted/components.d.ts
index a736e4a..6bd7b25 100644
--- a/fcb-fronted/components.d.ts
+++ b/fcb-fronted/components.d.ts
@@ -11,11 +11,9 @@ declare module 'vue' {
ElButton: typeof import('element-plus/es')['ElButton']
ElCard: typeof import('element-plus/es')['ElCard']
ElCol: typeof import('element-plus/es')['ElCol']
- ElDialog: typeof import('element-plus/es')['ElDialog']
ElDrawer: typeof import('element-plus/es')['ElDrawer']
ElIcon: typeof import('element-plus/es')['ElIcon']
ElInput: typeof import('element-plus/es')['ElInput']
- ElModal: typeof import('element-plus/es')['ElModal']
ElOption: typeof import('element-plus/es')['ElOption']
ElProgress: typeof import('element-plus/es')['ElProgress']
ElRadio: typeof import('element-plus/es')['ElRadio']
diff --git a/fcb-fronted/src/components/FileBox.vue b/fcb-fronted/src/components/FileBox.vue
index 566cfcf..49ddb11 100644
--- a/fcb-fronted/src/components/FileBox.vue
+++ b/fcb-fronted/src/components/FileBox.vue
@@ -50,7 +50,8 @@ const copyText = (text: any, style = 0) => {