diff --git a/facefusion/apis/core.py b/facefusion/apis/core.py index 0e3a5aea..8c78c575 100644 --- a/facefusion/apis/core.py +++ b/facefusion/apis/core.py @@ -12,6 +12,7 @@ from facefusion.apis.endpoints.capabilities import get_capabilities from facefusion.apis.endpoints.jobs import create_job, create_step, delete_job, delete_jobs, delete_step, get_job, get_jobs, update_job, update_jobs from facefusion.apis.endpoints.metrics import get_metrics, websocket_metrics from facefusion.apis.endpoints.ping import websocket_ping +from facefusion.apis.endpoints.root import get_root from facefusion.apis.endpoints.session import create_session, destroy_session, get_session, refresh_session from facefusion.apis.endpoints.state import get_state, set_state from facefusion.apis.endpoints.stream import delete_stream, post_stream, websocket_stream @@ -35,6 +36,7 @@ def create_api() -> Starlette: session_guard = Middleware(create_session_guard) routes =\ [ + Route('/', get_root, methods = [ 'GET' ]), Route('/session', create_session, methods = [ 'POST' ]), Route('/session', get_session, methods = [ 'GET' ], middleware = [ session_guard ]), Route('/session', refresh_session, methods = [ 'PUT' ]), diff --git a/tests/test_api_root.py b/tests/test_api_root.py new file mode 100644 index 00000000..843cb490 --- /dev/null +++ b/tests/test_api_root.py @@ -0,0 +1,28 @@ +from typing import Iterator + +import pytest +from starlette.testclient import TestClient + +from facefusion import metadata, session_manager +from facefusion.apis.core import create_api + + +@pytest.fixture(scope = 'function', autouse = True) +def before_each() -> None: + session_manager.API_SESSIONS.clear() + + +@pytest.fixture(scope = 'module') +def test_client() -> Iterator[TestClient]: + with TestClient(create_api()) as test_client: + yield test_client + + +def test_get_root(test_client : TestClient) -> None: + root_response = test_client.get('/') + root_body = root_response.json() + + assert root_body.get('name') == metadata.get('name') + assert root_body.get('description') == metadata.get('description') + assert root_body.get('version') == metadata.get('version') + assert root_response.status_code == 200