accounts/apps/gooyal_oauth2/validators.py
Sayyid Hamid Mahdavi be4df13c78 revoke token test
2026-04-07 09:47:00 +03:30

269 lines
No EOL
10 KiB
Python
Executable file

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')
else:
return result
def revoke_token(self, token, token_type_hint, request, *args, **kwargs):
return super().revoke_token(token, token_type_hint, request, *args, **kwargs)