From 9f284fbfff4724a1f9f3bc008a2d230801c6524d Mon Sep 17 00:00:00 2001 From: Sayyid Hamid Mahdavi Date: Wed, 8 Apr 2026 10:49:03 +0330 Subject: [PATCH] throttle --- apps/gooyal_oauth2/throttling.py | 48 ++++++++++++++++++++++++++++++ apps/gooyal_oauth2/validators.py | 50 +++++++++++++++----------------- 2 files changed, 72 insertions(+), 26 deletions(-) create mode 100644 apps/gooyal_oauth2/throttling.py diff --git a/apps/gooyal_oauth2/throttling.py b/apps/gooyal_oauth2/throttling.py new file mode 100644 index 0000000..6f5414a --- /dev/null +++ b/apps/gooyal_oauth2/throttling.py @@ -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 \ No newline at end of file diff --git a/apps/gooyal_oauth2/validators.py b/apps/gooyal_oauth2/validators.py index 6d47be1..0928ccc 100755 --- a/apps/gooyal_oauth2/validators.py +++ b/apps/gooyal_oauth2/validators.py @@ -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): """