This commit is contained in:
Sayyid Hamid Mahdavi 2026-04-08 10:49:03 +03:30
parent 5f0684d354
commit 9f284fbfff
2 changed files with 72 additions and 26 deletions

View 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

View file

@ -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):
"""