mirror of
https://github.com/facefusion/facefusion.git
synced 2026-09-15 20:15:28 +02:00
expose session id via context
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,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(
|
||||
|
||||
Reference in New Issue
Block a user