Files
AI-Profile-Router/platform/ltx2-studio/remote-app_factory.py
T

180 lines
7.0 KiB
Python

"""FastAPI app factory decoupled from runtime bootstrap side effects."""
from __future__ import annotations
import base64
import hmac
import os
from pathlib import Path
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any
from fastapi import FastAPI, Request
from fastapi.exceptions import RequestValidationError
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.responses import Response as StarletteResponse
from _routes._errors import HTTPError, build_http_error_response
from _routes.generation import router as generation_router
from _routes.hf_auth import router as hf_auth_router
from _routes.health import router as health_router
from _routes.ic_lora import router as ic_lora_router
from _routes.lora_catalog import router as lora_catalog_router
from _routes.image_gen import router as image_gen_router
from _routes.prompt_enhancement import router as prompt_enhancement_router
from _routes.models import router as models_router
from _routes.suggest_gap_prompt import router as suggest_gap_prompt_router
from _routes.retake import router as retake_router
from _routes.extend import router as extend_router
from _routes.runtime_policy import router as runtime_policy_router
from _routes.settings import router as settings_router
from api_types import HTTPErrorResponse
from logging_policy import log_http_error, log_unhandled_exception
from state import init_state_service
if TYPE_CHECKING:
from app_handler import AppHandler
DEFAULT_ALLOWED_ORIGINS: list[str] = [
"http://localhost:5173",
"http://127.0.0.1:5173",
]
DEFAULT_ERROR_RESPONSES: dict[int | str, dict[str, Any]] = {
"4XX": {
"model": HTTPErrorResponse,
"description": "Client Error",
},
"5XX": {
"model": HTTPErrorResponse,
"description": "Server Error",
},
}
def create_app(
*,
handler: "AppHandler",
allowed_origins: list[str] | None = None,
title: str = "LTX-2 Video Generation Server",
auth_token: str = "",
admin_token: str = "",
) -> FastAPI:
"""Create a configured FastAPI app bound to the provided handler."""
init_state_service(handler)
app = FastAPI(title=title, responses=DEFAULT_ERROR_RESPONSES)
remote_auth_token = os.environ.get("LTX_REMOTE_AUTH_TOKEN", "").strip()
remote_token_file = os.environ.get("LTX_REMOTE_AUTH_TOKEN_FILE", "").strip()
if remote_token_file:
remote_auth_token = Path(remote_token_file).read_text(encoding="utf-8").strip()
app.state.admin_token = admin_token # type: ignore[attr-defined]
app.add_middleware(
CORSMiddleware,
allow_origins=allowed_origins or DEFAULT_ALLOWED_ORIGINS,
allow_methods=["*"],
allow_headers=["*"],
)
@app.middleware("http")
async def _auth_middleware( # pyright: ignore[reportUnusedFunction]
request: Request,
call_next: Callable[[Request], Awaitable[StarletteResponse]],
) -> StarletteResponse:
if not auth_token:
return await call_next(request)
if request.method == "OPTIONS":
return await call_next(request)
if request.url.path == "/api/auth/huggingface/callback":
return await call_next(request)
def _token_matches(candidate: str) -> bool:
return hmac.compare_digest(candidate, auth_token) or (
bool(remote_auth_token) and hmac.compare_digest(candidate, remote_auth_token)
)
# WebSocket: check query param
if request.headers.get("upgrade", "").lower() == "websocket":
if _token_matches(request.query_params.get("token", "")):
return await call_next(request)
return JSONResponse(
status_code=401,
content=build_http_error_response(401, "Unauthorized").model_dump(),
)
# HTTP: Bearer or Basic auth
auth_header = request.headers.get("authorization", "")
if auth_header.startswith("Bearer ") and _token_matches(auth_header[7:]):
return await call_next(request)
if auth_header.startswith("Basic "):
try:
decoded = base64.b64decode(auth_header[6:]).decode()
_, _, password = decoded.partition(":")
if _token_matches(password):
return await call_next(request)
except Exception:
pass
return JSONResponse(
status_code=401,
content=build_http_error_response(401, "Unauthorized").model_dump(),
)
async def _route_http_error_handler(request: Request, exc: Exception) -> JSONResponse:
if isinstance(exc, HTTPError):
log_http_error(request, exc)
return JSONResponse(status_code=exc.status_code, content=exc.response.model_dump())
return JSONResponse(
status_code=500,
content=build_http_error_response(500, str(exc)).model_dump(),
)
async def _starlette_http_error_handler(request: Request, exc: Exception) -> JSONResponse:
if isinstance(exc, StarletteHTTPException):
return JSONResponse(
status_code=exc.status_code,
content=build_http_error_response(exc.status_code, exc.detail).model_dump(),
)
return JSONResponse(
status_code=500,
content=build_http_error_response(500, str(exc)).model_dump(),
)
async def _validation_error_handler(request: Request, exc: Exception) -> JSONResponse:
if isinstance(exc, RequestValidationError):
return JSONResponse(
status_code=422,
content=build_http_error_response(422, str(exc)).model_dump(),
)
return JSONResponse(
status_code=422,
content=build_http_error_response(422, str(exc)).model_dump(),
)
async def _route_generic_error_handler(request: Request, exc: Exception) -> JSONResponse:
log_unhandled_exception(request, exc)
return JSONResponse(
status_code=500,
content=build_http_error_response(500, str(exc)).model_dump(),
)
app.add_exception_handler(RequestValidationError, _validation_error_handler)
app.add_exception_handler(HTTPError, _route_http_error_handler)
app.add_exception_handler(StarletteHTTPException, _starlette_http_error_handler)
app.add_exception_handler(Exception, _route_generic_error_handler)
app.include_router(health_router)
app.include_router(generation_router)
app.include_router(models_router)
app.include_router(settings_router)
app.include_router(image_gen_router)
app.include_router(suggest_gap_prompt_router)
app.include_router(retake_router)
app.include_router(extend_router)
app.include_router(ic_lora_router)
app.include_router(lora_catalog_router)
app.include_router(prompt_enhancement_router)
app.include_router(runtime_policy_router)
app.include_router(hf_auth_router)
return app