mirror of
https://github.com/chanzuckerberg/cellxgene.git
synced 2026-10-06 06:08:12 +08:00
Handle the refresh token in oauth authentication (#1766)
* Handle the refresh token in oauth authentication
If the token has expired, then it can be refreshed to get a new token.
This is automatically handled by the server without the client being aware.
Also in the PR:
- refactor the auth_oauth.py file to more simply handle the save/restore of the token,
and the refresh token
- added an end2end test for oauth, which also tests refresh.
* adding python-jose and Authlib to requirements-dev.txt
They are needed in the auth_oauth test
This commit is contained in:
+174
-79
@@ -1,9 +1,10 @@
|
|||||||
from flask import session, request, redirect, current_app, has_request_context, g
|
from flask import session, request, redirect, current_app, after_this_request, has_request_context, g
|
||||||
from server.auth.auth import AuthTypeClientBase, AuthTypeFactory
|
from server.auth.auth import AuthTypeClientBase, AuthTypeFactory
|
||||||
from server.common.errors import AuthenticationError, ConfigurationError
|
from server.common.errors import AuthenticationError, ConfigurationError
|
||||||
from urllib.parse import urlencode
|
from urllib.parse import urlencode
|
||||||
from urllib.request import urlopen
|
|
||||||
import json
|
import json
|
||||||
|
import requests
|
||||||
|
import base64
|
||||||
|
|
||||||
# It is not required to have authlib or jose.
|
# It is not required to have authlib or jose.
|
||||||
# However, it is a configuration error to use this auth type if they are not installed.
|
# However, it is a configuration error to use this auth type if they are not installed.
|
||||||
@@ -20,10 +21,22 @@ except ModuleNotFoundError:
|
|||||||
missingimport.append("jose")
|
missingimport.append("jose")
|
||||||
|
|
||||||
|
|
||||||
|
class Tokens:
|
||||||
|
"""Simple class to represent the tokens that are saved/restored from the cookie"""
|
||||||
|
|
||||||
|
def __init__(self, access_token, id_token, refresh_token, expires_at):
|
||||||
|
self.access_token = access_token
|
||||||
|
self.id_token = id_token
|
||||||
|
self.refresh_token = refresh_token
|
||||||
|
self.expires_at = expires_at
|
||||||
|
if not (access_token and id_token and refresh_token and expires_at):
|
||||||
|
raise KeyError(str(self.__dict__))
|
||||||
|
|
||||||
|
|
||||||
class AuthTypeOAuth(AuthTypeClientBase):
|
class AuthTypeOAuth(AuthTypeClientBase):
|
||||||
"""An authentication type for oauth2 logins."""
|
"""An authentication type for oauth2 logins."""
|
||||||
|
|
||||||
CXG_ID_TOKEN = "id_token"
|
CXG_TOKENS = "auth_tokens"
|
||||||
|
|
||||||
def __init__(self, server_config):
|
def __init__(self, server_config):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@@ -46,8 +59,8 @@ class AuthTypeOAuth(AuthTypeClientBase):
|
|||||||
# any JSON Web Token (JWT) issued by the authorization server and signed using the RS256
|
# any JSON Web Token (JWT) issued by the authorization server and signed using the RS256
|
||||||
try:
|
try:
|
||||||
jwksloc = f"{self.api_base_url}/.well-known/jwks.json"
|
jwksloc = f"{self.api_base_url}/.well-known/jwks.json"
|
||||||
jwksurl = urlopen(jwksloc)
|
jwksurl = requests.get(jwksloc)
|
||||||
self.jwks = json.loads(jwksurl.read())
|
self.jwks = jwksurl.json()
|
||||||
except Exception:
|
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.api_base_url}, cannot access {jwksloc}")
|
||||||
|
|
||||||
@@ -87,48 +100,43 @@ class AuthTypeOAuth(AuthTypeClientBase):
|
|||||||
self.callback_base_url = f"http://{server_config.app__host}:{server_config.app__port}"
|
self.callback_base_url = f"http://{server_config.app__host}:{server_config.app__port}"
|
||||||
|
|
||||||
self.client = self.oauth.register(
|
self.client = self.oauth.register(
|
||||||
"oauth",
|
"auth0",
|
||||||
client_id=self.client_id,
|
client_id=self.client_id,
|
||||||
client_secret=self.client_secret,
|
client_secret=self.client_secret,
|
||||||
api_base_url=self.api_base_url,
|
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",
|
access_token_url=f"{self.api_base_url}/oauth/token",
|
||||||
authorize_url=f"{self.api_base_url}/authorize",
|
authorize_url=f"{self.api_base_url}/authorize",
|
||||||
client_kwargs={
|
client_kwargs={"scope": "openid profile email offline_access"},
|
||||||
"scope" : "openid profile email",
|
|
||||||
}
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def is_user_authenticated(self):
|
def is_user_authenticated(self):
|
||||||
try:
|
payload = self.get_userinfo()
|
||||||
payload = self.get_jwt_payload()
|
return payload is not None
|
||||||
return payload is not None
|
|
||||||
except AuthenticationError:
|
|
||||||
return False
|
|
||||||
|
|
||||||
def get_user_id(self):
|
def get_user_id(self):
|
||||||
payload = self.get_jwt_payload()
|
payload = self.get_userinfo()
|
||||||
if payload and payload.get("sub"):
|
if payload and payload.get("sub"):
|
||||||
return payload.get("sub")
|
return payload.get("sub")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def get_user_name(self):
|
def get_user_name(self):
|
||||||
payload = self.get_jwt_payload()
|
payload = self.get_userinfo()
|
||||||
if payload and payload.get("name"):
|
if payload and payload.get("name"):
|
||||||
return payload.get("name")
|
return payload.get("name")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def get_user_email(self):
|
def get_user_email(self):
|
||||||
payload = self.get_jwt_payload()
|
payload = self.get_userinfo()
|
||||||
if payload and payload.get("email"):
|
if payload and payload.get("email"):
|
||||||
return payload.get("email")
|
return payload.get("email")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def update_response(self, response):
|
def update_response(self, response):
|
||||||
response.cache_control.update(
|
response.cache_control.update(dict(public=True, max_age=0, no_store=True, no_cache=True, must_revalidate=True))
|
||||||
dict(public=True, max_age=0, no_store=True, no_cache=True, must_revalidate=True))
|
|
||||||
|
|
||||||
def login(self):
|
def login(self):
|
||||||
callbackurl = f'{self.callback_base_url}/oauth2/callback'
|
callbackurl = f"{self.callback_base_url}/oauth2/callback"
|
||||||
return_path = request.args.get("dataset", "")
|
return_path = request.args.get("dataset", "")
|
||||||
return_to = f"{self.callback_base_url}/{return_path}"
|
return_to = f"{self.callback_base_url}/{return_path}"
|
||||||
# save the return path in the session cookie, accessed in the callback function
|
# save the return path in the session cookie, accessed in the callback function
|
||||||
@@ -138,40 +146,83 @@ class AuthTypeOAuth(AuthTypeClientBase):
|
|||||||
return response
|
return response
|
||||||
|
|
||||||
def logout(self):
|
def logout(self):
|
||||||
params = {'returnTo' : self.callback_base_url, 'client_id' : self.client_id}
|
self.remove_tokens()
|
||||||
response = redirect(self.client.api_base_url + '/v2/logout?' + urlencode(params))
|
params = {"returnTo": self.callback_base_url, "client_id": self.client_id}
|
||||||
|
response = redirect(self.client.api_base_url + "/v2/logout?" + urlencode(params))
|
||||||
if self.session_cookie:
|
|
||||||
if self.CXG_ID_TOKEN in session:
|
|
||||||
del session[self.CXG_ID_TOKEN]
|
|
||||||
else:
|
|
||||||
response.set_cookie(self.cookie_params["key"], "", expires=0)
|
|
||||||
|
|
||||||
self.update_response(response)
|
self.update_response(response)
|
||||||
return response
|
return response
|
||||||
|
|
||||||
def callback(self):
|
def callback(self):
|
||||||
token = self.client.authorize_access_token()
|
data = self.client.authorize_access_token()
|
||||||
id_token = token.get("id_token")
|
tokens = Tokens(
|
||||||
|
access_token=data.get("access_token"),
|
||||||
|
id_token=data.get("id_token"),
|
||||||
|
refresh_token=data.get("refresh_token"),
|
||||||
|
expires_at=data.get("expires_at"),
|
||||||
|
)
|
||||||
|
self.save_tokens(tokens)
|
||||||
oauth_callback_redirect = session.pop("oauth_callback_redirect", "/")
|
oauth_callback_redirect = session.pop("oauth_callback_redirect", "/")
|
||||||
resp = redirect(oauth_callback_redirect)
|
response = redirect(oauth_callback_redirect)
|
||||||
|
self.update_response(response)
|
||||||
|
return response
|
||||||
|
|
||||||
|
def get_tokens(self):
|
||||||
|
"""Extract the tokens from the cookie, and store them in the flask global context"""
|
||||||
|
if "tokens" in g:
|
||||||
|
return g.tokens
|
||||||
|
|
||||||
|
try:
|
||||||
|
if self.session_cookie:
|
||||||
|
tokensdict = session.get(self.CXG_TOKENS)
|
||||||
|
if tokensdict:
|
||||||
|
g.tokens = Tokens(**tokensdict)
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
else:
|
||||||
|
value = request.cookies.get(self.cookie_params["key"])
|
||||||
|
value = base64.b64decode(value)
|
||||||
|
try:
|
||||||
|
tokensdict = json.loads(value)
|
||||||
|
g.tokens = Tokens(**tokensdict)
|
||||||
|
except (TypeError, KeyError, json.decoder.JSONDecodeError):
|
||||||
|
g.pop("tokens", None)
|
||||||
|
return None
|
||||||
|
|
||||||
|
except (TypeError, KeyError):
|
||||||
|
g.pop("tokens", None)
|
||||||
|
return None
|
||||||
|
|
||||||
|
return g.tokens
|
||||||
|
|
||||||
|
def save_tokens(self, tokens):
|
||||||
|
g.tokens = tokens
|
||||||
if self.session_cookie:
|
if self.session_cookie:
|
||||||
session[self.CXG_ID_TOKEN] = id_token
|
session[self.CXG_TOKENS] = tokens.__dict__
|
||||||
else:
|
else:
|
||||||
args = self.cookie_params.copy()
|
|
||||||
del args["key"]
|
|
||||||
try:
|
|
||||||
resp.set_cookie(
|
|
||||||
self.cookie_params["key"],
|
|
||||||
id_token,
|
|
||||||
**args)
|
|
||||||
g.token = id_token
|
|
||||||
except Exception as e:
|
|
||||||
raise AuthenticationError(f"unable to set_cookie {self.cookie_params}") from e
|
|
||||||
|
|
||||||
self.update_response(resp)
|
@after_this_request
|
||||||
return resp
|
def set_cookie(response):
|
||||||
|
args = self.cookie_params.copy()
|
||||||
|
value = base64.b64encode(json.dumps(tokens.__dict__).encode("utf-8"))
|
||||||
|
del args["key"]
|
||||||
|
try:
|
||||||
|
response.set_cookie(self.cookie_params["key"], value, **args)
|
||||||
|
except Exception as e:
|
||||||
|
raise AuthenticationError(f"unable to set_cookie {self.cookie_params}") from e
|
||||||
|
return response
|
||||||
|
|
||||||
|
def remove_tokens(self):
|
||||||
|
g.pop("tokens", None)
|
||||||
|
if self.session_cookie:
|
||||||
|
if self.CXG_TOKENS in session:
|
||||||
|
del session[self.CXG_TOKENS]
|
||||||
|
else:
|
||||||
|
|
||||||
|
@after_this_request
|
||||||
|
def remove_cookie(response):
|
||||||
|
response.set_cookie(self.cookie_params["key"], "", expires=0)
|
||||||
|
self.update_response(response)
|
||||||
|
return response
|
||||||
|
|
||||||
def get_login_url(self, data_adaptor):
|
def get_login_url(self, data_adaptor):
|
||||||
"""Return the url for the login route"""
|
"""Return the url for the login route"""
|
||||||
@@ -184,60 +235,104 @@ class AuthTypeOAuth(AuthTypeClientBase):
|
|||||||
"""Return the url for the logout route"""
|
"""Return the url for the logout route"""
|
||||||
return "/logout"
|
return "/logout"
|
||||||
|
|
||||||
def get_token(self):
|
def check_jwt_payload(self, id_token):
|
||||||
"""Function to return the token"""
|
|
||||||
if "token" in g:
|
|
||||||
return g.token
|
|
||||||
if self.session_cookie:
|
|
||||||
g.token = session.get(self.CXG_ID_TOKEN)
|
|
||||||
else:
|
|
||||||
g.token = request.cookies.get(self.cookie_params["key"])
|
|
||||||
|
|
||||||
return g.token
|
|
||||||
|
|
||||||
def get_jwt_payload(self):
|
|
||||||
if not has_request_context():
|
|
||||||
return None
|
|
||||||
|
|
||||||
token = self.get_token()
|
|
||||||
if token is None:
|
|
||||||
return None
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
unverified_header = jwt.get_unverified_header(token)
|
unverified_header = jwt.get_unverified_header(id_token)
|
||||||
except JWTError:
|
except JWTError:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
rsa_key = {}
|
rsa_key = {}
|
||||||
for key in self.jwks['keys']:
|
for key in self.jwks["keys"]:
|
||||||
if key['kid'] == unverified_header['kid']:
|
if key["kid"] == unverified_header["kid"]:
|
||||||
rsa_key = {
|
rsa_key = {
|
||||||
'kty': key['kty'],
|
"kty": key["kty"],
|
||||||
'kid': key['kid'],
|
"kid": key["kid"],
|
||||||
'use': key['use'],
|
"use": key["use"],
|
||||||
'n': key['n'],
|
"n": key.get("n"),
|
||||||
'e': key['e']
|
"e": key.get("e"),
|
||||||
}
|
}
|
||||||
if rsa_key:
|
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:
|
try:
|
||||||
payload = jwt.decode(
|
payload = jwt.decode(
|
||||||
token,
|
id_token,
|
||||||
rsa_key,
|
rsa_key,
|
||||||
algorithms=self.algorithms,
|
algorithms=self.algorithms,
|
||||||
audience=self.audience,
|
audience=self.audience,
|
||||||
issuer=self.api_base_url + "/"
|
issuer=self.api_base_url + "/",
|
||||||
|
options=options,
|
||||||
)
|
)
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
except JWTError as e:
|
|
||||||
raise AuthenticationError(f"invalid signature: {str(e)}")
|
|
||||||
except ExpiredSignatureError:
|
except ExpiredSignatureError:
|
||||||
# TODO, handle expired sessions by refreshing the token
|
# This exception is handled in get_userinfo
|
||||||
return None
|
raise
|
||||||
except JWTClaimsError as e:
|
except JWTClaimsError as e:
|
||||||
raise AuthenticationError(f"invalid claims {str(e)}")
|
raise AuthenticationError(f"invalid claims {str(e)}") from e
|
||||||
|
except JWTError as e:
|
||||||
|
raise AuthenticationError(f"invalid signature: {str(e)}") from e
|
||||||
|
|
||||||
raise AuthenticationError("Unable to find the appropriate key")
|
raise AuthenticationError("Unable to find the appropriate key")
|
||||||
|
|
||||||
|
def get_userinfo(self):
|
||||||
|
if not has_request_context():
|
||||||
|
return None
|
||||||
|
|
||||||
|
# check if the userinfo has been retrieved already in this request
|
||||||
|
if "userinfo" in g:
|
||||||
|
return g.get("userinfo")
|
||||||
|
|
||||||
|
# if there is no id_token, return None (user is not authenticated)
|
||||||
|
tokens = self.get_tokens()
|
||||||
|
if tokens is None or tokens.id_token is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
# check the jwt payload. This raises an AuthenticationError if the token is not valid.
|
||||||
|
# It the token has expired, we attempt to refresh the token
|
||||||
|
g.userinfo = self.check_jwt_payload(tokens.id_token)
|
||||||
|
return g.userinfo
|
||||||
|
|
||||||
|
except ExpiredSignatureError:
|
||||||
|
tokens = self.refresh_expired_token(tokens.refresh_token)
|
||||||
|
if tokens is None or tokens.id_token is None:
|
||||||
|
return None
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
g.userinfo = self.check_jwt_payload(tokens.id_token)
|
||||||
|
return g.userinfo
|
||||||
|
except JWTError as e:
|
||||||
|
raise AuthenticationError(f"error during token refresh: {str(e)}") from e
|
||||||
|
|
||||||
|
except AuthenticationError:
|
||||||
|
self.remove_tokens()
|
||||||
|
raise
|
||||||
|
|
||||||
|
def refresh_expired_token(self, refresh_token):
|
||||||
|
params = {
|
||||||
|
"grant_type": "refresh_token",
|
||||||
|
"client_id": self.client_id,
|
||||||
|
"refresh_token": refresh_token,
|
||||||
|
"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)
|
||||||
|
if request.status_code != 200:
|
||||||
|
# unable to refresh the token, log the user out
|
||||||
|
self.remove_tokens()
|
||||||
|
return None
|
||||||
|
data = request.json()
|
||||||
|
tokens = Tokens(
|
||||||
|
access_token=data.get("access_token"),
|
||||||
|
id_token=data.get("id_token"),
|
||||||
|
refresh_token=data.get("refresh_token", refresh_token),
|
||||||
|
expires_at=data.get("expires_at"),
|
||||||
|
)
|
||||||
|
self.save_tokens(tokens)
|
||||||
|
return tokens
|
||||||
|
|
||||||
|
|
||||||
AuthTypeFactory.register("oauth", AuthTypeOAuth)
|
AuthTypeFactory.register("oauth", AuthTypeOAuth)
|
||||||
|
|||||||
@@ -1,9 +1,11 @@
|
|||||||
|
Authlib>=0.14.3
|
||||||
black
|
black
|
||||||
bumpversion>=0.5
|
bumpversion>=0.5
|
||||||
parameterized>=0.7.0
|
|
||||||
pytest>=3.6.3
|
|
||||||
twine>=1.12.1
|
|
||||||
codecov>=2.0.15
|
codecov>=2.0.15
|
||||||
scanpy>=1.4.6
|
parameterized>=0.7.0
|
||||||
psycopg2==2.7.7
|
psycopg2==2.7.7
|
||||||
|
pytest>=3.6.3
|
||||||
|
python-jose>=3.2.0
|
||||||
|
scanpy>=1.4.6
|
||||||
|
twine>=1.12.1
|
||||||
-r requirements.txt
|
-r requirements.txt
|
||||||
|
|||||||
@@ -0,0 +1,176 @@
|
|||||||
|
import unittest
|
||||||
|
import random
|
||||||
|
import time
|
||||||
|
import base64
|
||||||
|
import json
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from flask import Flask, jsonify, make_response, request, redirect
|
||||||
|
from multiprocessing import Process
|
||||||
|
|
||||||
|
import jose
|
||||||
|
from server.common.app_config import AppConfig
|
||||||
|
from server.test import FIXTURES_ROOT, test_server
|
||||||
|
|
||||||
|
# This tests the oauth authentication type.
|
||||||
|
# This test starts a cellxgene server and a mock oauth server.
|
||||||
|
# API requests to login and logout and get the userinfo are made
|
||||||
|
# to the cellxgene server, which then sends requests to the mock
|
||||||
|
# oauth server.
|
||||||
|
|
||||||
|
# number of seconds that the oauth token is valid
|
||||||
|
TOKEN_EXPIRES = 5
|
||||||
|
|
||||||
|
# Create a mocked out oauth token, which servers all the endpoints needed by the oauth type.
|
||||||
|
mock_oauth_app = Flask("mock_oauth_app")
|
||||||
|
|
||||||
|
|
||||||
|
@mock_oauth_app.route("/authorize")
|
||||||
|
def authorize():
|
||||||
|
callback = request.args.get("redirect_uri")
|
||||||
|
state = request.args.get("state")
|
||||||
|
return redirect(callback + f"?code=fakecode&state={state}")
|
||||||
|
|
||||||
|
|
||||||
|
@mock_oauth_app.route("/oauth/token", methods=["POST"])
|
||||||
|
def token():
|
||||||
|
headers = dict(alg="RS256", kid="fake_kid")
|
||||||
|
payload = dict(name="fake_user", sub="fake_id", email="fake_user@email.com", email_verified=True)
|
||||||
|
jwt = jose.jwt.encode(claims=payload, key="mysecret", algorithm="HS256", headers=headers)
|
||||||
|
r = {
|
||||||
|
"access_token": f"access-{time.time()}",
|
||||||
|
"id_token": jwt,
|
||||||
|
"refresh_token": f"random-{time.time()}",
|
||||||
|
"scope": "openid profile email",
|
||||||
|
"expires_in": TOKEN_EXPIRES,
|
||||||
|
"token_type": "Bearer",
|
||||||
|
"expires_at": time.time() + TOKEN_EXPIRES,
|
||||||
|
}
|
||||||
|
return make_response(jsonify(r))
|
||||||
|
|
||||||
|
|
||||||
|
@mock_oauth_app.route("/v2/logout")
|
||||||
|
def logout():
|
||||||
|
return_to = request.args.get("returnTo")
|
||||||
|
return redirect(return_to)
|
||||||
|
|
||||||
|
|
||||||
|
@mock_oauth_app.route("/.well-known/jwks.json")
|
||||||
|
def jwks():
|
||||||
|
data = dict(alg="RS256", kty="RSA", use="sig", kid="fake_kid",)
|
||||||
|
return make_response(jsonify(dict(keys=[data])))
|
||||||
|
|
||||||
|
|
||||||
|
# The port that the mock oauth server will listen on
|
||||||
|
PORT = random.randint(10000, 12000)
|
||||||
|
|
||||||
|
|
||||||
|
# function to launch the mock oauth server
|
||||||
|
def launch_mock_oauth():
|
||||||
|
mock_oauth_app.run(port=PORT)
|
||||||
|
|
||||||
|
|
||||||
|
class AuthTest(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.dataset_dataroot = FIXTURES_ROOT
|
||||||
|
self.mock_oauth_process = Process(target=launch_mock_oauth)
|
||||||
|
self.mock_oauth_process.start()
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
self.mock_oauth_process.terminate()
|
||||||
|
|
||||||
|
def auth_flow(self, app_config, cookie_key=None):
|
||||||
|
|
||||||
|
with test_server(app_config=app_config) as server:
|
||||||
|
session = requests.Session()
|
||||||
|
|
||||||
|
# auth datasets
|
||||||
|
config = session.get(f"{server}/d/pbmc3k.cxg/api/v0.2/config").json()
|
||||||
|
userinfo = session.get(f"{server}/d/pbmc3k.cxg/api/v0.2/userinfo").json()
|
||||||
|
|
||||||
|
self.assertFalse(userinfo["userinfo"]["is_authenticated"])
|
||||||
|
self.assertIsNone(userinfo["userinfo"]["username"])
|
||||||
|
self.assertTrue(config["config"]["authentication"]["requires_client_login"])
|
||||||
|
self.assertTrue(config["config"]["parameters"]["annotations"])
|
||||||
|
|
||||||
|
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")
|
||||||
|
|
||||||
|
r = session.get(f"{server}/{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/")
|
||||||
|
config = session.get(f"{server}/d/pbmc3k.cxg/api/v0.2/config").json()
|
||||||
|
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")
|
||||||
|
self.assertTrue(config["config"]["parameters"]["annotations"])
|
||||||
|
|
||||||
|
if cookie_key:
|
||||||
|
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")
|
||||||
|
|
||||||
|
# let the token expire
|
||||||
|
time.sleep(TOKEN_EXPIRES + 1)
|
||||||
|
|
||||||
|
# check that refresh works
|
||||||
|
session.get(f"{server}/{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")
|
||||||
|
|
||||||
|
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")
|
||||||
|
|
||||||
|
self.assertNotEqual(access_token_before, access_token_after)
|
||||||
|
self.assertTrue(expires_at_after - expires_at_before > TOKEN_EXPIRES)
|
||||||
|
|
||||||
|
r = session.get(f"{server}/{logout_uri}")
|
||||||
|
# check that the logout redirect worked
|
||||||
|
self.assertEqual(r.history[0].status_code, 302)
|
||||||
|
self.assertEqual(r.url, f"{server}")
|
||||||
|
config = session.get(f"{server}/d/pbmc3k.cxg/api/v0.2/config").json()
|
||||||
|
userinfo = session.get(f"{server}/d/pbmc3k.cxg/api/v0.2/userinfo").json()
|
||||||
|
self.assertFalse(userinfo["userinfo"]["is_authenticated"])
|
||||||
|
self.assertIsNone(userinfo["userinfo"]["username"])
|
||||||
|
self.assertTrue(config["config"]["parameters"]["annotations"])
|
||||||
|
|
||||||
|
def test_auth_oauth_session(self):
|
||||||
|
# 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")
|
||||||
Reference in New Issue
Block a user