better api key validation

This commit is contained in:
henryruhs
2026-09-04 09:26:04 +02:00
parent 9a4c263390
commit 2dd3948893
3 changed files with 21 additions and 3 deletions
+2 -3
View File
@@ -1,4 +1,3 @@
import os
import secrets
from starlette.requests import Request
@@ -7,14 +6,14 @@ from starlette.status import HTTP_200_OK, HTTP_201_CREATED, HTTP_401_UNAUTHORIZE
from facefusion import session_context, session_manager, state_manager, translator
from facefusion.apis import asset_store
from facefusion.apis.session_helper import extract_access_token
from facefusion.apis.session_helper import extract_access_token, validate_api_key
from facefusion.filesystem import remove_directory
async def create_session(request : Request) -> JSONResponse:
body = await request.json()
if not body.get('api_key') or body.get('api_key') == os.getenv('FACEFUSION_API_KEY'):
if validate_api_key(body.get('api_key')):
session_id = secrets.token_urlsafe(16)
session = session_manager.create_session()
session_context.set_session_id(session_id)
+11
View File
@@ -1,3 +1,5 @@
import os
import secrets
from typing import Optional
from starlette.datastructures import Headers
@@ -7,6 +9,15 @@ from facefusion.apis.api_helper import get_sec_websocket_protocol
from facefusion.types import Token
def validate_api_key(api_key : Optional[str]) -> bool:
__api_key__ = os.getenv('FACEFUSION_API_KEY')
if api_key and __api_key__:
return secrets.compare_digest(api_key, __api_key__)
return bool(api_key) == bool(__api_key__)
def extract_access_token(scope : Scope) -> Optional[Token]:
if scope.get('type') == 'http':
auth_header = Headers(scope = scope).get('Authorization')
+8
View File
@@ -55,6 +55,14 @@ def test_create_session(test_client : TestClient) -> None:
assert create_session_response.status_code == 401
os.environ['FACEFUSION_API_KEY'] = 'TEST'
create_session_response = test_client.post('/session', json =
{
'client_version': metadata.get('version')
})
assert create_session_response.status_code == 401
os.environ['FACEFUSION_API_KEY'] = 'TEST'
create_session_response = test_client.post('/session', json =
{