diff --git a/facefusion/apis/core.py b/facefusion/apis/core.py index d8ca2f2c..9555ac89 100644 --- a/facefusion/apis/core.py +++ b/facefusion/apis/core.py @@ -8,6 +8,7 @@ from starlette.routing import Route, WebSocketRoute from facefusion.apis.endpoints.assets import delete_assets, get_asset, get_assets, upload_asset from facefusion.apis.endpoints.capabilities import get_capabilities +from facefusion.apis.endpoints.jobs import create_job from facefusion.apis.endpoints.metrics import get_metrics, websocket_metrics from facefusion.apis.endpoints.ping import websocket_ping from facefusion.apis.endpoints.session import create_session, destroy_session, get_session, refresh_session @@ -47,6 +48,7 @@ def create_api() -> Starlette: Route('/metrics', get_metrics, methods = [ 'GET' ], middleware = [ session_guard ]), Route('/stream', post_stream, methods = [ 'POST' ], middleware = [ session_guard ]), Route('/stream', delete_stream, methods = [ 'DELETE' ], name = 'delete_stream', middleware = [ session_guard ]), + Route('/jobs', create_job, methods = [ 'POST' ], middleware = [ session_guard ]), WebSocketRoute('/metrics', websocket_metrics, middleware = [ session_guard ]), WebSocketRoute('/ping', websocket_ping, middleware = [ session_guard ]), WebSocketRoute('/stream', websocket_stream, middleware = [ session_guard ]) diff --git a/facefusion/apis/endpoints/jobs.py b/facefusion/apis/endpoints/jobs.py new file mode 100644 index 00000000..8103cd17 --- /dev/null +++ b/facefusion/apis/endpoints/jobs.py @@ -0,0 +1,21 @@ +from starlette.requests import Request +from starlette.responses import JSONResponse +from starlette.status import HTTP_201_CREATED, HTTP_400_BAD_REQUEST + +from facefusion import translator +from facefusion.jobs import job_helper, job_manager + + +async def create_job(request : Request) -> JSONResponse: + job_id = job_helper.suggest_job_id() + + if job_manager.create_job(job_id): + return JSONResponse( + { + 'job_id': job_id + }, status_code = HTTP_201_CREATED) + + return JSONResponse( + { + 'message': translator.get('job_not_created', 'facefusion.apis') + }, status_code = HTTP_400_BAD_REQUEST) diff --git a/facefusion/apis/locales.py b/facefusion/apis/locales.py index 831be2f4..a514fdc7 100644 --- a/facefusion/apis/locales.py +++ b/facefusion/apis/locales.py @@ -10,6 +10,7 @@ LOCALES : Locales =\ 'invalid_refresh_token': 'invalid refresh token', 'source_asset_not_found': 'source asset not found', 'target_asset_not_found': 'target asset not found', - 'invalid_state_key': 'invalid state key' + 'invalid_state_key': 'invalid state key', + 'job_not_created': 'job not created' } } diff --git a/facefusion/core.py b/facefusion/core.py index 2d179bd5..4c27c00b 100755 --- a/facefusion/core.py +++ b/facefusion/core.py @@ -57,6 +57,8 @@ def route(args : Args) -> None: if state_manager.get_item('command') == 'api': if not common_pre_check() or not processors_pre_check() or not facefusion.apis.core.pre_check(): hard_exit(2) + if not job_manager.init_jobs(state_manager.get_jobs_path()): + hard_exit(1) logger.info(translator.get('api_started').format(host = state_manager.get_item('api_host'), port = state_manager.get_item('api_port')), __name__) uvicorn.run(facefusion.apis.core.create_api(), host = state_manager.get_item('api_host'), port = state_manager.get_item('api_port')) diff --git a/tests/test_api_jobs.py b/tests/test_api_jobs.py new file mode 100644 index 00000000..7e472941 --- /dev/null +++ b/tests/test_api_jobs.py @@ -0,0 +1,62 @@ +from typing import Iterator +from unittest.mock import patch + +import pytest +from starlette.testclient import TestClient + +from facefusion import metadata, session_manager +from facefusion.apis.core import create_api +from facefusion.jobs.job_manager import clear_jobs, find_job_ids, init_jobs +from .assert_helper import get_test_jobs_directory + + +@pytest.fixture(scope = 'function', autouse = True) +def before_each() -> None: + session_manager.SESSIONS.clear() + clear_jobs(get_test_jobs_directory()) + init_jobs(get_test_jobs_directory()) + + +@pytest.fixture(scope = 'module') +def test_client() -> Iterator[TestClient]: + with TestClient(create_api()) as test_client: + yield test_client + + +def test_create_job(test_client : TestClient) -> None: + create_job_response = test_client.post('/jobs') + + assert create_job_response.status_code == 401 + + create_session_response = test_client.post('/session', json = + { + 'client_version': metadata.get('version') + }) + create_session_body = create_session_response.json() + access_token = create_session_body.get('access_token') + + create_job_response = test_client.post('/jobs', headers = + { + 'Authorization': 'Bearer ' + access_token + }) + create_job_body = create_job_response.json() + + assert create_job_body.get('job_id') in find_job_ids('drafted') + assert create_job_response.status_code == 201 + + with patch('facefusion.jobs.job_helper.suggest_job_id', return_value = 'job-test-create-job'): + create_job_response = test_client.post('/jobs', headers = + { + 'Authorization': 'Bearer ' + access_token + }) + + assert create_job_response.status_code == 201 + + create_job_response = test_client.post('/jobs', headers = + { + 'Authorization': 'Bearer ' + access_token + }) + create_job_body = create_job_response.json() + + assert create_job_body.get('message') == 'job not created' + assert create_job_response.status_code == 400