From cf05ab0463081aa0fee695d8f7ca85fd918d8b8c Mon Sep 17 00:00:00 2001 From: Giacomo Vercesi Date: Thu, 28 Apr 2022 16:59:48 +0200 Subject: [PATCH] Substitute Flask with Starlette In the future we will need to use GraphQL subscriptions. This is done via websockets and is supported in Ariadne. However this support is limited to ASGI frameworks, which Flask isn't a part of. Startlette is a direct depencency of Ariadne, and all of Ariadne's features are fully integrated with Starlette, so the switch allows to drop some Flask integration cruft and streamline the revng.daemon package. --- python/CMakeLists.txt | 1 - python/requirements.txt | 4 +- python/revng/cli/daemon.py | 20 ++++-- python/revng/daemon/__init__.py | 87 ++++++++++++++++-------- python/revng/daemon/api.py | 74 -------------------- python/revng/daemon/demo_webpage.py | 12 ++-- python/revng/daemon/schema.py | 33 ++++----- python/revng/daemon/templates/index.html | 2 +- tests/daemon/test.py | 4 +- 9 files changed, 102 insertions(+), 135 deletions(-) delete mode 100644 python/revng/daemon/api.py 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(