more guards to prevent attacks

This commit is contained in:
henryruhs
2026-09-13 18:56:17 +02:00
parent 93a624f912
commit 19ba62fe1e
2 changed files with 25 additions and 6 deletions
+17
View File
@@ -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')
+8 -6
View File
@@ -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: