diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index e12ab60b..7aaced37 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 content_store, face_store, inference_manager, process_manager, rtc_store, session_context, session_manager, state_manager, translator, video_manager +from facefusion import content_store, face_store, inference_manager, process_manager, rtc_store, session_context, session_manager, state_manager, store_creator, translator, video_manager from facefusion.apis import asset_store from facefusion.apis.session_helper import validate_api_key from facefusion.apis.stream_manager import destroy_stream @@ -96,16 +96,23 @@ async def destroy_session(request : Request) -> JSONResponse: }, status_code = HTTP_404_NOT_FOUND) destroy_stream() - - asset_store.delete_assets() + video_manager.clear() session_manager.clear_session(session_id) - state_manager.clear() - content_store.clear() - face_store.clear() - inference_manager.clear() - video_manager.clear() - process_manager.clear() + stores =\ + [ + state_manager.STATE_SET, + asset_store.ASSET_STORE, + content_store.CONTENT_STORE, + face_store.FACE_STORE, + inference_manager.INFERENCE_POOL_STORE, + process_manager.PROCESS_STORE, + rtc_store.RTC_STORE, + video_manager.VIDEO_POOL_STORE + ] + + for store in stores: + store_creator.delete_content(store, session_id) return JSONResponse( { diff --git a/facefusion/core.py b/facefusion/core.py index ead2fe7a..9a28ff0e 100755 --- a/facefusion/core.py +++ b/facefusion/core.py @@ -326,7 +326,15 @@ 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 ]: + stores =\ + [ + state_manager.STATE_SET, + content_store.CONTENT_STORE, + process_manager.PROCESS_STORE, + video_manager.VIDEO_POOL_STORE + ] + + for store in stores: store_creator.delete_content(store, fork_id) return is_done diff --git a/facefusion/inference_manager.py b/facefusion/inference_manager.py index bf0ba658..e881d948 100644 --- a/facefusion/inference_manager.py +++ b/facefusion/inference_manager.py @@ -51,11 +51,12 @@ def get_inference_pool(module_name : str, model_names : List[str], model_source_ def find_inference_pool(inference_context : str) -> Optional[InferencePool]: - inference_pool_sets = list(INFERENCE_POOL_STORE.get('content_set').values()) + inference_pool_values = list(INFERENCE_POOL_STORE.get('content_set').values()) - for inference_pool_set in inference_pool_sets: + for inference_pool_set in inference_pool_values: if inference_pool_set.get(inference_context): return inference_pool_set.get(inference_context) + return None diff --git a/tests/test_api_session.py b/tests/test_api_session.py index c7efb2e5..eda629ee 100644 --- a/tests/test_api_session.py +++ b/tests/test_api_session.py @@ -8,7 +8,7 @@ from unittest.mock import patch import pytest from starlette.testclient import TestClient -from facefusion import metadata, process_manager, rtc, rtc_store, session_context, session_manager, state_manager +from facefusion import metadata, process_manager, rtc, rtc_store, session_context, session_manager, state_manager, store_creator from facefusion.apis import asset_store from facefusion.apis.core import create_api from facefusion.download import conditional_download @@ -277,8 +277,8 @@ def test_destroy_session(test_client : TestClient) -> None: }) assert session_manager.find_session_id(access_token) is None - assert asset_store.get_assets() == {} - assert rtc_store.has_peer() is False + assert store_creator.has_content(asset_store.ASSET_STORE, session_id) is False + assert store_creator.has_content(rtc_store.RTC_STORE, session_id) is False assert delete_session_response.status_code == 200 for asset_path in asset_paths: