diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index 7aaced37..77dc44af 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -17,9 +17,9 @@ async def create_session(request : Request) -> JSONResponse: if validate_api_key(body.get('api_key')): session_id = secrets.token_urlsafe(16) - session = session_manager.create_session() + session = session_manager.create_api_session() session_context.set_session_id(session_id) - session_manager.set_session(session_id, session) + session_manager.set_api_session(session_id, session) state_manager.init() asset_store.init() @@ -47,7 +47,7 @@ async def create_session(request : Request) -> JSONResponse: async def get_session(request : Request) -> JSONResponse: session_id = session_context.get_session_id() - session = session_manager.get_session(session_id) + session = session_manager.get_api_session(session_id) return JSONResponse( { @@ -61,10 +61,10 @@ async def get_session(request : Request) -> JSONResponse: async def refresh_session(request : Request) -> JSONResponse: body = await request.json() - for session_id, session in session_manager.SESSIONS.items(): - if session.get('refresh_token') == body.get('refresh_token') and session_manager.validate_session(session_id): - __session__ = session_manager.create_session() - session_manager.set_session(session_id, __session__) + for session_id, session in session_manager.API_SESSIONS.items(): + if session.get('refresh_token') == body.get('refresh_token') and session_manager.validate_api_session(session_id): + __session__ = session_manager.create_api_session() + session_manager.set_api_session(session_id, __session__) return JSONResponse( { @@ -97,7 +97,7 @@ async def destroy_session(request : Request) -> JSONResponse: destroy_stream() video_manager.clear() - session_manager.clear_session(session_id) + session_manager.clear_api_session(session_id) stores =\ [ diff --git a/facefusion/apis/middlewares/session.py b/facefusion/apis/middlewares/session.py index 9a935295..59e27148 100644 --- a/facefusion/apis/middlewares/session.py +++ b/facefusion/apis/middlewares/session.py @@ -11,10 +11,10 @@ def create_session_guard(app : ASGIApp) -> ASGIApp: access_token = extract_access_token(scope) if access_token: - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) if session_id: - if session_manager.validate_session(session_id): + if session_manager.validate_api_session(session_id): session_context.set_session_id(session_id) return await app(scope, receive, send) diff --git a/facefusion/core.py b/facefusion/core.py index 9a28ff0e..8ff040bc 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, store_creator, 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_context, 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 @@ -38,6 +38,9 @@ def cli() -> None: if state_manager.get_item('command'): logger.init(state_manager.get_item('log_level')) + session_id = session_context.get_session_id() + session_manager.set_cli_session(session_id, session_manager.create_cli_session()) + content_store.init() face_store.init() inference_manager.init() diff --git a/facefusion/session_manager.py b/facefusion/session_manager.py index 484bbe20..da556729 100644 --- a/facefusion/session_manager.py +++ b/facefusion/session_manager.py @@ -4,13 +4,14 @@ 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 +from facefusion.types import ApiSession, CliSession, SessionId -SESSIONS : Dict[SessionId, Session] = {} +API_SESSIONS : Dict[SessionId, ApiSession] = {} +CLI_SESSIONS : Dict[SessionId, CliSession] = {} -def create_session() -> Session: - session : Session =\ +def create_api_session() -> ApiSession: + api_session : ApiSession =\ { 'access_token': secrets.token_urlsafe(64), 'refresh_token': secrets.token_urlsafe(64), @@ -18,20 +19,25 @@ def create_session() -> Session: 'expires_at': datetime.now() + timedelta(minutes = 10) } - return session + return api_session + + +def create_cli_session() -> CliSession: + cli_session : CliSession =\ + { + 'owner_id': get_session_id(), + 'created_at': datetime.now() + } + + return cli_session def fork_session() -> SessionId: fork_id = secrets.token_urlsafe(16) - owner_id = get_session_id() - session : Session =\ - { - 'owner_id': owner_id, - 'created_at': datetime.now() - } + cli_session = create_cli_session() + set_cli_session(fork_id, cli_session) set_session_id(fork_id) - set_session(fork_id, session) return fork_id @@ -40,40 +46,57 @@ def join_session() -> None: fork_id = get_session_id() owner_id = resolve_owner_id() - clear_session(fork_id) + clear_cli_session(fork_id) set_session_id(owner_id) -def get_session(session_id : SessionId) -> Optional[Session]: - return SESSIONS.get(session_id) +def get_api_session(session_id : SessionId) -> Optional[ApiSession]: + return API_SESSIONS.get(session_id) -def find_session_id(access_token : str) -> Optional[SessionId]: - for session_id, session in SESSIONS.items(): - if session.get('access_token') == access_token: +def get_cli_session(session_id : SessionId) -> Optional[CliSession]: + return CLI_SESSIONS.get(session_id) + + +def find_api_session_id(access_token : str) -> Optional[SessionId]: + for session_id, api_session in API_SESSIONS.items(): + if api_session.get('access_token') == access_token: return session_id return None def resolve_owner_id() -> SessionId: session_id = get_session_id() - session = get_session(session_id) + cli_session = get_cli_session(session_id) - if session and session.get('owner_id'): - return session.get('owner_id') + if cli_session: + return cli_session.get('owner_id') return session_id -def set_session(session_id : SessionId, session : Session) -> None: - SESSIONS[session_id] = session +def set_api_session(session_id : SessionId, api_session : ApiSession) -> None: + API_SESSIONS[session_id] = api_session -def validate_session(session_id : SessionId) -> bool: - session = get_session(session_id) - return session and datetime.now() < session.get('expires_at') +def set_cli_session(session_id : SessionId, cli_session : CliSession) -> None: + CLI_SESSIONS[session_id] = cli_session -def clear_session(session_id : SessionId) -> None: - if session_id in SESSIONS: - del SESSIONS[session_id] +def validate_api_session(session_id : SessionId) -> bool: + api_session = get_api_session(session_id) + + if api_session: + return datetime.now() < api_session.get('expires_at') + + return False + + +def clear_api_session(session_id : SessionId) -> None: + if session_id in API_SESSIONS: + del API_SESSIONS[session_id] + + +def clear_cli_session(session_id : SessionId) -> None: + if session_id in CLI_SESSIONS: + del CLI_SESSIONS[session_id] diff --git a/facefusion/types.py b/facefusion/types.py index a49d8b55..c769d1f8 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -157,13 +157,17 @@ Content : TypeAlias = Dict[str, Any] Token : TypeAlias = str SessionId : TypeAlias = str -Session = TypedDict('Session', +ApiSession = TypedDict('ApiSession', { - 'access_token' : NotRequired[Token], - 'refresh_token' : NotRequired[Token], - 'owner_id' : NotRequired[SessionId], + 'access_token' : Token, + 'refresh_token' : Token, 'created_at' : datetime, - 'expires_at' : NotRequired[datetime] + 'expires_at' : datetime +}) +CliSession = TypedDict('CliSession', +{ + 'owner_id' : SessionId, + 'created_at' : datetime }) StoreContent : TypeAlias = Any diff --git a/tests/test_api_assets.py b/tests/test_api_assets.py index 246029b6..f9365d92 100644 --- a/tests/test_api_assets.py +++ b/tests/test_api_assets.py @@ -44,7 +44,7 @@ def before_each() -> Iterator[None]: state_manager.init_item('temp_path', tempfile.gettempdir()) state_manager.init_item('jobs_path', get_test_jobs_directory()) state_manager.init_item('temp_frame_format', 'png') - session_manager.SESSIONS.clear() + session_manager.API_SESSIONS.clear() asset_store.delete_assets() yield @@ -77,7 +77,7 @@ def test_upload_assets(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) with open(source_path, 'rb') as source_file: @@ -281,7 +281,7 @@ def test_delete_assets(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) source_path = get_test_example_file('source.jpg') @@ -351,7 +351,7 @@ def test_delete_asset(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) source_path = get_test_example_file('source.jpg') diff --git a/tests/test_api_capabilities.py b/tests/test_api_capabilities.py index dbeec217..822d89c3 100644 --- a/tests/test_api_capabilities.py +++ b/tests/test_api_capabilities.py @@ -36,7 +36,7 @@ def before_all() -> None: @pytest.fixture(scope = 'function', autouse = True) def before_each() -> None: - session_manager.SESSIONS.clear() + session_manager.API_SESSIONS.clear() @pytest.fixture(scope = 'module') diff --git a/tests/test_api_jobs.py b/tests/test_api_jobs.py index 223da51a..a527b828 100644 --- a/tests/test_api_jobs.py +++ b/tests/test_api_jobs.py @@ -32,7 +32,7 @@ def before_each() -> Iterator[None]: local_id = session_context.resolve_local_id() session_context.set_session_id(local_id) - session_manager.SESSIONS.clear() + session_manager.API_SESSIONS.clear() asset_store.delete_assets() state_manager.init_item('source_paths', [ get_test_example_file('source.jpg') ]) state_manager.init_item('target_path', get_test_example_file('target-240p.mp4')) @@ -63,7 +63,7 @@ def test_get_jobs(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) get_jobs_response = test_client.get('/jobs?status=invalid', headers = @@ -114,7 +114,7 @@ def test_get_job(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) get_job_response = test_client.get('/jobs/job-test-unknown', headers = @@ -165,7 +165,7 @@ def test_create_job(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) create_job_response = test_client.post('/jobs', headers = @@ -206,7 +206,7 @@ def test_submit_jobs(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) submit_jobs_response = test_client.patch('/jobs?action=invalid', headers = @@ -259,7 +259,7 @@ def test_submit_job(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) submit_job_response = test_client.patch('/jobs/job-test-submit-job?action=invalid', headers = @@ -312,7 +312,7 @@ def test_run_jobs(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) run_jobs_response = test_client.patch('/jobs?action=run', headers = @@ -360,7 +360,7 @@ def test_run_job(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) create_job('job-test-run-job') @@ -409,7 +409,7 @@ def test_retry_jobs(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) retry_jobs_response = test_client.patch('/jobs?action=retry', headers = @@ -459,7 +459,7 @@ def test_retry_job(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) retry_job_response = test_client.patch('/jobs/job-test-retry-job?action=retry', headers = @@ -509,7 +509,7 @@ def test_delete_jobs(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) delete_jobs_response = test_client.delete('/jobs', headers = @@ -544,7 +544,7 @@ def test_delete_job(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) delete_job_response = test_client.delete('/jobs/job-test-unknown', headers = @@ -581,7 +581,7 @@ def test_create_step(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) create_job('job-test-create-step') @@ -607,7 +607,7 @@ def test_create_step(test_client : TestClient) -> None: get_job_body = get_job_response.json() step_args = get_job_body.get('steps')[0].get('args') - access_session_id = session_manager.find_session_id(access_token) + access_session_id = session_manager.find_api_session_id(access_token) assert step_args ==\ { @@ -663,7 +663,7 @@ def test_delete_step(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) create_job('job-test-delete-step') diff --git a/tests/test_api_metrics.py b/tests/test_api_metrics.py index cacd54a3..fb7c393d 100644 --- a/tests/test_api_metrics.py +++ b/tests/test_api_metrics.py @@ -17,7 +17,7 @@ def before_all() -> None: @pytest.fixture(scope = 'function', autouse = True) def before_each() -> None: - session_manager.SESSIONS.clear() + session_manager.API_SESSIONS.clear() @pytest.fixture(scope = 'module') diff --git a/tests/test_api_ping.py b/tests/test_api_ping.py index 17562bdf..0c116685 100644 --- a/tests/test_api_ping.py +++ b/tests/test_api_ping.py @@ -16,7 +16,7 @@ def before_all() -> None: @pytest.fixture(scope = 'function', autouse = True) def before_each() -> None: - session_manager.SESSIONS.clear() + session_manager.API_SESSIONS.clear() @pytest.fixture(scope = 'module') diff --git a/tests/test_api_session.py b/tests/test_api_session.py index eda629ee..ab1359ae 100644 --- a/tests/test_api_session.py +++ b/tests/test_api_session.py @@ -13,7 +13,7 @@ from facefusion.apis import asset_store from facefusion.apis.core import create_api from facefusion.download import conditional_download from facefusion.libraries import datachannel as datachannel_module -from facefusion.types import RtcPeer, Session +from facefusion.types import ApiSession, RtcPeer from .assert_helper import get_test_example_file, get_test_examples_directory, get_test_jobs_directory @@ -37,7 +37,7 @@ def before_each() -> Iterator[None]: session_context.set_session_id(local_id) state_manager.init_item('temp_path', tempfile.gettempdir()) state_manager.init_item('jobs_path', get_test_jobs_directory()) - session_manager.SESSIONS.clear() + session_manager.API_SESSIONS.clear() asset_store.delete_assets() yield @@ -117,9 +117,9 @@ def test_get_session(test_client : TestClient) -> None: assert get_session_response.status_code == 200 - session_id = session_manager.find_session_id(create_session_body.get('access_token')) - session : Session = session_manager.get_session(session_id) - session_manager.set_session(session_id, + session_id = session_manager.find_api_session_id(create_session_body.get('access_token')) + session : ApiSession = session_manager.get_api_session(session_id) + session_manager.set_api_session(session_id, { 'access_token': session.get('access_token'), 'refresh_token': session.get('refresh_token'), @@ -150,6 +150,15 @@ def test_refresh_session(test_client : TestClient) -> None: assert refresh_session_response.status_code == 401 access_token = create_session_body.get('access_token') + session_id = session_manager.find_api_session_id(access_token) + session_context.set_session_id(session_id) + session_manager.fork_session() + + refresh_session_response = test_client.put('/session', json = {}) + + assert refresh_session_response.status_code == 401 + + session_manager.join_session() refresh_session_response = test_client.put('/session', json = { @@ -159,7 +168,7 @@ def test_refresh_session(test_client : TestClient) -> None: assert refresh_session_body.get('access_token') assert refresh_session_body.get('refresh_token') - assert session_manager.find_session_id(access_token) is None + assert session_manager.find_api_session_id(access_token) is None assert refresh_session_response.status_code == 200 refresh_session_response = test_client.put('/session', json = @@ -175,9 +184,9 @@ def test_refresh_session(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() - session_id = session_manager.find_session_id(create_session_body.get('access_token')) - session : Session = session_manager.get_session(session_id) - session_manager.set_session(session_id, + session_id = session_manager.find_api_session_id(create_session_body.get('access_token')) + session : ApiSession = session_manager.get_api_session(session_id) + session_manager.set_api_session(session_id, { 'access_token': session.get('access_token'), 'refresh_token': session.get('refresh_token'), @@ -199,7 +208,7 @@ def test_destroy_session(test_client : TestClient) -> None: 'client_version': metadata.get('version') }) access_token = create_session_response.json().get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) jobs_path = os.path.join(get_test_jobs_directory(), session_id) delete_session_response = test_client.delete('/session', headers = @@ -216,7 +225,7 @@ def test_destroy_session(test_client : TestClient) -> None: }) assert os.path.isdir(jobs_path) is False - assert session_manager.find_session_id(access_token) is None + assert session_manager.find_api_session_id(access_token) is None assert delete_session_response.status_code == 200 create_session_response = test_client.post('/session', json = @@ -224,7 +233,7 @@ def test_destroy_session(test_client : TestClient) -> None: 'client_version': metadata.get('version') }) access_token = create_session_response.json().get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) source_path = get_test_example_file('source.jpg') @@ -252,7 +261,7 @@ def test_destroy_session(test_client : TestClient) -> None: assert os.path.exists(asset_path) is True assert delete_session_response.json().get('message') == 'directory not removed' - assert session_manager.find_session_id(access_token) == session_id + assert session_manager.find_api_session_id(access_token) == session_id assert delete_session_response.status_code == 404 peer_connection = rtc.create_peer_connection() @@ -276,7 +285,7 @@ def test_destroy_session(test_client : TestClient) -> None: 'Authorization': 'Bearer ' + access_token }) - assert session_manager.find_session_id(access_token) is None + assert session_manager.find_api_session_id(access_token) is None 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 diff --git a/tests/test_api_state.py b/tests/test_api_state.py index f618ce62..2f670a2f 100644 --- a/tests/test_api_state.py +++ b/tests/test_api_state.py @@ -73,7 +73,7 @@ def before_each() -> Iterator[None]: local_id = session_context.resolve_local_id() session_context.set_session_id(local_id) - session_manager.SESSIONS.clear() + session_manager.API_SESSIONS.clear() asset_store.delete_assets() yield @@ -176,7 +176,7 @@ def test_select_source_assets(test_client : TestClient) -> None: create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) source_paths =\ [ @@ -226,7 +226,7 @@ def test_select_target_assets(test_client : TestClient) -> None: }) create_session_body = create_session_response.json() access_token = create_session_body.get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) target_path = get_test_example_file('target-240p.jpg') asset_id = asset_store.create_asset('target', target_path).get('id') diff --git a/tests/test_api_stream.py b/tests/test_api_stream.py index aaaf349c..f2aae03f 100644 --- a/tests/test_api_stream.py +++ b/tests/test_api_stream.py @@ -38,7 +38,7 @@ def before_each() -> Iterator[None]: local_id = session_context.resolve_local_id() session_context.set_session_id(local_id) - session_manager.SESSIONS.clear() + session_manager.API_SESSIONS.clear() asset_store.delete_assets() rtc_store.delete_peer() @@ -151,7 +151,7 @@ def test_delete_stream_video(test_client : TestClient) -> None: 'client_version': metadata.get('version') }) access_token = create_session_response.json().get('access_token') - session_id = session_manager.find_session_id(access_token) + session_id = session_manager.find_api_session_id(access_token) session_context.set_session_id(session_id) peer_connection = rtc.create_peer_connection() diff --git a/tests/test_session_manager.py b/tests/test_session_manager.py index 46df3534..426c14d0 100644 --- a/tests/test_session_manager.py +++ b/tests/test_session_manager.py @@ -1,11 +1,11 @@ import secrets -from datetime import datetime, timedelta +from datetime import timedelta from typing import Iterator 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 +from facefusion.session_manager import clear_api_session, clear_cli_session, create_api_session, create_cli_session, find_api_session_id, fork_session, get_api_session, get_cli_session, join_session, resolve_owner_id, set_api_session, set_cli_session, validate_api_session @pytest.fixture(scope = 'function', autouse = True) @@ -25,77 +25,112 @@ def test_fork_session() -> None: assert get_session_id() == fork_id assert resolve_owner_id() == local_id - 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 + assert get_cli_session(fork_id).get('owner_id') == local_id join_session() -def test_get_and_set_session() -> None: - session = create_session() - session_id = secrets.token_urlsafe(16) - - set_session(session_id, session) - - 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) - - set_session(session_id, session) - - assert validate_session(session_id) is True - - set_session(session_id, - { - 'access_token': session.get('access_token'), - 'refresh_token': session.get('refresh_token'), - 'created_at': session.get('created_at'), - 'expires_at': session.get('expires_at') - timedelta(hours = 1) - }) - - assert validate_session(session_id) is False - - -def test_clear_session() -> None: - session = create_session() - session_id = secrets.token_urlsafe(16) - - set_session(session_id, session) - - assert validate_session(session_id) is True - - clear_session(session_id) - - assert validate_session(session_id) is None - - def test_join_session() -> None: local_id = resolve_local_id() fork_id = fork_session() join_session() assert get_session_id() == local_id - assert get_session(fork_id) is None + assert get_cli_session(fork_id) is None + + +def test_get_and_set_api_session() -> None: + api_session = create_api_session() + session_id = secrets.token_urlsafe(16) + + set_api_session(session_id, api_session) + + assert get_api_session(session_id) == api_session + + clear_api_session(session_id) + + +def test_get_and_set_cli_session() -> None: + local_id = resolve_local_id() + cli_session = create_cli_session() + session_id = secrets.token_urlsafe(16) + + set_cli_session(session_id, cli_session) + + assert get_cli_session(session_id) == cli_session + assert get_cli_session(session_id).get('owner_id') == local_id + + clear_cli_session(session_id) + + +def test_find_api_session_id() -> None: + api_session = create_api_session() + session_id = secrets.token_urlsafe(16) + + set_api_session(session_id, api_session) + + assert find_api_session_id(api_session.get('access_token')) == session_id + assert find_api_session_id('INVALID') is None + + clear_api_session(session_id) + + +def test_resolve_owner_id() -> None: + local_id = resolve_local_id() + + assert resolve_owner_id() == local_id + + set_api_session('session-a', create_api_session()) + set_session_id('session-a') + + assert resolve_owner_id() == 'session-a' + + fork_session() + + assert resolve_owner_id() == 'session-a' + + join_session() + clear_api_session('session-a') + + +def test_validate_api_session() -> None: + api_session = create_api_session() + session_id = secrets.token_urlsafe(16) + + set_api_session(session_id, api_session) + + assert validate_api_session(session_id) is True + + set_api_session(session_id, + { + 'access_token': api_session.get('access_token'), + 'refresh_token': api_session.get('refresh_token'), + 'created_at': api_session.get('created_at'), + 'expires_at': api_session.get('expires_at') - timedelta(hours = 1) + }) + + assert validate_api_session(session_id) is False + + clear_api_session(session_id) + + assert validate_api_session(session_id) is False + + +def test_clear_api_session() -> None: + api_session = create_api_session() + session_id = secrets.token_urlsafe(16) + + set_api_session(session_id, api_session) + clear_api_session(session_id) + + assert get_api_session(session_id) is None + + +def test_clear_cli_session() -> None: + cli_session = create_cli_session() + session_id = secrets.token_urlsafe(16) + + set_cli_session(session_id, cli_session) + clear_cli_session(session_id) + + assert get_cli_session(session_id) is None diff --git a/tests/test_state_manager.py b/tests/test_state_manager.py index a8f6e8e3..3813841b 100644 --- a/tests/test_state_manager.py +++ b/tests/test_state_manager.py @@ -1,10 +1,10 @@ -from datetime import datetime from typing import Iterator import pytest -from facefusion import session_manager, store_creator +from facefusion import store_creator from facefusion.session_context import resolve_local_id, set_session_id +from facefusion.session_manager import fork_session, join_session from facefusion.state_manager import STATE_SET, clear, clone_state, get_item, get_state, init, init_item, set_item, set_state @@ -16,7 +16,6 @@ def before_each() -> Iterator[None]: clear() store_creator.delete_content(STATE_SET, 'session-a') store_creator.delete_content(STATE_SET, 'session-a1') - session_manager.clear_session('session-a1') yield @@ -53,17 +52,18 @@ def test_set_state() -> None: 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') + fork_id = fork_session() clone_state() set_item('video_memory_strategy', 'strict') assert get_state() == { 'video_memory_strategy': 'strict' } - set_session_id(resolve_local_id()) + join_session() assert get_state() == { 'video_memory_strategy': 'tolerant' } + store_creator.delete_content(STATE_SET, fork_id) + def test_clear() -> None: init_item('video_memory_strategy', 'tolerant')