From 864038b8f7eb26f5157e46265f66311c054bb377 Mon Sep 17 00:00:00 2001 From: henryruhs Date: Fri, 11 Sep 2026 21:26:07 +0200 Subject: [PATCH] expose session id via context --- facefusion/apis/endpoints/assets.py | 167 +++++++++++-------------- facefusion/apis/endpoints/jobs.py | 9 +- facefusion/apis/endpoints/session.py | 40 +++--- facefusion/apis/endpoints/state.py | 13 +- facefusion/apis/endpoints/stream.py | 30 ++--- facefusion/apis/middlewares/session.py | 3 +- 6 files changed, 107 insertions(+), 155 deletions(-) diff --git a/facefusion/apis/endpoints/assets.py b/facefusion/apis/endpoints/assets.py index fadcf74c..b229af67 100644 --- a/facefusion/apis/endpoints/assets.py +++ b/facefusion/apis/endpoints/assets.py @@ -5,80 +5,21 @@ from starlette.requests import Request from starlette.responses import FileResponse, JSONResponse, Response from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_400_BAD_REQUEST, HTTP_404_NOT_FOUND, HTTP_415_UNSUPPORTED_MEDIA_TYPE -from facefusion import session_context, session_manager, translator +from facefusion import session_context, 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.session_helper import extract_access_token from facefusion.filesystem import remove_file from facefusion.vision import is_vision_frames, to_strip_buffer async def get_assets(request : Request) -> Response: - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() + asset_set = asset_store.get_assets(session_id) + assets = [] - if session_id: - asset_set = asset_store.get_assets(session_id) - assets = [] - - if asset_set: - for asset in asset_set.values(): - assets.append( - { - 'id': asset.get('id'), - 'created_at': asset.get('created_at').isoformat(), - 'expires_at': asset.get('expires_at').isoformat(), - 'type': asset.get('type'), - 'media': asset.get('media'), - 'name': asset.get('name'), - 'format': asset.get('format'), - 'size': asset.get('size'), - 'metadata': asset.get('metadata') - }) - - return JSONResponse( - { - 'assets': assets - }, status_code = HTTP_200_OK) - - return JSONResponse( - { - 'message': translator.get('something_went_wrong', 'facefusion.apis') - }, status_code = HTTP_404_NOT_FOUND) - - -async def get_asset(request : Request) -> Response: - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) - asset_id = request.path_params.get('asset_id') - - if session_id and asset_id: - asset = asset_store.get_asset(session_id, asset_id) - - if asset: - if asset.get('media') in [ 'image', 'video' ] and request.query_params.get('action') == 'capture': - resolution = request.query_params.get('resolution') - 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 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') - - return Response(status_code = HTTP_400_BAD_REQUEST) - - if request.query_params.get('action') == 'download': - asset_path = asset.get('path') - - if os.path.exists(asset_path): - return FileResponse(asset_path, filename = asset.get('name')) - - return JSONResponse( + if asset_set: + for asset in asset_set.values(): + assets.append( { 'id': asset.get('id'), 'created_at': asset.get('created_at').isoformat(), @@ -89,7 +30,54 @@ async def get_asset(request : Request) -> Response: 'format': asset.get('format'), 'size': asset.get('size'), 'metadata': asset.get('metadata') - }, status_code = HTTP_200_OK) + }) + + return JSONResponse( + { + 'assets': assets + }, status_code = HTTP_200_OK) + + +async def get_asset(request : Request) -> Response: + session_id = session_context.get_session_id() + asset_id = request.path_params.get('asset_id') + asset = asset_store.get_asset(session_id, asset_id) + + if asset: + if asset.get('media') in [ 'image', 'video' ] and request.query_params.get('action') == 'capture': + resolution = request.query_params.get('resolution') + 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 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') + + return Response(status_code = HTTP_400_BAD_REQUEST) + + if request.query_params.get('action') == 'download': + asset_path = asset.get('path') + + if os.path.exists(asset_path): + return FileResponse(asset_path, filename = asset.get('name')) + + return JSONResponse( + { + 'id': asset.get('id'), + 'created_at': asset.get('created_at').isoformat(), + 'expires_at': asset.get('expires_at').isoformat(), + 'type': asset.get('type'), + 'media': asset.get('media'), + 'name': asset.get('name'), + 'format': asset.get('format'), + 'size': asset.get('size'), + 'metadata': asset.get('metadata') + }, status_code = HTTP_200_OK) return JSONResponse( { @@ -98,13 +86,10 @@ async def get_asset(request : Request) -> Response: async def upload_assets(request : Request) -> Response: - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() asset_type = request.query_params.get('type') - if session_id and asset_type in [ 'source', 'target' ]: - session_context.set_session_id(session_id) - + if asset_type in [ 'source', 'target' ]: form = await request.form() upload_files = form.getlist('file') @@ -135,36 +120,28 @@ async def upload_assets(request : Request) -> Response: async def delete_assets(request : Request) -> Response: - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() + asset_set = asset_store.get_assets(session_id) + asset_ids : List[str] = [] - if session_id: - asset_set = asset_store.get_assets(session_id) - asset_ids : List[str] = [] + if asset_set: + for asset in asset_set.values(): + if remove_file(asset.get('path')): + asset_ids.append(asset.get('id')) - if asset_set: - for asset in asset_set.values(): - if remove_file(asset.get('path')): - asset_ids.append(asset.get('id')) + for asset_id in asset_ids: + asset_store.delete_asset(session_id, asset_id) - for asset_id in asset_ids: - asset_store.delete_asset(session_id, asset_id) - - return Response(status_code = HTTP_200_OK) - - return Response(status_code = HTTP_404_NOT_FOUND) + return Response(status_code = HTTP_200_OK) async def delete_asset(request : Request) -> Response: - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() asset_id = request.path_params.get('asset_id') + asset = asset_store.get_asset(session_id, asset_id) - if session_id and asset_id: - asset = asset_store.get_asset(session_id, asset_id) - - if asset and remove_file(asset.get('path')): - asset_store.delete_asset(session_id, asset_id) - return Response(status_code = HTTP_200_OK) + if asset and remove_file(asset.get('path')): + asset_store.delete_asset(session_id, asset_id) + return Response(status_code = HTTP_200_OK) return Response(status_code = HTTP_404_NOT_FOUND) diff --git a/facefusion/apis/endpoints/jobs.py b/facefusion/apis/endpoints/jobs.py index 37749285..adb59cc2 100644 --- a/facefusion/apis/endpoints/jobs.py +++ b/facefusion/apis/endpoints/jobs.py @@ -8,9 +8,8 @@ from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_202_ACCEPTED, H import facefusion.choices import facefusion.core -from facefusion import args_helper, session_context, session_manager, state_manager, translator +from facefusion import args_helper, session_context, state_manager, translator from facefusion.apis import jobs_helper -from facefusion.apis.session_helper import extract_access_token from facefusion.filesystem import create_directory, get_file_extension, is_directory from facefusion.jobs import job_helper, job_manager, job_runner @@ -118,8 +117,7 @@ async def update_jobs(request : Request) -> JSONResponse: async def update_job(request : Request) -> JSONResponse: job_id = request.path_params.get('job_id') action = request.query_params.get('action') - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() if action == 'submit': if job_manager.submit_job(job_id): @@ -211,9 +209,6 @@ async def create_step(request : Request) -> JSONResponse: step_args['source_paths'] = state_manager.get_item('source_paths') if state_manager.get_item('target_path'): - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) - session_context.set_session_id(session_id) temp_path = state_manager.get_temp_path() step_args['target_path'] = state_manager.get_item('target_path') diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index 7ee4ff2d..a8a3e833 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -6,7 +6,7 @@ from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_401_UNAUTHORIZE from facefusion import content_store, process_manager, session_context, session_manager, state_manager, translator from facefusion.apis import asset_store -from facefusion.apis.session_helper import extract_access_token, validate_api_key +from facefusion.apis.session_helper import validate_api_key from facefusion.filesystem import is_directory, remove_directory @@ -35,8 +35,7 @@ async def create_session(request : Request) -> JSONResponse: async def get_session(request : Request) -> JSONResponse: - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() session = session_manager.get_session(session_id) return JSONResponse( @@ -69,31 +68,22 @@ async def refresh_session(request : Request) -> JSONResponse: async def destroy_session(request : Request) -> JSONResponse: - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) - - if session_id: - session_context.set_session_id(session_id) - temp_path = state_manager.get_temp_path() - - if is_directory(temp_path) and not remove_directory(temp_path): - return JSONResponse( - { - 'message': translator.get('directory_not_removed', 'facefusion.apis') - }, status_code = HTTP_404_NOT_FOUND) - - asset_store.delete_assets(session_id) - session_manager.clear_session(session_id) - - content_store.clear() - process_manager.clear() + session_id = session_context.get_session_id() + temp_path = state_manager.get_temp_path() + if is_directory(temp_path) and not remove_directory(temp_path): return JSONResponse( { - 'message': translator.get('ok', 'facefusion.apis') - }, status_code = HTTP_200_OK) + 'message': translator.get('directory_not_removed', 'facefusion.apis') + }, status_code = HTTP_404_NOT_FOUND) + + asset_store.delete_assets(session_id) + session_manager.clear_session(session_id) + + content_store.clear() + process_manager.clear() return JSONResponse( { - 'message': translator.get('something_went_wrong', 'facefusion.apis') - }, status_code = HTTP_401_UNAUTHORIZED) + 'message': translator.get('ok', 'facefusion.apis') + }, status_code = HTTP_200_OK) diff --git a/facefusion/apis/endpoints/state.py b/facefusion/apis/endpoints/state.py index 1a1a1c5a..3cf985c9 100644 --- a/facefusion/apis/endpoints/state.py +++ b/facefusion/apis/endpoints/state.py @@ -2,9 +2,8 @@ from starlette.requests import Request from starlette.responses import JSONResponse, Response from starlette.status import HTTP_200_OK, HTTP_400_BAD_REQUEST, HTTP_404_NOT_FOUND, HTTP_422_UNPROCESSABLE_CONTENT -from facefusion import args_helper, capability_store, session_manager, state_manager, translator +from facefusion import args_helper, capability_store, session_context, state_manager, translator from facefusion.apis import asset_store -from facefusion.apis.session_helper import extract_access_token async def get_state(request : Request) -> JSONResponse: @@ -50,10 +49,9 @@ async def set_state(request : Request) -> Response: async def select_source(request : Request) -> JSONResponse: body = await request.json() asset_ids = body.get('asset_ids') - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() - if isinstance(asset_ids, list) and session_id: + if isinstance(asset_ids, list): source_paths = [] for asset_id in asset_ids: @@ -76,10 +74,9 @@ async def select_source(request : Request) -> JSONResponse: async def select_target(request : Request) -> JSONResponse: body = await request.json() asset_id = body.get('asset_id') - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() - if isinstance(asset_id, str) and session_id: + if isinstance(asset_id, str): asset = asset_store.get_asset(session_id, asset_id) if asset: diff --git a/facefusion/apis/endpoints/stream.py b/facefusion/apis/endpoints/stream.py index b3bfef14..6fbc239b 100644 --- a/facefusion/apis/endpoints/stream.py +++ b/facefusion/apis/endpoints/stream.py @@ -3,17 +3,13 @@ from starlette.responses import Response from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_404_NOT_FOUND, HTTP_409_CONFLICT from starlette.websockets import WebSocket, WebSocketState -from facefusion import rtc_store, session_context, session_manager +from facefusion import rtc_store, session_context from facefusion.apis.api_helper import get_sec_websocket_protocol -from facefusion.apis.session_helper import extract_access_token from facefusion.apis.stream_manager import destroy_stream, process_image, process_video async def websocket_stream(websocket : WebSocket) -> None: subprotocol = get_sec_websocket_protocol(websocket.scope) - access_token = extract_access_token(websocket.scope) - session_id = session_manager.find_session_id(access_token) - session_context.set_session_id(session_id) await websocket.accept(subprotocol = subprotocol) await process_image(websocket) @@ -27,29 +23,25 @@ async def post_stream(request : Request) -> Response: { 'Location': request.url_for('delete_stream').path } - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) - session_context.set_session_id(session_id) + session_id = session_context.get_session_id() - if session_id: - if not rtc_store.has_peer(session_id): - sdp_offer = await request.body() - sdp_answer = process_video(session_id, sdp_offer.decode()) + if not rtc_store.has_peer(session_id): + sdp_offer = await request.body() + sdp_answer = process_video(session_id, sdp_offer.decode()) - if sdp_answer: - return Response(sdp_answer, status_code = HTTP_201_CREATED, media_type = 'application/sdp', headers = headers) + if sdp_answer: + return Response(sdp_answer, status_code = HTTP_201_CREATED, media_type = 'application/sdp', headers = headers) - else: - return Response(status_code = HTTP_409_CONFLICT) + else: + return Response(status_code = HTTP_409_CONFLICT) return Response(status_code = HTTP_404_NOT_FOUND) async def delete_stream(request : Request) -> Response: - access_token = extract_access_token(request.scope) - session_id = session_manager.find_session_id(access_token) + session_id = session_context.get_session_id() - if session_id and destroy_stream(session_id): + if destroy_stream(session_id): return Response(status_code = HTTP_200_OK) return Response(status_code = HTTP_404_NOT_FOUND) diff --git a/facefusion/apis/middlewares/session.py b/facefusion/apis/middlewares/session.py index 57327e27..9a935295 100644 --- a/facefusion/apis/middlewares/session.py +++ b/facefusion/apis/middlewares/session.py @@ -2,7 +2,7 @@ from starlette.responses import JSONResponse from starlette.status import HTTP_401_UNAUTHORIZED, HTTP_426_UPGRADE_REQUIRED from starlette.types import ASGIApp, Receive, Scope, Send -from facefusion import session_manager, translator +from facefusion import session_context, session_manager, translator from facefusion.apis.session_helper import extract_access_token @@ -15,6 +15,7 @@ def create_session_guard(app : ASGIApp) -> ASGIApp: if session_id: if session_manager.validate_session(session_id): + session_context.set_session_id(session_id) return await app(scope, receive, send) response = JSONResponse(