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.
This commit is contained in:
John Chilton
2021-01-27 19:55:48 -05:00
parent 352c488aaf
commit c46ff97123
6 changed files with 166 additions and 34 deletions
+28 -21
View File
@@ -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."""
+15 -10
View File
@@ -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 ======
# ======================================
+31
View File
@@ -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(
+1 -1
View File
@@ -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)
+7 -2
View File
@@ -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,
)
@@ -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