# # This file is distributed under the MIT License. See LICENSE.md for details. # import asyncio import json import os from contextlib import asynccontextmanager, suppress from functools import wraps from typing import Mapping from starlette.applications import Starlette from starlette.datastructures import UploadFile from starlette.middleware import Middleware from starlette.middleware.cors import CORSMiddleware from starlette.middleware.gzip import GZipMiddleware from starlette.requests import Request from starlette.responses import JSONResponse, PlainTextResponse from starlette.responses import Response as StarletteResponse from starlette.routing import Route, WebSocketRoute from starlette.websockets import WebSocket, WebSocketDisconnect from revng.pypeline.utils import PypelineException from revng.pypeline.utils.logger import pypeline_logger from .daemon import Daemon, Response from .exceptions import DaemonException, MalformedRequestError from .notification_broker import WebSocketStream from .notification_broker.local_broker import LocalNotificationBroker # Global instances notification_broker = LocalNotificationBroker() # This is initialized by the `lifespan` function below, once the actual event # loop has been created, otherwise this creates another event loop and an # exception is thrown since the two loops don't match. shutdown_begun: asyncio.Event | None = None def daemon_exception_handler(request: Request, exc: Exception) -> JSONResponse: assert isinstance(exc, DaemonException) """Handle BasicHTTPException and return JSON response""" return JSONResponse(content=exc.body, status_code=exc.code) def pypeline_exception_handler(request: Request, exc: Exception) -> JSONResponse: assert isinstance(exc, PypelineException) """Handle BasicHTTPException and return JSON response""" return JSONResponse(content={"message": str(exc)}, status_code=500) def get_project_id(headers: Mapping[str, str]) -> str | None: """Extract project ID from headers, return none if missing""" return headers.get("x-project-id") async def invalidation_websocket(websocket: WebSocket): """Handle WebSocket connections for invalidations""" assert shutdown_begun is not None await websocket.accept() subscriber = None pending = None try: project_id = get_project_id(websocket.headers) subscriber = await notification_broker.subscribe(project_id, WebSocketStream(websocket)) _, pending = await asyncio.wait( ( asyncio.create_task(shutdown_begun.wait()), asyncio.create_task(subscriber.listen_for_messages()), ), return_when=asyncio.FIRST_COMPLETED, ) except WebSocketDisconnect: pass except Exception as e: pypeline_logger.log(f"Uncaught exception: {str(e)}") with suppress(RuntimeError): await websocket.close(code=500, reason=f"Internal server error: {str(e)}") finally: # Clean up unfinished tasks if pending is not None: for task in pending: task.cancel() # Clean up the subscription if subscriber is not None: await notification_broker.unsubscribe(subscriber) def prepare_endpoint(func): """A decorator that abstracts the boilerplate needed to adapt the http data the agnostic daemon implementation.""" @wraps(func) async def wrapper(request: Request) -> StarletteResponse: project_id = get_project_id(request.headers) # Prepare the data dictionary with the common attributes we extract # from the headers data = { "project_id": project_id, # TODO add auth token forwarding } response: Response = await func(request, data) # Forward any websocket notification for notification in response.notifications: await notification_broker.notify(project_id, json.dumps(notification)) # Convert the daemon response based on the body type: # * a JSON response if the body is a dict or list # * a binary response if the body is bytes # * throw an error for everything else if isinstance(response.body, (dict, list)): return JSONResponse( status_code=response.code, content=response.body, headers=response.headers, ) elif isinstance(response.body, bytes): return StarletteResponse( status_code=response.code, media_type=response.content_type, content=response.body, headers=response.headers, ) else: raise ValueError(f"Unknown response body type: {type(response.body)}") return wrapper def make_starlette(daemon: Daemon) -> Starlette: # This doesn't use `prepare_endpoint` as it doesn't need the project_id nor token async def pipeline_endpoint(request: Request) -> JSONResponse: """Get pipeline information""" response = daemon.get_pipeline() return JSONResponse( status_code=response.code, content=response.body, headers=response.headers, ) @prepare_endpoint async def epoch_endpoint(request: Request, data: dict) -> Response: """Get epoch information for a project""" return await daemon.get_epoch(data) @prepare_endpoint async def model_endpoint(request: Request, data: dict) -> Response: """Get model data for a project""" return await daemon.get_model(data) @prepare_endpoint async def put_file_endpoint(request: Request, data: dict) -> Response: """Put a file in storage""" async with request.form() as form: if not isinstance(form["file"], UploadFile): raise MalformedRequestError('"file" parameter is not a file') file: UploadFile = form["file"] return await daemon.put_file( {"name": file.filename, "contents": await file.read(), **data} ) @prepare_endpoint async def artifact_endpoint(request: Request, data: dict) -> Response: """Process artifact requests""" return await daemon.artifact({**await request.json(), **data}) @prepare_endpoint async def analysis_endpoint(request: Request, data: dict) -> Response: """Process analysis requests""" return await daemon.analyze({**await request.json(), **data}) async def status(request): return PlainTextResponse("OK") # Define routes routes = [ Route("/api/epoch", epoch_endpoint, methods=["GET"]), Route("/api/pipeline", pipeline_endpoint, methods=["GET"]), Route("/api/model", model_endpoint, methods=["GET"]), Route("/api/put-file", put_file_endpoint, methods=["POST"]), Route("/api/artifact", artifact_endpoint, methods=["POST"]), Route("/api/analysis", analysis_endpoint, methods=["POST"]), WebSocketRoute("/api/subscribe", invalidation_websocket), Route("/status", status, methods=["GET"]), ] origins: list[str] = [] if "REVNG_ORIGINS" in os.environ: origins = os.environ["REVNG_ORIGINS"].split(",") expose_headers: list[str] = ["x-pypeline-configuration-hash"] if "REVNG_EXPOSE_HEADERS" in os.environ: expose_headers.extend(os.environ["REVNG_EXPOSE_HEADERS"].split(",")) @asynccontextmanager async def lifespan(app): global shutdown_begun shutdown_begun = asyncio.Event() yield # Create the Starlette application return Starlette( debug=False, routes=routes, exception_handlers={ DaemonException: daemon_exception_handler, PypelineException: pypeline_exception_handler, }, middleware=[ Middleware( # type: ignore CORSMiddleware, # type: ignore allow_origins=origins, expose_headers=expose_headers, allow_methods=["*"], allow_headers=["*"], ), Middleware(GZipMiddleware, minimum_size=1024), # type: ignore ], lifespan=lifespan, )