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")
|
log = logging.getLogger("oauth2_provider")
|
||||||
|
|
||||||
|
|
||||||
UserModel = get_user_model()
|
UserModel = get_user_model()
|
||||||
Application = get_application_model()
|
Application = get_application_model()
|
||||||
AccessToken = get_access_token_model()
|
AccessToken = get_access_token_model()
|
||||||
|
|
@ -32,7 +31,6 @@ Grant = get_grant_model()
|
||||||
RefreshToken = get_refresh_token_model()
|
RefreshToken = get_refresh_token_model()
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class OAuth2Validator(BaseOAuth2Validator): # pylint: disable=w0223
|
class OAuth2Validator(BaseOAuth2Validator): # pylint: disable=w0223
|
||||||
def validate_user(self, username, password, client, request, *args, **kwargs):
|
def validate_user(self, username, password, client, request, *args, **kwargs):
|
||||||
auth_fields = getattr(request, 'auth_fields', 'username:password').split(':')
|
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 _get_token_from_authentication_server
|
||||||
def validate_bearer_token(self,token, scopes, request):
|
# def validate_bearer_token(self, token, scopes, request):
|
||||||
result = super().validate_bearer_token(token, scopes, request)
|
# result = super().validate_bearer_token(token, scopes, request)
|
||||||
if result:
|
# if result:
|
||||||
application = request.client
|
# application = request.client
|
||||||
if hasattr(request, "max_allowed_session"):
|
# if hasattr(request, "max_allowed_session"):
|
||||||
max_allowed_session = request.max_allowed_session
|
# max_allowed_session = request.max_allowed_session
|
||||||
else:
|
# else:
|
||||||
max_allowed_session = application.max_allowed_session
|
# max_allowed_session = application.max_allowed_session
|
||||||
|
#
|
||||||
if max_allowed_session :
|
# if max_allowed_session:
|
||||||
session_count = AccessToken.objects.filter(application=application, user=request.user).count()
|
# session_count = AccessToken.objects.filter(application=application, user=request.user).count()
|
||||||
if max_allowed_session >= session_count:
|
# if max_allowed_session >= session_count:
|
||||||
return result
|
# return result
|
||||||
else:
|
# else:
|
||||||
active_tokens = AccessToken.objects.filter(
|
# active_tokens = AccessToken.objects.filter(
|
||||||
application=application, user=request.user
|
# application=application, user=request.user
|
||||||
).order_by("created")[:max_allowed_session].values_list("token", flat=True)
|
# ).order_by("created")[:max_allowed_session].values_list("token", flat=True)
|
||||||
if token in active_tokens:
|
# if token in active_tokens:
|
||||||
return result
|
# return result
|
||||||
else:
|
# else:
|
||||||
raise ServiceUnavailable(code='max_allowed_session_reached')
|
# raise ServiceUnavailable(code='max_allowed_session_reached')
|
||||||
|
#
|
||||||
else:
|
# else:
|
||||||
return result
|
# return result
|
||||||
|
|
||||||
def revoke_token(self, token, token_type_hint, request, *args, **kwargs):
|
def revoke_token(self, token, token_type_hint, request, *args, **kwargs):
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue