From fb21d1d1b51b298f6ccf320cba82e2148f702523 Mon Sep 17 00:00:00 2001 From: Henry Ruhs Date: Sat, 12 Sep 2026 11:02:14 +0200 Subject: [PATCH] refactor state manager (#1237) * refactor state manager * refactor state manager * refactor state manager * refactor state manager --- facefusion/apis/endpoints/session.py | 2 + facefusion/apis/stream_manager.py | 38 +++++++++++---- facefusion/core.py | 2 + facefusion/processors/types.py | 3 +- facefusion/state_manager.py | 46 ++++++++++++------ facefusion/thread_helper.py | 8 ++++ facefusion/types.py | 1 - facefusion/workflows/core.py | 6 +-- facefusion/workflows/to_video.py | 6 +-- tests/test_api_assets.py | 2 + tests/test_api_jobs.py | 2 + tests/test_api_session.py | 2 + tests/test_api_state.py | 2 + tests/test_api_stream.py | 1 + tests/test_api_stream_audio.py | 1 + tests/test_api_stream_manager.py | 23 ++++++--- tests/test_api_stream_video.py | 1 + tests/test_audio.py | 4 +- tests/test_cli_age_modifier.py | 4 +- tests/test_cli_background_remover.py | 4 +- tests/test_cli_batch_runner.py | 4 +- tests/test_cli_expression_restorer.py | 4 +- tests/test_cli_face_debugger.py | 4 +- tests/test_cli_face_editor.py | 4 +- tests/test_cli_face_enhancer.py | 4 +- tests/test_cli_face_swapper.py | 4 +- tests/test_cli_frame_colorizer.py | 4 +- tests/test_cli_frame_enhancer.py | 4 +- tests/test_cli_job_manager.py | 4 +- tests/test_cli_job_runner.py | 4 +- tests/test_cli_lip_syncer.py | 4 +- tests/test_cli_output_fps.py | 4 +- tests/test_cli_output_scale.py | 4 +- tests/test_codec_aom_decoder.py | 1 + tests/test_codec_aom_encoder.py | 1 + tests/test_codec_opus_decoder.py | 1 + tests/test_codec_opus_encoder.py | 1 + tests/test_codec_vpx_decoder.py | 1 + tests/test_codec_vpx_encoder.py | 1 + tests/test_config.py | 1 + tests/test_face_creator.py | 2 + tests/test_face_detector.py | 2 + tests/test_face_tracker.py | 2 + tests/test_ffmpeg.py | 2 + tests/test_ffprobe.py | 4 +- tests/test_inference_manager.py | 1 + tests/test_job_runner.py | 4 +- tests/test_library_aom.py | 1 + tests/test_library_datachannel.py | 1 + tests/test_library_opus.py | 1 + tests/test_library_vpx.py | 1 + tests/test_process_manager.py | 10 +++- tests/test_rtc.py | 1 + tests/test_session_context.py | 17 +++++-- tests/test_state_manager.py | 68 ++++++++++++++++++++------- tests/test_temp_helper.py | 1 + tests/test_thread_helper.py | 22 +++++++++ tests/test_video_manager.py | 2 + tests/test_vision.py | 4 +- tests/test_workflow.py | 2 + 60 files changed, 285 insertions(+), 80 deletions(-) create mode 100644 tests/test_thread_helper.py diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index a8a3e833..7761aeec 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -19,6 +19,7 @@ async def create_session(request : Request) -> JSONResponse: session_context.set_session_id(session_id) session_manager.set_session(session_id, session) + state_manager.init() content_store.init() process_manager.init() @@ -80,6 +81,7 @@ async def destroy_session(request : Request) -> JSONResponse: asset_store.delete_assets(session_id) session_manager.clear_session(session_id) + state_manager.clear() content_store.clear() process_manager.clear() diff --git a/facefusion/apis/stream_manager.py b/facefusion/apis/stream_manager.py index 3ac9101a..f94571c7 100644 --- a/facefusion/apis/stream_manager.py +++ b/facefusion/apis/stream_manager.py @@ -1,13 +1,14 @@ import ctypes import threading from collections.abc import AsyncIterator -from concurrent.futures import Future, ThreadPoolExecutor +from concurrent.futures import Future +from contextvars import copy_context from queue import Queue from typing import Optional, Tuple from starlette.websockets import WebSocket -from facefusion import content_store, rtc, rtc_store, session_context, state_manager, streamer +from facefusion import content_store, rtc, rtc_store, state_manager, streamer, thread_helper from facefusion.apis.stream_audio import receive_audio_frames, run_audio_encode_loop from facefusion.apis.stream_video import receive_video_frames, run_video_encode_loop from facefusion.content_analyser import analyse_frame @@ -94,7 +95,11 @@ def process_video(session_id : SessionId, sdp_offer : SdpOffer) -> Optional[SdpA content_store.clear() rtc_store.set_peer(session_id, rtc_peer) - threading.Thread(target = run_peer_loop, args = (session_id, rtc_peer), daemon = True).start() + threading.Thread( + target = copy_context().run, + args = (run_peer_loop, session_id, rtc_peer), + daemon = True + ).start() return sdp_answer @@ -104,20 +109,35 @@ def process_video(session_id : SessionId, sdp_offer : SdpOffer) -> Optional[SdpA def run_peer_loop(session_id : SessionId, rtc_peer : RtcPeer) -> None: - session_context.set_session_id(session_id) execution_thread_count = state_manager.get_item('execution_thread_count') video_queue : Queue[Tuple[Time, Future[BufferPack]]] = Queue(maxsize = execution_thread_count) audio_queue : Queue[Tuple[Time, AudioFrame]] = Queue(maxsize = execution_thread_count * 10) - video_executor = ThreadPoolExecutor(max_workers = execution_thread_count, initializer = session_context.set_session_id, initargs = (session_id,)) + video_executor = thread_helper.create_executor(execution_thread_count) - video_receiver_thread = threading.Thread(target = receive_video_frames, args = (rtc_peer.get('video'), video_queue, video_executor), daemon = True) - video_encoder_thread = threading.Thread(target = run_video_encode_loop, args = (rtc_peer, video_queue), daemon = True) + video_receiver_thread = threading.Thread( + target = copy_context().run, + args = (receive_video_frames, rtc_peer.get('video'), video_queue, video_executor), + daemon = True + ) + video_encoder_thread = threading.Thread( + target = copy_context().run, + args = (run_video_encode_loop, rtc_peer, video_queue), + daemon = True + ) video_receiver_thread.start() video_encoder_thread.start() if rtc_peer.get('audio'): - audio_receiver_thread = threading.Thread(target = receive_audio_frames, args = (rtc_peer.get('audio'), audio_queue), daemon = True) - audio_encoder_thread = threading.Thread(target = run_audio_encode_loop, args = (rtc_peer, audio_queue), daemon = True) + audio_receiver_thread = threading.Thread( + target = copy_context().run, + args = (receive_audio_frames, rtc_peer.get('audio'), audio_queue), + daemon = True + ) + audio_encoder_thread = threading.Thread( + target = copy_context().run, + args = (run_audio_encode_loop, rtc_peer, audio_queue), + daemon = True + ) audio_receiver_thread.start() audio_encoder_thread.start() audio_receiver_thread.join() diff --git a/facefusion/core.py b/facefusion/core.py index 329ed5fd..222b1076 100755 --- a/facefusion/core.py +++ b/facefusion/core.py @@ -25,6 +25,8 @@ from facefusion.workflows.core import detect_workflow_mode def cli() -> None: + state_manager.init() + if pre_check(): signal.signal(signal.SIGINT, signal_exit) program = create_program() diff --git a/facefusion/processors/types.py b/facefusion/processors/types.py index a172f8f1..547fc9be 100644 --- a/facefusion/processors/types.py +++ b/facefusion/processors/types.py @@ -2,7 +2,7 @@ from typing import Any, Callable, Dict, Tuple, TypeAlias from numpy.typing import NDArray -from facefusion.types import AppContext, Mask, VisionFrame +from facefusion.types import Mask, VisionFrame LivePortraitPitch : TypeAlias = float LivePortraitYaw : TypeAlias = float @@ -17,7 +17,6 @@ LivePortraitTranslation : TypeAlias = NDArray[Any] ProcessorStateValue : TypeAlias = Any ProcessorStateKey : TypeAlias = str ProcessorState : TypeAlias = Dict[ProcessorStateKey, ProcessorStateValue] -ProcessorStateSet : TypeAlias = Dict[AppContext, ProcessorState] ApplyStateItem : TypeAlias = Callable[[ProcessorStateKey, ProcessorStateValue], None] diff --git a/facefusion/state_manager.py b/facefusion/state_manager.py index 5f2d2b6c..cd058009 100644 --- a/facefusion/state_manager.py +++ b/facefusion/state_manager.py @@ -1,21 +1,39 @@ import os +from copy import deepcopy from typing import Union -from facefusion.app_context import detect_app_context -from facefusion.processors.types import ProcessorState, ProcessorStateKey, ProcessorStateSet -from facefusion.session_context import get_session_id -from facefusion.types import Args, State, StateKey, StateSet, StateValue +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.types import Args, State, StateKey, StateValue, Store -STATE_SET : Union[StateSet, ProcessorStateSet] =\ -{ - 'api': {}, #type:ignore[assignment] - 'cli': {} #type:ignore[assignment] -} +STATE_SET : Store = store_creator.create_store({}) + + +def init() -> None: + session_id = get_session_id() + local_id = resolve_local_id() + + if session_id == local_id: + store_creator.init_content(STATE_SET, session_id) + else: + store_creator.set_content(STATE_SET, session_id, deepcopy(store_creator.get_content(STATE_SET, local_id))) def get_state() -> Union[State, ProcessorState]: - app_context = detect_app_context() - return STATE_SET.get(app_context) + session_id = get_session_id() + + return store_creator.get_content(STATE_SET, session_id) + + +def set_state(state : Union[State, ProcessorState]) -> None: + session_id = get_session_id() + store_creator.set_content(STATE_SET, session_id, state) + + +def clear() -> None: + session_id = get_session_id() + store_creator.init_content(STATE_SET, session_id) def collect_state(args : Args) -> Union[State, ProcessorState]: @@ -27,8 +45,7 @@ def collect_state(args : Args) -> Union[State, ProcessorState]: def init_item(key : Union[StateKey, ProcessorStateKey], value : StateValue) -> None: - STATE_SET['api'][key] = value #type:ignore[literal-required] - STATE_SET['cli'][key] = value #type:ignore[literal-required] + get_state()[key] = value #type:ignore[literal-required] def get_item(key : Union[StateKey, ProcessorStateKey]) -> StateValue: @@ -36,8 +53,7 @@ def get_item(key : Union[StateKey, ProcessorStateKey]) -> StateValue: def set_item(key : Union[StateKey, ProcessorStateKey], value : StateValue) -> None: - app_context = detect_app_context() - STATE_SET[app_context][key] = value #type:ignore[literal-required] + get_state()[key] = value #type:ignore[literal-required] def clear_item(key : Union[StateKey, ProcessorStateKey]) -> None: diff --git a/facefusion/thread_helper.py b/facefusion/thread_helper.py index 9bd1b7b5..9d5919f3 100644 --- a/facefusion/thread_helper.py +++ b/facefusion/thread_helper.py @@ -1,9 +1,11 @@ import threading +from concurrent.futures import ThreadPoolExecutor from contextlib import nullcontext from typing import ContextManager, Union from facefusion.common_helper import is_linux, is_windows from facefusion.execution import has_execution_provider +from facefusion.session_context import get_session_id, set_session_id THREAD_LOCK : threading.Lock = threading.Lock() THREAD_SEMAPHORE : threading.Semaphore = threading.Semaphore() @@ -22,3 +24,9 @@ def conditional_thread_semaphore() -> Union[threading.Semaphore, ContextManager[ if is_windows() and has_execution_provider('directml') or is_linux() and has_execution_provider('migraphx') or is_linux() and has_execution_provider('rocm'): return THREAD_SEMAPHORE return NULL_CONTEXT + + +def create_executor(max_workers : int) -> ThreadPoolExecutor: + session_id = get_session_id() + + return ThreadPoolExecutor(max_workers = max_workers, initializer = set_session_id, initargs = (session_id,)) diff --git a/facefusion/types.py b/facefusion/types.py index 6c3ace26..303844ae 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -676,6 +676,5 @@ State = TypedDict('State', 'job_status' : JobStatus, 'step_index' : int }) -StateSet : TypeAlias = Dict[AppContext, State] ApplyStateItem : TypeAlias = Callable[[StateKey, StateValue], None] diff --git a/facefusion/workflows/core.py b/facefusion/workflows/core.py index 4d394c0d..bf47ce18 100644 --- a/facefusion/workflows/core.py +++ b/facefusion/workflows/core.py @@ -1,10 +1,10 @@ from collections import deque -from concurrent.futures import Future, ThreadPoolExecutor +from concurrent.futures import Future from typing import Deque, List import numpy -from facefusion import cli_progress, logger, process_manager, state_manager, translator +from facefusion import cli_progress, logger, process_manager, state_manager, thread_helper, translator from facefusion.audio import create_empty_audio_frame, get_audio_frame, get_voice_frame from facefusion.common_helper import get_first from facefusion.filesystem import filter_audio_paths, get_file_extension, has_audio, has_image, has_video @@ -159,7 +159,7 @@ def process_frames() -> ErrorCode: progress.set_title(translator.get('processing')) progress.count(temp_frame_set) - with ThreadPoolExecutor(max_workers = state_manager.get_item('execution_thread_count')) as executor: + with thread_helper.create_executor(state_manager.get_item('execution_thread_count')) as executor: futures : Deque[Future[bool]] = deque() for frame_index, temp_frame_path in temp_frame_set.items(): diff --git a/facefusion/workflows/to_video.py b/facefusion/workflows/to_video.py index 83236b15..6cc5f2a6 100644 --- a/facefusion/workflows/to_video.py +++ b/facefusion/workflows/to_video.py @@ -1,11 +1,11 @@ from collections import deque -from concurrent.futures import Future, ThreadPoolExecutor +from concurrent.futures import Future from typing import Deque import cv2 import numpy -from facefusion import cli_progress, content_analyser, ffmpeg, logger, process_manager, state_manager, translator, video_manager +from facefusion import cli_progress, content_analyser, ffmpeg, logger, process_manager, state_manager, thread_helper, translator, video_manager from facefusion.common_helper import get_first, get_middle from facefusion.filesystem import filter_audio_paths, is_video from facefusion.media_helper import restrict_trim_frame @@ -81,7 +81,7 @@ def process_memory_frames() -> ErrorCode: read_static_video_frame(state_manager.get_item('target_path'), state_manager.get_item('reference_frame_index')) - with ThreadPoolExecutor(max_workers = state_manager.get_item('execution_thread_count')) as executor: + with thread_helper.create_executor(state_manager.get_item('execution_thread_count')) as executor: futures : Deque[Future[VisionFrame]] = deque() for frame_index in temp_frame_range: diff --git a/tests/test_api_assets.py b/tests/test_api_assets.py index bf90bfb7..d39eca3d 100644 --- a/tests/test_api_assets.py +++ b/tests/test_api_assets.py @@ -14,6 +14,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_api_jobs.py b/tests/test_api_jobs.py index 922d0461..364c92d2 100644 --- a/tests/test_api_jobs.py +++ b/tests/test_api_jobs.py @@ -16,6 +16,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + create_program() conditional_download(get_test_examples_directory(), diff --git a/tests/test_api_session.py b/tests/test_api_session.py index 1e81adb7..ed4ed44e 100644 --- a/tests/test_api_session.py +++ b/tests/test_api_session.py @@ -17,6 +17,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_api_state.py b/tests/test_api_state.py index e1de5833..9c012c1f 100644 --- a/tests/test_api_state.py +++ b/tests/test_api_state.py @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() program = ArgumentParser() capability_store.register_capability_set( diff --git a/tests/test_api_stream.py b/tests/test_api_stream.py index 13c9da1f..a12de75f 100644 --- a/tests/test_api_stream.py +++ b/tests/test_api_stream.py @@ -17,6 +17,7 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('execution_device_ids', [ 0 ]) state_manager.init_item('execution_providers', [ 'cpu' ]) state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) diff --git a/tests/test_api_stream_audio.py b/tests/test_api_stream_audio.py index 6e04d0da..2a97dd10 100644 --- a/tests/test_api_stream_audio.py +++ b/tests/test_api_stream_audio.py @@ -20,6 +20,7 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) state_manager.init_item('processors', []) diff --git a/tests/test_api_stream_manager.py b/tests/test_api_stream_manager.py index cb3afebb..c49b459a 100644 --- a/tests/test_api_stream_manager.py +++ b/tests/test_api_stream_manager.py @@ -1,22 +1,25 @@ import ctypes import threading +from contextvars import copy_context +from typing import Iterator from unittest.mock import AsyncMock, patch import pytest -from facefusion import rtc, rtc_store, state_manager +from facefusion import rtc, rtc_store, state_manager, store_creator from facefusion.apis.stream_manager import destroy_stream, process_image, process_video, receive_vision_frames, run_peer_loop from facefusion.common_helper import is_linux, is_windows from facefusion.download import conditional_download from facefusion.hash_helper import create_hash from facefusion.libraries import datachannel as datachannel_module -from facefusion.session_context import set_session_id +from facefusion.session_context import resolve_local_id, set_session_id from facefusion.types import RtcPeer, SessionId, VideoCodec from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) state_manager.init_item('execution_thread_count', 8) state_manager.init_item('processors', []) @@ -30,9 +33,15 @@ def before_all() -> None: @pytest.fixture(scope = 'function', autouse = True) -def before_each() -> None: +def before_each() -> Iterator[None]: + local_id = resolve_local_id() + rtc_store.clear() + yield + + set_session_id(local_id) + @pytest.mark.anyio async def test_process_image() -> None: @@ -148,17 +157,19 @@ def test_run_peer_loop(video_codec : VideoCodec, payload_type : int, session_id } rtc_store.set_peer(session_id, rtc_peer) + store_creator.set_content(state_manager.STATE_SET, session_id, state_manager.get_state()) + set_session_id(session_id) assert rtc_store.has_peer(session_id) is True - with patch('facefusion.apis.stream_manager.ThreadPoolExecutor') as thread_pool_executor_mock: + with patch('facefusion.thread_helper.ThreadPoolExecutor') as thread_pool_executor_mock: with patch('facefusion.apis.stream_manager.receive_video_frames'): with patch('facefusion.apis.stream_manager.run_video_encode_loop'): - thread = threading.Thread(target = run_peer_loop, args = (session_id, rtc_peer), daemon = True) + thread = threading.Thread(target = copy_context().run, args = (run_peer_loop, session_id, rtc_peer), daemon = True) thread.start() thread.join(timeout = 5.0) - thread_pool_executor_mock.assert_called_once_with(max_workers = 8, initializer = set_session_id, initargs = (session_id,)) + thread_pool_executor_mock.assert_called_once_with(max_workers = 8, initializer = set_session_id, initargs = tuple([ session_id ])) assert rtc_store.has_peer(session_id) is False diff --git a/tests/test_api_stream_video.py b/tests/test_api_stream_video.py index c10f73b3..5d0db57d 100644 --- a/tests/test_api_stream_video.py +++ b/tests/test_api_stream_video.py @@ -24,6 +24,7 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) state_manager.init_item('execution_thread_count', 8) state_manager.init_item('processors', []) diff --git a/tests/test_audio.py b/tests/test_audio.py index c7c92c3a..3f2ae491 100644 --- a/tests/test_audio.py +++ b/tests/test_audio.py @@ -2,7 +2,7 @@ import pytest from pytest import approx -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.audio import detect_audio_duration, get_audio_frame, read_static_audio, restrict_trim_audio_frame from facefusion.download import conditional_download from .assert_helper import get_test_example_file, get_test_examples_directory @@ -10,6 +10,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_age_modifier.py b/tests/test_cli_age_modifier.py index 3683fc25..e0828ca0 100644 --- a/tests/test_cli_age_modifier.py +++ b/tests/test_cli_age_modifier.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_background_remover.py b/tests/test_cli_background_remover.py index a38b2d04..8420aa05 100644 --- a/tests/test_cli_background_remover.py +++ b/tests/test_cli_background_remover.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_batch_runner.py b/tests/test_cli_batch_runner.py index 5be66936..6a3a5947 100644 --- a/tests/test_cli_batch_runner.py +++ b/tests/test_cli_batch_runner.py @@ -3,7 +3,7 @@ import sys import pytest -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs 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 @@ -11,6 +11,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_expression_restorer.py b/tests/test_cli_expression_restorer.py index e19d280e..5a52cb99 100644 --- a/tests/test_cli_expression_restorer.py +++ b/tests/test_cli_expression_restorer.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_face_debugger.py b/tests/test_cli_face_debugger.py index f55d0b4a..ec2a4282 100644 --- a/tests/test_cli_face_debugger.py +++ b/tests/test_cli_face_debugger.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_face_editor.py b/tests/test_cli_face_editor.py index 42cde089..f4bc1899 100644 --- a/tests/test_cli_face_editor.py +++ b/tests/test_cli_face_editor.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_face_enhancer.py b/tests/test_cli_face_enhancer.py index 78605460..d75656ba 100644 --- a/tests/test_cli_face_enhancer.py +++ b/tests/test_cli_face_enhancer.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_face_swapper.py b/tests/test_cli_face_swapper.py index 582098b8..168486ac 100644 --- a/tests/test_cli_face_swapper.py +++ b/tests/test_cli_face_swapper.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_frame_colorizer.py b/tests/test_cli_frame_colorizer.py index 8b9a3b2b..49cd62b8 100644 --- a/tests/test_cli_frame_colorizer.py +++ b/tests/test_cli_frame_colorizer.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_frame_enhancer.py b/tests/test_cli_frame_enhancer.py index 0c557402..8b9ec324 100644 --- a/tests/test_cli_frame_enhancer.py +++ b/tests/test_cli_frame_enhancer.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_job_manager.py b/tests/test_cli_job_manager.py index a52ab3d3..bb9488a8 100644 --- a/tests/test_cli_job_manager.py +++ b/tests/test_cli_job_manager.py @@ -4,7 +4,7 @@ import sys import pytest -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_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 @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_job_runner.py b/tests/test_cli_job_runner.py index 4d422cd0..10112b47 100644 --- a/tests/test_cli_job_runner.py +++ b/tests/test_cli_job_runner.py @@ -4,7 +4,7 @@ import sys import pytest -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_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 @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_lip_syncer.py b/tests/test_cli_lip_syncer.py index 19681eb9..fe303c3b 100644 --- a/tests/test_cli_lip_syncer.py +++ b/tests/test_cli_lip_syncer.py @@ -4,7 +4,7 @@ import sys import pytest import facefusion.choices -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import WorkflowStrategy @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_output_fps.py b/tests/test_cli_output_fps.py index aec6b65f..d2b4b4c8 100644 --- a/tests/test_cli_output_fps.py +++ b/tests/test_cli_output_fps.py @@ -4,7 +4,7 @@ import sys import numpy import pytest -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import Fps, WorkflowStrategy @@ -14,6 +14,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_cli_output_scale.py b/tests/test_cli_output_scale.py index 176d1a24..36c7360e 100644 --- a/tests/test_cli_output_scale.py +++ b/tests/test_cli_output_scale.py @@ -3,7 +3,7 @@ import sys import pytest -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.jobs.job_manager import clear_jobs, init_jobs from facefusion.types import Resolution, Scale @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_codec_aom_decoder.py b/tests/test_codec_aom_decoder.py index 1de82c04..cf11a648 100644 --- a/tests/test_codec_aom_decoder.py +++ b/tests/test_codec_aom_decoder.py @@ -16,6 +16,7 @@ from facefusion.vision import read_video_frame @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) conditional_download(get_test_examples_directory(), [ 'https://github.com/facefusion/facefusion-assets/releases/download/examples-3.0.0/target-240p.mp4' ]) diff --git a/tests/test_codec_aom_encoder.py b/tests/test_codec_aom_encoder.py index 811a432d..3215bd17 100644 --- a/tests/test_codec_aom_encoder.py +++ b/tests/test_codec_aom_encoder.py @@ -15,6 +15,7 @@ from facefusion.vision import read_video_frame @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) conditional_download(get_test_examples_directory(), [ 'https://github.com/facefusion/facefusion-assets/releases/download/examples-3.0.0/target-240p.mp4' ]) diff --git a/tests/test_codec_opus_decoder.py b/tests/test_codec_opus_decoder.py index d38eab71..4b2b5f41 100644 --- a/tests/test_codec_opus_decoder.py +++ b/tests/test_codec_opus_decoder.py @@ -16,6 +16,7 @@ from facefusion.libraries import opus as opus_module @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) conditional_download(get_test_examples_directory(), [ 'https://github.com/facefusion/facefusion-assets/releases/download/examples-3.0.0/source.mp3' ]) diff --git a/tests/test_codec_opus_encoder.py b/tests/test_codec_opus_encoder.py index 5a3dc474..9d42f4ec 100644 --- a/tests/test_codec_opus_encoder.py +++ b/tests/test_codec_opus_encoder.py @@ -15,6 +15,7 @@ from facefusion.libraries import opus as opus_module @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) conditional_download(get_test_examples_directory(), [ 'https://github.com/facefusion/facefusion-assets/releases/download/examples-3.0.0/source.mp3' ]) diff --git a/tests/test_codec_vpx_decoder.py b/tests/test_codec_vpx_decoder.py index bb436143..00aff33e 100644 --- a/tests/test_codec_vpx_decoder.py +++ b/tests/test_codec_vpx_decoder.py @@ -17,6 +17,7 @@ from facefusion.vision import read_video_frame @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) conditional_download(get_test_examples_directory(), [ 'https://github.com/facefusion/facefusion-assets/releases/download/examples-3.0.0/target-240p.mp4' ]) diff --git a/tests/test_codec_vpx_encoder.py b/tests/test_codec_vpx_encoder.py index 17deacbd..4176a54e 100644 --- a/tests/test_codec_vpx_encoder.py +++ b/tests/test_codec_vpx_encoder.py @@ -16,6 +16,7 @@ from facefusion.vision import read_video_frame @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) conditional_download(get_test_examples_directory(), [ 'https://github.com/facefusion/facefusion-assets/releases/download/examples-3.0.0/target-240p.mp4' ]) diff --git a/tests/test_config.py b/tests/test_config.py index ba821508..85073cf8 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -5,6 +5,7 @@ from facefusion import config, state_manager @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('config_path', 'facefusion.ini') config_parser = config.get_static_config_parser() config_parser.read_dict( diff --git a/tests/test_face_creator.py b/tests/test_face_creator.py index d6e59977..8d76bab3 100644 --- a/tests/test_face_creator.py +++ b/tests/test_face_creator.py @@ -12,6 +12,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_face_detector.py b/tests/test_face_detector.py index 80898992..c73afad4 100644 --- a/tests/test_face_detector.py +++ b/tests/test_face_detector.py @@ -11,6 +11,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_face_tracker.py b/tests/test_face_tracker.py index f9465229..ce4e2ee6 100644 --- a/tests/test_face_tracker.py +++ b/tests/test_face_tracker.py @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + conditional_download(get_test_examples_directory(), [ 'https://github.com/facefusion/facefusion-assets/releases/download/examples-3.0.0/target-240p.mp4' diff --git a/tests/test_ffmpeg.py b/tests/test_ffmpeg.py index 8c7f88c1..1f166889 100644 --- a/tests/test_ffmpeg.py +++ b/tests/test_ffmpeg.py @@ -17,6 +17,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() state_manager.init_item('temp_path', tempfile.gettempdir()) state_manager.init_item('temp_frame_format', 'png') diff --git a/tests/test_ffprobe.py b/tests/test_ffprobe.py index 7f3d1517..31d44201 100644 --- a/tests/test_ffprobe.py +++ b/tests/test_ffprobe.py @@ -1,7 +1,7 @@ import pytest -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.ffprobe import extract_audio_metadata, extract_video_metadata from .assert_helper import get_test_example_file, get_test_examples_directory @@ -9,6 +9,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), diff --git a/tests/test_inference_manager.py b/tests/test_inference_manager.py index 6e7e1e60..13de7066 100644 --- a/tests/test_inference_manager.py +++ b/tests/test_inference_manager.py @@ -11,6 +11,7 @@ from facefusion.inference_manager import get_inference_pool, resolve_static_infe @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('execution_device_ids', [ 0 ]) state_manager.init_item('execution_providers', [ 'cpu' ]) state_manager.init_item('download_providers', [ 'github' ]) diff --git a/tests/test_job_runner.py b/tests/test_job_runner.py index 2698afed..96474d48 100644 --- a/tests/test_job_runner.py +++ b/tests/test_job_runner.py @@ -2,7 +2,7 @@ import os import pytest -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.filesystem import copy_file, create_directory, get_file_extension from facefusion.jobs.job_manager import add_step, clear_jobs, create_job, init_jobs, move_job_file, submit_job, submit_jobs @@ -13,6 +13,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_library_aom.py b/tests/test_library_aom.py index f43370f3..76d65476 100644 --- a/tests/test_library_aom.py +++ b/tests/test_library_aom.py @@ -8,6 +8,7 @@ from facefusion.libraries import aom as aom_module @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) aom_module.pre_check() diff --git a/tests/test_library_datachannel.py b/tests/test_library_datachannel.py index 6790d2fd..568b9c21 100644 --- a/tests/test_library_datachannel.py +++ b/tests/test_library_datachannel.py @@ -8,6 +8,7 @@ from facefusion.libraries import datachannel as datachannel_module @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) datachannel_module.pre_check() diff --git a/tests/test_library_opus.py b/tests/test_library_opus.py index 5cca63a6..f149de9a 100644 --- a/tests/test_library_opus.py +++ b/tests/test_library_opus.py @@ -8,6 +8,7 @@ from facefusion.libraries import opus as opus_module @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) opus_module.pre_check() diff --git a/tests/test_library_vpx.py b/tests/test_library_vpx.py index af49b61e..9b8477c9 100644 --- a/tests/test_library_vpx.py +++ b/tests/test_library_vpx.py @@ -8,6 +8,7 @@ from facefusion.libraries import vpx as vpx_module @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) vpx_module.pre_check() diff --git a/tests/test_process_manager.py b/tests/test_process_manager.py index c3ae86d3..808f6d06 100644 --- a/tests/test_process_manager.py +++ b/tests/test_process_manager.py @@ -1,3 +1,5 @@ +from typing import Iterator + import pytest from facefusion import session_context @@ -5,13 +7,19 @@ from facefusion.process_manager import clear, end, get_state, init, is_pending, @pytest.fixture(scope = 'function', autouse = True) -def before_each() -> None: +def before_each() -> Iterator[None]: + local_id = session_context.resolve_local_id() + session_context.set_session_id('session-a') clear() session_context.set_session_id('session-b') clear() session_context.set_session_id('session-a') + yield + + session_context.set_session_id(local_id) + def test_init() -> None: set_state('processing') diff --git a/tests/test_rtc.py b/tests/test_rtc.py index 40c0c372..e11259f0 100644 --- a/tests/test_rtc.py +++ b/tests/test_rtc.py @@ -10,6 +10,7 @@ from facefusion.types import RtcPeer, VideoCodec @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) datachannel_module.pre_check() diff --git a/tests/test_session_context.py b/tests/test_session_context.py index ba4ee854..06f0ce36 100644 --- a/tests/test_session_context.py +++ b/tests/test_session_context.py @@ -1,16 +1,25 @@ +from typing import Iterator + 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: +def before_each() -> Iterator[None]: local_id = resolve_local_id() + + set_session_id(local_id) + + yield + set_session_id(local_id) def test_get_session_id() -> None: - assert get_session_id() == resolve_local_id() + local_id = resolve_local_id() + + assert get_session_id() == local_id set_session_id('session-a') @@ -18,4 +27,6 @@ def test_get_session_id() -> None: def test_resolve_local_id() -> None: - assert resolve_local_id() == resolve_local_id() + local_id = resolve_local_id() + + assert resolve_local_id() == local_id diff --git a/tests/test_state_manager.py b/tests/test_state_manager.py index d2ea313f..1f1cfcce 100644 --- a/tests/test_state_manager.py +++ b/tests/test_state_manager.py @@ -1,35 +1,67 @@ -from typing import Union +from typing import Iterator import pytest -from facefusion.processors.types import ProcessorState -from facefusion.state_manager import STATE_SET, get_item, init_item, set_item -from facefusion.types import AppContext, State - - -def get_state(app_context : AppContext) -> Union[State, ProcessorState]: - return STATE_SET.get(app_context) - - -def clear_state(app_context : AppContext) -> None: - STATE_SET[app_context] = {} #type:ignore[typeddict-item] +from facefusion import 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 @pytest.fixture(scope = 'function', autouse = True) -def before_each() -> None: - clear_state('cli') - clear_state('api') +def before_each() -> Iterator[None]: + local_id = resolve_local_id() + + set_session_id(local_id) + clear() + store_creator.delete_content(STATE_SET, 'session-a') + + yield + + set_session_id(local_id) + + +def test_init() -> None: + init_item('video_memory_strategy', 'tolerant') + set_session_id('session-a') + + assert get_state() is None + + init() + 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_get_state() -> None: + init_item('video_memory_strategy', 'tolerant') + + assert get_state() == { 'video_memory_strategy': 'tolerant' } + + +def test_set_state() -> None: + set_state({ 'video_memory_strategy': 'strict' }) + + assert get_state() == { 'video_memory_strategy': 'strict' } + + +def test_clear() -> None: + init_item('video_memory_strategy', 'tolerant') + clear() + + assert get_state() == {} def test_init_item() -> None: init_item('video_memory_strategy', 'tolerant') - assert get_state('cli').get('video_memory_strategy') == 'tolerant' - assert get_state('api').get('video_memory_strategy') == 'tolerant' + assert get_state().get('video_memory_strategy') == 'tolerant' def test_get_item_and_set_item() -> None: set_item('video_memory_strategy', 'tolerant') assert get_item('video_memory_strategy') == 'tolerant' - assert get_state('api').get('video_memory_strategy') is None diff --git a/tests/test_temp_helper.py b/tests/test_temp_helper.py index 7be93708..8768caa4 100644 --- a/tests/test_temp_helper.py +++ b/tests/test_temp_helper.py @@ -11,6 +11,7 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() state_manager.init_item('temp_path', tempfile.gettempdir()) state_manager.init_item('temp_frame_format', 'png') diff --git a/tests/test_thread_helper.py b/tests/test_thread_helper.py new file mode 100644 index 00000000..9e369d83 --- /dev/null +++ b/tests/test_thread_helper.py @@ -0,0 +1,22 @@ +from typing import Iterator + +import pytest + +from facefusion.session_context import get_session_id, resolve_local_id, set_session_id +from facefusion.thread_helper import create_executor + + +@pytest.fixture(scope = 'function', autouse = True) +def before_each() -> Iterator[None]: + local_id = resolve_local_id() + + set_session_id('session-a') + + yield + + set_session_id(local_id) + + +def test_create_executor() -> None: + with create_executor(1) as executor: + assert executor.submit(get_session_id).result() == 'session-a' diff --git a/tests/test_video_manager.py b/tests/test_video_manager.py index c133c371..0ce7f757 100644 --- a/tests/test_video_manager.py +++ b/tests/test_video_manager.py @@ -15,6 +15,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_vision.py b/tests/test_vision.py index f5d208a4..a697eb66 100644 --- a/tests/test_vision.py +++ b/tests/test_vision.py @@ -2,7 +2,7 @@ import numpy import pytest -from facefusion import ffmpeg, ffmpeg_builder, process_manager +from facefusion import ffmpeg, ffmpeg_builder, process_manager, state_manager from facefusion.download import conditional_download from facefusion.vision import calculate_histogram_difference, count_video_frame_total, detect_image_resolution, detect_video_duration, detect_video_fps, detect_video_resolution, match_frame_color, normalize_resolution, pack_resolution, predict_video_frame_total, read_image, read_video_frame, resolve_extract_frame_index, resolve_target_frame_index, restrict_image_resolution, restrict_trim_video_frame, restrict_video_fps, restrict_video_resolution, scale_resolution, select_video_frames, unpack_resolution, write_image from .assert_helper import get_test_example_file, get_test_examples_directory, get_test_output_path, prepare_test_output_directory @@ -10,6 +10,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + process_manager.start() conditional_download(get_test_examples_directory(), [ diff --git a/tests/test_workflow.py b/tests/test_workflow.py index aede08d5..a23dde68 100644 --- a/tests/test_workflow.py +++ b/tests/test_workflow.py @@ -8,6 +8,8 @@ from .assert_helper import get_test_example_file, get_test_examples_directory @pytest.fixture(scope = 'module', autouse = True) def before_all() -> None: + state_manager.init() + conditional_download(get_test_examples_directory(), [ 'https://github.com/facefusion/facefusion-assets/releases/download/examples-3.0.0/source.jpg',