diff --git a/server/auth/auth_oauth.py b/server/auth/auth_oauth.py index ca2adf95..bb224e21 100644 --- a/server/auth/auth_oauth.py +++ b/server/auth/auth_oauth.py @@ -122,13 +122,19 @@ class AuthTypeOAuth(AuthTypeClientBase): return payload.get("email") 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): callbackurl = f'{self.callback_base_url}/oauth2/callback' return_path = request.args.get("dataset", "") return_to = f"{self.callback_base_url}/{return_path}" # save the return path in the session cookie, accessed in the callback function 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): if self.session_cookie: @@ -138,12 +144,15 @@ class AuthTypeOAuth(AuthTypeClientBase): @after_this_request def remove_cookie(response): response.set_cookie(self.cookie_params["key"], "", expires=0) + self.update_response() return response return_path = request.args.get("dataset", "") return_to = f"{self.callback_base_url}/{return_path}" 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): token = self.client.authorize_access_token() @@ -165,6 +174,7 @@ class AuthTypeOAuth(AuthTypeClientBase): except Exception as e: raise AuthenticationError(f"unable to set_cookie {self.cookie_params}") from e + self.update_response(resp) return resp def get_login_url(self, data_adaptor): diff --git a/server/eb/app.py b/server/eb/app.py index 64645fc3..2e451e77 100644 --- a/server/eb/app.py +++ b/server/eb/app.py @@ -31,7 +31,7 @@ except Exception: 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() 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: var = get_secret_value_response["SecretString"] secret = json.loads(var) - return secret.get(secret_key) + return secret except Exception: logging.critical("Caught exception during get_secret_key", exc_info=True) sys.exit(1) @@ -48,6 +48,46 @@ def get_secret_key(region_name, secret_name, secret_key): 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): def __init__(self, app_config): super().__init__(app_config) @@ -158,30 +198,14 @@ try: logging.info("Configuration from CXG_DATAROOT") app_config.update_server_config(multi_dataset__dataroot=dataroot) - secret_name = os.getenv("CXG_AWS_SECRET_NAME") - if secret_name: - # 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) + # update from secret manager + handle_config_from_secret(app_config) # features are unsupported in the current hosted server app_config.update_default_dataset_config( user_annotations__enable=False, embeddings__enable_reembedding=False, ) app_config.update_server_config(multi_dataset__allowed_matrix_types=["cxg"],) - app_config.complete_config(logging.info) if not app_config.server_config.app__flask_secret_key: