diff --git a/facefusion.ini b/facefusion.ini index 68657359..acdf8064 100644 --- a/facefusion.ini +++ b/facefusion.ini @@ -150,6 +150,7 @@ benchmark_cycle_count = api_host = api_port = api_security_strategy = +api_session_limit = [execution] execution_device_ids = diff --git a/facefusion/apis/endpoints/session.py b/facefusion/apis/endpoints/session.py index 77dc44af..95fc3354 100644 --- a/facefusion/apis/endpoints/session.py +++ b/facefusion/apis/endpoints/session.py @@ -2,7 +2,7 @@ import secrets 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 starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_401_UNAUTHORIZED, HTTP_404_NOT_FOUND, HTTP_503_SERVICE_UNAVAILABLE 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 @@ -14,30 +14,37 @@ from facefusion.jobs import job_manager async def create_session(request : Request) -> JSONResponse: body = await request.json() + session_limit = state_manager.get_item('api_session_limit') if validate_api_key(body.get('api_key')): - session_id = secrets.token_urlsafe(16) - session = session_manager.create_api_session() - session_context.set_session_id(session_id) - session_manager.set_api_session(session_id, session) + if session_manager.count_api_sessions() < session_limit: + session_id = secrets.token_urlsafe(16) + session = session_manager.create_api_session() + session_context.set_session_id(session_id) + session_manager.set_api_session(session_id, session) - state_manager.init() - asset_store.init() - content_store.init() - face_store.init() - inference_manager.init() - video_manager.init() - process_manager.init() - rtc_store.init() + state_manager.init() + asset_store.init() + content_store.init() + face_store.init() + inference_manager.init() + video_manager.init() + process_manager.init() + rtc_store.init() - jobs_path = state_manager.get_jobs_path() - job_manager.init_jobs(jobs_path) + jobs_path = state_manager.get_jobs_path() + job_manager.init_jobs(jobs_path) + + return JSONResponse( + { + 'access_token': session.get('access_token'), + 'refresh_token': session.get('refresh_token') + }, status_code = HTTP_201_CREATED) return JSONResponse( { - 'access_token': session.get('access_token'), - 'refresh_token': session.get('refresh_token') - }, status_code = HTTP_201_CREATED) + 'message': translator.get('session_limit_reached', 'facefusion.apis') + }, status_code = HTTP_503_SERVICE_UNAVAILABLE) return JSONResponse( { diff --git a/facefusion/apis/locales.py b/facefusion/apis/locales.py index 13e626e9..21da94e5 100644 --- a/facefusion/apis/locales.py +++ b/facefusion/apis/locales.py @@ -9,9 +9,10 @@ LOCALES : Locales =\ 'directory_not_removed': 'directory not removed', 'invalid_access_token': 'invalid access token', 'invalid_refresh_token': 'invalid refresh token', + 'session_limit_reached': 'session limit reached', + 'invalid_state_key': 'invalid state key', 'source_asset_not_found': 'source asset not found', 'target_asset_not_found': 'target asset not found', - 'invalid_state_key': 'invalid state key', 'invalid_job_status': 'invalid job status', 'invalid_job_action': 'invalid job action', 'job_not_found': 'job not found', diff --git a/facefusion/args_helper.py b/facefusion/args_helper.py index 59332581..574f0139 100644 --- a/facefusion/args_helper.py +++ b/facefusion/args_helper.py @@ -83,6 +83,7 @@ def apply_args(args : Args, apply_state_item : ApplyStateItem) -> None: apply_state_item('api_host', args.get('api_host')) apply_state_item('api_port', args.get('api_port')) apply_state_item('api_security_strategy', args.get('api_security_strategy')) + apply_state_item('api_session_limit', args.get('api_session_limit')) apply_state_item('video_memory_strategy', args.get('video_memory_strategy')) apply_state_item('log_level', args.get('log_level')) apply_state_item('halt_on_error', args.get('halt_on_error')) diff --git a/facefusion/choices.py b/facefusion/choices.py index 35de62a6..40d490a3 100755 --- a/facefusion/choices.py +++ b/facefusion/choices.py @@ -162,6 +162,7 @@ progress_action_set : ProgressActionSet =\ job_statuses : List[JobStatus] = list(get_args(JobStatus)) +api_session_limit_range : Sequence[int] = create_int_range(1, 100, 1) benchmark_cycle_count_range : Sequence[int] = create_int_range(1, 10, 1) execution_thread_count_range : Sequence[int] = create_int_range(1, 32, 1) face_detector_margin_range : Sequence[int] = create_int_range(0, 100, 1) diff --git a/facefusion/locales.py b/facefusion/locales.py index b81ae5e0..396d1342 100644 --- a/facefusion/locales.py +++ b/facefusion/locales.py @@ -161,6 +161,7 @@ LOCALES : Locales =\ 'api_host': 'specify the API host', 'api_port': 'specify the API port', 'api_security_strategy': 'specify the API security strategy used for sanitizing uploaded assets', + 'api_session_limit': 'specify the maximum amount of concurrent sessions', 'execution_device_ids': 'specify the devices used for processing', 'execution_providers': 'inference using different providers (choices: {choices}, ...)', 'execution_thread_count': 'specify the amount of parallel threads while processing', diff --git a/facefusion/program.py b/facefusion/program.py index 8cee080e..b1fceaac 100755 --- a/facefusion/program.py +++ b/facefusion/program.py @@ -947,6 +947,14 @@ def create_api_program() -> ArgumentParser: default = config.get_str_value('api', 'api_security_strategy', 'strict'), choices = facefusion.choices.api_security_strategies ) + group_api.add_argument( + '--api-session-limit', + help = translator.get('help.api_session_limit'), + type = int, + default = config.get_int_value('api', 'api_session_limit', '1'), + choices = facefusion.choices.api_session_limit_range, + metavar = create_int_metavar(facefusion.choices.api_session_limit_range) + ) return program diff --git a/facefusion/session_manager.py b/facefusion/session_manager.py index da556729..a66d3815 100644 --- a/facefusion/session_manager.py +++ b/facefusion/session_manager.py @@ -65,6 +65,16 @@ def find_api_session_id(access_token : str) -> Optional[SessionId]: return None +def count_api_sessions() -> int: + session_total = 0 + + for session_id in API_SESSIONS: + if validate_api_session(session_id): + session_total += 1 + + return session_total + + def resolve_owner_id() -> SessionId: session_id = get_session_id() cli_session = get_cli_session(session_id) diff --git a/facefusion/types.py b/facefusion/types.py index c769d1f8..a0cc19cb 100755 --- a/facefusion/types.py +++ b/facefusion/types.py @@ -597,6 +597,7 @@ StateKey = Literal\ 'api_port', 'api_key', 'api_security_strategy', + 'api_session_limit', 'job_id', 'job_status', 'step_index' @@ -673,6 +674,7 @@ State = TypedDict('State', 'api_port' : int, 'api_key' : str, 'api_security_strategy' : ApiSecurityStrategy, + 'api_session_limit' : int, 'job_id' : str, 'job_status' : JobStatus, 'step_index' : int diff --git a/tests/test_api_assets.py b/tests/test_api_assets.py index f9365d92..7be82545 100644 --- a/tests/test_api_assets.py +++ b/tests/test_api_assets.py @@ -43,6 +43,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()) + state_manager.init_item('api_session_limit', 10) state_manager.init_item('temp_frame_format', 'png') session_manager.API_SESSIONS.clear() asset_store.delete_assets() diff --git a/tests/test_api_jobs.py b/tests/test_api_jobs.py index 8f724383..2ad3a18c 100644 --- a/tests/test_api_jobs.py +++ b/tests/test_api_jobs.py @@ -38,6 +38,7 @@ def before_each() -> Iterator[None]: state_manager.init_item('target_path', get_test_example_file('target-240p.mp4')) state_manager.init_item('temp_path', get_test_jobs_directory()) state_manager.init_item('jobs_path', get_test_jobs_directory()) + state_manager.init_item('api_session_limit', 10) clear_jobs(get_test_jobs_directory()) init_jobs(state_manager.get_jobs_path()) diff --git a/tests/test_api_metrics.py b/tests/test_api_metrics.py index fb7c393d..18113ddb 100644 --- a/tests/test_api_metrics.py +++ b/tests/test_api_metrics.py @@ -13,6 +13,7 @@ from .assert_helper import get_test_jobs_directory def before_all() -> None: state_manager.init() state_manager.init_item('jobs_path', get_test_jobs_directory()) + state_manager.init_item('api_session_limit', 10) @pytest.fixture(scope = 'function', autouse = True) diff --git a/tests/test_api_ping.py b/tests/test_api_ping.py index 0c116685..df597445 100644 --- a/tests/test_api_ping.py +++ b/tests/test_api_ping.py @@ -12,6 +12,7 @@ from .assert_helper import get_test_jobs_directory def before_all() -> None: state_manager.init() state_manager.init_item('jobs_path', get_test_jobs_directory()) + state_manager.init_item('api_session_limit', 10) @pytest.fixture(scope = 'function', autouse = True) diff --git a/tests/test_api_session.py b/tests/test_api_session.py index ab1359ae..ec4c8062 100644 --- a/tests/test_api_session.py +++ b/tests/test_api_session.py @@ -37,6 +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()) + state_manager.init_item('api_session_limit', 10) session_manager.API_SESSIONS.clear() asset_store.delete_assets() diff --git a/tests/test_api_state.py b/tests/test_api_state.py index 2f670a2f..87802b11 100644 --- a/tests/test_api_state.py +++ b/tests/test_api_state.py @@ -15,6 +15,7 @@ from .assert_helper import get_test_example_file, get_test_examples_directory, g def before_all() -> None: state_manager.init() state_manager.init_item('jobs_path', get_test_jobs_directory()) + state_manager.init_item('api_session_limit', 10) process_manager.start() program = ArgumentParser() diff --git a/tests/test_api_stream.py b/tests/test_api_stream.py index f2aae03f..d1539cb7 100644 --- a/tests/test_api_stream.py +++ b/tests/test_api_stream.py @@ -23,6 +23,7 @@ def before_all() -> None: state_manager.init_item('download_providers', [ 'github', 'huggingface' ]) state_manager.init_item('temp_path', tempfile.gettempdir()) state_manager.init_item('jobs_path', get_test_jobs_directory()) + state_manager.init_item('api_session_limit', 10) state_manager.init_item('processors', []) pre_check()