diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index d7414dcd..8ab5dc27 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -4,7 +4,7 @@ from starlette.requests import Request from starlette.responses import JSONResponse from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_401_UNAUTHORIZED, HTTP_404_NOT_FOUND -from facefusion import session_context, session_manager, state_manager, translator +from facefusion import 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.filesystem import is_directory, remove_directory @@ -18,6 +18,7 @@ async def create_session(request : Request) -> JSONResponse: session = session_manager.create_session() session_context.set_session_id(session_id) session_manager.set_session(session_id, session) + process_manager.init_process_state() return JSONResponse( { @@ -80,6 +81,7 @@ async def destroy_session(request : Request) -> JSONResponse: }, status_code = HTTP_404_NOT_FOUND) asset_store.delete_assets(session_id) + process_manager.clear_process_state() session_manager.clear_session(session_id) return JSONResponse( diff --git a/facefusion/core.py b/facefusion/core.py index 1ca676f9..09249e27 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, logger, state_manager, translator +from facefusion import args_helper, benchmarker, cli_helper, content_analyser, content_store, hash_helper, logger, process_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 @@ -35,6 +35,7 @@ def cli() -> None: if state_manager.get_item('command'): logger.init(state_manager.get_item('log_level')) + process_manager.init_process_state() route(args) else: program.print_help() diff --git a/facefusion/process_manager.py b/facefusion/process_manager.py index f5377fa9..ce897fad 100644 --- a/facefusion/process_manager.py +++ b/facefusion/process_manager.py @@ -1,16 +1,24 @@ -from facefusion.types import ProcessState +from facefusion import store_creator +from facefusion.session_context import get_session_id +from facefusion.types import ProcessState, Store -PROCESS_STATE : ProcessState = 'pending' +PROCESS_STORE : Store = store_creator.create_store('pending') + + +def init_process_state() -> None: + store_creator.init_content(PROCESS_STORE, get_session_id()) def get_process_state() -> ProcessState: - return PROCESS_STATE + return store_creator.get_content(PROCESS_STORE, get_session_id()) def set_process_state(process_state : ProcessState) -> None: - global PROCESS_STATE + store_creator.set_content(PROCESS_STORE, get_session_id(), process_state) - PROCESS_STATE = process_state + +def clear_process_state() -> None: + store_creator.delete_content(PROCESS_STORE, get_session_id()) def is_checking() -> bool: diff --git a/facefusion/session_context.py b/facefusion/session_context.py index d85584fe..b1952cc1 100644 --- a/facefusion/session_context.py +++ b/facefusion/session_context.py @@ -1,4 +1,7 @@ +import hashlib +import uuid from contextvars import ContextVar +from functools import lru_cache from typing import Optional from facefusion.types import SessionId @@ -6,13 +9,14 @@ from facefusion.types import SessionId SESSION_ID : ContextVar[Optional[SessionId]] = ContextVar('SESSION_ID', default = None) +def get_session_id() -> SessionId: + return SESSION_ID.get() or resolve_local_id() + + def set_session_id(session_id : SessionId) -> None: SESSION_ID.set(session_id) -def get_session_id() -> Optional[SessionId]: - return SESSION_ID.get() - - -def clear_session_id() -> None: - SESSION_ID.set(None) +@lru_cache() +def resolve_local_id() -> SessionId: + return hashlib.sha1(str(uuid.getnode()).encode()).hexdigest() diff --git a/facefusion/session_manager.py b/facefusion/session_manager.py index 06284147..906c2bd5 100644 --- a/facefusion/session_manager.py +++ b/facefusion/session_manager.py @@ -37,10 +37,9 @@ def set_session(session_id : SessionId, session : Session) -> None: def validate_session(session_id : SessionId) -> bool: session = get_session(session_id) - return session and datetime.now() <= session.get('expires_at') + return session and datetime.now() < session.get('expires_at') def clear_session(session_id : SessionId) -> None: if session_id in SESSIONS: del SESSIONS[session_id] - diff --git a/facefusion/state_manager.py b/facefusion/state_manager.py index f9b9e7d7..5f2d2b6c 100644 --- a/facefusion/state_manager.py +++ b/facefusion/state_manager.py @@ -48,15 +48,11 @@ def get_jobs_path() -> str: jobs_path = get_item('jobs_path') session_id = get_session_id() - if session_id: - return os.path.join(jobs_path, session_id) - return jobs_path + return os.path.join(jobs_path, session_id) def get_temp_path() -> str: temp_path = get_item('temp_path') session_id = get_session_id() - if session_id: - return os.path.join(temp_path, session_id) - return temp_path + return os.path.join(temp_path, session_id) diff --git a/facefusion/store_creator.py b/facefusion/store_creator.py new file mode 100644 index 00000000..23773dbd --- /dev/null +++ b/facefusion/store_creator.py @@ -0,0 +1,34 @@ +from copy import deepcopy + +from facefusion.types import SessionId, Store, StoreContent + + +def create_store(content : StoreContent) -> Store: + store : Store =\ + { + '__init__': content, + 'content_set': {} + } + + return store + + +def init_content(store : Store, session_id : SessionId) -> None: + store['content_set'][session_id] = deepcopy(store.get('__init__')) + + +def has_content(store : Store, session_id : SessionId) -> bool: + return session_id in store.get('content_set') + + +def get_content(store : Store, session_id : SessionId) -> StoreContent: + return store.get('content_set').get(session_id) + + +def set_content(store : Store, session_id : SessionId, content : StoreContent) -> None: + store['content_set'][session_id] = content + + +def delete_content(store : Store, session_id : SessionId) -> None: + if has_content(store, session_id): + del store['content_set'][session_id] diff --git a/facefusion/types.py b/facefusion/types.py index 16fa9de4..6c3ace26 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -165,6 +165,13 @@ Session = TypedDict('Session', 'expires_at' : datetime }) +StoreContent : TypeAlias = Any +Store = TypedDict('Store', +{ + '__init__' : StoreContent, + 'content_set' : Dict[SessionId, StoreContent] +}) + Command : TypeAlias = str CommandSet : TypeAlias = Dict[str, List[Command]] diff --git a/tests/assert_helper.py b/tests/assert_helper.py index 82c07a7e..64588e6f 100644 --- a/tests/assert_helper.py +++ b/tests/assert_helper.py @@ -2,6 +2,7 @@ import os import tempfile from facefusion.filesystem import are_images, create_directory, is_directory, is_file, remove_directory, resolve_file_paths +from facefusion.session_context import resolve_local_id from facefusion.types import JobStatus @@ -10,7 +11,9 @@ def is_test_job_file(file_path : str, job_status : JobStatus) -> bool: def get_test_job_file(file_path : str, job_status : JobStatus) -> str: - return os.path.join(get_test_jobs_directory(), job_status, file_path) + jobs_path = os.path.join(get_test_jobs_directory(), resolve_local_id()) + + return os.path.join(jobs_path, job_status, file_path) def get_test_jobs_directory() -> str: diff --git a/tests/test_cli_job_manager.py b/tests/test_cli_job_manager.py index 3c88231f..91a0d1e0 100644 --- a/tests/test_cli_job_manager.py +++ b/tests/test_cli_job_manager.py @@ -1,3 +1,4 @@ +import os import subprocess import sys @@ -6,6 +7,7 @@ import pytest from facefusion import ffmpeg, ffmpeg_builder, process_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, count_step_total, init_jobs +from facefusion.session_context import resolve_local_id from .assert_helper import get_test_example_file, get_test_examples_directory, get_test_jobs_directory, get_test_output_path, is_test_job_file @@ -32,8 +34,10 @@ def before_all() -> None: @pytest.fixture(scope = 'function', autouse = True) def before_each() -> None: + jobs_path = os.path.join(get_test_jobs_directory(), resolve_local_id()) + clear_jobs(get_test_jobs_directory()) - init_jobs(get_test_jobs_directory()) + init_jobs(jobs_path) def test_job_list() -> None: diff --git a/tests/test_cli_job_runner.py b/tests/test_cli_job_runner.py index 0a96dc23..9bafc311 100644 --- a/tests/test_cli_job_runner.py +++ b/tests/test_cli_job_runner.py @@ -1,3 +1,4 @@ +import os import subprocess import sys @@ -6,6 +7,7 @@ import pytest from facefusion import ffmpeg, ffmpeg_builder, process_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs, move_job_file, set_steps_status +from facefusion.session_context import resolve_local_id from .assert_helper import get_test_example_file, get_test_examples_directory, get_test_jobs_directory, get_test_output_path, is_test_output_file, prepare_test_output_directory @@ -32,8 +34,11 @@ def before_all() -> None: @pytest.fixture(scope = 'function', autouse = True) def before_each() -> None: + jobs_path = os.path.join(get_test_jobs_directory(), resolve_local_id()) + clear_jobs(get_test_jobs_directory()) - init_jobs(get_test_jobs_directory()) + init_jobs(jobs_path) + prepare_test_output_directory() diff --git a/tests/test_process_manager.py b/tests/test_process_manager.py index 85e64645..b013487a 100644 --- a/tests/test_process_manager.py +++ b/tests/test_process_manager.py @@ -1,4 +1,43 @@ -from facefusion.process_manager import end, is_pending, is_processing, is_stopping, set_process_state, start, stop +import pytest + +from facefusion.process_manager import clear_process_state, end, get_process_state, init_process_state, is_pending, is_processing, is_stopping, set_process_state, start, stop +from facefusion.session_context import set_session_id + + +@pytest.fixture(scope = 'function', autouse = True) +def before_each() -> None: + set_session_id('session-a') + clear_process_state() + set_session_id('session-b') + clear_process_state() + set_session_id('session-a') + + +def test_init_process_state() -> None: + assert get_process_state() is None + + init_process_state() + + assert get_process_state() == 'pending' + + +def test_get_process_state() -> None: + set_process_state('processing') + set_session_id('session-b') + set_process_state('stopping') + + assert get_process_state() == 'stopping' + + set_session_id('session-a') + + assert get_process_state() == 'processing' + + +def test_clear_process_state() -> None: + set_process_state('processing') + clear_process_state() + + assert get_process_state() is None def test_start() -> None: diff --git a/tests/test_session_context.py b/tests/test_session_context.py new file mode 100644 index 00000000..ae3877d5 --- /dev/null +++ b/tests/test_session_context.py @@ -0,0 +1,20 @@ +import pytest + +from facefusion.session_context import get_session_id, resolve_local_id, set_session_id + + +@pytest.fixture(scope = 'function', autouse = True) +def before_each() -> None: + set_session_id(resolve_local_id()) + + +def test_get_session_id() -> None: + assert get_session_id() == resolve_local_id() + + set_session_id('session-a') + + assert get_session_id() == 'session-a' + + +def test_resolve_local_id() -> None: + assert resolve_local_id() == resolve_local_id() diff --git a/tests/test_store_creator.py b/tests/test_store_creator.py new file mode 100644 index 00000000..9df287e5 --- /dev/null +++ b/tests/test_store_creator.py @@ -0,0 +1,64 @@ +from facefusion.store_creator import create_store, delete_content, get_content, has_content, init_content, set_content + + +def test_create_store() -> None: + store = create_store({ 'total': 0 }) + + assert store.get('__init__') == { 'total': 0 } + assert store.get('content_set') == {} + + +def test_init_content() -> None: + store = create_store({ 'total': 0 }) + + init_content(store, 'session-a') + + assert store.get('content_set').get('session-a') == { 'total': 0 } + + +def test_has_content() -> None: + store = create_store({ 'total': 0 }) + + assert has_content(store, 'session-a') is False + + init_content(store, 'session-a') + + assert has_content(store, 'session-a') is True + assert has_content(store, 'session-b') is False + + +def test_get_content() -> None: + store = create_store({ 'total': 0 }) + + init_content(store, 'session-a') + + assert get_content(store, 'session-a') == { 'total': 0 } + assert get_content(store, 'session-b') is None + + get_content(store, 'session-a')['total'] += 1 + + assert get_content(store, 'session-a') == { 'total': 1 } + + +def test_set_content() -> None: + store = create_store({ 'total': 0 }) + + set_content(store, 'session-a', { 'total': 3 }) + + assert get_content(store, 'session-a') == { 'total': 3 } + assert has_content(store, 'session-b') is False + + +def test_delete_content() -> None: + store = create_store({ 'total': 0 }) + + init_content(store, 'session-a') + init_content(store, 'session-b') + delete_content(store, 'session-a') + + assert list(store.get('content_set').keys()) == [ 'session-b' ] + + delete_content(store, 'session-a') + + assert list(store.get('content_set').keys()) == [ 'session-b' ] +