增加独立配置文件

This commit is contained in:
veoco
2022-12-11 21:07:39 +08:00
parent 24dc5fbd38
commit 5f5156a262
2 changed files with 62 additions and 48 deletions
+34 -48
View File
@@ -4,6 +4,7 @@ import uuid
import threading import threading
import random import random
import asyncio import asyncio
from pathlib import Path
from fastapi import FastAPI, Depends, UploadFile, Form, File from fastapi import FastAPI, Depends, UploadFile, Form, File
from starlette.requests import Request from starlette.requests import Request
@@ -13,12 +14,17 @@ from starlette.staticfiles import StaticFiles
from sqlalchemy import or_, select, update, delete from sqlalchemy import or_, select, update, delete
from sqlalchemy.ext.asyncio.session import AsyncSession from sqlalchemy.ext.asyncio.session import AsyncSession
import settings
from database import get_session, Codes, init_models, engine from database import get_session, Codes, init_models, engine
app = FastAPI() app = FastAPI(debug=settings.DEBUG)
if not os.path.exists('./static'):
os.makedirs('./static') DATA_ROOT = Path(settings.DATA_ROOT)
app.mount("/static", StaticFiles(directory="static"), name="static") if not DATA_ROOT.exists():
DATA_ROOT.mkdir(parents=True)
STATIC_URL = settings.STATIC_URL
app.mount(STATIC_URL, StaticFiles(directory=DATA_ROOT), name="static")
@app.on_event('startup') @app.on_event('startup')
@@ -27,35 +33,14 @@ async def startup():
asyncio.create_task(delete_expire_files()) asyncio.create_task(delete_expire_files())
############################################
# 需要修改的参数
# 允许错误次数
error_count = 5
# 禁止分钟数
error_minute = 10
# 后台地址
admin_address = 'admin'
# 管理密码
admin_password = 'admin'
# 文件大小限制 10M
file_size_limit = 1024 * 1024 * 10
# 系统标题
title = '文件快递柜'
# 系统描述
description = 'FileCodeBox,文件快递柜,口令传送箱,匿名口令分享文本,文件,图片,视频,音频,压缩包等文件'
# 系统关键字
keywords = 'FileCodeBox,文件快递柜,口令传送箱,匿名口令分享文本,文件,图片,视频,音频,压缩包等文件'
############################################
index_html = open('templates/index.html', 'r', encoding='utf-8').read() \ index_html = open('templates/index.html', 'r', encoding='utf-8').read() \
.replace('{{title}}', title) \ .replace('{{title}}', settings.TITLE) \
.replace('{{description}}', description) \ .replace('{{description}}', settings.DESCRIPTION) \
.replace('{{keywords}}', keywords) .replace('{{keywords}}', settings.KEYWORDS)
admin_html = open('templates/admin.html', 'r', encoding='utf-8').read() \ admin_html = open('templates/admin.html', 'r', encoding='utf-8').read() \
.replace('{{title}}', title) \ .replace('{{title}}', settings.TITLE) \
.replace('{{description}}', description) \ .replace('{{description}}', settings.DESCRIPTION) \
.replace('{{keywords}}', keywords) .replace('{{keywords}}', settings.KEYWORDS)
error_ip_count = {} error_ip_count = {}
@@ -63,7 +48,7 @@ error_ip_count = {}
def delete_file(files): def delete_file(files):
for file in files: for file in files:
if file['type'] != 'text': if file['type'] != 'text':
os.remove('.' + file['text']) os.remove(DATA_ROOT / file['text'].lstrip(STATIC_URL+'/'))
async def delete_expire_files(): async def delete_expire_files():
@@ -92,25 +77,26 @@ def get_file_name(key, ext, file):
now = datetime.datetime.now() now = datetime.datetime.now()
file_bytes = file.file.read() file_bytes = file.file.read()
size = len(file_bytes) size = len(file_bytes)
if size > file_size_limit: if size > settings.FILE_SIZE_LIMIT:
return size, '', '', '' return size, '', '', ''
path = f'./static/upload/{now.year}/{now.month}/{now.day}/' path = DATA_ROOT / f"upload/{now.year}/{now.month}/{now.day}/"
name = f'{key}.{ext}' name = f'{key}.{ext}'
if not os.path.exists(path): if not path.exists():
os.makedirs(path) path.mkdir(parents=True)
with open(f'{os.path.join(path, name)}', 'wb') as f: filepath = path / name
with open(filepath, 'wb') as f:
f.write(file_bytes) f.write(file_bytes)
return size, path[1:] + name, file.content_type, file.filename return size, f"{STATIC_URL}/{filepath.relative_to(DATA_ROOT)}", file.content_type, file.filename
@app.get(f'/{admin_address}') @app.get(f'/{settings.ADMIN_ADDRESS}')
async def admin(): async def admin():
return HTMLResponse(admin_html) return HTMLResponse(admin_html)
@app.post(f'/{admin_address}') @app.post(f'/{settings.ADMIN_ADDRESS}')
async def admin_post(request: Request, s: AsyncSession = Depends(get_session)): async def admin_post(request: Request, s: AsyncSession = Depends(get_session)):
if request.headers.get('pwd') == admin_password: if request.headers.get('pwd') == settings.ADMIN_PASSWORD:
query = select(Codes) query = select(Codes)
codes = (await s.execute(query)).scalars().all() codes = (await s.execute(query)).scalars().all()
return {'code': 200, 'msg': '查询成功', 'data': codes} return {'code': 200, 'msg': '查询成功', 'data': codes}
@@ -118,9 +104,9 @@ async def admin_post(request: Request, s: AsyncSession = Depends(get_session)):
return {'code': 404, 'msg': '密码错误'} return {'code': 404, 'msg': '密码错误'}
@app.delete(f'/{admin_address}') @app.delete(f'/{settings.ADMIN_ADDRESS}')
async def admin_delete(request: Request, code: str, s: AsyncSession = Depends(get_session)): async def admin_delete(request: Request, code: str, s: AsyncSession = Depends(get_session)):
if request.headers.get('pwd') == admin_password: if request.headers.get('pwd') == settings.ADMIN_PASSWORD:
query = select(Codes).where(Codes.code == code) query = select(Codes).where(Codes.code == code)
file = (await s.execute(query)).scalars().first() file = (await s.execute(query)).scalars().first()
await asyncio.to_thread(delete_file, [{'type': file.type, 'text': file.text}]) await asyncio.to_thread(delete_file, [{'type': file.type, 'text': file.text}])
@@ -139,8 +125,8 @@ async def index():
def check_ip(ip): def check_ip(ip):
# 检查ip是否被禁止 # 检查ip是否被禁止
if ip in error_ip_count: if ip in error_ip_count:
if error_ip_count[ip]['count'] >= error_count: if error_ip_count[ip]['count'] >= settings.ERROR_COUNT:
if error_ip_count[ip]['time'] + datetime.timedelta(minutes=error_minute) > datetime.datetime.now(): if error_ip_count[ip]['time'] + datetime.timedelta(minutes=settings.ERROR_MINUTE) > datetime.datetime.now():
return False return False
else: else:
error_ip_count.pop(ip) error_ip_count.pop(ip)
@@ -162,7 +148,7 @@ async def get_file(code: str, s: AsyncSession = Depends(get_session)):
if info.type == 'text': if info.type == 'text':
return {'code': code, 'msg': '查询成功', 'data': info.text} return {'code': code, 'msg': '查询成功', 'data': info.text}
else: else:
return FileResponse('.' + info.text, filename=info.name) return FileResponse(DATA_ROOT / info.text.lstrip(STATIC_URL+'/'), filename=info.name)
else: else:
return {'code': 404, 'msg': '口令不存在'} return {'code': 404, 'msg': '口令不存在'}
@@ -175,7 +161,7 @@ async def index(request: Request, code: str, s: AsyncSession = Depends(get_sessi
query = select(Codes).where(Codes.code == code) query = select(Codes).where(Codes.code == code)
info = (await s.execute(query)).scalars().first() info = (await s.execute(query)).scalars().first()
if not info: if not info:
return {'code': 404, 'msg': f'取件码错误,错误{error_count - ip_error(ip)}次将被禁止10分钟'} return {'code': 404, 'msg': f'取件码错误,错误{settings.ERROR_COUNT - ip_error(ip)}次将被禁止10分钟'}
if info.exp_time < datetime.datetime.now() or info.count == 0: if info.exp_time < datetime.datetime.now() or info.count == 0:
threading.Thread(target=delete_file, args=([{'type': info.type, 'text': info.text}],)).start() threading.Thread(target=delete_file, args=([{'type': info.type, 'text': info.text}],)).start()
await s.delete(info) await s.delete(info)
@@ -215,7 +201,7 @@ async def share(text: str = Form(default=None), style: str = Form(default='2'),
key = uuid.uuid4().hex key = uuid.uuid4().hex
if file: if file:
size, _text, _type, name = get_file_name(key, file.filename.split('.')[-1], file) size, _text, _type, name = get_file_name(key, file.filename.split('.')[-1], file)
if size > file_size_limit: if size > settings.FILE_SIZE_LIMIT:
return {'code': 404, 'msg': '文件过大'} return {'code': 404, 'msg': '文件过大'}
else: else:
size, _text, _type, name = len(text), text, 'text', '文本分享' size, _text, _type, name = len(text), text, 'text', '文本分享'
+28
View File
@@ -0,0 +1,28 @@
from starlette.config import Config
config = Config(".env")
DEBUG = config('DEBUG', cast=bool, default=False)
DATABASE_URL = config('DATABASE_URL', cast=str, default="sqlite+aiosqlite:///database.db")
DATA_ROOT = config('DATA_ROOT', cast=str, default="./static")
STATIC_URL = config('STATIC_URL', cast=str, default="/static")
ERROR_COUNT = config('ERROR_COUNT', cast=int, default=5)
ERROR_MINUTE = config('ERROR_MINUTE', cast=int, default=10)
ADMIN_ADDRESS = config('ADMIN_ADDRESS', cast=str, default="admin")
ADMIN_PASSWORD = config('ADMIN_ADDRESS', cast=str, default="admin")
FILE_SIZE_LIMIT = config('FILE_SIZE_LIMIT', cast=int, default=1024 * 1024 * 10)
TITLE = config('TITLE', cast=str, default="文件快递柜")
DESCRIPTION = config('DESCRIPTION', cast=str, default="FileCodeBox,文件快递柜,口令传送箱,匿名口令分享文本,文件,图片,视频,音频,压缩包等文件")
KEYWORDS = config('TITLE', cast=str, default="FileCodeBox,文件快递柜,口令传送箱,匿名口令分享文本,文件,图片,视频,音频,压缩包等文件")