import itertools import logging import time from posixpath import join from urllib.request import urlopen import jwt try: import structlog except ImportError: structlog = None import requests from flask import ( _request_ctx_stack, request, abort, redirect, has_request_context, ) from . import models, utils, claims from .domains import get_allowed_root_domain class AuthManager: """ Flask extension to manage auth process. Wires up into request-response cycle. """ # Using nonzero leeway by default because PyJWT is using iss validation # in last versions and it may cause much more issues than no leeway for # expiration because of non synced time on servers DEFAULT_TOKEN_LEEWAY = 15 DEFAULT_COOKIES_MAX_AGE = 60 * 60 * 24 * 30 def __init__( self, app=None, public_key=None, login_url=None, public_key_url=None, refresh_token_url=None, encoding_algorithm=None, related_claims_namespace=None, bearer_token_cookie_name=None, refresh_token_cookie_name=None, token_cookie_max_age=None, token_leeway=None, logger=None, ): self.app = app self._login_url = login_url self._public_key_url = public_key_url self._refresh_token_url = refresh_token_url self._public_key = None self._encoding_algorithm = encoding_algorithm self._related_claims_namespace = related_claims_namespace self._bearer_token_cookie_name = bearer_token_cookie_name self._refresh_token_cookie_name = refresh_token_cookie_name self._token_cookie_max_age = token_cookie_max_age self._token_leeway = token_leeway self._logger = logger or logging.getLogger(__name__) if app is not None: self.init_app(app) def init_app(self, app): self._login_url = self._login_url or app.config.get("ATLAS_LOGIN_URL") self._public_key_url = self._public_key_url or app.config.get( "ATLAS_PUBLIC_KEY_URL" ) self._refresh_token_url = self._refresh_token_url or app.config.get( "ATLAS_REFRESH_TOKEN_URL" ) self._encoding_algorithm = self._encoding_algorithm or app.config.get( "ATLAS_ENCODING_ALGORITHM" ) self._related_claims_namespace = ( self._related_claims_namespace or app.config.get("ATLAS_RELATED_CLAIMS_NAMESPACE") ) self._bearer_token_cookie_name = ( self._bearer_token_cookie_name or app.config.get("ATLAS_BEARER_TOKEN_COOKIE_NAME") or "dna_bearer_token" ) self._refresh_token_cookie_name = ( self._refresh_token_cookie_name or app.config.get("ATLAS_REFRESH_TOKEN_COOKIE_NAME") or "dna_refresh_token" ) self._token_cookie_max_age = ( self._token_cookie_max_age or app.config.get("ATLAS_TOKEN_COOKIES_MAX_AGE") or self.DEFAULT_COOKIES_MAX_AGE ) self._token_leeway = ( self._token_leeway or app.config.get("ATLAS_TOKEN_LEEWAY") or self.DEFAULT_TOKEN_LEEWAY ) app.auth_manager = self app.after_request(self._load_user) app.after_request(self._refresh_token) def _load_user(self, response=None): """ Hook to load user from cookie token. Makes signature verification and decoding. """ # check if already processed from utils.current_user if has_request_context() and hasattr(_request_ctx_stack.top, "user"): return response anonymous_user = models.AnonymousUser() token = self._get_token() if not token: self._update_request_context_with_user(anonymous_user) return response try: claimset = jwt.decode( token, self.public_key, self._encoding_algorithm, **self._get_validation_params(token), ) except jwt.PyJWTError as e: self._logger.warning(f"Token decoding error: {e}") self._update_request_context_with_user(anonymous_user) return response if not isinstance(claimset, dict): self._logger.warning("Corrupted claimset") self._update_request_context_with_user(anonymous_user) return response sub = claimset.get("sub") if not sub: self._logger.warning("Empty sub in valid token") self._update_request_context_with_user(anonymous_user) return response email = claimset.get("email", "") name = claimset.get("name", "") user = models.User( sub, email, name, self._get_assigned_claims(claimset) ) self._update_request_context_with_user(user) return response def _get_token(self): """ Fetch token from one of the following sources: - auth header - cookie """ token = None # case with token in header, e.g. Authorization: Bearer auth_header = request.headers.get("Authorization", "") if auth_header and auth_header.startswith("Bearer "): token = auth_header.replace("Bearer ", "") # case with token in cookie if not token: token = request.cookies.get(self._bearer_token_cookie_name) return token def _refresh_token(self, response=None): bearer_token = request.cookies.get(self._bearer_token_cookie_name) refresh_token = request.cookies.get(self._refresh_token_cookie_name) if ( not bearer_token or not refresh_token or not self._refresh_token_url ): return response try: payload = jwt.decode( bearer_token, options={"verify_signature": False} ) except Exception: return response exp = payload.get("exp") if not isinstance(exp, int): return response if int(time.time()) < exp - 60: return response try: atlas_resp = requests.post( self._refresh_token_url, data={ "refresh_token": refresh_token, "resource_group": self._related_claims_namespace, }, ) except Exception as e: self._logger.error(f"Error on token refresh: {e}") return response if atlas_resp.status_code != 200: self._logger.warning( f"Failed to refresh the token: {atlas_resp.content}" ) return response data = atlas_resp.json() bearer_token = data.get("access_token", "") refresh_token = data.get("refresh_token", "") if not bearer_token or not refresh_token: self._logger.warning(f"Received no tokens: {atlas_resp.content}") return response domain = get_allowed_root_domain() domain_cookie = f".{domain}" if domain else None response.set_cookie( self._bearer_token_cookie_name, bearer_token, httponly=True, domain=domain_cookie, secure=True, max_age=self._token_cookie_max_age, ) response.set_cookie( self._refresh_token_cookie_name, refresh_token, httponly=True, domain=domain_cookie, secure=True, max_age=self._token_cookie_max_age, ) return response @property def public_key(self): if self._public_key: return self._public_key public_key = urlopen(self._public_key_url) # nosec public_key = public_key.read() self._public_key = public_key.decode() return self._public_key def _get_validation_params(self, token): """ We may have 2 kinds of tokens: - generic - without 'aud' claim, suitable for usage with all projects in a scope of end user authorization - m2m - with 'aud' claim, suitable for projects specified in 'aud' in a scope of server to server authorization This function returns params to pass the validation for both cases. We need this, as pyjwt lib has hard validation of `aud` claim on token decoding. """ params = {} try: claimset = jwt.decode(token, options={"verify_signature": False}) except jwt.PyJWTError as e: self._logger.warning(f"Token decoding error: {e}") return params if "aud" in claimset: audience = list( itertools.chain( *[ c.Values.list() for c in claims.registered_claims if c.is_reserved and c.path == "aud" ] ) ) params["audience"] = audience try: leeway = int(self._token_leeway) except ValueError: leeway = 0 self._logger.warning(f"Invalid leeway value: {self._token_leeway}") params["leeway"] = leeway return params def _get_assigned_claims(self, claimset): """Get claims, related to current application.""" assigned_claims = [] for claim in claims.registered_claims: full_path = ( claim.path if claim.is_reserved else join(self._related_claims_namespace, claim.path) ) for key, value in claimset.items(): if full_path != key: continue # appending claims for allowed values try: assigned_claims.append(claim.from_token_value(value)) except ValueError: pass return assigned_claims def _update_request_context_with_user(self, user): ctx = _request_ctx_stack.top ctx.user = user if structlog: structlog.contextvars.bind_contextvars(user=user) def unauthorized(self, login_redirect=True): """Redirect lo login if not authenticated or abort with 401.""" if utils.current_user.is_authenticated or not login_redirect: abort(401) return redirect(f"{self._login_url}?next={request.url}")