import base64 import binascii import logging from copy import deepcopy from datetime import datetime, timedelta from urllib.parse import unquote_plus import requests # import service_clients from django.contrib.auth import get_user_model from django.utils import timezone from django.utils.timezone import make_aware from oauth2_provider.exceptions import FatalClientError from oauth2_provider.models import get_access_token_model, get_application_model, get_id_token_model, get_grant_model, \ get_refresh_token_model from oauth2_provider.oauth2_validators import OAuth2Validator as BaseOAuth2Validator from requests import Session from utils.exceptions import ServiceUnavailable from .settings import oauth2_settings from django.conf import settings from django.db import router, transaction log = logging.getLogger("oauth2_provider") UserModel = get_user_model() Application = get_application_model() AccessToken = get_access_token_model() IDToken = get_id_token_model() Grant = get_grant_model() RefreshToken = get_refresh_token_model() class OAuth2Validator(BaseOAuth2Validator): # pylint: disable=w0223 def validate_user(self, username, password, client, request, *args, **kwargs): auth_fields = getattr(request, 'auth_fields', 'username:password').split(':') if len(auth_fields) != 2: return False user_field, pass_field = auth_fields if user_field not in ['phone_number', 'username', 'email']: return False if pass_field not in ['password', 'otp', 'ott']: return False if not username or not password: return False user = UserModel.objects.filter(**{user_field: username}).first() if not user: return False if not user.check_auth(pass_field, password): return False if user.is_active: request.user = user return True return False def get_additional_claims(self, request): return { "given_name": request.user.first_name, "family_name": request.user.last_name, "name": ' '.join([request.user.first_name, request.user.last_name]), "preferred_username": request.user.username, "email": request.user.email, } def get_claim_dict(self, request): if self._get_additional_claims_is_request_agnostic(): claims = {"sub": lambda r: str(r.user.pk)} else: claims = {"sub": str(request.user.pk)} # https://openid.net/specs/openid-connect-core-1_0.html#StandardClaims if self._get_additional_claims_is_request_agnostic(): add = self.get_additional_claims() else: add = self.get_additional_claims(request) claims.update(add) return claims def get_oidc_issuer_endpoint(self, request): return oauth2_settings.oidc_issuer(request) def save_token(self, token, request, *args, **kwargs): """Persist the token with a token type specific method. Currently, only save_bearer_token is supported. :param token: A (Bearer) token dict. :param request: OAuthlib request. :type request: oauthlib.common.Request """ return self.save_bearer_token(token, request, *args, **kwargs) def save_bearer_token(self, token, request, *args, **kwargs): """ Save access and refresh token. Override _save_bearer_token and not this function when adding custom logic for the storing of these token. This allows the transaction logic to be separate from the token handling. """ # Use the AccessToken's database instead of making the assumption it is in 'default'. with transaction.atomic(using=router.db_for_write(AccessToken)): return self._save_bearer_token(token, request, *args, **kwargs) def _save_bearer_token(self, token, request, *args, **kwargs): """ Save access and refresh token. If refresh token is issued, remove or reuse old refresh token as in rfc:`6`. @see: https://rfc-editor.org/rfc/rfc6749.html#section-6 """ if "scope" not in token: raise FatalClientError("Failed to renew access token: missing scope") # expires_in is passed to Server on initialization # custom server class can have logic to override this expires = timezone.now() + timedelta( seconds=token.get( "expires_in", oauth2_settings.ACCESS_TOKEN_EXPIRE_SECONDS, ) ) if request.grant_type == "client_credentials": request.user = None # This comes from OAuthLib: # https://github.com/idan/oauthlib/blob/1.0.3/oauthlib/oauth2/rfc6749/tokens.py#L267 # Its value is either a new random code; or if we are reusing # refresh tokens, then it is the same value that the request passed in # (stored in `request.refresh_token`) refresh_token_code = token.get("refresh_token", None) if refresh_token_code: # an instance of `RefreshToken` that matches the old refresh code. # Set on the request in `validate_refresh_token` refresh_token_instance = getattr(request, "refresh_token_instance", None) # If we are to reuse tokens, and we can: do so if ( not self.rotate_refresh_token(request) and isinstance(refresh_token_instance, RefreshToken) and refresh_token_instance.access_token ): access_token = AccessToken.objects.select_for_update().get( pk=refresh_token_instance.access_token.pk ) access_token.user = request.user access_token.scope = token["scope"] access_token.expires = expires access_token.token = token["access_token"] access_token.application = request.client access_token.save() # else create fresh with access & refresh tokens else: # revoke existing tokens if possible to allow reuse of grant if isinstance(refresh_token_instance, RefreshToken): # First, to ensure we don't have concurrency issues, we refresh the refresh token # from the db while acquiring a lock on it # We also put it in the "request cache" refresh_token_instance = RefreshToken.objects.select_for_update().get( pk=refresh_token_instance.pk ) request.refresh_token_instance = refresh_token_instance previous_access_token = AccessToken.objects.filter( source_refresh_token=refresh_token_instance ).first() try: refresh_token_instance.revoke() except (AccessToken.DoesNotExist, RefreshToken.DoesNotExist): pass else: setattr(request, "refresh_token_instance", None) else: previous_access_token = None # If the refresh token has already been used to create an # access token (ie it's within the grace period), return that # access token if not previous_access_token: access_token = self._create_access_token( expires, request, token, source_refresh_token=refresh_token_instance, ) self._create_refresh_token( request, refresh_token_code, access_token, refresh_token_instance ) else: # make sure that the token data we're returning matches # the existing token token["access_token"] = previous_access_token.token token["refresh_token"] = ( RefreshToken.objects.filter(access_token=previous_access_token).first().token ) token["scope"] = previous_access_token.scope # No refresh token should be created, just access token else: self._create_access_token(expires, request, token) def _create_access_token(self, expires, request, token, source_refresh_token=None): id_token = token.get("id_token", None) if id_token: id_token = self._load_id_token(id_token) headers = {} for header, value in dict(request.headers).items(): if type(value) in [str, bool, int, float, tuple, list]: print(f'unserializable header: {header} -> {value}') headers[header] = value return AccessToken.objects.create( user=request.user, scope=token["scope"], expires=expires, token=token["access_token"], id_token=id_token, application=request.client, source_refresh_token=source_refresh_token, detail={'headers': headers}, ) # def _get_token_from_authentication_server def validate_bearer_token(self,token, scopes, request): result = super().validate_bearer_token(token, scopes, request) if result: application = request.client if hasattr(request, "max_allowed_session"): max_allowed_session = request.max_allowed_session else: max_allowed_session = application.max_allowed_session if max_allowed_session : session_count = AccessToken.objects.filter(application=application, user=request.user).count() if max_allowed_session >= session_count: return result else: active_tokens = AccessToken.objects.filter( application=application, user=request.user ).order_by("created")[:max_allowed_session].values_list("token", flat=True) if token in active_tokens: return result else: raise ServiceUnavailable(code='max_allowed_session_reached')