Lan
2026-01-07 19:50:02 +08:00
parent ebbf08f06e
commit 830dd65f6a
4 changed files with 246 additions and 155 deletions
+14
View File
@@ -0,0 +1,14 @@
from tortoise import connections
async def add_save_path_to_uploadchunk():
conn = connections.get("default")
await conn.execute_script(
"""
ALTER TABLE uploadchunk ADD COLUMN save_path VARCHAR(512) NULL;
"""
)
async def migrate():
await add_save_path_to_uploadchunk()
+5 -1
View File
@@ -49,6 +49,7 @@ class UploadChunk(models.Model):
file_size = fields.BigIntField() file_size = fields.BigIntField()
chunk_size = fields.IntField() chunk_size = fields.IntField()
file_name = fields.CharField(max_length=255) file_name = fields.CharField(max_length=255)
save_path = fields.CharField(max_length=512, null=True)
created_at = fields.DatetimeField(auto_now_add=True) created_at = fields.DatetimeField(auto_now_add=True)
completed = fields.BooleanField(default=False) completed = fields.BooleanField(default=False)
@@ -66,6 +67,7 @@ class KeyValue(Model):
class PresignUploadSession(models.Model): class PresignUploadSession(models.Model):
"""预签名上传会话模型""" """预签名上传会话模型"""
id = fields.IntField(pk=True) id = fields.IntField(pk=True)
upload_id = fields.CharField(max_length=36, unique=True, index=True) upload_id = fields.CharField(max_length=36, unique=True, index=True)
file_name = fields.CharField(max_length=255) file_name = fields.CharField(max_length=255)
@@ -85,4 +87,6 @@ class PresignUploadSession(models.Model):
file_codes_pydantic = pydantic_model_creator(FileCodes, name="FileCodes") file_codes_pydantic = pydantic_model_creator(FileCodes, name="FileCodes")
upload_chunk_pydantic = pydantic_model_creator(UploadChunk, name="UploadChunk") upload_chunk_pydantic = pydantic_model_creator(UploadChunk, name="UploadChunk")
key_value_pydantic = pydantic_model_creator(KeyValue, name="KeyValue") key_value_pydantic = pydantic_model_creator(KeyValue, name="KeyValue")
presign_upload_session_pydantic = pydantic_model_creator(PresignUploadSession, name="PresignUploadSession") presign_upload_session_pydantic = pydantic_model_creator(
PresignUploadSession, name="PresignUploadSession"
)
+155 -84
View File
@@ -5,13 +5,25 @@ import uuid
from datetime import timedelta from datetime import timedelta
from urllib.parse import unquote from urllib.parse import unquote
from typing import Optional, Tuple, Union
from fastapi import APIRouter, Form, UploadFile, File, Depends, HTTPException from fastapi import APIRouter, Form, UploadFile, File, Depends, HTTPException
from starlette import status from starlette import status
from apps.admin.dependencies import share_required_login from apps.admin.dependencies import share_required_login
from apps.base.models import FileCodes, UploadChunk, PresignUploadSession from apps.base.models import FileCodes, UploadChunk, PresignUploadSession
from apps.base.schemas import SelectFileModel, InitChunkUploadModel, CompleteUploadModel, PresignUploadInitRequest from apps.base.schemas import (
from apps.base.utils import get_expire_info, get_file_path_name, ip_limit, get_chunk_file_path_name SelectFileModel,
InitChunkUploadModel,
CompleteUploadModel,
PresignUploadInitRequest,
)
from apps.base.utils import (
get_expire_info,
get_file_path_name,
ip_limit,
get_chunk_file_path_name,
)
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
@@ -25,7 +37,9 @@ class FileUploadService:
"""统一的文件上传服务""" """统一的文件上传服务"""
@staticmethod @staticmethod
async def generate_file_path(file_name: str, upload_id: str = None) -> tuple[str, str, str, str, str]: async def generate_file_path(
file_name: str, upload_id: Optional[str] = None
) -> tuple[str, str, str, str, str]:
"""统一的路径生成""" """统一的路径生成"""
today = datetime.datetime.now() today = datetime.datetime.now()
storage_path = settings.storage_path.strip("/") storage_path = settings.storage_path.strip("/")
@@ -44,10 +58,12 @@ class FileUploadService:
file_path: str, file_path: str,
expire_value: int, expire_value: int,
expire_style: str, expire_style: str,
**extra_fields **extra_fields,
) -> str: ) -> str:
"""统一创建FileCodes记录,返回code""" """统一创建FileCodes记录,返回code"""
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
)
prefix, suffix = os.path.splitext(file_name) prefix, suffix = os.path.splitext(file_name)
await FileCodes.create( await FileCodes.create(
@@ -60,7 +76,7 @@ class FileUploadService:
expired_at=expired_at, expired_at=expired_at,
expired_count=expired_count, expired_count=expired_count,
used_count=used_count, used_count=used_count,
**extra_fields **extra_fields,
) )
return code return code
@@ -68,8 +84,7 @@ class FileUploadService:
async def validate_file_size(file: UploadFile, max_size: int) -> int: async def validate_file_size(file: UploadFile, max_size: int) -> int:
size = file.size size = file.size
if size is None: if size is None:
# 读取流计算大小,保持指针复位 await file.seek(0, 2) # type: ignore[arg-type]
await file.seek(0, 2)
size = file.file.tell() size = file.file.tell()
await file.seek(0) await file.seek(0)
if size > max_size: if size > max_size:
@@ -122,7 +137,9 @@ async def share_file(
file_size = await validate_file_size(file, settings.uploadSize) file_size = await validate_file_size(file, settings.uploadSize)
if expire_style not in settings.expireStyle: if expire_style not in settings.expireStyle:
raise HTTPException(status_code=400, detail="过期时间类型错误") raise HTTPException(status_code=400, detail="过期时间类型错误")
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
)
path, suffix, prefix, uuid_file_name, save_path = await get_file_path_name(file) path, suffix, prefix, uuid_file_name, save_path = await get_file_path_name(file)
file_storage: FileStorageInterface = storages[settings.file_storage]() file_storage: FileStorageInterface = storages[settings.file_storage]()
await file_storage.save_file(file, save_path) await file_storage.save_file(file, save_path)
@@ -141,7 +158,9 @@ async def share_file(
return APIResponse(detail={"code": code, "name": file.filename}) return APIResponse(detail={"code": code, "name": file.filename})
async def get_code_file_by_code(code, check=True): async def get_code_file_by_code(
code: str, check: bool = True
) -> Tuple[bool, Union[FileCodes, str]]:
file_code = await FileCodes.filter(code=code).first() file_code = await FileCodes.filter(code=code).first()
if not file_code: if not file_code:
return False, "文件不存在" return False, "文件不存在"
@@ -150,7 +169,7 @@ async def get_code_file_by_code(code, check=True):
return True, file_code return True, file_code
async def update_file_usage(file_code): async def update_file_usage(file_code: FileCodes) -> None:
file_code.used_count += 1 file_code.used_count += 1
if file_code.expired_count > 0: if file_code.expired_count > 0:
file_code.expired_count -= 1 file_code.expired_count -= 1
@@ -165,6 +184,7 @@ async def get_code_file(code: str, ip: str = Depends(ip_limit["error"])):
ip_limit["error"].add_ip(ip) ip_limit["error"].add_ip(ip)
return APIResponse(code=404, detail=file_code) return APIResponse(code=404, detail=file_code)
assert isinstance(file_code, FileCodes)
await update_file_usage(file_code) await update_file_usage(file_code)
return await file_storage.get_file_response(file_code) return await file_storage.get_file_response(file_code)
@@ -177,6 +197,7 @@ async def select_file(data: SelectFileModel, ip: str = Depends(ip_limit["error"]
ip_limit["error"].add_ip(ip) ip_limit["error"].add_ip(ip)
return APIResponse(code=404, detail=file_code) return APIResponse(code=404, detail=file_code)
assert isinstance(file_code, FileCodes)
await update_file_usage(file_code) await update_file_usage(file_code)
return APIResponse( return APIResponse(
detail={ detail={
@@ -201,6 +222,7 @@ async def download_file(key: str, code: str, ip: str = Depends(ip_limit["error"]
has, file_code = await get_code_file_by_code(code, False) has, file_code = await get_code_file_by_code(code, False)
if not has: if not has:
return APIResponse(code=404, detail="文件不存在") return APIResponse(code=404, detail="文件不存在")
assert isinstance(file_code, FileCodes)
return ( return (
APIResponse(detail=file_code.text) APIResponse(detail=file_code.text)
if file_code.text if file_code.text
@@ -219,8 +241,7 @@ async def init_chunk_upload(data: InitChunkUploadModel):
if max_possible_size > settings.uploadSize: if max_possible_size > settings.uploadSize:
max_size_mb = settings.uploadSize / (1024 * 1024) max_size_mb = settings.uploadSize / (1024 * 1024)
raise HTTPException( raise HTTPException(
status_code=403, status_code=403, detail=f"文件大小超过限制,最大为 {max_size_mb:.2f} MB"
detail=f"文件大小超过限制,最大为 {max_size_mb:.2f} MB"
) )
# # 秒传检查 # # 秒传检查
@@ -247,21 +268,25 @@ async def init_chunk_upload(data: InitChunkUploadModel):
).first() ).first()
if existing_session: if existing_session:
# 复用已有会话,获取已上传的分片列表 if not existing_session.save_path:
await UploadChunk.filter(upload_id=existing_session.upload_id).delete()
else:
uploaded_chunks = await UploadChunk.filter( uploaded_chunks = await UploadChunk.filter(
upload_id=existing_session.upload_id, upload_id=existing_session.upload_id, completed=True
completed=True ).values_list("chunk_index", flat=True)
).values_list('chunk_index', flat=True) return APIResponse(
return APIResponse(detail={ detail={
"existed": False, "existed": False,
"upload_id": existing_session.upload_id, "upload_id": existing_session.upload_id,
"chunk_size": existing_session.chunk_size, "chunk_size": existing_session.chunk_size,
"total_chunks": existing_session.total_chunks, "total_chunks": existing_session.total_chunks,
"uploaded_chunks": list(uploaded_chunks) "uploaded_chunks": list(uploaded_chunks),
}) }
)
# 创建新的上传会话 # 创建新的上传会话
upload_id = uuid.uuid4().hex upload_id = uuid.uuid4().hex
_, _, _, _, save_path = await get_chunk_file_path_name(data.file_name, upload_id)
await UploadChunk.create( await UploadChunk.create(
upload_id=upload_id, upload_id=upload_id,
chunk_index=-1, chunk_index=-1,
@@ -270,17 +295,23 @@ async def init_chunk_upload(data: InitChunkUploadModel):
chunk_size=data.chunk_size, chunk_size=data.chunk_size,
chunk_hash=data.file_hash, chunk_hash=data.file_hash,
file_name=data.file_name, file_name=data.file_name,
save_path=save_path,
) )
return APIResponse(detail={ return APIResponse(
detail={
"existed": False, "existed": False,
"upload_id": upload_id, "upload_id": upload_id,
"chunk_size": data.chunk_size, "chunk_size": data.chunk_size,
"total_chunks": total_chunks, "total_chunks": total_chunks,
"uploaded_chunks": [] "uploaded_chunks": [],
}) }
)
@chunk_api.post("/upload/chunk/{upload_id}/{chunk_index}", dependencies=[Depends(share_required_login)]) @chunk_api.post(
"/upload/chunk/{upload_id}/{chunk_index}",
dependencies=[Depends(share_required_login)],
)
async def upload_chunk( async def upload_chunk(
upload_id: str, upload_id: str,
chunk_index: int, chunk_index: int,
@@ -297,12 +328,12 @@ async def upload_chunk(
# 检查是否已上传(支持断点续传) # 检查是否已上传(支持断点续传)
existing_chunk = await UploadChunk.filter( existing_chunk = await UploadChunk.filter(
upload_id=upload_id, upload_id=upload_id, chunk_index=chunk_index, completed=True
chunk_index=chunk_index,
completed=True
).first() ).first()
if existing_chunk: if existing_chunk:
return APIResponse(detail={"chunk_hash": existing_chunk.chunk_hash, "skipped": True}) return APIResponse(
detail={"chunk_hash": existing_chunk.chunk_hash, "skipped": True}
)
# 读取分片数据并计算哈希 # 读取分片数据并计算哈希
chunk_data = await chunk.read() chunk_data = await chunk.read()
@@ -312,47 +343,49 @@ async def upload_chunk(
if chunk_size > chunk_info.chunk_size: if chunk_size > chunk_info.chunk_size:
raise HTTPException( raise HTTPException(
status.HTTP_400_BAD_REQUEST, status.HTTP_400_BAD_REQUEST,
detail=f"分片大小超过声明值: 最大 {chunk_info.chunk_size}, 实际 {chunk_size}" detail=f"分片大小超过声明值: 最大 {chunk_info.chunk_size}, 实际 {chunk_size}",
) )
# 计算已上传分片数,校验累计大小不超限(用分片数 * chunk_size 估算) # 计算已上传分片数,校验累计大小不超限(用分片数 * chunk_size 估算)
uploaded_count = await UploadChunk.filter( uploaded_count = await UploadChunk.filter(
upload_id=upload_id, upload_id=upload_id, completed=True
completed=True
).count() ).count()
# 已上传分片的最大可能大小 + 当前分片 # 已上传分片的最大可能大小 + 当前分片
max_uploaded_size = uploaded_count * chunk_info.chunk_size + chunk_size max_uploaded_size = uploaded_count * chunk_info.chunk_size + chunk_size
if max_uploaded_size > settings.uploadSize: if max_uploaded_size > settings.uploadSize:
max_size_mb = settings.uploadSize / (1024 * 1024) max_size_mb = settings.uploadSize / (1024 * 1024)
raise HTTPException( raise HTTPException(
status_code=403, status_code=403, detail=f"累计上传大小超过限制,最大为 {max_size_mb:.2f} MB"
detail=f"累计上传大小超过限制,最大为 {max_size_mb:.2f} MB"
) )
chunk_hash = hashlib.sha256(chunk_data).hexdigest() chunk_hash = hashlib.sha256(chunk_data).hexdigest()
# 获取文件路径 save_path = chunk_info.save_path
_, _, _, _, save_path = await get_chunk_file_path_name(chunk_info.file_name, upload_id)
# 保存分片到存储 # 保存分片到存储
storage = storages[settings.file_storage]() storage = storages[settings.file_storage]()
try: try:
await storage.save_chunk(upload_id, chunk_index, chunk_data, chunk_hash, save_path) await storage.save_chunk(
upload_id, chunk_index, chunk_data, chunk_hash, save_path
)
except Exception as e: except Exception as e:
raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"分片保存失败: {str(e)}") raise HTTPException(
status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"分片保存失败: {str(e)}"
)
# 更新或创建分片记录(保存成功后再记录) # 更新或创建分片记录(保存成功后再记录)
await UploadChunk.update_or_create( await UploadChunk.update_or_create(
upload_id=upload_id, upload_id=upload_id,
chunk_index=chunk_index, chunk_index=chunk_index,
defaults={ defaults={
'chunk_hash': chunk_hash, "chunk_hash": chunk_hash,
'completed': True, "completed": True,
'file_size': chunk_info.file_size, "file_size": chunk_info.file_size,
'total_chunks': chunk_info.total_chunks, "total_chunks": chunk_info.total_chunks,
'chunk_size': chunk_info.chunk_size, "chunk_size": chunk_info.chunk_size,
'file_name': chunk_info.file_name "file_name": chunk_info.file_name,
} "save_path": chunk_info.save_path,
},
) )
return APIResponse(detail={"chunk_hash": chunk_hash}) return APIResponse(detail={"chunk_hash": chunk_hash})
@@ -364,15 +397,14 @@ async def cancel_upload(upload_id: str):
if not chunk_info: if not chunk_info:
raise HTTPException(status.HTTP_404_NOT_FOUND, detail="上传会话不存在") raise HTTPException(status.HTTP_404_NOT_FOUND, detail="上传会话不存在")
# 获取文件路径 save_path = chunk_info.save_path
_, _, _, _, save_path = await get_chunk_file_path_name(chunk_info.file_name, upload_id)
# 清理存储中的临时文件 # 清理存储中的临时文件
storage = storages[settings.file_storage]() storage = storages[settings.file_storage]()
if save_path:
try: try:
await storage.clean_chunks(upload_id, save_path) await storage.clean_chunks(upload_id, save_path)
except Exception as e: except Exception as e:
# 记录错误但不阻止删除数据库记录
pass pass
# 清理数据库记录 # 清理数据库记录
@@ -381,7 +413,9 @@ async def cancel_upload(upload_id: str):
return APIResponse(detail={"message": "上传已取消"}) return APIResponse(detail={"message": "上传已取消"})
@chunk_api.get("/upload/status/{upload_id}", dependencies=[Depends(share_required_login)]) @chunk_api.get(
"/upload/status/{upload_id}", dependencies=[Depends(share_required_login)]
)
async def get_upload_status(upload_id: str): async def get_upload_status(upload_id: str):
"""获取上传状态""" """获取上传状态"""
chunk_info = await UploadChunk.filter(upload_id=upload_id, chunk_index=-1).first() chunk_info = await UploadChunk.filter(upload_id=upload_id, chunk_index=-1).first()
@@ -390,23 +424,28 @@ async def get_upload_status(upload_id: str):
# 获取已上传的分片列表 # 获取已上传的分片列表
uploaded_chunks = await UploadChunk.filter( uploaded_chunks = await UploadChunk.filter(
upload_id=upload_id, upload_id=upload_id, completed=True
completed=True ).values_list("chunk_index", flat=True)
).values_list('chunk_index', flat=True)
return APIResponse(detail={ return APIResponse(
detail={
"upload_id": upload_id, "upload_id": upload_id,
"file_name": chunk_info.file_name, "file_name": chunk_info.file_name,
"file_size": chunk_info.file_size, "file_size": chunk_info.file_size,
"chunk_size": chunk_info.chunk_size, "chunk_size": chunk_info.chunk_size,
"total_chunks": chunk_info.total_chunks, "total_chunks": chunk_info.total_chunks,
"uploaded_chunks": list(uploaded_chunks), "uploaded_chunks": list(uploaded_chunks),
"progress": len(uploaded_chunks) / chunk_info.total_chunks * 100 "progress": len(uploaded_chunks) / chunk_info.total_chunks * 100,
}) }
)
@chunk_api.post("/upload/complete/{upload_id}", dependencies=[Depends(share_required_login)]) @chunk_api.post(
async def complete_upload(upload_id: str, data: CompleteUploadModel, ip: str = Depends(ip_limit["upload"])): "/upload/complete/{upload_id}", dependencies=[Depends(share_required_login)]
)
async def complete_upload(
upload_id: str, data: CompleteUploadModel, ip: str = Depends(ip_limit["upload"])
):
# 获取上传基本信息 # 获取上传基本信息
chunk_info = await UploadChunk.filter(upload_id=upload_id, chunk_index=-1).first() chunk_info = await UploadChunk.filter(upload_id=upload_id, chunk_index=-1).first()
if not chunk_info: if not chunk_info:
@@ -415,8 +454,7 @@ async def complete_upload(upload_id: str, data: CompleteUploadModel, ip: str = D
storage = storages[settings.file_storage]() storage = storages[settings.file_storage]()
# 验证所有分片 # 验证所有分片
completed_chunks_list = await UploadChunk.filter( completed_chunks_list = await UploadChunk.filter(
upload_id=upload_id, upload_id=upload_id, completed=True
completed=True
).all() ).all()
if len(completed_chunks_list) != chunk_info.total_chunks: if len(completed_chunks_list) != chunk_info.total_chunks:
raise HTTPException(status.HTTP_400_BAD_REQUEST, detail="分片不完整") raise HTTPException(status.HTTP_400_BAD_REQUEST, detail="分片不完整")
@@ -424,8 +462,8 @@ async def complete_upload(upload_id: str, data: CompleteUploadModel, ip: str = D
# 用分片数 * chunk_size 校验最大可能大小 # 用分片数 * chunk_size 校验最大可能大小
max_total_size = len(completed_chunks_list) * chunk_info.chunk_size max_total_size = len(completed_chunks_list) * chunk_info.chunk_size
if max_total_size > settings.uploadSize: if max_total_size > settings.uploadSize:
# 清理已上传的分片 save_path = chunk_info.save_path
_, _, _, _, save_path = await get_chunk_file_path_name(chunk_info.file_name, upload_id) if save_path:
try: try:
await storage.clean_chunks(upload_id, save_path) await storage.clean_chunks(upload_id, save_path)
except Exception: except Exception:
@@ -433,18 +471,20 @@ async def complete_upload(upload_id: str, data: CompleteUploadModel, ip: str = D
await UploadChunk.filter(upload_id=upload_id).delete() await UploadChunk.filter(upload_id=upload_id).delete()
max_size_mb = settings.uploadSize / (1024 * 1024) max_size_mb = settings.uploadSize / (1024 * 1024)
raise HTTPException( raise HTTPException(
status_code=403, status_code=403, detail=f"实际上传大小超过限制,最大为 {max_size_mb:.2f} MB"
detail=f"实际上传大小超过限制,最大为 {max_size_mb:.2f} MB"
) )
# 获取文件路径 save_path = chunk_info.save_path
path, suffix, prefix, _, save_path = await get_chunk_file_path_name(chunk_info.file_name, upload_id) path = os.path.dirname(save_path) if save_path else ""
prefix, suffix = os.path.splitext(chunk_info.file_name)
try: try:
# 合并文件并计算哈希 # 合并文件并计算哈希
_, file_hash = await storage.merge_chunks(upload_id, chunk_info, save_path) _, file_hash = await storage.merge_chunks(upload_id, chunk_info, save_path)
# 创建文件记录 # 创建文件记录
expired_at, expired_count, used_count, code = await get_expire_info(data.expire_value, data.expire_style) expired_at, expired_count, used_count, code = await get_expire_info(
data.expire_value, data.expire_style
)
await FileCodes.create( await FileCodes.create(
code=code, code=code,
file_hash=file_hash, # 使用合并后计算的哈希 file_hash=file_hash, # 使用合并后计算的哈希
@@ -457,7 +497,7 @@ async def complete_upload(upload_id: str, data: CompleteUploadModel, ip: str = D
file_path=path, file_path=path,
uuid_file_name=f"{prefix}{suffix}", uuid_file_name=f"{prefix}{suffix}",
prefix=prefix, prefix=prefix,
suffix=suffix suffix=suffix,
) )
# 清理临时文件 # 清理临时文件
await storage.clean_chunks(upload_id, save_path) await storage.clean_chunks(upload_id, save_path)
@@ -473,7 +513,9 @@ async def complete_upload(upload_id: str, data: CompleteUploadModel, ip: str = D
await storage.clean_chunks(upload_id, save_path) await storage.clean_chunks(upload_id, save_path)
except Exception: except Exception:
pass pass
raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"文件合并失败: {str(e)}") raise HTTPException(
status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"文件合并失败: {str(e)}"
)
# ============ 预签名上传API ============ # ============ 预签名上传API ============
@@ -482,7 +524,9 @@ presign_api = APIRouter(prefix="/presign", tags=["预签名上传"])
PRESIGN_SESSION_EXPIRES = 900 # 15分钟 PRESIGN_SESSION_EXPIRES = 900 # 15分钟
async def _get_valid_session(upload_id: str, expected_mode: str = None) -> PresignUploadSession: async def _get_valid_session(
upload_id: str, expected_mode: Optional[str] = None
) -> PresignUploadSession:
"""获取并验证会话""" """获取并验证会话"""
session = await PresignUploadSession.filter(upload_id=upload_id).first() session = await PresignUploadSession.filter(upload_id=upload_id).first()
if not session: if not session:
@@ -496,18 +540,27 @@ async def _get_valid_session(upload_id: str, expected_mode: str = None) -> Presi
@presign_api.post("/upload/init", dependencies=[Depends(share_required_login)]) @presign_api.post("/upload/init", dependencies=[Depends(share_required_login)])
async def presign_upload_init(data: PresignUploadInitRequest, ip: str = Depends(ip_limit["upload"])): async def presign_upload_init(
data: PresignUploadInitRequest, ip: str = Depends(ip_limit["upload"])
):
"""初始化预签名上传,S3返回直传URL,其他存储返回代理URL""" """初始化预签名上传,S3返回直传URL,其他存储返回代理URL"""
if data.file_size > settings.uploadSize: if data.file_size > settings.uploadSize:
raise HTTPException(403, f"文件大小超过限制,最大为 {settings.uploadSize / (1024*1024):.2f} MB") raise HTTPException(
403,
f"文件大小超过限制,最大为 {settings.uploadSize / (1024 * 1024):.2f} MB",
)
if data.expire_style not in settings.expireStyle: if data.expire_style not in settings.expireStyle:
raise HTTPException(400, "过期时间类型错误") raise HTTPException(400, "过期时间类型错误")
upload_id = uuid.uuid4().hex upload_id = uuid.uuid4().hex
path, _, _, filename, save_path = await FileUploadService.generate_file_path(data.file_name, upload_id) path, _, _, filename, save_path = await FileUploadService.generate_file_path(
data.file_name, upload_id
)
storage: FileStorageInterface = storages[settings.file_storage]() storage: FileStorageInterface = storages[settings.file_storage]()
presigned_url = await storage.generate_presigned_upload_url(save_path, PRESIGN_SESSION_EXPIRES) presigned_url = await storage.generate_presigned_upload_url(
save_path, PRESIGN_SESSION_EXPIRES
)
mode = "direct" if presigned_url else "proxy" mode = "direct" if presigned_url else "proxy"
upload_url = presigned_url or f"/api/presign/upload/proxy/{upload_id}" upload_url = presigned_url or f"/api/presign/upload/proxy/{upload_id}"
@@ -524,16 +577,22 @@ async def presign_upload_init(data: PresignUploadInitRequest, ip: str = Depends(
) )
ip_limit["upload"].add_ip(ip) ip_limit["upload"].add_ip(ip)
return APIResponse(detail={ return APIResponse(
detail={
"upload_id": upload_id, "upload_id": upload_id,
"upload_url": upload_url, "upload_url": upload_url,
"mode": mode, "mode": mode,
"expires_in": PRESIGN_SESSION_EXPIRES, "expires_in": PRESIGN_SESSION_EXPIRES,
}) }
)
@presign_api.put("/upload/proxy/{upload_id}", dependencies=[Depends(share_required_login)]) @presign_api.put(
async def presign_upload_proxy(upload_id: str, file: UploadFile = File(...), ip: str = Depends(ip_limit["upload"])): "/upload/proxy/{upload_id}", dependencies=[Depends(share_required_login)]
)
async def presign_upload_proxy(
upload_id: str, file: UploadFile = File(...), ip: str = Depends(ip_limit["upload"])
):
"""代理模式上传,服务器转存到存储后端""" """代理模式上传,服务器转存到存储后端"""
session = await _get_valid_session(upload_id, expected_mode="proxy") session = await _get_valid_session(upload_id, expected_mode="proxy")
@@ -548,8 +607,11 @@ async def presign_upload_proxy(upload_id: str, file: UploadFile = File(...), ip:
raise HTTPException(500, f"文件保存失败: {str(e)}") raise HTTPException(500, f"文件保存失败: {str(e)}")
code = await FileUploadService.create_file_record( code = await FileUploadService.create_file_record(
session.file_name, file_size, os.path.dirname(session.save_path), session.file_name,
session.expire_value, session.expire_style file_size,
os.path.dirname(session.save_path),
session.expire_value,
session.expire_style,
) )
await session.delete() await session.delete()
@@ -557,7 +619,9 @@ async def presign_upload_proxy(upload_id: str, file: UploadFile = File(...), ip:
return APIResponse(detail={"code": code, "name": session.file_name}) return APIResponse(detail={"code": code, "name": session.file_name})
@presign_api.post("/upload/confirm/{upload_id}", dependencies=[Depends(share_required_login)]) @presign_api.post(
"/upload/confirm/{upload_id}", dependencies=[Depends(share_required_login)]
)
async def presign_upload_confirm(upload_id: str, ip: str = Depends(ip_limit["upload"])): async def presign_upload_confirm(upload_id: str, ip: str = Depends(ip_limit["upload"])):
"""直传确认,客户端完成S3直传后调用获取分享码""" """直传确认,客户端完成S3直传后调用获取分享码"""
session = await _get_valid_session(upload_id, expected_mode="direct") session = await _get_valid_session(upload_id, expected_mode="direct")
@@ -567,8 +631,11 @@ async def presign_upload_confirm(upload_id: str, ip: str = Depends(ip_limit["upl
raise HTTPException(404, "文件未上传或上传失败") raise HTTPException(404, "文件未上传或上传失败")
code = await FileUploadService.create_file_record( code = await FileUploadService.create_file_record(
session.file_name, session.file_size, os.path.dirname(session.save_path), session.file_name,
session.expire_value, session.expire_style session.file_size,
os.path.dirname(session.save_path),
session.expire_value,
session.expire_style,
) )
await session.delete() await session.delete()
@@ -576,14 +643,17 @@ async def presign_upload_confirm(upload_id: str, ip: str = Depends(ip_limit["upl
return APIResponse(detail={"code": code, "name": session.file_name}) return APIResponse(detail={"code": code, "name": session.file_name})
@presign_api.get("/upload/status/{upload_id}", dependencies=[Depends(share_required_login)]) @presign_api.get(
"/upload/status/{upload_id}", dependencies=[Depends(share_required_login)]
)
async def presign_upload_status(upload_id: str): async def presign_upload_status(upload_id: str):
"""查询上传会话状态""" """查询上传会话状态"""
session = await PresignUploadSession.filter(upload_id=upload_id).first() session = await PresignUploadSession.filter(upload_id=upload_id).first()
if not session: if not session:
raise HTTPException(404, "上传会话不存在") raise HTTPException(404, "上传会话不存在")
return APIResponse(detail={ return APIResponse(
detail={
"upload_id": session.upload_id, "upload_id": session.upload_id,
"file_name": session.file_name, "file_name": session.file_name,
"file_size": session.file_size, "file_size": session.file_size,
@@ -591,7 +661,8 @@ async def presign_upload_status(upload_id: str):
"created_at": session.created_at.isoformat(), "created_at": session.created_at.isoformat(),
"expires_at": session.expires_at.isoformat(), "expires_at": session.expires_at.isoformat(),
"is_expired": await session.is_expired(), "is_expired": await session.is_expired(),
}) }
)
@presign_api.delete("/upload/{upload_id}", dependencies=[Depends(share_required_login)]) @presign_api.delete("/upload/{upload_id}", dependencies=[Depends(share_required_login)])
+14 -12
View File
@@ -48,36 +48,38 @@ async def delete_expire_files():
async def clean_incomplete_uploads(): async def clean_incomplete_uploads():
"""清理超时未完成的分片上传""" """清理超时未完成的分片上传"""
file_storage: FileStorageInterface = storages[settings.file_storage]() file_storage: FileStorageInterface = storages[settings.file_storage]()
# 默认 24 小时未完成的上传视为过期 expire_hours = getattr(settings, "chunk_expire_hours", 24)
expire_hours = getattr(settings, 'chunk_expire_hours', 24)
while True: while True:
try: try:
expire_time = datetime.datetime.now() - datetime.timedelta(hours=expire_hours) expire_time = datetime.datetime.now() - datetime.timedelta(
# 查找所有过期的上传会话(chunk_index=-1 的记录) hours=expire_hours
)
expired_sessions = await UploadChunk.filter( expired_sessions = await UploadChunk.filter(
chunk_index=-1, chunk_index=-1, created_at__lt=expire_time
created_at__lt=expire_time
).all() ).all()
for session in expired_sessions: for session in expired_sessions:
try: try:
# 获取分片存储路径 save_path = session.save_path
if not save_path:
_, _, _, _, save_path = await get_chunk_file_path_name( _, _, _, _, save_path = await get_chunk_file_path_name(
session.file_name, session.upload_id session.file_name, session.upload_id
) )
# 清理存储中的临时文件
await file_storage.clean_chunks(session.upload_id, save_path) await file_storage.clean_chunks(session.upload_id, save_path)
except Exception as e: except Exception as e:
logging.error(f"清理分片文件失败 upload_id={session.upload_id}: {e}") logging.error(
f"清理分片文件失败 upload_id={session.upload_id}: {e}"
)
try: try:
# 删除该会话的所有数据库记录
await UploadChunk.filter(upload_id=session.upload_id).delete() await UploadChunk.filter(upload_id=session.upload_id).delete()
logging.info(f"已清理过期上传会话 upload_id={session.upload_id}") logging.info(f"已清理过期上传会话 upload_id={session.upload_id}")
except Exception as e: except Exception as e:
logging.error(f"删除分片记录失败 upload_id={session.upload_id}: {e}") logging.error(
f"删除分片记录失败 upload_id={session.upload_id}: {e}"
)
except Exception as e: except Exception as e:
logging.error(f"清理未完成上传任务异常: {e}") logging.error(f"清理未完成上传任务异常: {e}")
finally: finally:
await asyncio.sleep(3600) # 每小时执行一次 await asyncio.sleep(3600)