Merge pull request #460 from vastsa/dev

feat: 优化代码结构
This commit is contained in:
Lan
2026-01-27 10:52:06 +08:00
committed by GitHub
4 changed files with 259 additions and 189 deletions
+2 -2
View File
@@ -39,9 +39,9 @@ COPY --from=frontend-builder /build/fronted-2023/dist ./themes/2023
RUN pip install --no-cache-dir -r requirements.txt RUN pip install --no-cache-dir -r requirements.txt
# 环境变量配置 # 环境变量配置
ENV HOST="::" \ ENV HOST="0.0.0.0" \
PORT=12345 \ PORT=12345 \
WORKERS=4 \ WORKERS=1 \
LOG_LEVEL="info" LOG_LEVEL="info"
EXPOSE 12345 EXPOSE 12345
+29 -37
View File
@@ -50,7 +50,7 @@ def verify_token(token: str) -> dict:
).digest() ).digest()
expected_signature_b64 = base64.b64encode(expected_signature).decode() expected_signature_b64 = base64.b64encode(expected_signature).decode()
if signature_b64 != expected_signature_b64: if not hmac.compare_digest(signature_b64, expected_signature_b64):
raise ValueError("无效的签名") raise ValueError("无效的签名")
# 解码payload # 解码payload
@@ -65,39 +65,41 @@ def verify_token(token: str) -> dict:
raise ValueError(f"token验证失败: {str(e)}") raise ValueError(f"token验证失败: {str(e)}")
def _extract_bearer_token(authorization: str) -> str:
if not authorization or not authorization.startswith("Bearer "):
raise HTTPException(status_code=401, detail="未授权或授权校验失败")
token = authorization.split(" ", 1)[1].strip()
if not token:
raise HTTPException(status_code=401, detail="未授权或授权校验失败")
return token
def _require_admin_payload(authorization: str) -> dict:
token = _extract_bearer_token(authorization)
try:
payload = verify_token(token)
except ValueError as e:
raise HTTPException(status_code=401, detail=str(e))
if not payload.get("is_admin", False):
raise HTTPException(status_code=401, detail="未授权或授权校验失败")
return payload
ADMIN_PUBLIC_ENDPOINTS = {("POST", "/admin/login")}
async def admin_required( async def admin_required(
authorization: str = Header(default=None), request: Request = None authorization: str = Header(default=None), request: Request = None
): ):
""" """
验证管理员权限 验证管理员权限
""" """
try: if request and (request.method, request.url.path) in ADMIN_PUBLIC_ENDPOINTS:
if not authorization or not authorization.startswith("Bearer "): return None
is_admin = False return _require_admin_payload(authorization)
else:
try:
token = authorization.split(" ")[1]
payload = verify_token(token)
is_admin = payload.get("is_admin", False)
except ValueError as e:
is_admin = False
if request.url.path.startswith("/share/"):
if not settings.openUpload and not is_admin:
raise HTTPException(
status_code=403, detail="本站未开启游客上传,如需上传请先登录后台"
)
else:
if not is_admin:
raise HTTPException(status_code=401, detail="未授权或授权校验失败")
return is_admin
except ValueError as e:
raise HTTPException(status_code=401, detail=str(e))
async def share_required_login( async def share_required_login(authorization: str = Header(default=None)):
authorization: str = Header(default=None), request: Request = None
):
""" """
验证分享上传权限 验证分享上传权限
@@ -109,21 +111,11 @@ async def share_required_login(
:return: 验证结果 :return: 验证结果
""" """
if not settings.openUpload: if not settings.openUpload:
try:
if not authorization or not authorization.startswith("Bearer "): if not authorization or not authorization.startswith("Bearer "):
raise HTTPException( raise HTTPException(
status_code=403, detail="本站未开启游客上传,如需上传请先登录后台" status_code=403, detail="本站未开启游客上传,如需上传请先登录后台"
) )
_require_admin_payload(authorization)
token = authorization.split(" ")[1]
try:
payload = verify_token(token)
if not payload.get("is_admin", False):
raise HTTPException(status_code=401, detail="未授权或授权校验失败")
except ValueError as e:
raise HTTPException(status_code=401, detail=str(e))
except Exception as e:
raise HTTPException(status_code=401, detail="认证失败:" + str(e))
return True return True
+4 -11
View File
@@ -19,7 +19,9 @@ from apps.admin.dependencies import create_token
from core.settings import settings from core.settings import settings
from core.utils import get_now, verify_password from core.utils import get_now, verify_password
admin_api = APIRouter(prefix="/admin", tags=["管理"]) admin_api = APIRouter(
prefix="/admin", tags=["管理"], dependencies=[Depends(admin_required)]
)
@admin_api.post("/login") @admin_api.post("/login")
@@ -32,7 +34,7 @@ async def login(data: LoginData):
@admin_api.get("/dashboard") @admin_api.get("/dashboard")
async def dashboard(admin: bool = Depends(admin_required)): async def dashboard():
all_codes = await FileCodes.all() all_codes = await FileCodes.all()
all_size = str(sum([code.size for code in all_codes])) all_size = str(sum([code.size for code in all_codes]))
sys_start = await KeyValue.filter(key="sys_start").first() sys_start = await KeyValue.filter(key="sys_start").first()
@@ -61,7 +63,6 @@ async def dashboard(admin: bool = Depends(admin_required)):
async def file_delete( async def file_delete(
data: IDData, data: IDData,
file_service: FileService = Depends(get_file_service), file_service: FileService = Depends(get_file_service),
admin: bool = Depends(admin_required),
): ):
await file_service.delete_file(data.id) await file_service.delete_file(data.id)
return APIResponse() return APIResponse()
@@ -73,7 +74,6 @@ async def file_list(
size: int = 10, size: int = 10,
keyword: str = "", keyword: str = "",
file_service: FileService = Depends(get_file_service), file_service: FileService = Depends(get_file_service),
admin: bool = Depends(admin_required),
): ):
files, total = await file_service.list_files(page, size, keyword) files, total = await file_service.list_files(page, size, keyword)
return APIResponse( return APIResponse(
@@ -89,7 +89,6 @@ async def file_list(
@admin_api.get("/config/get") @admin_api.get("/config/get")
async def get_config( async def get_config(
config_service: ConfigService = Depends(get_config_service), config_service: ConfigService = Depends(get_config_service),
admin: bool = Depends(admin_required),
): ):
return APIResponse(detail=config_service.get_config()) return APIResponse(detail=config_service.get_config())
@@ -98,7 +97,6 @@ async def get_config(
async def update_config( async def update_config(
data: dict, data: dict,
config_service: ConfigService = Depends(get_config_service), config_service: ConfigService = Depends(get_config_service),
admin: bool = Depends(admin_required),
): ):
data.pop("themesChoices") data.pop("themesChoices")
await config_service.update_config(data) await config_service.update_config(data)
@@ -109,7 +107,6 @@ async def update_config(
async def file_download( async def file_download(
id: int, id: int,
file_service: FileService = Depends(get_file_service), file_service: FileService = Depends(get_file_service),
admin: bool = Depends(admin_required),
): ):
file_content = await file_service.download_file(id) file_content = await file_service.download_file(id)
return file_content return file_content
@@ -118,7 +115,6 @@ async def file_download(
@admin_api.get("/local/lists") @admin_api.get("/local/lists")
async def get_local_lists( async def get_local_lists(
local_file_service: LocalFileService = Depends(get_local_file_service), local_file_service: LocalFileService = Depends(get_local_file_service),
admin: bool = Depends(admin_required),
): ):
files = await local_file_service.list_files() files = await local_file_service.list_files()
return APIResponse(detail=files) return APIResponse(detail=files)
@@ -128,7 +124,6 @@ async def get_local_lists(
async def delete_local_file( async def delete_local_file(
item: DeleteItem, item: DeleteItem,
local_file_service: LocalFileService = Depends(get_local_file_service), local_file_service: LocalFileService = Depends(get_local_file_service),
admin: bool = Depends(admin_required),
): ):
result = await local_file_service.delete_file(item.filename) result = await local_file_service.delete_file(item.filename)
return APIResponse(detail=result) return APIResponse(detail=result)
@@ -138,7 +133,6 @@ async def delete_local_file(
async def share_local_file( async def share_local_file(
item: ShareItem, item: ShareItem,
file_service: FileService = Depends(get_file_service), file_service: FileService = Depends(get_file_service),
admin: bool = Depends(admin_required),
): ):
share_info = await file_service.share_local_file(item) share_info = await file_service.share_local_file(item)
return APIResponse(detail=share_info) return APIResponse(detail=share_info)
@@ -147,7 +141,6 @@ async def share_local_file(
@admin_api.patch("/file/update") @admin_api.patch("/file/update")
async def update_file( async def update_file(
data: UpdateFileData, data: UpdateFileData,
admin: bool = Depends(admin_required),
): ):
file_code = await FileCodes.filter(id=data.id).first() file_code = await FileCodes.filter(id=data.id).first()
if not file_code: if not file_code:
+125 -40
View File
@@ -4,6 +4,8 @@
# @Software: PyCharm # @Software: PyCharm
import base64 import base64
import hashlib import hashlib
import os
import tempfile
from core.logger import logger from core.logger import logger
import shutil import shutil
from typing import Optional from typing import Optional
@@ -14,7 +16,6 @@ import aiohttp
import asyncio import asyncio
from pathlib import Path from pathlib import Path
import datetime import datetime
import io
import re import re
import aioboto3 import aioboto3
from botocore.config import Config from botocore.config import Config
@@ -285,11 +286,12 @@ class S3FileStorage(FileStorageInterface):
region_name=self.region_name, region_name=self.region_name,
config=Config(signature_version=self.signature_version), config=Config(signature_version=self.signature_version),
) as s3: ) as s3:
await s3.put_object( # 使用 upload_fileobj 流式上传,避免将整个文件加载到内存
Bucket=self.bucket_name, await s3.upload_fileobj(
Key=save_path, file.file,
Body=await file.read(), self.bucket_name,
ContentType=file.content_type, save_path,
ExtraArgs={"ContentType": file.content_type or "application/octet-stream"},
) )
async def delete_file(self, file_code: FileCodes): async def delete_file(self, file_code: FileCodes):
@@ -425,11 +427,10 @@ class S3FileStorage(FileStorageInterface):
async def merge_chunks(self, upload_id: str, chunk_info: UploadChunk, save_path: str) -> tuple[str, str]: async def merge_chunks(self, upload_id: str, chunk_info: UploadChunk, save_path: str) -> tuple[str, str]:
""" """
合并 S3 上的分片文件 合并 S3 上的分片文件
由于分片是独立对象存储的,需要下载后合并再上传 使用 S3 的 multipart upload API 实现流式合并,避免内存问题
""" """
file_sha256 = hashlib.sha256() file_sha256 = hashlib.sha256()
chunk_dir = str(Path(save_path).parent / "chunks" / upload_id) chunk_dir = str(Path(save_path).parent / "chunks" / upload_id)
merged_data = io.BytesIO()
async with self.session.client( async with self.session.client(
's3', 's3',
@@ -438,7 +439,17 @@ class S3FileStorage(FileStorageInterface):
region_name=self.region_name, region_name=self.region_name,
config=Config(signature_version=self.signature_version), config=Config(signature_version=self.signature_version),
) as s3: ) as s3:
# 按顺序读取并验证每个分片 # 创建 multipart upload
mpu = await s3.create_multipart_upload(
Bucket=self.bucket_name,
Key=save_path,
ContentType='application/octet-stream'
)
mpu_id = mpu['UploadId']
parts = []
try:
# 按顺序读取、验证并上传每个分片
for i in range(chunk_info.total_chunks): for i in range(chunk_info.total_chunks):
chunk_key = f"{chunk_dir}/{i}.part" chunk_key = f"{chunk_dir}/{i}.part"
chunk_record = await UploadChunk.filter(upload_id=upload_id, chunk_index=i).first() chunk_record = await UploadChunk.filter(upload_id=upload_id, chunk_index=i).first()
@@ -459,16 +470,38 @@ class S3FileStorage(FileStorageInterface):
raise ValueError(f"分片{i}哈希不匹配: 期望 {chunk_record.chunk_hash}, 实际 {current_hash}") raise ValueError(f"分片{i}哈希不匹配: 期望 {chunk_record.chunk_hash}, 实际 {current_hash}")
file_sha256.update(chunk_data) file_sha256.update(chunk_data)
merged_data.write(chunk_data)
# 上传合并后的文件 # 上传分片到 multipart upload
merged_data.seek(0) part_response = await s3.upload_part(
await s3.put_object(
Bucket=self.bucket_name, Bucket=self.bucket_name,
Key=save_path, Key=save_path,
Body=merged_data.getvalue(), UploadId=mpu_id,
ContentType='application/octet-stream' PartNumber=i + 1, # S3 part numbers start at 1
Body=chunk_data
) )
parts.append({
'PartNumber': i + 1,
'ETag': part_response['ETag']
})
# 释放内存
del chunk_data
# 完成 multipart upload
await s3.complete_multipart_upload(
Bucket=self.bucket_name,
Key=save_path,
UploadId=mpu_id,
MultipartUpload={'Parts': parts}
)
except Exception as e:
# 出错时取消 multipart upload
await s3.abort_multipart_upload(
Bucket=self.bucket_name,
Key=save_path,
UploadId=mpu_id
)
raise e
return save_path, file_sha256.hexdigest() return save_path, file_sha256.hexdigest()
@@ -762,11 +795,16 @@ class OneDriveFileStorage(FileStorageInterface):
current_folder.upload(filename, data).execute_query() current_folder.upload(filename, data).execute_query()
async def merge_chunks(self, upload_id: str, chunk_info: UploadChunk, save_path: str) -> tuple[str, str]: async def merge_chunks(self, upload_id: str, chunk_info: UploadChunk, save_path: str) -> tuple[str, str]:
"""合并 OneDrive 上的分片文件""" """合并 OneDrive 上的分片文件,使用临时文件避免内存问题"""
file_sha256 = hashlib.sha256() file_sha256 = hashlib.sha256()
chunk_dir = str(Path(save_path).parent / "chunks" / upload_id) chunk_dir = str(Path(save_path).parent / "chunks" / upload_id)
merged_data = io.BytesIO()
# 使用临时文件存储合并数据
with tempfile.NamedTemporaryFile(delete=False) as temp_file:
temp_path = temp_file.name
try:
async with aiofiles.open(temp_path, 'wb') as out_file:
for i in range(chunk_info.total_chunks): for i in range(chunk_info.total_chunks):
chunk_path = f"{chunk_dir}/{i}.part" chunk_path = f"{chunk_dir}/{i}.part"
chunk_record = await UploadChunk.filter(upload_id=upload_id, chunk_index=i).first() chunk_record = await UploadChunk.filter(upload_id=upload_id, chunk_index=i).first()
@@ -783,11 +821,17 @@ class OneDriveFileStorage(FileStorageInterface):
raise ValueError(f"分片{i}哈希不匹配: 期望 {chunk_record.chunk_hash}, 实际 {current_hash}") raise ValueError(f"分片{i}哈希不匹配: 期望 {chunk_record.chunk_hash}, 实际 {current_hash}")
file_sha256.update(chunk_data) file_sha256.update(chunk_data)
merged_data.write(chunk_data) await out_file.write(chunk_data)
del chunk_data # 释放内存
# 上传合并后的文件 # 读取临时文件并上传
merged_data.seek(0) async with aiofiles.open(temp_path, 'rb') as f:
await asyncio.to_thread(self._upload_merged, save_path, merged_data.getvalue()) merged_content = await f.read()
await asyncio.to_thread(self._upload_merged, save_path, merged_content)
finally:
# 清理临时文件
if os.path.exists(temp_path):
os.unlink(temp_path)
return save_path, file_sha256.hexdigest() return save_path, file_sha256.hexdigest()
@@ -845,7 +889,9 @@ class OpenDALFileStorage(FileStorageInterface):
) )
async def save_file(self, file: UploadFile, save_path: str): async def save_file(self, file: UploadFile, save_path: str):
await self.operator.write(save_path, file.file.read()) # 使用 asyncio.to_thread 避免阻塞事件循环
content = await asyncio.to_thread(file.file.read)
await self.operator.write(save_path, content)
async def delete_file(self, file_code: FileCodes): async def delete_file(self, file_code: FileCodes):
await self.operator.delete(await file_code.get_file_path()) await self.operator.delete(await file_code.get_file_path())
@@ -915,11 +961,16 @@ class OpenDALFileStorage(FileStorageInterface):
await self.operator.write(chunk_path, chunk_data) await self.operator.write(chunk_path, chunk_data)
async def merge_chunks(self, upload_id: str, chunk_info: UploadChunk, save_path: str) -> tuple[str, str]: async def merge_chunks(self, upload_id: str, chunk_info: UploadChunk, save_path: str) -> tuple[str, str]:
"""合并 OpenDAL 存储上的分片文件""" """合并 OpenDAL 存储上的分片文件,使用临时文件避免内存问题"""
file_sha256 = hashlib.sha256() file_sha256 = hashlib.sha256()
chunk_dir = str(Path(save_path).parent / "chunks" / upload_id) chunk_dir = str(Path(save_path).parent / "chunks" / upload_id)
merged_data = io.BytesIO()
# 使用临时文件存储合并数据
with tempfile.NamedTemporaryFile(delete=False) as temp_file:
temp_path = temp_file.name
try:
async with aiofiles.open(temp_path, 'wb') as out_file:
for i in range(chunk_info.total_chunks): for i in range(chunk_info.total_chunks):
chunk_path = f"{chunk_dir}/{i}.part" chunk_path = f"{chunk_dir}/{i}.part"
chunk_record = await UploadChunk.filter(upload_id=upload_id, chunk_index=i).first() chunk_record = await UploadChunk.filter(upload_id=upload_id, chunk_index=i).first()
@@ -936,11 +987,17 @@ class OpenDALFileStorage(FileStorageInterface):
raise ValueError(f"分片{i}哈希不匹配: 期望 {chunk_record.chunk_hash}, 实际 {current_hash}") raise ValueError(f"分片{i}哈希不匹配: 期望 {chunk_record.chunk_hash}, 实际 {current_hash}")
file_sha256.update(chunk_data) file_sha256.update(chunk_data)
merged_data.write(chunk_data) await out_file.write(chunk_data)
del chunk_data # 释放内存
# 写入合并后的文件 # 读取临时文件并写入存储
merged_data.seek(0) async with aiofiles.open(temp_path, 'rb') as f:
await self.operator.write(save_path, merged_data.getvalue()) merged_content = await f.read()
await self.operator.write(save_path, merged_content)
finally:
# 清理临时文件
if os.path.exists(temp_path):
os.unlink(temp_path)
return save_path, file_sha256.hexdigest() return save_path, file_sha256.hexdigest()
@@ -1033,7 +1090,7 @@ class WebDAVFileStorage(FileStorageInterface):
current_path = current_path.parent current_path = current_path.parent
async def save_file(self, file: UploadFile, save_path: str): async def save_file(self, file: UploadFile, save_path: str):
"""保存文件(自动创建目录)""" """保存文件(自动创建目录,流式上传"""
path_obj = Path(save_path) path_obj = Path(save_path)
directory_path = str(path_obj.parent) directory_path = str(path_obj.parent)
# 提取原始文件名并进行清理 # 提取原始文件名并进行清理
@@ -1044,13 +1101,23 @@ class WebDAVFileStorage(FileStorageInterface):
try: try:
# 先创建目录结构 # 先创建目录结构
await self._mkdir_p(directory_path) await self._mkdir_p(directory_path)
# 上传文件 # 上传文件(流式)
url = self._build_url(safe_save_path) url = self._build_url(safe_save_path)
async def file_sender():
"""流式读取文件内容"""
chunk_size = 256 * 1024 # 256KB chunks
while True:
chunk = await asyncio.to_thread(file.file.read, chunk_size)
if not chunk:
break
yield chunk
async with aiohttp.ClientSession(auth=self.auth) as session: async with aiohttp.ClientSession(auth=self.auth) as session:
content = await file.read()
async with session.put( async with session.put(
url, data=content, headers={ url,
"Content-Type": file.content_type} data=file_sender(),
headers={"Content-Type": file.content_type or "application/octet-stream"}
) as resp: ) as resp:
if resp.status not in (200, 201, 204): if resp.status not in (200, 201, 204):
content = await resp.text() content = await resp.text()
@@ -1161,14 +1228,19 @@ class WebDAVFileStorage(FileStorageInterface):
async def merge_chunks(self, upload_id: str, chunk_info: UploadChunk, save_path: str) -> tuple[str, str]: async def merge_chunks(self, upload_id: str, chunk_info: UploadChunk, save_path: str) -> tuple[str, str]:
""" """
合并 WebDAV 上的分片文件 合并 WebDAV 上的分片文件
由于大多数 WebDAV 服务器不支持 PATCH 追加,这里下载所有分片后合并上传 使用临时文件避免内存问题
""" """
file_sha256 = hashlib.sha256() file_sha256 = hashlib.sha256()
chunk_dir = str(Path(save_path).parent / "chunks" / upload_id) chunk_dir = str(Path(save_path).parent / "chunks" / upload_id)
merged_data = io.BytesIO()
# 使用临时文件存储合并数据,避免内存问题
with tempfile.NamedTemporaryFile(delete=False) as temp_file:
temp_path = temp_file.name
try:
async with aiohttp.ClientSession(auth=self.auth) as session: async with aiohttp.ClientSession(auth=self.auth) as session:
# 按顺序读取并验证每个分片 # 按顺序读取并验证每个分片,写入临时文件
async with aiofiles.open(temp_path, 'wb') as out_file:
for i in range(chunk_info.total_chunks): for i in range(chunk_info.total_chunks):
chunk_path = f"{chunk_dir}/{i}.part" chunk_path = f"{chunk_dir}/{i}.part"
chunk_url = self._build_url(chunk_path) chunk_url = self._build_url(chunk_path)
@@ -1190,22 +1262,35 @@ class WebDAVFileStorage(FileStorageInterface):
raise ValueError(f"分片{i}哈希不匹配: 期望 {chunk_record.chunk_hash}, 实际 {current_hash}") raise ValueError(f"分片{i}哈希不匹配: 期望 {chunk_record.chunk_hash}, 实际 {current_hash}")
file_sha256.update(chunk_data) file_sha256.update(chunk_data)
merged_data.write(chunk_data) await out_file.write(chunk_data)
del chunk_data # 释放内存
# 确保目标目录存在 # 确保目标目录存在
output_dir = str(Path(save_path).parent) output_dir = str(Path(save_path).parent)
await self._mkdir_p(output_dir) await self._mkdir_p(output_dir)
# 上传合并后的文件 # 流式上传合并后的文件
output_url = self._build_url(save_path) output_url = self._build_url(save_path)
merged_data.seek(0)
async with session.put(output_url, data=merged_data.getvalue()) as resp: async def file_sender():
async with aiofiles.open(temp_path, 'rb') as f:
while True:
chunk = await f.read(256 * 1024)
if not chunk:
break
yield chunk
async with session.put(output_url, data=file_sender()) as resp:
if resp.status not in (200, 201, 204): if resp.status not in (200, 201, 204):
content = await resp.text() content = await resp.text()
raise HTTPException( raise HTTPException(
status_code=resp.status, status_code=resp.status,
detail=f"合并文件上传失败: {content[:200]}" detail=f"合并文件上传失败: {content[:200]}"
) )
finally:
# 清理临时文件
if os.path.exists(temp_path):
os.unlink(temp_path)
return save_path, file_sha256.hexdigest() return save_path, file_sha256.hexdigest()