diff --git a/facefusion/core.py b/facefusion/core.py index bb43d402..ead2fe7a 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, face_store, hash_helper, inference_manager, logger, process_manager, session_manager, state_manager, translator, video_manager +from facefusion import args_helper, benchmarker, cli_helper, content_analyser, content_store, face_store, hash_helper, inference_manager, logger, process_manager, session_manager, state_manager, store_creator, translator, video_manager 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 @@ -311,7 +311,7 @@ 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() + fork_id = session_manager.fork_session() state_manager.clone_state() video_manager.init() @@ -326,6 +326,9 @@ def process_step(job_id : str, step_index : int, step_args : Args) -> bool: is_done = common_pre_check() and processors_pre_check() and conditional_process() == 0 session_manager.join_session() + for store in [ state_manager.STATE_SET, content_store.CONTENT_STORE, process_manager.PROCESS_STORE, video_manager.VIDEO_POOL_STORE ]: + store_creator.delete_content(store, fork_id) + return is_done diff --git a/facefusion/session_manager.py b/facefusion/session_manager.py index 3a437847..484bbe20 100644 --- a/facefusion/session_manager.py +++ b/facefusion/session_manager.py @@ -22,25 +22,25 @@ def create_session() -> Session: def fork_session() -> SessionId: + fork_id = secrets.token_urlsafe(16) 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) + set_session_id(fork_id) + set_session(fork_id, session) - return session_id + return fork_id def join_session() -> None: - session_id = get_session_id() + fork_id = get_session_id() owner_id = resolve_owner_id() - clear_session(session_id) + clear_session(fork_id) set_session_id(owner_id) diff --git a/tests/test_session_manager.py b/tests/test_session_manager.py index dd18cfaa..46df3534 100644 --- a/tests/test_session_manager.py +++ b/tests/test_session_manager.py @@ -21,13 +21,13 @@ def before_each() -> Iterator[None]: def test_fork_session() -> None: local_id = resolve_local_id() - session_id = fork_session() + fork_id = fork_session() - assert get_session_id() == session_id + assert get_session_id() == fork_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 + assert get_session(fork_id).get('owner_id') == local_id + assert get_session(fork_id).get('access_token') is None + assert get_session(fork_id).get('expires_at') is None join_session() @@ -94,8 +94,8 @@ def test_clear_session() -> None: def test_join_session() -> None: local_id = resolve_local_id() - session_id = fork_session() + fork_id = fork_session() join_session() assert get_session_id() == local_id - assert get_session(session_id) is None + assert get_session(fork_id) is None