diff --git a/python/CMakeLists.txt b/python/CMakeLists.txt index d23c39bfa..f44787b21 100644 --- a/python/CMakeLists.txt +++ b/python/CMakeLists.txt @@ -170,7 +170,6 @@ endforeach() # set(REVNG_DAEMON_MODULE_FILES revng/daemon/__init__.py - revng/daemon/api.py revng/daemon/demo_webpage.py revng/daemon/schema.py revng/daemon/schema.graphql.tpl diff --git a/python/requirements.txt b/python/requirements.txt index c67adf43b..9d6a59fc0 100644 --- a/python/requirements.txt +++ b/python/requirements.txt @@ -11,11 +11,11 @@ grandiso cffi # revng.daemon -Flask -Werkzeug ariadne aiodataloader +hypercorn Jinja2 +python-multipart xdg # revng.daemon - tests diff --git a/python/revng/cli/daemon.py b/python/revng/cli/daemon.py index 8f64872c9..474cdf5d7 100644 --- a/python/revng/cli/daemon.py +++ b/python/revng/cli/daemon.py @@ -29,11 +29,23 @@ class DaemonCommand(Command): **os.environ, "REVNG_ANALYSIS_LIBRARIES": ":".join(libraries), "REVNG_PIPELINES": ",".join(pipelines), - "FLASK_RUN_PORT": port, - "FLASK_APP": "revng.daemon", - "FLASK_ENV": "development", + "STARLETTE_DEBUG": "1", } - return run([py_executable, "-m", "flask", "run", "--no-reload"], options, env, True) + return run( + [ + py_executable, + "-m", + "hypercorn", + "-b", + f"127.0.0.1:{port}", + "-k", + "asyncio", + "revng.daemon:app", + ], + options, + env, + True, + ) commands_registry.register_command(DaemonCommand()) diff --git a/python/revng/daemon/__init__.py b/python/revng/daemon/__init__.py index 92073316e..a0d38a138 100644 --- a/python/revng/daemon/__init__.py +++ b/python/revng/daemon/__init__.py @@ -2,57 +2,84 @@ # This file is distributed under the MIT License. See LICENSE.md for details. # -import atexit import logging -import secrets from pathlib import Path from typing import Optional -from flask import Flask, g +from ariadne.asgi import GraphQL +from ariadne.contrib.tracing.apollotracing import ApolloTracingExtension +from starlette.applications import Starlette +from starlette.config import Config +from starlette.middleware import Middleware +from starlette.middleware.base import BaseHTTPMiddleware +from starlette.requests import Request +from starlette.responses import PlainTextResponse from revng.api._capi import initialize as capi_initialize -from .api import api_blueprint -from .demo_webpage import demo_blueprint +from .demo_webpage import demo_page from .schema import SchemafulManager from .util import project_workdir workdir: Path = project_workdir() manager: Optional[SchemafulManager] = None +startup_done = False + +config = Config() +DEBUG = config("STARLETTE_DEBUG", cast=bool, default=False) -app = Flask(__name__) -app.register_blueprint(api_blueprint) -app.register_blueprint(demo_blueprint) +class ManagerMiddleware(BaseHTTPMiddleware): + async def dispatch(self, request: Request, call_next): + assert manager, "Manager not initialized" -app.secret_key = secrets.token_hex(16) + request.scope["workdir"] = workdir + request.scope["manager"] = manager + + # Safety checks + assert request.scope["workdir"] is not None + assert request.scope["workdir"].exists() and request.scope["workdir"].is_dir() + assert request.scope["manager"] is not None + + return await call_next(request) -def cleanup(): +async def status(request): + if startup_done: + return PlainTextResponse("OK") + else: + return PlainTextResponse("KO", 503) + + +def startup(): + global manager, startup_done + capi_initialize() + manager = SchemafulManager(workdir=str(workdir.resolve())) + app.mount( + "/graphql", + GraphQL( + manager.schema, + context_value={"manager": manager, "workdir": workdir}, + extensions=[ApolloTracingExtension], + debug=DEBUG, + ), + ) + startup_done = True + + +def shutdown(): if manager is not None: store_result = manager.store_containers() if not store_result: logging.warning("Failed to store manager's containers") -atexit.register(cleanup) +app = Starlette( + debug=DEBUG, + middleware=[Middleware(ManagerMiddleware)], + on_startup=[startup], + on_shutdown=[shutdown], +) - -@app.before_first_request -def init(): - global manager - capi_initialize() - manager = SchemafulManager(workdir=str(workdir.resolve())) - - -@app.before_request -def init_global_object(): - assert manager, "Manager not initialized" - - g.workdir = workdir - g.manager = manager - - # Safety checks - assert g.workdir is not None - assert g.workdir.exists() and g.workdir.is_dir() - assert g.manager is not None +app.add_route("/", demo_page, ["GET"]) +app.add_route("/status", status, ["GET"]) diff --git a/python/revng/daemon/api.py b/python/revng/daemon/api.py deleted file mode 100644 index 34b7c2f2a..000000000 --- a/python/revng/daemon/api.py +++ /dev/null @@ -1,74 +0,0 @@ -# -# This file is distributed under the MIT License. See LICENSE.md for details. -# - -import json -from typing import TYPE_CHECKING, cast - -from flask import Blueprint, current_app, jsonify, request - -from ariadne import graphql -from ariadne.constants import PLAYGROUND_HTML -from ariadne.contrib.tracing.apollotracing import ApolloTracingExtension -from ariadne.file_uploads import FilesDict, combine_multipart_data - -from .schema import SchemafulManager - -if TYPE_CHECKING: - from flask.ctx import _AppCtxGlobals - - class FlaskGlobals(_AppCtxGlobals): - manager: SchemafulManager - - g = FlaskGlobals() -else: - from flask import g - - -api_blueprint = Blueprint("api", __name__) - - -def json_response(data=None): - if data is None: - data = {} - return jsonify(data) - - -def json_error(message, http_code=404): - data = { - "error": message, - } - return jsonify(data), http_code - - -# When navigating to http:///graphql show a graphql sandbox -@api_blueprint.route("/graphql", methods=["GET"]) -def graphql_playground(): - return PLAYGROUND_HTML, 200 - - -@api_blueprint.route("/graphql", methods=["POST"]) -async def graphql_server(): - if request.content_type.startswith("multipart/form-data;"): - operations = request.form.get("operations") - req_map = request.form.get("map") - if operations is not None and req_map is not None: - request_files = cast(FilesDict, request.files) - data = combine_multipart_data( - json.loads(operations), json.loads(req_map), request_files - ) - else: - return json_error("Invalid form data") - else: - data = request.get_json() - - success, result = await graphql( - g.manager.schema, - data, - context_value={"g": g, "request": request}, - extensions=[ApolloTracingExtension], - debug=current_app.config["DEBUG"], - ) - - status_code = 200 if success else 400 - return jsonify(result), status_code diff --git a/python/revng/daemon/demo_webpage.py b/python/revng/daemon/demo_webpage.py index b738178cc..8c47fa251 100644 --- a/python/revng/daemon/demo_webpage.py +++ b/python/revng/daemon/demo_webpage.py @@ -2,11 +2,13 @@ # This file is distributed under the MIT License. See LICENSE.md for details. # -from flask import Blueprint, render_template +from pathlib import Path -demo_blueprint = Blueprint("demo", __name__) +from starlette.templating import Jinja2Templates + +module_dir = Path(__file__).parent.resolve() +templates = Jinja2Templates(directory=module_dir / "templates") -@demo_blueprint.route("/") -def index(): - return render_template("index.html") +async def demo_page(request): + return templates.TemplateResponse("index.html", {"request": request}) diff --git a/python/revng/daemon/schema.py b/python/revng/daemon/schema.py index 4e7b9923b..2a1aec2cc 100644 --- a/python/revng/daemon/schema.py +++ b/python/revng/daemon/schema.py @@ -10,6 +10,7 @@ from typing import Dict, List, Optional from ariadne import MutationType, ObjectType, QueryType, make_executable_schema, upload_scalar from graphql.type.schema import GraphQLSchema from jinja2 import Environment, FileSystemLoader +from starlette.datastructures import UploadFile from revng.api.manager import Manager from revng.api.rank import Rank @@ -31,28 +32,28 @@ async def resolve_root(_, info): @query.field("produce") async def resolve_produce(obj, info, *, step, container, target_list, only_if_ready=False): - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] targets = target_list.split(",") return manager.produce_target(step, targets, container, only_if_ready) @query.field("produce_artifacts") async def resolve_produce_artifacts(obj, info, *, step, paths=None, only_if_ready=False): - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] target_paths = paths.split(",") if paths is not None else None return manager.produce_target(step, target_paths, only_if_ready=only_if_ready) @query.field("step") async def resolve_step(_, info, *, name): - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] step = manager.get_step(name) return step.as_dict() if step is not None else {} @query.field("container") async def resolve_container(_, info, *, name, step): - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] container_id = manager.get_container_with_name(name) step = manager.get_step(step) if step is None or container_id is None: @@ -63,7 +64,7 @@ async def resolve_container(_, info, *, name, step): @query.field("targets") async def resolve_targets(_, info, *, pathspec): - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] targets = manager.get_all_targets() result = [ { @@ -81,16 +82,16 @@ async def resolve_targets(_, info, *, pathspec): @mutation.field("upload_b64") async def resolve_upload_b64(_, info, *, input: str, container: str): # noqa: A002 - g = info.context["g"] - g.manager.set_input(container, b64decode(input)) + manager: Manager = info.context["manager"] + manager.set_input(container, b64decode(input)) logging.info(f"Saved file for container {container}") return True @mutation.field("upload_file") -async def resolve_upload_file(_, info, *, file, container: str): - g = info.context["g"] - g.manager.set_input(container, file.read()) +async def resolve_upload_file(_, info, *, file: UploadFile, container: str): + manager: Manager = info.context["manager"] + manager.set_input(container, await file.read()) logging.info(f"Saved file for container {container}") return True @@ -102,19 +103,19 @@ async def resolve_ranks(_, info): @info.field("kinds") async def resolve_root_kinds(_, info): - manager = info.context["g"].manager + manager = info.context["manager"] return [k.as_dict() for k in manager.kinds()] @info.field("model") async def resolve_root_model(_, info): - manager = info.context["g"].manager + manager = info.context["manager"] return manager.get_model() @info.field("steps") async def resolve_root_steps(_, info): - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] return [s.as_dict() for s in manager.steps()] @@ -123,7 +124,7 @@ async def resolve_step_containers(step_obj, info): if "containers" in step_obj: return step_obj["containers"] - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] step = manager.get_step(step_obj["name"]) if step is None: return [] @@ -136,7 +137,7 @@ async def resolve_container_targets(container_obj, info): if "targets" in container_obj: return container_obj["targets"] - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] targets = manager.get_targets(container_obj["_step"], container_obj["name"]) return [t.as_dict() for t in targets] @@ -218,7 +219,7 @@ class BindableGen: @staticmethod def gen_step_handle(step: Step): def rank_step_handle(obj, info, *, only_if_ready=False): - manager: Manager = info.context["g"].manager + manager: Manager = info.context["manager"] return manager.produce_target(step.name, obj["_target"], only_if_ready=only_if_ready) return rank_step_handle diff --git a/python/revng/daemon/templates/index.html b/python/revng/daemon/templates/index.html index 818ea2705..e5b7b9dee 100644 --- a/python/revng/daemon/templates/index.html +++ b/python/revng/daemon/templates/index.html @@ -11,7 +11,7 @@
Welcome to revng!
- Working directory is {{ g.workdir }} + Working directory is {{ request.scope.workdir }}
diff --git a/tests/daemon/test.py b/tests/daemon/test.py index b6eef699c..ac38a9ef5 100755 --- a/tests/daemon/test.py +++ b/tests/daemon/test.py @@ -26,7 +26,7 @@ def print_fd(fd: int): def check_server_up(port: int): for _ in range(10): try: - req = urlopen(f"http://127.0.0.1:{port}/", timeout=1.0) + req = urlopen(f"http://127.0.0.1:{port}/status", timeout=1.0) if req.code == 200: return sleep(1.0) @@ -54,7 +54,7 @@ def client(pytestconfig: Config, request) -> Generator[Client, None, None]: raise e binary = pytestconfig.getoption("binary") - transport = RequestsHTTPTransport(f"http://127.0.0.1:{port}/graphql") + transport = RequestsHTTPTransport(f"http://127.0.0.1:{port}/graphql/") gql_client = Client(transport=transport, fetch_schema_from_transport=True) upload_q = gql(