ai-gateway / app.py
basyx's picture
Upload 61 files
0287e4e verified
Raw
History Blame Contribute Delete
8.57 kB
"""AI Gateway FastAPI application."""
from __future__ import annotations
import logging
import time
from contextlib import asynccontextmanager
from dataclasses import dataclass
from uuid import uuid4
import gradio as gr
from fastapi import FastAPI, HTTPException, Request
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse, RedirectResponse
from fastapi.staticfiles import StaticFiles
from fastmcp.utilities.lifespan import combine_lifespans
from gradio.context import LocalContext
from config import Settings, get_settings
from core.errors import GatewayError
from core.loader import ModelLoader
from core.manager import AIService, TaskManager
from core.runtime import initialize_zerogpu
from core.workflow import WorkflowService
from gateway_mcp.auth import AuthenticationError, authenticate_headers
from gateway_mcp.server import MCPServerBundle, create_mcp_server
from gradio_ui import create_gradio_ui
from routes import audio, health, image, video, workflow
from routes.schemas import ErrorResponse
from utils.files import OutputManager
from utils.logger import configure_logging, request_id_context
logger = logging.getLogger(__name__)
@dataclass(slots=True)
class GatewayServices:
"""Process-wide services made available to route dependencies."""
settings: Settings
outputs: OutputManager
loader: ModelLoader
tasks: TaskManager
workflows: WorkflowService
ai: AIService
mcp: MCPServerBundle
def create_app(settings: Settings | None = None) -> FastAPI:
"""Create an application instance, allowing isolated test settings."""
resolved_settings = settings or get_settings()
configure_logging(resolved_settings.log_level)
outputs = OutputManager(resolved_settings)
outputs.initialize()
loader = ModelLoader(resolved_settings)
tasks = TaskManager(resolved_settings, loader)
workflows = WorkflowService(resolved_settings, tasks, outputs)
ai = AIService(resolved_settings, tasks, loader, outputs, workflows)
mcp = create_mcp_server(ai)
gradio_ui = create_gradio_ui(ai)
services = GatewayServices(
resolved_settings, outputs, loader, tasks, workflows, ai, mcp
)
@asynccontextmanager
async def lifespan(application: FastAPI):
application.state.gateway = services
removed = outputs.cleanup_stale_tmp()
logger.info("gateway starting", extra={"stale_tmp_removed": removed})
initialize_zerogpu()
await tasks.start()
try:
yield
finally:
await tasks.stop()
loader.close()
logger.info("gateway stopped")
application = FastAPI(
title=resolved_settings.app_name,
version=resolved_settings.app_version,
description="A serialized, multi-model media inference gateway.",
lifespan=combine_lifespans(lifespan, mcp.app.lifespan),
)
application.state.gateway = services
application.state.gateway_mcp = mcp
@application.middleware("http")
async def request_context(request: Request, call_next):
request_id = uuid4().hex
request.state.request_id = request_id
token = request_id_context.set(request_id)
gradio_token = LocalContext.request.set(gr.Request(request))
started = time.perf_counter()
try:
public_paths = {"/", "/health", "/docs", "/openapi.json", "/redoc"}
gradio_path = request.url.path == "/ui" or request.url.path.startswith(
"/ui/"
)
if request.url.path not in public_paths and not gradio_path:
try:
identity = authenticate_headers(
request.headers,
resolved_settings,
mcp_request=(
request.url.path == "/mcp"
or request.url.path.startswith("/mcp/")
),
)
except AuthenticationError:
body = ErrorResponse(
error="Invalid or missing API credentials",
code="unauthorized",
request_id=request_id,
)
response = JSONResponse(status_code=401, content=body.model_dump())
else:
request.state.auth_user = identity.user
response = await call_next(request)
else:
request.state.auth_user = "anonymous"
response = await call_next(request)
finally:
elapsed = round(time.perf_counter() - started, 3)
logger.info(
"request completed",
extra={
"method": request.method,
"path": request.url.path,
"status_code": getattr(locals().get("response"), "status_code", 500),
"execution_time": elapsed,
},
)
request_id_context.reset(token)
LocalContext.request.reset(gradio_token)
response.headers["X-Request-ID"] = request_id
return response
@application.exception_handler(GatewayError)
async def gateway_error_handler(request: Request, exc: GatewayError) -> JSONResponse:
logger.warning(
"gateway request failed",
extra={"error_code": exc.code, "status_code": exc.status_code},
)
body = ErrorResponse(
error=exc.message,
code=exc.code,
request_id=getattr(request.state, "request_id", "-"),
)
return JSONResponse(status_code=exc.status_code, content=body.model_dump())
@application.exception_handler(RequestValidationError)
async def validation_error_handler(
request: Request, exc: RequestValidationError
) -> JSONResponse:
errors = "; ".join(
f"{'.'.join(str(item) for item in error['loc'])}: {error['msg']}"
for error in exc.errors()
)
logger.warning(
"request validation failed",
extra={"error_code": "validation_error", "status_code": 422},
)
body = ErrorResponse(
error=errors,
code="validation_error",
request_id=getattr(request.state, "request_id", "-"),
)
return JSONResponse(status_code=422, content=body.model_dump())
@application.exception_handler(HTTPException)
async def http_error_handler(request: Request, exc: HTTPException) -> JSONResponse:
body = ErrorResponse(
error=str(exc.detail),
code="http_error",
request_id=getattr(request.state, "request_id", "-"),
)
return JSONResponse(status_code=exc.status_code, content=body.model_dump())
@application.exception_handler(Exception)
async def unexpected_error_handler(request: Request, exc: Exception) -> JSONResponse:
logger.exception("unhandled request error")
body = ErrorResponse(
error="Internal server error",
code="internal_error",
request_id=getattr(request.state, "request_id", "-"),
)
return JSONResponse(status_code=500, content=body.model_dump())
application.include_router(health.router)
application.include_router(image.router)
application.include_router(video.router)
application.include_router(audio.router)
application.include_router(workflow.router)
application.mount("/mcp", mcp.app, name="mcp")
@application.get("/", include_in_schema=False)
@application.get("/ui", include_in_schema=False)
async def ui_entry(request: Request) -> RedirectResponse:
query = f"?{request.url.query}" if request.url.query else ""
return RedirectResponse(url=f"/ui/{query}")
application.mount(
"/output",
StaticFiles(directory=resolved_settings.output_folder, check_dir=False),
name="output",
)
return gr.mount_gradio_app(
application,
gradio_ui,
path="/ui",
server_name=resolved_settings.host,
server_port=resolved_settings.port,
allowed_paths=[str(outputs.root)],
show_error=False,
max_file_size=f"{resolved_settings.max_upload_mb}mb",
ssr_mode=False,
mcp_server=True,
)
app = create_app()
if __name__ == "__main__":
import uvicorn
settings = get_settings()
uvicorn.run(app, host=settings.host, port=settings.port, proxy_headers=True)