269 lines
No EOL
10 KiB
Python
Executable file
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) |