expose session id via context

This commit is contained in:
henryruhs
2026-09-11 21:26:07 +02:00
parent d96580b34c
commit 864038b8f7
6 changed files with 107 additions and 155 deletions
+72 -95
View File
@@ -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)
+2 -7
View File
@@ -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')
+15 -25
View File
@@ -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)
+5 -8
View File
@@ -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:
+11 -19
View File
@@ -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)
+2 -1
View File
@@ -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(