diff --git a/facefusion/apis/asset_helper.py b/facefusion/apis/asset_helper.py index 1045d8e4..7cf68359 100644 --- a/facefusion/apis/asset_helper.py +++ b/facefusion/apis/asset_helper.py @@ -61,6 +61,23 @@ def validate_asset_files(upload_files : List[UploadFile]) -> bool: return True +def validate_frame_resolution(resolution : str) -> bool: + resolution_limit = 4096 + frame_width, frame_height = unpack_resolution(resolution) + + return frame_width < resolution_limit and frame_height < resolution_limit + + +def validate_file_size(upload_files : List[UploadFile]) -> bool: + size_limit = 512 * 1024 * 1024 + + for upload_file in upload_files: + if upload_file.size > size_limit: + return False + + return True + + async def save_asset_files(upload_files : List[UploadFile]) -> List[str]: asset_paths : List[str] = [] api_security_strategy = state_manager.get_item('api_security_strategy') diff --git a/facefusion/apis/endpoints/assets.py b/facefusion/apis/endpoints/assets.py index 639417bd..69356241 100644 --- a/facefusion/apis/endpoints/assets.py +++ b/facefusion/apis/endpoints/assets.py @@ -7,7 +7,7 @@ from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_400_BAD_REQUEST from facefusion import translator from facefusion.apis import asset_store -from facefusion.apis.asset_helper import capture_asset_faces, capture_asset_frames, save_asset_files, validate_asset_files +from facefusion.apis.asset_helper import capture_asset_faces, capture_asset_frames, save_asset_files, validate_asset_files, validate_frame_resolution, validate_file_size from facefusion.filesystem import remove_file from facefusion.vision import is_vision_frames, to_strip_buffer @@ -47,11 +47,13 @@ async def get_asset(request : Request) -> Response: frame_indexes = request.query_params.getlist('frame_index') vision_frames = [] - if request.query_params.get('subject') == 'frame': - vision_frames = capture_asset_frames(asset, frame_indexes, resolution) #type:ignore[arg-type] + if validate_frame_resolution(resolution): - if request.query_params.get('subject') == 'face': - vision_frames = capture_asset_faces(asset, frame_indexes, resolution) #type:ignore[arg-type] + if request.query_params.get('subject') == 'frame': + vision_frames = capture_asset_frames(asset, frame_indexes, resolution) #type:ignore[arg-type] + + if request.query_params.get('subject') == 'face': + vision_frames = capture_asset_faces(asset, frame_indexes, resolution) #type:ignore[arg-type] if is_vision_frames(vision_frames): return Response(content = to_strip_buffer(vision_frames), media_type = 'image/jpeg') @@ -90,7 +92,7 @@ async def upload_assets(request : Request) -> Response: form = await request.form() upload_files = form.getlist('file') - if upload_files and validate_asset_files(upload_files): + if upload_files and validate_file_size(upload_files) and validate_asset_files(upload_files): asset_paths = await save_asset_files(upload_files) if asset_paths: