From c46ff97123d35904e822476134e9f01ac3deb668 Mon Sep 17 00:00:00 2001 From: John Chilton Date: Wed, 27 Jan 2021 19:39:06 -0500 Subject: [PATCH] FastAPI/ASGI CORS handling. - Some tests to clarify existing behavior (at least where they don't diverge). Including frameworks enhancements to allow integration tests to force on server type or the other. - In WSGI-land return a 400 if an invalid origin is specified (to bring inline with CORSMiddleware). - In ASGI-land, override CORSMiddleware to use Galaxy's logic for checking an origin against the config object. - When building a WSGI webapp for FastAPI, skip CORS handling at that level so it can be handled by FastAPI app. --- lib/galaxy/webapps/base/webapp.py | 49 ++++++----- lib/galaxy/webapps/galaxy/buildapp.py | 25 +++--- lib/galaxy/webapps/galaxy/fast_app.py | 31 +++++++ lib/galaxy/webapps/galaxy/fast_factory.py | 2 +- lib/galaxy_test/driver/driver_util.py | 9 +- test/integration/test_web_framework_config.py | 84 +++++++++++++++++++ 6 files changed, 166 insertions(+), 34 deletions(-) create mode 100644 test/integration/test_web_framework_config.py diff --git a/lib/galaxy/webapps/base/webapp.py b/lib/galaxy/webapps/base/webapp.py index ead9ccee1f1..50e9c07ac6f 100644 --- a/lib/galaxy/webapps/base/webapp.py +++ b/lib/galaxy/webapps/base/webapp.py @@ -9,6 +9,7 @@ import socket import string import time from http.cookies import CookieError +from typing import Any, Dict from urllib.parse import urlparse import mako.lookup @@ -154,13 +155,36 @@ class WebApplication(base.WebApplication): return T(app) +def config_allows_origin(origin_raw, config): + # boil origin header down to hostname + origin = urlparse(origin_raw).hostname + + # singular match + def matches_allowed_origin(origin, allowed_origin): + if isinstance(allowed_origin, str): + return origin == allowed_origin + match = allowed_origin.match(origin) + return match and match.group() == origin + + # localhost uses no origin header (== null) + if not origin: + return False + + # check for '*' or compare to list of allowed + for allowed_origin in config.allowed_origin_hostnames: + if allowed_origin == '*' or matches_allowed_origin(origin, allowed_origin): + return True + + return False + + class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryContext): """ Encapsulates web transaction specific state for the Galaxy application (specifically the user's "cookie" session and history) """ - def __init__(self, environ, app, webapp, session_cookie=None): + def __init__(self, environ: Dict[str, Any], app, webapp, session_cookie=None) -> None: self._app = app self.webapp = webapp self.user_manager = UserManager(app) @@ -303,28 +327,11 @@ class GalaxyWebTransaction(base.DefaultWebTransaction, context.ProvidesHistoryCo if not origin_header: return - # singular match - def matches_allowed_origin(origin, allowed_origin): - if isinstance(allowed_origin, str): - return origin == allowed_origin - match = allowed_origin.match(origin) - return match and match.group() == origin - - # check for '*' or compare to list of allowed - def is_allowed_origin(origin): - # localhost uses no origin header (== null) - if not origin: - return False - for allowed_origin in self.app.config.allowed_origin_hostnames: - if allowed_origin == '*' or matches_allowed_origin(origin, allowed_origin): - return True - return False - - # boil origin header down to hostname - origin = urlparse(origin_header).hostname # check against the list of allowed strings/regexp hostnames, echo original if cleared - if is_allowed_origin(origin): + if config_allows_origin(origin_header, self.app.config): self.set_cors_origin(origin=origin_header) + else: + self.response.status = 400 def get_user(self): """Return the current user if logged in or None.""" diff --git a/lib/galaxy/webapps/galaxy/buildapp.py b/lib/galaxy/webapps/galaxy/buildapp.py index 7b7873cacd3..ae2fad7279a 100644 --- a/lib/galaxy/webapps/galaxy/buildapp.py +++ b/lib/galaxy/webapps/galaxy/buildapp.py @@ -38,9 +38,12 @@ def app_factory(*args, **kwargs): return app_pair(*args, **kwargs)[0] -def app_pair(global_conf, load_app_kwds=None, **kwargs): +def app_pair(global_conf, load_app_kwds=None, wsgi_preflight=True, **kwargs): """ Return a wsgi application serving the root object and the Galaxy application. + + When creating an app for asgi, set wsgi_preflight to False to allow FastAPI + middleware to handle CORS options, etc.. """ load_app_kwds = load_app_kwds or {} kwargs = load_app_properties( @@ -107,6 +110,15 @@ def app_pair(global_conf, load_app_kwds=None, **kwargs): # TODO: Refactor above routes into external method to allow testing in # isolation as well. populate_api_routes(webapp, app) + if wsgi_preflight: + # API OPTIONS RESPONSE + webapp.mapper.connect( + 'options', + '/api/{path_info:.*?}', + controller='authenticate', + action='options', + conditions={'method': ['OPTIONS']}, + ) # CLIENTSIDE ROUTES # The following are routes that are handled completely on the clientside. @@ -455,8 +467,8 @@ def populate_api_routes(webapp, app): webapp.mapper.connect('/api/workflows/menu', action='set_workflow_menu', controller="workflows", conditions=dict(method=["PUT"])) webapp.mapper.connect('/api/workflows/{id}/refactor', action='refactor', controller="workflows", conditions=dict(method=["PUT"])) webapp.mapper.resource('workflow', 'workflows', path_prefix='/api') - webapp.mapper.connect('/api/licenses', controller='licenses', action='index') - webapp.mapper.connect('/api/licenses/{id}', controller='licenses', action='get') + webapp.mapper.connect('/api/licenses', controller='licenses', action='index', conditions=dict(method="GET")) + webapp.mapper.connect('/api/licenses/{id}', controller='licenses', action='get', conditions=dict(method="GET")) webapp.mapper.resource_with_deleted('history', 'histories', path_prefix='/api') webapp.mapper.connect('/api/histories/{history_id}/citations', action='citations', controller="histories") webapp.mapper.connect('/api/histories/{id}/sharing', action='sharing', controller="histories", conditions=dict(method=["GET", "POST"])) @@ -708,13 +720,6 @@ def populate_api_routes(webapp, app): action='get_api_key', conditions=dict(method=["GET"])) - # API OPTIONS RESPONSE - webapp.mapper.connect('options', - '/api/{path_info:.*?}', - controller='authenticate', - action='options', - conditions={'method': ['OPTIONS']}) - # ====================================== # ====== DISPLAY APPLICATIONS API ====== # ====================================== diff --git a/lib/galaxy/webapps/galaxy/fast_app.py b/lib/galaxy/webapps/galaxy/fast_app.py index 88daeaa1f7b..82ab083e740 100644 --- a/lib/galaxy/webapps/galaxy/fast_app.py +++ b/lib/galaxy/webapps/galaxy/fast_app.py @@ -2,6 +2,7 @@ from fastapi import FastAPI, Request from fastapi.exceptions import RequestValidationError from fastapi.middleware.wsgi import WSGIMiddleware from fastapi.responses import JSONResponse +from starlette.middleware.cors import CORSMiddleware from starlette.responses import Response from galaxy.exceptions import MessageException @@ -10,6 +11,8 @@ from galaxy.web.framework.decorators import ( api_error_message, validation_error_to_message_exception ) +from galaxy.webapps.base.webapp import config_allows_origin + # https://fastapi.tiangolo.com/tutorial/metadata/#metadata-for-tags api_tags_metadata = [ @@ -40,6 +43,16 @@ api_tags_metadata = [ ] +class GalaxyCORSMiddleware(CORSMiddleware): + + def __init__(self, *args, **kwds): + self.config = kwds.pop("config") + super().__init__(*args, **kwds) + + def is_allowed_origin(self, origin: str) -> bool: + return config_allows_origin(origin, self.config) + + def add_exception_handler( app: FastAPI ) -> None: @@ -72,6 +85,24 @@ def add_galaxy_middleware(app: FastAPI, gx_app): response.headers['X-Frame-Options'] = x_frame_options return response + if gx_app.config.get('allowed_origin_hostnames', None): + app.add_middleware( + GalaxyCORSMiddleware, + config=gx_app.config, + allow_headers=["*"], + allow_methods=["*"], + max_age=600, + ) + else: + + # handle CORS preflight requests - synchronize with wsgi behavior. + @app.options('/api/{rest_of_path:path}') + async def preflight_handler(request: Request, rest_of_path: str) -> Response: + response = Response() + response.headers['Access-Control-Allow-Headers'] = '*' + response.headers['Access-Control-Max-Age'] = '600' + return response + def initialize_fast_app(gx_webapp, gx_app): app = FastAPI( diff --git a/lib/galaxy/webapps/galaxy/fast_factory.py b/lib/galaxy/webapps/galaxy/fast_factory.py index a717e5dcc8a..824325a730d 100644 --- a/lib/galaxy/webapps/galaxy/fast_factory.py +++ b/lib/galaxy/webapps/galaxy/fast_factory.py @@ -78,5 +78,5 @@ def factory(): global_conf = {} if config_is_ini(config_file): global_conf["__file__"] = config_file - gx_webapp, gx_app = app_pair(global_conf=global_conf, load_app_kwds=kwds) + gx_webapp, gx_app = app_pair(global_conf=global_conf, load_app_kwds=kwds, wsgi_preflight=False) return initialize_fast_app(gx_webapp, gx_app) diff --git a/lib/galaxy_test/driver/driver_util.py b/lib/galaxy_test/driver/driver_util.py index 5417d6dbcd5..f1b52d91be8 100644 --- a/lib/galaxy_test/driver/driver_util.py +++ b/lib/galaxy_test/driver/driver_util.py @@ -970,6 +970,11 @@ class GalaxyTestDriver(TestDriver): use_uwsgi = True self.use_uwsgi = use_uwsgi + if getattr(config_object, "use_uvicorn", USE_UVICORN): + self.else_use_uvicorn = True + else: + self.else_use_uvicorn = False + # Allow controlling the log format log_format = os.environ.get('GALAXY_TEST_LOG_FORMAT', None) if not log_format and use_uwsgi: @@ -1061,11 +1066,11 @@ class GalaxyTestDriver(TestDriver): tempdir=tempdir, config_object=config_object, ) - elif USE_UVICORN: + elif self.else_use_uvicorn: self.app = build_galaxy_app(galaxy_config) server_wrapper = launch_uvicorn( self.app, - buildapp.app_factory, + lambda *args, **kwd: buildapp.app_factory(*args, wsgi_preflight=False, **kwd), galaxy_config, config_object=config_object, ) diff --git a/test/integration/test_web_framework_config.py b/test/integration/test_web_framework_config.py new file mode 100644 index 00000000000..6892fc98e36 --- /dev/null +++ b/test/integration/test_web_framework_config.py @@ -0,0 +1,84 @@ +"""Integration tests for framework configuration code.""" +from requests import options + +from galaxy_test.driver import integration_util + + +class BaseWebFrameworkTestCase(integration_util.IntegrationTestCase): + + def _options(self, headers=None): + url = self._api_url("licenses") + options_response = options(url, headers=headers or {}) + return options_response + + +class CorsDefaultIntegrationTestCase(BaseWebFrameworkTestCase): + use_uvicorn = True + + def test_options(self): + headers = { + "Access-Control-Request-Method": "GET", + "origin": "http://192.168.0.101:8083", + } + options_response = self._options(headers) + assert options_response.status_code == 200 + assert 'access-control-allow-origin' not in options_response.headers + + def test_origin_not_allowed_default(self): + headers = { + "Access-Control-Request-Method": "GET", + "Access-Control-Request-Headers": "Authorization", + "origin": "http://192.168.0.101:8083", + } + options_response = self._options(headers) + assert options_response.status_code == 200 + assert 'access-control-allow-origin' not in options_response.headers + + +class AllowOriginIntegrationTestCase(BaseWebFrameworkTestCase): + use_uvicorn = True + + @classmethod + def handle_galaxy_config_kwds(cls, config): + config["allowed_origin_hostnames"] = "192.168.0.101,/.*.galaxyproject.org/" + + def test_origin_allowed_if_configured(self): + headers = { + "Access-Control-Request-Method": "GET", + "origin": "http://192.168.0.101:8083", + "Access-Control-Request-Headers": "Authorization", + } + options_response = self._options(headers) + options_response.raise_for_status() + assert 'access-control-allow-origin' in options_response.headers + assert options_response.headers['access-control-allow-origin'] == "http://192.168.0.101:8083" + assert options_response.headers['access-control-max-age'] == "600" + + def test_origin_allowed_if_configured_via_regex(self): + headers = { + "Access-Control-Request-Method": "GET", + "origin": "http://rna.galaxyproject.org", + "Access-Control-Request-Headers": "Authorization", + } + options_response = self._options(headers) + options_response.raise_for_status() + assert 'access-control-allow-origin' in options_response.headers + assert options_response.headers['access-control-allow-origin'] == "http://rna.galaxyproject.org" + assert options_response.headers['access-control-max-age'] == "600" + + def test_origin_not_allowed_if_not_in_configured_list(self): + headers = { + "Access-Control-Request-Method": "GET", + "origin": "http://192.168.0.102:8083", # swapped ip by one + "Access-Control-Request-Headers": "Authorization", + } + options_response = self._options(headers) + assert options_response.status_code == 400 + + +class AllowOriginPasteIntegrationTestCase(AllowOriginIntegrationTestCase): + use_uvicorn = False + + +class CorsDefaultPasteIntegrationTestCase(CorsDefaultIntegrationTestCase): + use_uvicorn = False