mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
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:
@@ -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."""
|
||||
|
||||
@@ -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 ======
|
||||
# ======================================
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user