import logging from urllib.request import urlopen import jwt from django.conf import settings from django.contrib.auth import get_user_model from django.contrib.auth.backends import BaseBackend from core_images.exceptions import AccessDeniedError, BearerTokenExpiredError UserModel = get_user_model() logger = logging.getLogger(__name__) session = {} class AtlasAuthBackend(BaseBackend): ALGORITHMS = "RS256" ALLOWED_CLAIMS = {"core-images/role": "admin"} def authenticate( self, request, token=None, refresh_token=None, response=None, **kwargs ): if not token: return None try: rsa_key = self.get_public_key() except Exception as e: logger.error(f"Unable to fetch the Atlas public key: {e}") return None try: payload = jwt.decode( token, rsa_key, algorithms=self.ALGORITHMS, ) except jwt.exceptions.ExpiredSignatureError as e: logger.info(f"Refresh token process: {e}") raise BearerTokenExpiredError except jwt.PyJWTError as e: logger.warning(f"Token validation error: {e}") return None username = payload.get("email") if not username: return None if not self.has_allowed_claims(payload): raise AccessDeniedError user, created = UserModel._default_manager.update_or_create( **{UserModel.USERNAME_FIELD: username}, defaults={ "first_name": payload.get("given_name", ""), "last_name": payload.get("family_name", ""), "is_staff": True, "is_active": True, "is_superuser": True, }, ) if created: user.set_unusable_password() user.save() return user def get_user(self, user_id): try: user = UserModel._default_manager.get(pk=user_id) except UserModel.DoesNotExist: return None return user def get_public_key(self): """Fetch public key from shared resource on the network.""" key = session.get("public_key") if key: return key public_key = urlopen(settings.ATLAS_PUBLIC_KEY_URL) # nosec public_key = public_key.read() key = public_key.decode() session["public_key"] = key return key def has_allowed_claims(self, payload): """Checking permissions with valid token.""" if not payload: return False for claim_name, claim_value in self.ALLOWED_CLAIMS.items(): if claim_value in payload.get(claim_name, []): return True return False