258 lines
10 KiB
Python
258 lines
10 KiB
Python
import logging
|
|
import time
|
|
from datetime import timedelta
|
|
|
|
import requests
|
|
from django.core.exceptions import ImproperlyConfigured
|
|
from django.db.models import JSONField
|
|
from oauth2_provider.models import AbstractApplication, AbstractAccessToken, AbstractGrant, AbstractRefreshToken, \
|
|
AbstractIDToken, get_access_token_model, get_refresh_token_model, get_id_token_model, get_grant_model
|
|
from oauth2_provider.scopes import get_scopes_backend
|
|
from oauth2_provider.settings import oauth2_settings
|
|
import uuid
|
|
from contextlib import suppress
|
|
|
|
from django.conf import settings
|
|
from django.db import models, router, transaction
|
|
from django.utils import timezone
|
|
|
|
from utils.models import BaseModel
|
|
from .settings import oauth2_settings
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class Application(AbstractApplication, BaseModel):
|
|
"""
|
|
Application model for use with Django OAuth Toolkit that allows the scopes
|
|
available to an application to be restricted on a per-application basis.
|
|
"""
|
|
id = None
|
|
name = models.CharField(max_length=255, blank=False, unique=True)
|
|
user = models.ForeignKey(
|
|
settings.AUTH_USER_MODEL,
|
|
related_name="%(app_label)s_%(class)s",
|
|
on_delete=models.PROTECT
|
|
)
|
|
allowed_scope = models.TextField(blank=True)
|
|
|
|
avatar = models.ImageField(upload_to='avatars', null=True, blank=True)
|
|
max_allowed_session = models.PositiveIntegerField(default=1)
|
|
|
|
|
|
@property
|
|
def allowed_scopes(self):
|
|
"""
|
|
Returns the set of allowed scope names for this application.
|
|
"""
|
|
all_scopes = set(get_scopes_backend().get_all_scopes().keys())
|
|
app_scopes = set(self.allowed_scope.split())
|
|
return app_scopes.intersection(all_scopes)
|
|
|
|
@property
|
|
def allowed_scopes_queryset(self):
|
|
allowed_scopes = self.allowed_scopes
|
|
return Scope.objects.filter(name__in=allowed_scopes).order_by('name')
|
|
|
|
def allows_grant_type(self, *grant_types):
|
|
# Assume, for this example, that self.authorization_grant_type is set to self.GRANT_AUTHORIZATION_CODE
|
|
return bool(set([self.authorization_grant_type, self.GRANT_CLIENT_CREDENTIALS]) & set(grant_types))
|
|
|
|
|
|
class Scope(BaseModel):
|
|
"""
|
|
Django model for an OAuth scope.
|
|
"""
|
|
#: The application that created the scope
|
|
# NOTE: This is not used to limit access to the scope in any way - we want the
|
|
# scope to be available to other applications in order to request access
|
|
# to the resource it protects!
|
|
application = models.ForeignKey(
|
|
oauth2_settings.APPLICATION_MODEL,
|
|
models.CASCADE,
|
|
# This field is nullable because it is only set for scopes created by
|
|
# external resource servers, which have a corresponding OAuth application
|
|
# record on the authorisation server
|
|
blank=True, null=True,
|
|
help_text='The application to which the scope belongs.',
|
|
related_name='scopes'
|
|
)
|
|
#: The name of the scope
|
|
name = models.CharField(
|
|
max_length=255,
|
|
unique=True,
|
|
help_text='The name of the scope.'
|
|
)
|
|
#: A brief description of the scope
|
|
description = models.TextField(
|
|
help_text='A brief description of the scope. This text is displayed '
|
|
'to users when authorising access for the scope.'
|
|
)
|
|
#: Indicates if the scope should be included in the default scopes
|
|
is_default = models.BooleanField(
|
|
default=False,
|
|
help_text='Indicates if this scope should be included in the default scopes.'
|
|
)
|
|
|
|
@property
|
|
def final_name(self):
|
|
return self.name
|
|
|
|
@property
|
|
def final_description(self):
|
|
return self.description
|
|
|
|
@classmethod
|
|
def register(cls, name, description, is_default=False):
|
|
"""
|
|
Registers a scope with the given values. It always creates an instance in
|
|
the local database, but if this resource server has an external authorisation
|
|
server, it will also register the scope there.
|
|
|
|
Returns ``True`` on success. Should raise on failure.
|
|
"""
|
|
endpoint = settings.RESOURCE_SERVER_REGISTER_SCOPE_URL
|
|
if endpoint:
|
|
# If the endpoint is set, make the callout to the authz server
|
|
token = "Bearer {}".format(oauth2_settings.RESOURCE_SERVER_AUTH_TOKEN)
|
|
# Let any failures bubble up
|
|
# The idea is to call this method during deployment as a post-migrate
|
|
# hook, so we want failures to halt the deployment
|
|
response = requests.post(
|
|
endpoint,
|
|
json={
|
|
'name': name,
|
|
'description': description,
|
|
'is_default': is_default
|
|
},
|
|
headers={"Authorization": token}
|
|
)
|
|
# Raise the exception for anything other than 20x responses
|
|
response.raise_for_status()
|
|
# Always create/update the scope record locally
|
|
_ = Scope.objects.update_or_create(
|
|
name=name,
|
|
defaults={'description': description, 'is_default': is_default}
|
|
)
|
|
return True
|
|
|
|
|
|
class Grant(AbstractGrant, BaseModel):
|
|
id = None
|
|
application = models.ForeignKey(
|
|
oauth2_settings.APPLICATION_MODEL, on_delete=models.CASCADE, related_name='grants'
|
|
)
|
|
|
|
|
|
class RefreshToken(BaseModel, AbstractRefreshToken):
|
|
id = None
|
|
application = models.ForeignKey(
|
|
oauth2_settings.APPLICATION_MODEL, on_delete=models.CASCADE, related_name='refresh_tokens')
|
|
|
|
# TODO: submit merge request
|
|
def revoke(self):
|
|
"""
|
|
Mark this refresh token revoked and revoke related access token
|
|
"""
|
|
access_token_model = get_access_token_model()
|
|
access_token_database = router.db_for_write(access_token_model)
|
|
refresh_token_model = get_refresh_token_model()
|
|
|
|
# Use the access_token_database instead of making the assumption it is in 'default'.
|
|
with transaction.atomic(using=access_token_database):
|
|
token = refresh_token_model.objects.select_for_update().filter(pk=self.pk, revoked__isnull=True)
|
|
if not token:
|
|
return
|
|
self = list(token)[0]
|
|
|
|
with suppress(access_token_model.DoesNotExist):
|
|
access_token_model.objects.get(pk=self.access_token_id).revoke()
|
|
|
|
self.access_token = None
|
|
self.revoked = timezone.now()
|
|
self.save()
|
|
|
|
|
|
class AccessToken(AbstractAccessToken):
|
|
id = None
|
|
uuid = models.UUIDField(primary_key=True, editable=False, default=uuid.uuid4, unique=True, db_index=True)
|
|
detail = JSONField(null=True, blank=True)
|
|
application = models.ForeignKey(
|
|
oauth2_settings.APPLICATION_MODEL, on_delete=models.CASCADE, blank=True, null=True, related_name='access_tokens'
|
|
)
|
|
|
|
|
|
class IDToken(AbstractIDToken):
|
|
id = None
|
|
uuid = models.UUIDField(primary_key=True, editable=False, default=uuid.uuid4, unique=True, db_index=True)
|
|
|
|
|
|
def clear_expired():
|
|
def batch_delete(queryset, query):
|
|
CLEAR_EXPIRED_TOKENS_BATCH_SIZE = oauth2_settings.CLEAR_EXPIRED_TOKENS_BATCH_SIZE
|
|
CLEAR_EXPIRED_TOKENS_BATCH_INTERVAL = oauth2_settings.CLEAR_EXPIRED_TOKENS_BATCH_INTERVAL
|
|
current_no = start_no = queryset.count()
|
|
|
|
while current_no:
|
|
flat_queryset = queryset.values_list("pk", flat=True)[:CLEAR_EXPIRED_TOKENS_BATCH_SIZE]
|
|
batch_length = flat_queryset.count()
|
|
queryset.model.objects.filter(pk__in=list(flat_queryset)).delete()
|
|
logger.debug(f"{batch_length} tokens deleted, {current_no - batch_length} left")
|
|
queryset = queryset.model.objects.filter(query)
|
|
time.sleep(CLEAR_EXPIRED_TOKENS_BATCH_INTERVAL)
|
|
current_no = queryset.count()
|
|
|
|
stop_no = queryset.model.objects.filter(query).count()
|
|
deleted = start_no - stop_no
|
|
return deleted
|
|
|
|
now = timezone.now()
|
|
refresh_expire_at = None
|
|
access_token_model = get_access_token_model()
|
|
refresh_token_model = get_refresh_token_model()
|
|
id_token_model = get_id_token_model()
|
|
grant_model = get_grant_model()
|
|
REFRESH_TOKEN_EXPIRE_SECONDS = oauth2_settings.REFRESH_TOKEN_EXPIRE_SECONDS
|
|
|
|
if REFRESH_TOKEN_EXPIRE_SECONDS:
|
|
if not isinstance(REFRESH_TOKEN_EXPIRE_SECONDS, timedelta):
|
|
try:
|
|
REFRESH_TOKEN_EXPIRE_SECONDS = timedelta(seconds=REFRESH_TOKEN_EXPIRE_SECONDS)
|
|
except TypeError:
|
|
e = "REFRESH_TOKEN_EXPIRE_SECONDS must be either a timedelta or seconds"
|
|
raise ImproperlyConfigured(e)
|
|
refresh_expire_at = now - REFRESH_TOKEN_EXPIRE_SECONDS
|
|
|
|
if refresh_expire_at:
|
|
revoked_query = models.Q(revoked__lt=refresh_expire_at)
|
|
revoked = refresh_token_model.objects.filter(revoked_query)
|
|
|
|
revoked_deleted_no = batch_delete(revoked, revoked_query)
|
|
logger.info("%s Revoked refresh tokens deleted", revoked_deleted_no)
|
|
|
|
expired_query = models.Q(access_token__expires__lt=refresh_expire_at)
|
|
expired = refresh_token_model.objects.filter(expired_query)
|
|
|
|
expired_deleted_no = batch_delete(expired, expired_query)
|
|
logger.info("%s Expired refresh tokens deleted", expired_deleted_no)
|
|
else:
|
|
logger.info("refresh_expire_at is %s. No refresh tokens deleted.", refresh_expire_at)
|
|
|
|
access_token_query = models.Q(refresh_token__isnull=True, expires__lt=now)
|
|
access_tokens = access_token_model.objects.filter(access_token_query)
|
|
|
|
access_tokens_delete_no = batch_delete(access_tokens, access_token_query)
|
|
logger.info("%s Expired access tokens deleted", access_tokens_delete_no)
|
|
|
|
id_token_query = models.Q(access_token__isnull=True, expires__lt=now)
|
|
id_tokens = id_token_model.objects.filter(id_token_query)
|
|
|
|
id_tokens_delete_no = batch_delete(id_tokens, id_token_query)
|
|
logger.info("%s Expired ID tokens deleted", id_tokens_delete_no)
|
|
|
|
grants_query = models.Q(expires__lt=now)
|
|
grants = grant_model.objects.filter(grants_query)
|
|
|
|
grants_deleted_no = batch_delete(grants, grants_query)
|
|
logger.info("%s Expired grant tokens deleted", grants_deleted_no)
|