diff --git a/server/app/app.py b/server/app/app.py
index 5f47e235..72378caf 100644
--- a/server/app/app.py
+++ b/server/app/app.py
@@ -2,6 +2,9 @@ import datetime
import logging
from functools import wraps
from http import HTTPStatus
+from urllib.parse import urlparse
+import hashlib
+import os
from flask import Flask, redirect, current_app, make_response, render_template, abort, Blueprint, request, \
send_from_directory
@@ -84,10 +87,12 @@ def dataset_index(url_dataroot=None, dataset=None):
cache_manager = current_app.matrix_data_cache_manager
with cache_manager.data_adaptor(url_dataroot, location, app_config) as data_adaptor:
data_adaptor.set_uri_path(f"{url_dataroot}/{dataset}")
- dataset_title = app_config.get_title(data_adaptor)
- return render_template(
- "index.html", datasetTitle=dataset_title, SCRIPTS=scripts, INLINE_SCRIPTS=inline_scripts
- )
+ args = {
+ "SCRIPTS" : scripts,
+ "INLINE_SCRIPTS" : inline_scripts
+ }
+ return render_template("index.html", **args)
+
except DatasetAccessError as e:
return common_rest.abort_and_log(
e.status_code, f"Invalid dataset {dataset}: {e.message}", loglevel=logging.INFO, include_exc_info=True
@@ -179,9 +184,9 @@ def dataroot_test_index():
data += f"
Logged in as {auth.get_user_id()} / {auth.get_user_name()} / {auth.get_user_email()}
"
if auth.requires_client_login():
if server_config.auth.is_user_authenticated():
- data += "Logout
"
+ data += f"Logout
"
else:
- data += "Login
"
+ data += f"Login
"
datasets = []
for dataroot_dict in server_config.multi_dataset__dataroot.values():
@@ -329,6 +334,28 @@ def get_api_resources(bp_api, url_dataroot=None):
return api
+def handle_api_base_url(app, app_config):
+ """If an api_base_url is provided, then an inline script is generated to
+ handle the new API prefix"""
+ api_base_url = app_config.server_config.get_api_base_url()
+ if not api_base_url:
+ return
+
+ if api_base_url.endswith("/"):
+ api_base_url = api_base_url[:-1]
+
+ sha256 = hashlib.sha256(api_base_url.encode()).hexdigest()
+ script_name = f"api_base_url-{sha256}.js"
+ script_path = os.path.join(app.root_path, "../common/web/templates", script_name)
+ with open(script_path, "w") as fout:
+ fout.write("window.CELLXGENE.API.prefix = `" + api_base_url + "${location.pathname}api/`;\n")
+
+ dataset_configs = [app_config.default_dataset_config] + list(app_config.dataroot_config.values())
+ for dataset_config in dataset_configs:
+ inline_scripts = dataset_config.app__inline_scripts
+ inline_scripts.append(script_name)
+
+
class Server:
@staticmethod
def _before_adding_routes(app, app_config):
@@ -337,6 +364,7 @@ class Server:
def __init__(self, app_config):
self.app = Flask(__name__, static_folder=None)
+ handle_api_base_url(self.app, app_config)
self._before_adding_routes(self.app, app_config)
self.app.json_encoder = Float32JSONEncoder
server_config = app_config.server_config
@@ -353,6 +381,12 @@ class Server:
self.app.register_blueprint(webbp)
api_version = "/api/v0.2"
+ api_base_url = server_config.get_api_base_url()
+ api_path = "/"
+ if api_base_url:
+ parse = urlparse(api_base_url)
+ api_path = parse.path
+
if app_config.is_multi_dataset():
# NOTE: These routes only allow the dataset to be in the directory
# of the dataroot, and not a subdirectory. We may want to change
@@ -360,7 +394,8 @@ class Server:
for dataroot_dict in server_config.multi_dataset__dataroot.values():
url_dataroot = dataroot_dict["base_url"]
bp_api = Blueprint(
- f"api_dataset_{url_dataroot}", __name__, url_prefix=f"/{url_dataroot}/" + api_version
+ f"api_dataset_{url_dataroot}", __name__,
+ url_prefix=f"{api_path}/{url_dataroot}/" + api_version
)
resources = get_api_resources(bp_api, url_dataroot)
self.app.register_blueprint(resources.blueprint)
@@ -378,7 +413,7 @@ class Server:
)
else:
- bp_api = Blueprint("api", __name__, url_prefix=api_version)
+ bp_api = Blueprint("api", __name__, url_prefix=f"{api_path}{api_version}")
resources = get_api_resources(bp_api)
self.app.register_blueprint(resources.blueprint)
self.app.add_url_rule(
diff --git a/server/auth/auth_oauth.py b/server/auth/auth_oauth.py
index 428d0048..2b4ddf45 100644
--- a/server/auth/auth_oauth.py
+++ b/server/auth/auth_oauth.py
@@ -1,7 +1,7 @@
from flask import session, request, redirect, current_app, after_this_request, has_request_context, g
from server.auth.auth import AuthTypeClientBase, AuthTypeFactory
from server.common.errors import AuthenticationError, ConfigurationError
-from urllib.parse import urlencode
+from urllib.parse import urlencode, urlparse
import json
import requests
import base64
@@ -45,13 +45,20 @@ class AuthTypeOAuth(AuthTypeClientBase):
if missingimport:
raise ConfigurationError(f"oauth requires these modules: {', '.join(missingimport)}")
self.algorithms = ["RS256"]
- self.api_base_url = server_config.authentication__params_oauth__api_base_url
+ self.oauth_api_base_url = server_config.authentication__params_oauth__oauth_api_base_url
self.client_id = server_config.authentication__params_oauth__client_id
self.client_secret = server_config.authentication__params_oauth__client_secret
- self.callback_base_url = server_config.authentication__params_oauth__callback_base_url
self.session_cookie = server_config.authentication__params_oauth__session_cookie
self.cookie_params = server_config.authentication__params_oauth__cookie
+ self.jwt_decode_options = server_config.authentication__params_oauth__jwt_decode_options
+
self._validate_cookie_params()
+ self._validate_jwt_decode_options()
+
+ self.api_base_url = server_config.get_api_base_url()
+ self.web_base_url = server_config.get_web_base_url()
+ if self.api_base_url is None:
+ raise ConfigurationError("oauth requires the app__api_base_url to be set")
# set the audience
self.audience = self.client_id
@@ -60,11 +67,13 @@ class AuthTypeOAuth(AuthTypeClientBase):
# The JSON Web Key Set (JWKS) is a set of keys which contains the public keys used to verify
# any JSON Web Token (JWT) issued by the authorization server and signed using the RS256
try:
- jwksloc = f"{self.api_base_url}/.well-known/jwks.json"
+ jwksloc = f"{self.oauth_api_base_url}/.well-known/jwks.json"
jwksurl = requests.get(jwksloc)
self.jwks = jwksurl.json()
except Exception:
- raise ConfigurationError(f"error in oauth, api_url_base: {self.api_base_url}, cannot access {jwksloc}")
+ raise ConfigurationError(
+ f"error in oauth, api_url_base: {self.oauth_api_base_url}, cannot access {jwksloc}"
+ )
def _validate_cookie_params(self):
"""check the cookie_params, and raise a ConfigurationError if there is something wrong"""
@@ -81,6 +90,20 @@ class AuthTypeOAuth(AuthTypeClientBase):
if "key" not in keys:
raise ConfigurationError("must have a key (name) in the cookie params")
+ def _validate_jwt_decode_options(self):
+ """check the jwt_decode_options, and raise a ConfigurationError if there is something wrong"""
+ if self.jwt_decode_options is None:
+ self.jwt_decode_options = {}
+ return
+
+ valid_keys = {
+ "verify_signature", "verify_aud", "verify_iat", "verify_exp", "verify_nbf", "verify_iss",
+ "verify_sub", "verify_jti", "verify_at_hash", "leeway"}
+ keys = set(self.jwt_decode_options.keys())
+ unknown = keys - valid_keys
+ if unknown:
+ raise ConfigurationError(f"unexpected key in jwt_decode_options: {', '.join(unknown)}")
+
def is_valid_authentication_type(self):
return True
@@ -88,27 +111,22 @@ class AuthTypeOAuth(AuthTypeClientBase):
return True
def add_url_rules(self, app):
- app.add_url_rule("/login", "login", self.login, methods=["GET"])
- app.add_url_rule("/logout", "logout", self.logout, methods=["GET"])
- app.add_url_rule("/oauth2/callback", "callback", self.callback, methods=["GET"])
+ parse = urlparse(self.api_base_url)
+ app.add_url_rule(f"{parse.path}/login", "login", self.login, methods=["GET"])
+ app.add_url_rule(f"{parse.path}/logout", "logout", self.logout, methods=["GET"])
+ app.add_url_rule(f"{parse.path}/oauth2/callback", "callback", self.callback, methods=["GET"])
def complete_setup(self, flask_app):
self.oauth = OAuth(flask_app)
- if self.callback_base_url is None:
- # In this case, assume the server is running on the same host as the client,
- # and the oauth provider has been configured
- # with a callback that understands a localhost callback (e.g. A http://localhost:5005).
- server_config = flask_app.app_config.server_config
- self.callback_base_url = f"http://{server_config.app__host}:{server_config.app__port}"
self.client = self.oauth.register(
"auth0",
client_id=self.client_id,
client_secret=self.client_secret,
- api_base_url=self.api_base_url,
- refresh_token_url=f"{self.api_base_url}/oauth/token",
- access_token_url=f"{self.api_base_url}/oauth/token",
- authorize_url=f"{self.api_base_url}/authorize",
+ api_base_url=self.oauth_api_base_url,
+ refresh_token_url=f"{self.oauth_api_base_url}/oauth/token",
+ access_token_url=f"{self.oauth_api_base_url}/oauth/token",
+ authorize_url=f"{self.oauth_api_base_url}/authorize",
client_kwargs={"scope": "openid profile email offline_access"},
)
@@ -138,9 +156,9 @@ class AuthTypeOAuth(AuthTypeClientBase):
response.cache_control.update(dict(public=True, max_age=0, no_store=True, no_cache=True, must_revalidate=True))
def login(self):
- callbackurl = f"{self.callback_base_url}/oauth2/callback"
+ callbackurl = f"{self.api_base_url}/oauth2/callback"
return_path = request.args.get("dataset", "")
- return_to = f"{self.callback_base_url}/{return_path}"
+ return_to = f"{self.web_base_url}/{return_path}/"
# save the return path in the session cookie, accessed in the callback function
session["oauth_callback_redirect"] = return_to
response = self.client.authorize_redirect(redirect_uri=callbackurl)
@@ -149,7 +167,7 @@ class AuthTypeOAuth(AuthTypeClientBase):
def logout(self):
self.remove_tokens()
- params = {"returnTo": self.callback_base_url, "client_id": self.client_id}
+ params = {"returnTo": self.web_base_url, "client_id": self.client_id}
response = redirect(self.client.api_base_url + "/v2/logout?" + urlencode(params))
self.update_response(response)
return response
@@ -228,14 +246,14 @@ class AuthTypeOAuth(AuthTypeClientBase):
def get_login_url(self, data_adaptor):
"""Return the url for the login route"""
- if current_app.app_config.is_multi_dataset():
- return f"/login?dataset={data_adaptor.uri_path}/"
+ if data_adaptor and current_app.app_config.is_multi_dataset():
+ return f"{self.api_base_url}/login?dataset={data_adaptor.uri_path}/"
else:
- return "/login"
+ return f"{self.api_base_url}/login"
def get_logout_url(self, data_adaptor):
"""Return the url for the logout route"""
- return "/logout"
+ return f"{self.api_base_url}/logout"
def check_jwt_payload(self, id_token):
try:
@@ -254,18 +272,14 @@ class AuthTypeOAuth(AuthTypeClientBase):
"e": key.get("e"),
}
if rsa_key:
- options = {}
- if not rsa_key["n"] or not rsa_key["e"]:
- # this is a mock auth server, do not validate
- options = {"verify_signature": False, "verify_iss": False}
try:
payload = jwt.decode(
id_token,
rsa_key,
algorithms=self.algorithms,
audience=self.audience,
- issuer=self.api_base_url + "/",
- options=options,
+ issuer=self.oauth_api_base_url + "/",
+ options=self.jwt_decode_options,
)
return payload
@@ -321,7 +335,7 @@ class AuthTypeOAuth(AuthTypeClientBase):
"client_secret": self.client_secret,
}
headers = {"content-type": "application/x-www-form-urlencoded"}
- request = requests.post(f"{self.api_base_url}/oauth/token", urlencode(params), headers=headers)
+ request = requests.post(f"{self.oauth_api_base_url}/oauth/token", urlencode(params), headers=headers)
if request.status_code != 200:
# unable to refresh the token, log the user out
self.remove_tokens()
diff --git a/server/cli/launch.py b/server/cli/launch.py
index d46161d4..d2a95f2a 100644
--- a/server/cli/launch.py
+++ b/server/cli/launch.py
@@ -296,7 +296,7 @@ class CliLaunchServer(Server):
"application/octet-stream",
]
Compress(app)
- if app_config.server_config.app__debug:
+ if app_config.server_config.app__cors_supports_credentials or app_config.server_config.app__debug:
CORS(app, supports_credentials=True)
diff --git a/server/common/app_config.py b/server/common/app_config.py
index ae67eeaf..487d70bb 100644
--- a/server/common/app_config.py
+++ b/server/common/app_config.py
@@ -417,6 +417,7 @@ class ServerConfig(BaseConfig):
dictval_cases = [
("app", "csp_directives"),
("authentication", "params_oauth", "cookie"),
+ ("authentication", "params_oauth", "jwt_decode_options"),
("adaptor", "cxg_adaptor", "tiledb_ctx"),
("multi_dataset", "dataroot"),
]
@@ -434,13 +435,18 @@ class ServerConfig(BaseConfig):
self.app__generate_cache_control_headers = dc["app"]["generate_cache_control_headers"]
self.app__server_timing_headers = dc["app"]["server_timing_headers"]
self.app__csp_directives = dc["app"]["csp_directives"]
+ self.app__cors_supports_credentials = dc["app"]["cors_supports_credentials"]
+ self.app__api_base_url = dc["app"]["api_base_url"]
+ self.app__web_base_url = dc["app"]["web_base_url"]
self.authentication__type = dc["authentication"]["type"]
- self.authentication__params_oauth__api_base_url = dc["authentication"]["params_oauth"]["api_base_url"]
+ self.authentication__params_oauth__oauth_api_base_url = dc["authentication"]["params_oauth"][
+ "oauth_api_base_url"
+ ]
self.authentication__params_oauth__client_id = dc["authentication"]["params_oauth"]["client_id"]
self.authentication__params_oauth__client_secret = dc["authentication"]["params_oauth"]["client_secret"]
- self.authentication__params_oauth__callback_base_url = \
- dc["authentication"]["params_oauth"]["callback_base_url"]
+ self.authentication__params_oauth__jwt_decode_options = dc["authentication"]["params_oauth"][
+ "jwt_decode_options"]
self.authentication__params_oauth__session_cookie = dc["authentication"]["params_oauth"]["session_cookie"]
self.authentication__params_oauth__cookie = dc["authentication"]["params_oauth"]["cookie"]
@@ -500,7 +506,10 @@ class ServerConfig(BaseConfig):
self.check_attr("app__flask_secret_key", (type(None), str))
self.check_attr("app__generate_cache_control_headers", bool)
self.check_attr("app__server_timing_headers", bool)
+ self.check_attr("app__cors_supports_credentials", bool)
self.check_attr("app__csp_directives", (type(None), dict))
+ self.check_attr("app__api_base_url", (type(None), str))
+ self.check_attr("app__web_base_url", (type(None), str))
if self.app__port:
try:
@@ -549,15 +558,18 @@ class ServerConfig(BaseConfig):
elif not isinstance(v, str):
raise ConfigurationError("CSP directive value must be a string or list of strings.")
+ if self.app__web_base_url is None:
+ self.app__web_base_url = self.app__api_base_url
+
def handle_authentication(self, context):
self.check_attr("authentication__type", (type(None), str))
# oauth
ptypes = str if self.authentication__type == "oauth" else (type(None), str)
- self.check_attr("authentication__params_oauth__api_base_url", ptypes)
+ self.check_attr("authentication__params_oauth__oauth_api_base_url", ptypes)
self.check_attr("authentication__params_oauth__client_id", ptypes)
self.check_attr("authentication__params_oauth__client_secret", ptypes)
- self.check_attr("authentication__params_oauth__callback_base_url", (type(None), str))
+ self.check_attr("authentication__params_oauth__jwt_decode_options", (type(None), dict))
self.check_attr("authentication__params_oauth__session_cookie", bool)
if self.authentication__params_oauth__session_cookie:
@@ -743,6 +755,18 @@ class ServerConfig(BaseConfig):
return False
return value > limit_value
+ def get_api_base_url(self):
+ if self.app__api_base_url == "local":
+ return f"http://{self.app__host}:{self.app__port}"
+ return self.app__api_base_url
+
+ def get_web_base_url(self):
+ if self.app__web_base_url == "local":
+ return f"http://{self.app__host}:{self.app__port}"
+ if self.app__web_base_url is None:
+ return self.get_api_base_url()
+ return self.app__web_base_url
+
class DatasetConfig(BaseConfig):
"""Manages the config attribute associated with a dataset."""
@@ -769,7 +793,7 @@ class DatasetConfig(BaseConfig):
self.user_annotations__ontology__obo_location = dc["user_annotations"]["ontology"]["obo_location"]
self.user_annotations__hosted_tiledb_array__db_uri = dc["user_annotations"]["hosted_tiledb_array"]["db_uri"]
self.user_annotations__hosted_tiledb_array__hosted_file_directory = \
- dc["user_annotations"]["hosted_tiledb_array"]["hosted_file_directory"] # noqa E501
+ dc["user_annotations"][ "hosted_tiledb_array" ][ "hosted_file_directory" ] # noqa E501
self.embeddings__names = dc["embeddings"]["names"]
self.embeddings__enable_reembedding = dc["embeddings"]["enable_reembedding"]
diff --git a/server/common/default_config.py b/server/common/default_config.py
index f473dd9f..3b58322f 100644
--- a/server/common/default_config.py
+++ b/server/common/default_config.py
@@ -14,6 +14,25 @@ server:
server_timing_headers: false
csp_directives: null
+ # CORS: Cross Origin Resource Sharing. If true, this allow users to make
+ # authenticated requests. This allows cookies and credentials to be submitted
+ # across domains
+ cors_supports_credentials: false
+
+ # By default, cellxgene will serve api requests from the same base url as the webpage.
+ # In general api_base_url and web_base_url will not need to be set.
+ # There are two reasons to set these parameters:
+ # 1. Oauth authentication is used; the oauth server will redirect back to the api_base_url after login,
+ # which then redirects back to the web_base_url. If the web_base_url is not set, it will default to
+ # the api_base_url. If oauth authentication is used, the api_base_url must be set.
+ # For a local test (where the server runs on "http://localhost:"), then the api_base_url may be
+ # set to the string "local".
+ # 2. The cellxgene deploymnent is in an environment where the webpage and api have
+ # different base urls. In this case both api_base_url and web_base_url must be set.
+ # It is up to the server admin to ensure that the networking is setup correctly for this environment.
+ api_base_url: null
+ web_base_url: null
+
authentication:
# The authentication types may be "none", "session", "oauth"
# none: No authentication support, features like user_annotations must not be enabled.
@@ -22,16 +41,17 @@ server:
type: session
params_oauth:
- # url to the auth server
- api_base_url: null
+ # url to the oauth server
+ oauth_api_base_url: null
# client_id of this app
client_id: null
# the client_secret known to the auth server and this app
client_secret: null
- # cellxgene server location;
- # the browser will be redirected to locations relative to this location during login and logout.
- # A value of None, indicates the client and server are on the localhost. http://localhost: will be used.
- callback_base_url: null
+ # jwt_decode_options, to specify non default decode options define
+ # jwt_decode_options to be a dictionary with key/values described by
+ # the options parameter of the jose.jwt.decode function:
+ # (https://python-jose.readthedocs.io/en/latest/jwt/api.html)
+ jwt_decode_options: null
# if true, the jwt containing the id_token is stored in a session cookie
session_cookie: true
diff --git a/server/test/__init__.py b/server/test/__init__.py
index 12159cc4..8ed6e3e7 100644
--- a/server/test/__init__.py
+++ b/server/test/__init__.py
@@ -131,16 +131,24 @@ def start_test_server(command_line_args=[], app_config=None):
where the server can be accessed within the context, and is terminated when
the context is exited.
- The port is automatically set using find_available_port.
+ The port is automatically set using find_available_port, unless passed in as a command line arg.
The verbose flag is automatically set to True.
If an app_config is provided, then this function writes a temporary
yaml config file, which this server will read and parse.
"""
- start = random.randint(DEFAULT_SERVER_PORT, 2 ** 16 - 1)
- port = int(os.environ.get("CXG_SERVER_PORT", start))
- port = find_available_port("localhost", port)
- command = ["cellxgene", "--no-upgrade-check", "launch", "--verbose", "--port=%d" % port] + command_line_args
+ command = ["cellxgene", "--no-upgrade-check", "launch", "--verbose"]
+ if "-p" in command_line_args:
+ port = int(command_line_args[command_line_args.index("-p") + 1])
+ elif "--port" in command_line_args:
+ port = int(command_line_args[command_line_args.index("--port") + 1])
+ else:
+ start = random.randint(DEFAULT_SERVER_PORT, 2 ** 16 - 1)
+ port = int(os.environ.get("CXG_SERVER_PORT", start))
+ port = find_available_port("localhost", port)
+ command += ["--port=%d" % port]
+
+ command += command_line_args
tempdir = None
if app_config:
diff --git a/server/test/unit/auth/test_oauth.py b/server/test/unit/auth/test_oauth.py
index 413a1084..42734a7b 100644
--- a/server/test/unit/auth/test_oauth.py
+++ b/server/test/unit/auth/test_oauth.py
@@ -19,7 +19,7 @@ from server.test import FIXTURES_ROOT, test_server
# oauth server.
# number of seconds that the oauth token is valid
-TOKEN_EXPIRES = 5
+TOKEN_EXPIRES = 2
# Create a mocked out oauth token, which servers all the endpoints needed by the oauth type.
mock_oauth_app = Flask("mock_oauth_app")
@@ -34,17 +34,19 @@ def authorize():
@mock_oauth_app.route("/oauth/token", methods=["POST"])
def token():
+ now = time.time()
+ expires_at = now + TOKEN_EXPIRES
headers = dict(alg="RS256", kid="fake_kid")
- payload = dict(name="fake_user", sub="fake_id", email="fake_user@email.com", email_verified=True)
+ payload = dict(name="fake_user", sub="fake_id", email="fake_user@email.com", email_verified=True, exp=expires_at)
jwt = jose.jwt.encode(claims=payload, key="mysecret", algorithm="HS256", headers=headers)
r = {
- "access_token": f"access-{time.time()}",
+ "access_token": f"access-{now}",
"id_token": jwt,
- "refresh_token": f"random-{time.time()}",
+ "refresh_token": f"random-{now}",
"scope": "openid profile email",
"expires_in": TOKEN_EXPIRES,
"token_type": "Bearer",
- "expires_at": time.time() + TOKEN_EXPIRES,
+ "expires_at": expires_at
}
return make_response(jsonify(r))
@@ -81,6 +83,19 @@ class AuthTest(unittest.TestCase):
def auth_flow(self, app_config, cookie_key=None):
+ app_config.update_server_config(
+ app__api_base_url="local",
+ authentication__type="oauth",
+ authentication__params_oauth__oauth_api_base_url=f"http://localhost:{PORT}",
+ authentication__params_oauth__client_id="mock_client_id",
+ authentication__params_oauth__client_secret="mock_client_secret",
+ authentication__params_oauth__jwt_decode_options={
+ "verify_signature": False, "verify_iss": False
+ })
+
+ app_config.update_server_config(multi_dataset__dataroot=self.dataset_dataroot)
+ app_config.complete_config()
+
with test_server(app_config=app_config) as server:
session = requests.Session()
@@ -96,10 +111,10 @@ class AuthTest(unittest.TestCase):
login_uri = config["config"]["authentication"]["login"]
logout_uri = config["config"]["authentication"]["logout"]
- self.assertEqual(login_uri, "/login?dataset=d/pbmc3k.cxg/")
- self.assertEqual(logout_uri, "/logout")
+ self.assertEqual(login_uri, f"{server}/login?dataset=d/pbmc3k.cxg/")
+ self.assertEqual(logout_uri, f"{server}/logout")
- r = session.get(f"{server}/{login_uri}")
+ r = session.get(login_uri)
# check that the login redirect worked
self.assertEqual(r.history[0].status_code, 302)
self.assertEqual(r.url, f"{server}/d/pbmc3k.cxg/")
@@ -113,13 +128,13 @@ class AuthTest(unittest.TestCase):
cookie = session.cookies.get(cookie_key)
token = json.loads(base64.b64decode(cookie))
access_token_before = token.get("access_token")
- expires_at_before = token.get("expires_at")
+ id_token_before = token.get("id_token")
# let the token expire
time.sleep(TOKEN_EXPIRES + 1)
# check that refresh works
- session.get(f"{server}/{login_uri}")
+ session.get(login_uri)
userinfo = session.get(f"{server}/d/pbmc3k.cxg/api/v0.2/userinfo").json()
self.assertTrue(userinfo["userinfo"]["is_authenticated"])
self.assertEqual(userinfo["userinfo"]["username"], "fake_user")
@@ -127,12 +142,12 @@ class AuthTest(unittest.TestCase):
cookie = session.cookies.get(cookie_key)
token = json.loads(base64.b64decode(cookie))
access_token_after = token.get("access_token")
- expires_at_after = token.get("expires_at")
+ id_token_after = token.get("id_token")
self.assertNotEqual(access_token_before, access_token_after)
- self.assertTrue(expires_at_after - expires_at_before > TOKEN_EXPIRES)
+ self.assertNotEqual(id_token_before, id_token_after)
- r = session.get(f"{server}/{logout_uri}")
+ r = session.get(logout_uri)
# check that the logout redirect worked
self.assertEqual(r.history[0].status_code, 302)
self.assertEqual(r.url, f"{server}")
@@ -146,31 +161,16 @@ class AuthTest(unittest.TestCase):
# test with session cookies
app_config = AppConfig()
app_config.update_server_config(
- authentication__type="oauth",
- authentication__params_oauth__api_base_url=f"http://localhost:{PORT}",
- authentication__params_oauth__client_id="mock_client_id",
- authentication__params_oauth__client_secret="mock_client_secret",
authentication__params_oauth__session_cookie=True,
)
-
- app_config.update_server_config(multi_dataset__dataroot=self.dataset_dataroot)
- app_config.complete_config()
-
self.auth_flow(app_config)
def test_auth_oauth_cookie(self):
# test with specified cookie
app_config = AppConfig()
app_config.update_server_config(
- authentication__type="oauth",
- authentication__params_oauth__api_base_url=f"http://localhost:{PORT}",
- authentication__params_oauth__client_id="mock_client_id",
- authentication__params_oauth__client_secret="mock_client_secret",
authentication__params_oauth__session_cookie=False,
authentication__params_oauth__cookie=dict(key="test_cxguser", httponly=True, max_age=60),
)
- app_config.update_server_config(multi_dataset__dataroot=self.dataset_dataroot)
- app_config.complete_config()
-
self.auth_flow(app_config, "test_cxguser")
diff --git a/server/test/unit/common/test_app_config.py b/server/test/unit/common/test_app_config.py
index a299bf51..543fe6da 100644
--- a/server/test/unit/common/test_app_config.py
+++ b/server/test/unit/common/test_app_config.py
@@ -7,6 +7,7 @@ import requests
from server.common.app_config import AppConfig
from server.common.errors import ConfigurationError
+from server.common.utils.utils import find_available_port
from server.test import PROJECT_ROOT, test_server, FIXTURES_ROOT
@@ -19,46 +20,46 @@ def mockenv(**envvars):
class AppConfigTest(unittest.TestCase):
def test_update(self):
- c = AppConfig()
- c.update_server_config(app__verbose=True, multi_dataset__dataroot="datadir")
- v = c.server_config.changes_from_default()
- self.assertCountEqual(v, [("app__verbose", True, False), ("multi_dataset__dataroot", "datadir", None)])
+ config = AppConfig()
+ config.update_server_config(app__verbose=True, multi_dataset__dataroot="datadir")
+ vars = config.server_config.changes_from_default()
+ self.assertCountEqual(vars, [("app__verbose", True, False), ("multi_dataset__dataroot", "datadir", None)])
- c = AppConfig()
- c.update_default_dataset_config(app__scripts=(), app__inline_scripts=())
- v = c.server_config.changes_from_default()
- self.assertCountEqual(v, [])
+ config = AppConfig()
+ config.update_default_dataset_config(app__scripts=(), app__inline_scripts=())
+ vars = config.server_config.changes_from_default()
+ self.assertCountEqual(vars, [])
- c = AppConfig()
- c.update_default_dataset_config(app__scripts=[], app__inline_scripts=[])
- v = c.default_dataset_config.changes_from_default()
- self.assertCountEqual(v, [])
+ config = AppConfig()
+ config.update_default_dataset_config(app__scripts=[], app__inline_scripts=[])
+ vars = config.default_dataset_config.changes_from_default()
+ self.assertCountEqual(vars, [])
- c = AppConfig()
- c.update_default_dataset_config(app__scripts=("a", "b"), app__inline_scripts=["c", "d"])
- v = c.default_dataset_config.changes_from_default()
- self.assertCountEqual(v, [("app__scripts", ["a", "b"], []), ("app__inline_scripts", ["c", "d"], [])])
+ config = AppConfig()
+ config.update_default_dataset_config(app__scripts=("a", "b"), app__inline_scripts=["c", "d"])
+ vars = config.default_dataset_config.changes_from_default()
+ self.assertCountEqual(vars, [("app__scripts", ["a", "b"], []), ("app__inline_scripts", ["c", "d"], [])])
def test_multi_dataset(self):
- c = AppConfig()
+ config = AppConfig()
# test for illegal url_dataroots
for illegal in ("../b", "!$*", "\\n", "", "(bad)"):
- c.update_server_config(
+ config.update_server_config(
multi_dataset__dataroot={"tag": {"base_url": illegal, "dataroot": "{PROJECT_ROOT}/example-dataset"}}
)
with self.assertRaises(ConfigurationError):
- c.complete_config()
+ config.complete_config()
# test for legal url_dataroots
for legal in ("d", "this.is-okay_", "a/b"):
- c.update_server_config(
+ config.update_server_config(
multi_dataset__dataroot={"tag": {"base_url": legal, "dataroot": "{PROJECT_ROOT}/example-dataset"}}
)
- c.complete_config()
+ config.complete_config()
# test that multi dataroots work end to end
- c.update_server_config(
+ config.update_server_config(
multi_dataset__dataroot=dict(
s1=dict(dataroot=f"{PROJECT_ROOT}/example-dataset", base_url="set1/1/2"),
s2=dict(dataroot=f"{FIXTURES_ROOT}", base_url="set2"),
@@ -67,46 +68,46 @@ class AppConfigTest(unittest.TestCase):
)
# Change this default to test if the dataroot overrides below work.
- c.update_default_dataset_config(app__about_legal_tos="tos_default.html")
+ config.update_default_dataset_config(app__about_legal_tos="tos_default.html")
# specialize the configs for set1
- c.add_dataroot_config(
+ config.add_dataroot_config(
"s1", user_annotations__enable=False, diffexp__enable=True, app__about_legal_tos="tos_set1.html"
)
# specialize the configs for set2
- c.add_dataroot_config(
+ config.add_dataroot_config(
"s2", user_annotations__enable=True, diffexp__enable=False, app__about_legal_tos="tos_set2.html"
)
# no specializations for set3 (they get the default dataset config)
- c.complete_config()
+ config.complete_config()
- with test_server(app_config=c) as server:
+ with test_server(app_config=config) as server:
session = requests.Session()
- r = session.get(f"{server}/set1/1/2/pbmc3k.h5ad/api/v0.2/config")
- data_config = r.json()
+ response = session.get(f"{server}/set1/1/2/pbmc3k.h5ad/api/v0.2/config")
+ data_config = response.json()
assert data_config["config"]["displayNames"]["dataset"] == "pbmc3k"
assert data_config["config"]["parameters"]["annotations"] is False
assert data_config["config"]["parameters"]["disable-diffexp"] is False
assert data_config["config"]["parameters"]["about_legal_tos"] == "tos_set1.html"
- r = session.get(f"{server}/set2/pbmc3k.cxg/api/v0.2/config")
- data_config = r.json()
+ response = session.get(f"{server}/set2/pbmc3k.cxg/api/v0.2/config")
+ data_config = response.json()
assert data_config["config"]["displayNames"]["dataset"] == "pbmc3k"
assert data_config["config"]["parameters"]["annotations"] is True
assert data_config["config"]["parameters"]["about_legal_tos"] == "tos_set2.html"
- r = session.get(f"{server}/set3/pbmc3k.cxg/api/v0.2/config")
- data_config = r.json()
+ response = session.get(f"{server}/set3/pbmc3k.cxg/api/v0.2/config")
+ data_config = response.json()
assert data_config["config"]["displayNames"]["dataset"] == "pbmc3k"
assert data_config["config"]["parameters"]["annotations"] is True
assert data_config["config"]["parameters"]["disable-diffexp"] is False
assert data_config["config"]["parameters"]["about_legal_tos"] == "tos_default.html"
- r = session.get(f"{server}/health")
- assert r.json()["status"] == "pass"
+ response = session.get(f"{server}/health")
+ assert response.json()["status"] == "pass"
@mockenv(CXG_AWS_SECRET_NAME="TESTING", CXG_AWS_SECRET_REGION_NAME="TEST_REGION")
@patch('server.common.aws_secret_utils.get_secret_key')
@@ -133,3 +134,23 @@ class AppConfigTest(unittest.TestCase):
self.assertEqual(config.server_config.app__flask_secret_key, "mock_flask_secret")
self.assertEqual(config.server_config.authentication__params_oauth__client_secret, "mock_oauth_secret")
self.assertEqual(config.default_dataset_config.user_annotations__hosted_tiledb_array__db_uri, "mock_db_uri")
+
+ def test_api_base_url(self):
+
+ # test the api_base_url feature, and that it can contain a path
+ config = AppConfig()
+ backend_port = find_available_port("localhost", 10000)
+ config.update_server_config(
+ app__api_base_url=f"http://localhost:{backend_port}/additional/path/before/dataroot",
+ multi_dataset__dataroot=f"{PROJECT_ROOT}/example-dataset"
+ )
+
+ config.complete_config()
+
+ with test_server(["-p", str(backend_port)], app_config=config) as server:
+ session = requests.Session()
+ self.assertEqual(server, f"http://localhost:{backend_port}")
+ response = session.get(f"{server}/additional/path/before/dataroot/d/pbmc3k.h5ad/api/v0.2/config")
+ self.assertEqual(response.status_code, 200)
+ data_config = response.json()
+ self.assertEqual(data_config["config"]["displayNames"]["dataset"], "pbmc3k")