throttle
This commit is contained in:
parent
5f0684d354
commit
9f284fbfff
2 changed files with 72 additions and 26 deletions
48
apps/gooyal_oauth2/throttling.py
Normal file
48
apps/gooyal_oauth2/throttling.py
Normal file
|
|
@ -0,0 +1,48 @@
|
|||
from rest_framework.throttling import UserRateThrottle
|
||||
|
||||
from apps.gooyal_oauth2.models import AccessToken
|
||||
|
||||
def get_application(request):
|
||||
try:
|
||||
application = request.auth.application
|
||||
except:
|
||||
application = None
|
||||
return application
|
||||
|
||||
|
||||
class TokenLimitThrottle:
|
||||
def allow_request(self, request, view):
|
||||
application = get_application(request)
|
||||
|
||||
if application:
|
||||
if hasattr(view, "max_allowed_session"):
|
||||
max_allowed_session = view.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 True
|
||||
else:
|
||||
active_tokens = AccessToken.objects.filter(
|
||||
application=application, user=request.user
|
||||
).order_by("created")[:max_allowed_session].values_list("token", flat=True)
|
||||
if request.auth.token in active_tokens:
|
||||
return True
|
||||
else:
|
||||
False
|
||||
# raise ServiceUnavailable(code='max_allowed_session_reached')
|
||||
|
||||
else:
|
||||
return True
|
||||
|
||||
def wait(self):
|
||||
"""
|
||||
Optionally, return a recommended number of seconds to wait before
|
||||
the next request.
|
||||
"""
|
||||
return None
|
||||
|
||||
|
||||
from rest_framework.exceptions import Throttled
|
||||
|
|
@ -23,7 +23,6 @@ from django.db import router, transaction
|
|||
|
||||
log = logging.getLogger("oauth2_provider")
|
||||
|
||||
|
||||
UserModel = get_user_model()
|
||||
Application = get_application_model()
|
||||
AccessToken = get_access_token_model()
|
||||
|
|
@ -32,7 +31,6 @@ 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(':')
|
||||
|
|
@ -240,30 +238,30 @@ class OAuth2Validator(BaseOAuth2Validator): # pylint: disable=w0223
|
|||
)
|
||||
|
||||
# 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 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):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue