Files
Giacomo Vercesi a91d87315a pypeline daemon: add /status endpoint
Add a `/status` endpoint to the daemon that allows verifying that the
daemon is up.
2026-04-10 11:45:12 +02:00

221 lines
8.1 KiB
Python

#
# 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,
)