diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index 59f09770..ad1ba5bf 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -7,6 +7,7 @@ from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_401_UNAUTHORIZE from facefusion import content_store, inference_manager, process_manager, session_context, session_manager, state_manager, translator from facefusion.apis import asset_store from facefusion.apis.session_helper import validate_api_key +from facefusion.apis.stream_manager import destroy_stream from facefusion.filesystem import is_directory, remove_directory @@ -79,6 +80,8 @@ async def destroy_session(request : Request) -> JSONResponse: 'message': translator.get('directory_not_removed', 'facefusion.apis') }, status_code = HTTP_404_NOT_FOUND) + destroy_stream(session_id) + asset_store.delete_assets(session_id) session_manager.clear_session(session_id) diff --git a/facefusion/core.py b/facefusion/core.py index 36317590..311ea781 100755 --- a/facefusion/core.py +++ b/facefusion/core.py @@ -8,7 +8,7 @@ from time import time import uvicorn import facefusion.apis.core -from facefusion import args_helper, benchmarker, cli_helper, content_analyser, content_store, hash_helper, inference_manager, logger, process_manager, state_manager, translator +from facefusion import args_helper, benchmarker, cli_helper, content_analyser, content_store, hash_helper, inference_manager, logger, process_manager, session_manager, state_manager, translator from facefusion.args_helper import apply_args from facefusion.download import conditional_download_hashes, conditional_download_sources from facefusion.exit_helper import hard_exit, signal_exit @@ -309,16 +309,20 @@ def process_batch(args : Args) -> ErrorCode: def process_step(job_id : str, step_index : int, step_args : Args) -> bool: step_total = job_manager.count_step_total(job_id) + session_manager.fork_session() + state_manager.clone_state() + cli_args = args_helper.extract_cli_args(state_manager.get_state()) args = cli_args.copy() args.update(step_args) apply_args(args, state_manager.set_item) logger.info(translator.get('processing_step').format(step_current = step_index + 1, step_total = step_total), __name__) - if common_pre_check() and processors_pre_check(): - error_code = conditional_process() - return error_code == 0 - return False + + is_done = common_pre_check() and processors_pre_check() and conditional_process() == 0 + session_manager.join_session() + + return is_done def conditional_process() -> ErrorCode: diff --git a/facefusion/inference_manager.py b/facefusion/inference_manager.py index 7bd9ca59..bf0ba658 100644 --- a/facefusion/inference_manager.py +++ b/facefusion/inference_manager.py @@ -11,7 +11,7 @@ from facefusion.common_helper import is_windows from facefusion.execution import create_inference_providers, get_onnxruntime_version, has_execution_provider from facefusion.exit_helper import fatal_exit from facefusion.filesystem import get_file_name, is_file -from facefusion.session_context import get_session_id +from facefusion.session_manager import resolve_owner_id from facefusion.time_helper import calculate_end_time from facefusion.types import DownloadSet, ExecutionProvider, InferencePool, InferenceProvider, Store @@ -19,16 +19,16 @@ INFERENCE_POOL_STORE : Store = store_creator.create_store({}) def init() -> None: - session_id = get_session_id() - store_creator.init_content(INFERENCE_POOL_STORE, session_id) + owner_id = resolve_owner_id() + store_creator.init_content(INFERENCE_POOL_STORE, owner_id) def get_inference_pool(module_name : str, model_names : List[str], model_source_set : DownloadSet) -> InferencePool: while process_manager.is_checking(): sleep(0.5) - session_id = get_session_id() - inference_pool_set = store_creator.get_content(INFERENCE_POOL_STORE, session_id) + owner_id = resolve_owner_id() + inference_pool_set = store_creator.get_content(INFERENCE_POOL_STORE, owner_id) execution_device_ids = state_manager.get_item('execution_device_ids') execution_providers = state_manager.get_item('execution_providers') has_arena_leak = has_execution_provider('cuda') and get_onnxruntime_version() > (1, 24, 4) @@ -51,7 +51,9 @@ def get_inference_pool(module_name : str, model_names : List[str], model_source_ def find_inference_pool(inference_context : str) -> Optional[InferencePool]: - for inference_pool_set in INFERENCE_POOL_STORE.get('content_set').values(): + inference_pool_sets = list(INFERENCE_POOL_STORE.get('content_set').values()) + + for inference_pool_set in inference_pool_sets: if inference_pool_set.get(inference_context): return inference_pool_set.get(inference_context) return None @@ -70,8 +72,8 @@ def create_inference_pool(model_source_set : DownloadSet, inference_providers : def clear_inference_pool(module_name : str, model_names : List[str]) -> None: - session_id = get_session_id() - inference_pool_set = store_creator.get_content(INFERENCE_POOL_STORE, session_id) + owner_id = resolve_owner_id() + inference_pool_set = store_creator.get_content(INFERENCE_POOL_STORE, owner_id) execution_device_ids = state_manager.get_item('execution_device_ids') execution_providers = state_manager.get_item('execution_providers') @@ -86,8 +88,8 @@ def clear_inference_pool(module_name : str, model_names : List[str]) -> None: def clear() -> None: - session_id = get_session_id() - store_creator.init_content(INFERENCE_POOL_STORE, session_id) + owner_id = resolve_owner_id() + store_creator.init_content(INFERENCE_POOL_STORE, owner_id) def create_inference_session(model_path : str, inference_providers : List[InferenceProvider]) -> InferenceSession: diff --git a/facefusion/session_manager.py b/facefusion/session_manager.py index 906c2bd5..3a437847 100644 --- a/facefusion/session_manager.py +++ b/facefusion/session_manager.py @@ -3,6 +3,7 @@ from datetime import datetime, timedelta from typing import Dict from typing import Optional +from facefusion.session_context import get_session_id, set_session_id from facefusion.types import Session, SessionId SESSIONS : Dict[SessionId, Session] = {} @@ -20,6 +21,29 @@ def create_session() -> Session: return session +def fork_session() -> SessionId: + owner_id = get_session_id() + session_id = secrets.token_urlsafe(16) + session : Session =\ + { + 'owner_id': owner_id, + 'created_at': datetime.now() + } + + set_session_id(session_id) + set_session(session_id, session) + + return session_id + + +def join_session() -> None: + session_id = get_session_id() + owner_id = resolve_owner_id() + + clear_session(session_id) + set_session_id(owner_id) + + def get_session(session_id : SessionId) -> Optional[Session]: return SESSIONS.get(session_id) @@ -31,6 +55,16 @@ def find_session_id(access_token : str) -> Optional[SessionId]: return None +def resolve_owner_id() -> SessionId: + session_id = get_session_id() + session = get_session(session_id) + + if session and session.get('owner_id'): + return session.get('owner_id') + + return session_id + + def set_session(session_id : SessionId, session : Session) -> None: SESSIONS[session_id] = session diff --git a/facefusion/state_manager.py b/facefusion/state_manager.py index cd058009..3f52d71b 100644 --- a/facefusion/state_manager.py +++ b/facefusion/state_manager.py @@ -5,6 +5,7 @@ from typing import Union from facefusion import store_creator from facefusion.processors.types import ProcessorState, ProcessorStateKey from facefusion.session_context import get_session_id, resolve_local_id +from facefusion.session_manager import resolve_owner_id from facefusion.types import Args, State, StateKey, StateValue, Store STATE_SET : Store = store_creator.create_store({}) @@ -31,6 +32,12 @@ def set_state(state : Union[State, ProcessorState]) -> None: store_creator.set_content(STATE_SET, session_id, state) +def clone_state() -> None: + session_id = get_session_id() + owner_id = resolve_owner_id() + store_creator.set_content(STATE_SET, session_id, deepcopy(store_creator.get_content(STATE_SET, owner_id))) + + def clear() -> None: session_id = get_session_id() store_creator.init_content(STATE_SET, session_id) @@ -62,13 +69,13 @@ def clear_item(key : Union[StateKey, ProcessorStateKey]) -> None: def get_jobs_path() -> str: jobs_path = get_item('jobs_path') - session_id = get_session_id() + owner_id = resolve_owner_id() - return os.path.join(jobs_path, session_id) + return os.path.join(jobs_path, owner_id) def get_temp_path() -> str: temp_path = get_item('temp_path') - session_id = get_session_id() + owner_id = resolve_owner_id() - return os.path.join(temp_path, session_id) + return os.path.join(temp_path, owner_id) diff --git a/facefusion/types.py b/facefusion/types.py index 667f5f65..10c673a9 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -159,10 +159,11 @@ Token : TypeAlias = str SessionId : TypeAlias = str Session = TypedDict('Session', { - 'access_token' : Token, - 'refresh_token' : Token, + 'access_token' : NotRequired[Token], + 'refresh_token' : NotRequired[Token], + 'owner_id' : NotRequired[SessionId], 'created_at' : datetime, - 'expires_at' : datetime + 'expires_at' : NotRequired[datetime] }) StoreContent : TypeAlias = Any diff --git a/tests/test_api_session.py b/tests/test_api_session.py index ed4ed44e..5bdacfd4 100644 --- a/tests/test_api_session.py +++ b/tests/test_api_session.py @@ -1,3 +1,4 @@ +import ctypes import os import tempfile from datetime import timedelta @@ -7,11 +8,12 @@ from unittest.mock import patch import pytest from starlette.testclient import TestClient -from facefusion import metadata, process_manager, session_manager, state_manager +from facefusion import metadata, process_manager, rtc, rtc_store, session_manager, state_manager from facefusion.apis import asset_store from facefusion.apis.core import create_api from facefusion.download import conditional_download -from facefusion.types import Session +from facefusion.libraries import datachannel as datachannel_module +from facefusion.types import RtcPeer, Session from .assert_helper import get_test_example_file, get_test_examples_directory @@ -19,6 +21,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory def before_all() -> None: state_manager.init() + datachannel_module.pre_check() + process_manager.start() conditional_download(get_test_examples_directory(), [ @@ -238,6 +242,21 @@ def test_destroy_session(test_client : TestClient) -> None: assert session_manager.find_session_id(access_token) == session_id assert delete_session_response.status_code == 404 + peer_connection = rtc.create_peer_connection() + rtc_peer : RtcPeer =\ + { + 'peer_connection': peer_connection, + 'video': + { + 'sender_track': rtc.add_video_track(peer_connection, 'sendonly', 'vp8', 96), + 'receiver_track': 0, + 'codec': 'vp8' + }, + 'sender_bitrate': ctypes.c_uint(0), + 'receiver_bitrate': ctypes.c_uint(0) + } + rtc_store.set_peer(session_id, rtc_peer) + delete_session_response = test_client.delete('/session', headers = { 'Authorization': 'Bearer ' + access_token @@ -245,6 +264,7 @@ def test_destroy_session(test_client : TestClient) -> None: assert session_manager.find_session_id(access_token) is None assert asset_store.get_assets(session_id) is None + assert rtc_store.has_peer(session_id) is False assert delete_session_response.status_code == 200 for asset_path in asset_paths: diff --git a/tests/test_session_manager.py b/tests/test_session_manager.py index 59e94938..dd18cfaa 100644 --- a/tests/test_session_manager.py +++ b/tests/test_session_manager.py @@ -1,7 +1,35 @@ import secrets -from datetime import timedelta +from datetime import datetime, timedelta +from typing import Iterator -from facefusion.session_manager import clear_session, create_session, get_session, set_session, validate_session +import pytest + +from facefusion.session_context import get_session_id, resolve_local_id, set_session_id +from facefusion.session_manager import clear_session, create_session, fork_session, get_session, join_session, resolve_owner_id, set_session, validate_session + + +@pytest.fixture(scope = 'function', autouse = True) +def before_each() -> Iterator[None]: + local_id = resolve_local_id() + + set_session_id(local_id) + + yield + + set_session_id(local_id) + + +def test_fork_session() -> None: + local_id = resolve_local_id() + session_id = fork_session() + + assert get_session_id() == session_id + assert resolve_owner_id() == local_id + assert get_session(session_id).get('owner_id') == local_id + assert get_session(session_id).get('access_token') is None + assert get_session(session_id).get('expires_at') is None + + join_session() def test_get_and_set_session() -> None: @@ -13,6 +41,25 @@ def test_get_and_set_session() -> None: assert get_session(session_id) == session +def test_resolve_owner_id() -> None: + local_id = resolve_local_id() + + assert resolve_owner_id() == local_id + + set_session('session-a', create_session()) + set_session('session-a1', { 'owner_id': 'session-a', 'created_at': datetime.now() }) + set_session_id('session-a') + + assert resolve_owner_id() == 'session-a' + + set_session_id('session-a1') + + assert resolve_owner_id() == 'session-a' + + clear_session('session-a1') + clear_session('session-a') + + def test_validate_session() -> None: session = create_session() session_id = secrets.token_urlsafe(16) @@ -43,3 +90,12 @@ def test_clear_session() -> None: clear_session(session_id) assert validate_session(session_id) is None + + +def test_join_session() -> None: + local_id = resolve_local_id() + session_id = fork_session() + join_session() + + assert get_session_id() == local_id + assert get_session(session_id) is None diff --git a/tests/test_state_manager.py b/tests/test_state_manager.py index 1f1cfcce..a8f6e8e3 100644 --- a/tests/test_state_manager.py +++ b/tests/test_state_manager.py @@ -1,10 +1,11 @@ +from datetime import datetime from typing import Iterator import pytest -from facefusion import store_creator +from facefusion import session_manager, store_creator from facefusion.session_context import resolve_local_id, set_session_id -from facefusion.state_manager import STATE_SET, clear, get_item, get_state, init, init_item, set_item, set_state +from facefusion.state_manager import STATE_SET, clear, clone_state, get_item, get_state, init, init_item, set_item, set_state @pytest.fixture(scope = 'function', autouse = True) @@ -14,6 +15,8 @@ def before_each() -> Iterator[None]: set_session_id(local_id) clear() store_creator.delete_content(STATE_SET, 'session-a') + store_creator.delete_content(STATE_SET, 'session-a1') + session_manager.clear_session('session-a1') yield @@ -48,6 +51,20 @@ def test_set_state() -> None: assert get_state() == { 'video_memory_strategy': 'strict' } +def test_clone_state() -> None: + init_item('video_memory_strategy', 'tolerant') + session_manager.set_session('session-a1', { 'owner_id': resolve_local_id(), 'created_at': datetime.now() }) + set_session_id('session-a1') + clone_state() + set_item('video_memory_strategy', 'strict') + + assert get_state() == { 'video_memory_strategy': 'strict' } + + set_session_id(resolve_local_id()) + + assert get_state() == { 'video_memory_strategy': 'tolerant' } + + def test_clear() -> None: init_item('video_memory_strategy', 'tolerant') clear()