import json import typing from urllib.request import urlopen from flask import request from jose import jwt from auth_api import config from auth_api.errors import AuthError, Error from auth_api.logs import logger from auth_api.providers.base import BaseAuthProvider from auth_api.users import User class Auth0AuthProvider(BaseAuthProvider): def __init__(self): self.keys = [] def authorize(self) -> User: """ Validates the token, returns the user or None (for m2m cases), or throws an exception if the request is not authorized. """ token = self.get_token_auth_header() try: unverified_header = jwt.get_unverified_header(token) except jwt.JWTError: logger.error(f"Code: invalid_token \n" "Description: Error decoding token headers \n" f"Token: {token}") raise AuthError(Error("invalid_token", "Error decoding token headers")) rsa_key = {} for key in self.get_keys(): if key["kid"] == unverified_header["kid"]: rsa_key = key if rsa_key: try: payload = jwt.decode( token, rsa_key, algorithms=config.ALGORITHMS, audience=config.API_AUDIENCE, issuer=config.AUTH0_DOMAIN, options={"verify_aud": bool(config.API_AUDIENCE)}, ) except jwt.ExpiredSignatureError: raise AuthError(Error("token_expired", "Token is expired")) except jwt.JWTClaimsError: raise AuthError( Error( "invalid_claims", "Incorrect claims, please check the audience and issuer", ) ) except Exception: raise AuthError(Error("invalid_header", "Unable to parse authentication token.")) return User(payload) raise AuthError(Error("invalid_header", "Unable to find appropriate key")) def get_keys(self) -> list: if self.keys: return self.keys jsonurl = urlopen(config.AUTH0_DOMAIN[0] + ".well-known/jwks.json") # nosec jwks = json.loads(jsonurl.read()) self.keys = jwks["keys"] return self.keys @staticmethod def get_token_auth_header(): """Obtains the Access Token from the Authorization Header.""" auth = request.headers.get("Authorization", None) if not auth: raise AuthError(Error("authorization_header_missing", "Authorization header is expected")) parts = auth.split() if parts[0].lower() != "bearer": logger.error(f"Unable to decode token: {auth}, \n" f"Token parts: {parts}") raise AuthError(Error("invalid_header", "Authorization header must start with Bearer")) elif len(parts) == 1: raise AuthError(Error("invalid_header", "Token not found")) elif len(parts) > 2: raise AuthError(Error("invalid_header", "Authorization header must be Bearer token")) return parts[1]