mirror of
https://github.com/facefusion/facefusion.git
synced 2026-09-15 12:05:27 +02:00
introduce api session limit
This commit is contained in:
@@ -150,6 +150,7 @@ benchmark_cycle_count =
|
||||
api_host =
|
||||
api_port =
|
||||
api_security_strategy =
|
||||
api_session_limit =
|
||||
|
||||
[execution]
|
||||
execution_device_ids =
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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'))
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user