mirror of
https://github.com/chanzuckerberg/cellxgene.git
synced 2026-10-09 07:38:12 +08:00
Update hosted app to get the oauth client secret from the secret manager (#1713)
* Update the hosted app to get the oauth client secret from the secret manager * fix to eb app, and set no cache on oauth endpoints
This commit is contained in:
@@ -122,13 +122,19 @@ class AuthTypeOAuth(AuthTypeClientBase):
|
|||||||
return payload.get("email")
|
return payload.get("email")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def update_response(self, response):
|
||||||
|
response.cache_control.update(
|
||||||
|
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
|
||||||
session["oauth_callback_redirect"] = return_to
|
session["oauth_callback_redirect"] = return_to
|
||||||
return self.client.authorize_redirect(redirect_uri=callbackurl)
|
response = self.client.authorize_redirect(redirect_uri=callbackurl)
|
||||||
|
self.update_response(response)
|
||||||
|
return response
|
||||||
|
|
||||||
def logout(self):
|
def logout(self):
|
||||||
if self.session_cookie:
|
if self.session_cookie:
|
||||||
@@ -138,12 +144,15 @@ class AuthTypeOAuth(AuthTypeClientBase):
|
|||||||
@after_this_request
|
@after_this_request
|
||||||
def remove_cookie(response):
|
def remove_cookie(response):
|
||||||
response.set_cookie(self.cookie_params["key"], "", expires=0)
|
response.set_cookie(self.cookie_params["key"], "", expires=0)
|
||||||
|
self.update_response()
|
||||||
return response
|
return response
|
||||||
|
|
||||||
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}"
|
||||||
params = {'returnTo' : return_to, 'client_id' : self.client_id}
|
params = {'returnTo' : return_to, 'client_id' : self.client_id}
|
||||||
return redirect(self.client.api_base_url + '/v2/logout?' + urlencode(params))
|
response = redirect(self.client.api_base_url + '/v2/logout?' + urlencode(params))
|
||||||
|
self.update_response()
|
||||||
|
return response
|
||||||
|
|
||||||
def callback(self):
|
def callback(self):
|
||||||
token = self.client.authorize_access_token()
|
token = self.client.authorize_access_token()
|
||||||
@@ -165,6 +174,7 @@ class AuthTypeOAuth(AuthTypeClientBase):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise AuthenticationError(f"unable to set_cookie {self.cookie_params}") from e
|
raise AuthenticationError(f"unable to set_cookie {self.cookie_params}") from e
|
||||||
|
|
||||||
|
self.update_response(resp)
|
||||||
return resp
|
return resp
|
||||||
|
|
||||||
def get_login_url(self, data_adaptor):
|
def get_login_url(self, data_adaptor):
|
||||||
|
|||||||
+44
-20
@@ -31,7 +31,7 @@ except Exception:
|
|||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
|
|
||||||
|
|
||||||
def get_secret_key(region_name, secret_name, secret_key):
|
def get_secret_key(region_name, secret_name):
|
||||||
session = boto3.session.Session()
|
session = boto3.session.Session()
|
||||||
client = session.client(service_name="secretsmanager", region_name=region_name)
|
client = session.client(service_name="secretsmanager", region_name=region_name)
|
||||||
|
|
||||||
@@ -40,7 +40,7 @@ def get_secret_key(region_name, secret_name, secret_key):
|
|||||||
if "SecretString" in get_secret_value_response:
|
if "SecretString" in get_secret_value_response:
|
||||||
var = get_secret_value_response["SecretString"]
|
var = get_secret_value_response["SecretString"]
|
||||||
secret = json.loads(var)
|
secret = json.loads(var)
|
||||||
return secret.get(secret_key)
|
return secret
|
||||||
except Exception:
|
except Exception:
|
||||||
logging.critical("Caught exception during get_secret_key", exc_info=True)
|
logging.critical("Caught exception during get_secret_key", exc_info=True)
|
||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
@@ -48,6 +48,46 @@ def get_secret_key(region_name, secret_name, secret_key):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def handle_config_from_secret(app_config):
|
||||||
|
"""Update configuration from the secret manager"""
|
||||||
|
secret_name = os.getenv("CXG_AWS_SECRET_NAME")
|
||||||
|
if not secret_name:
|
||||||
|
return
|
||||||
|
|
||||||
|
# need to find the secret manager region.
|
||||||
|
# 1. from CXG_AWS_SECRET_REGION_NAME
|
||||||
|
# 2. discover from dataroot location (if on s3)
|
||||||
|
# 3. discover from config file location (if on s3)
|
||||||
|
secret_region_name = os.getenv("CXG_AWS_SECRET_REGION_NAME")
|
||||||
|
if secret_region_name is None:
|
||||||
|
secret_region_name = discover_s3_region_name(app_config.multi_dataset__dataroot)
|
||||||
|
if not secret_region_name:
|
||||||
|
secret_region_name = discover_s3_region_name(config_file)
|
||||||
|
if not secret_region_name:
|
||||||
|
logging.error("Could not determine the AWS Secret Manager region")
|
||||||
|
sys.exit(1)
|
||||||
|
|
||||||
|
secrets = get_secret_key(secret_region_name, secret_name)
|
||||||
|
if not secrets:
|
||||||
|
return
|
||||||
|
|
||||||
|
keyattrs = (
|
||||||
|
("flask_secret_key", "app__flask_secret_key"),
|
||||||
|
("oauth_client_secret", "authentication__params_oauth__client_secret")
|
||||||
|
)
|
||||||
|
|
||||||
|
for key, attr in keyattrs:
|
||||||
|
curval = getattr(app_config.server_config, attr)
|
||||||
|
if curval:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# replace the attr with the secret if it is not set
|
||||||
|
val = secrets.get(key)
|
||||||
|
if val:
|
||||||
|
logging.error(f"set {attr} from secret")
|
||||||
|
app_config.update_server_config(**{attr : val})
|
||||||
|
|
||||||
|
|
||||||
class WSGIServer(Server):
|
class WSGIServer(Server):
|
||||||
def __init__(self, app_config):
|
def __init__(self, app_config):
|
||||||
super().__init__(app_config)
|
super().__init__(app_config)
|
||||||
@@ -158,30 +198,14 @@ try:
|
|||||||
logging.info("Configuration from CXG_DATAROOT")
|
logging.info("Configuration from CXG_DATAROOT")
|
||||||
app_config.update_server_config(multi_dataset__dataroot=dataroot)
|
app_config.update_server_config(multi_dataset__dataroot=dataroot)
|
||||||
|
|
||||||
secret_name = os.getenv("CXG_AWS_SECRET_NAME")
|
# update from secret manager
|
||||||
if secret_name:
|
handle_config_from_secret(app_config)
|
||||||
# need to find the secret manager region.
|
|
||||||
# 1. from CXG_AWS_SECRET_REGION_NAME
|
|
||||||
# 2. discover from dataroot location (if on s3)
|
|
||||||
# 3. discover from config file location (if on s3)
|
|
||||||
secret_region_name = os.getenv("CXG_AWS_SECRET_REGION_NAME")
|
|
||||||
if secret_region_name is None:
|
|
||||||
secret_region_name = discover_s3_region_name(app_config.multi_dataset__dataroot)
|
|
||||||
if not secret_region_name:
|
|
||||||
secret_region_name = discover_s3_region_name(config_file)
|
|
||||||
if not secret_region_name:
|
|
||||||
logging.error("Could not determine the AWS Secret Manager region")
|
|
||||||
sys.exit(1)
|
|
||||||
|
|
||||||
flask_secret_key = get_secret_key(secret_region_name, secret_name, 'flask_secret_key')
|
|
||||||
app_config.update_server_config(app__flask_secret_key=flask_secret_key)
|
|
||||||
|
|
||||||
# features are unsupported in the current hosted server
|
# features are unsupported in the current hosted server
|
||||||
app_config.update_default_dataset_config(
|
app_config.update_default_dataset_config(
|
||||||
user_annotations__enable=False, embeddings__enable_reembedding=False,
|
user_annotations__enable=False, embeddings__enable_reembedding=False,
|
||||||
)
|
)
|
||||||
app_config.update_server_config(multi_dataset__allowed_matrix_types=["cxg"],)
|
app_config.update_server_config(multi_dataset__allowed_matrix_types=["cxg"],)
|
||||||
|
|
||||||
app_config.complete_config(logging.info)
|
app_config.complete_config(logging.info)
|
||||||
|
|
||||||
if not app_config.server_config.app__flask_secret_key:
|
if not app_config.server_config.app__flask_secret_key:
|
||||||
|
|||||||
Reference in New Issue
Block a user