From 014c674961eb5a397e19048f6c9bfa18a5c4c1ad Mon Sep 17 00:00:00 2001 From: Sayyid Hamid Mahdavi Date: Sun, 22 Mar 2026 17:10:59 +0330 Subject: [PATCH] store headers in access token detail --- apps/gooyal_oauth2/validators.py | 167 ++++++++++++++++++++++++++++++- 1 file changed, 165 insertions(+), 2 deletions(-) diff --git a/apps/gooyal_oauth2/validators.py b/apps/gooyal_oauth2/validators.py index 797cddf..4a1f0d7 100755 --- a/apps/gooyal_oauth2/validators.py +++ b/apps/gooyal_oauth2/validators.py @@ -1,22 +1,35 @@ import base64 import binascii import logging +from copy import deepcopy from datetime import datetime, timedelta from urllib.parse import unquote_plus import requests # import service_clients from django.contrib.auth import get_user_model +from django.utils import timezone from django.utils.timezone import make_aware -from oauth2_provider.models import get_access_token_model +from oauth2_provider.exceptions import FatalClientError +from oauth2_provider.models import get_access_token_model, get_application_model, get_id_token_model, get_grant_model, \ + get_refresh_token_model from oauth2_provider.oauth2_validators import OAuth2Validator as BaseOAuth2Validator +from requests import Session + from .settings import oauth2_settings from django.conf import settings +from django.db import router, transaction log = logging.getLogger("oauth2_provider") -AccessTokenModel = get_access_token_model() + UserModel = get_user_model() +Application = get_application_model() +AccessToken = get_access_token_model() +IDToken = get_id_token_model() +Grant = get_grant_model() +RefreshToken = get_refresh_token_model() + class OAuth2Validator(BaseOAuth2Validator): # pylint: disable=w0223 @@ -77,3 +90,153 @@ class OAuth2Validator(BaseOAuth2Validator): # pylint: disable=w0223 def get_oidc_issuer_endpoint(self, request): return oauth2_settings.oidc_issuer(request) + + def save_token(self, token, request, *args, **kwargs): + """Persist the token with a token type specific method. + + Currently, only save_bearer_token is supported. + + :param token: A (Bearer) token dict. + :param request: OAuthlib request. + :type request: oauthlib.common.Request + """ + return self.save_bearer_token(token, request, *args, **kwargs) + + def save_bearer_token(self, token, request, *args, **kwargs): + """ + Save access and refresh token. + + Override _save_bearer_token and not this function when adding custom logic + for the storing of these token. This allows the transaction logic to be + separate from the token handling. + """ + # Use the AccessToken's database instead of making the assumption it is in 'default'. + with transaction.atomic(using=router.db_for_write(AccessToken)): + return self._save_bearer_token(token, request, *args, **kwargs) + + def _save_bearer_token(self, token, request, *args, **kwargs): + """ + Save access and refresh token. + + If refresh token is issued, remove or reuse old refresh token as in rfc:`6`. + + @see: https://rfc-editor.org/rfc/rfc6749.html#section-6 + """ + + if "scope" not in token: + raise FatalClientError("Failed to renew access token: missing scope") + + # expires_in is passed to Server on initialization + # custom server class can have logic to override this + expires = timezone.now() + timedelta( + seconds=token.get( + "expires_in", + oauth2_settings.ACCESS_TOKEN_EXPIRE_SECONDS, + ) + ) + + if request.grant_type == "client_credentials": + request.user = None + + # This comes from OAuthLib: + # https://github.com/idan/oauthlib/blob/1.0.3/oauthlib/oauth2/rfc6749/tokens.py#L267 + # Its value is either a new random code; or if we are reusing + # refresh tokens, then it is the same value that the request passed in + # (stored in `request.refresh_token`) + refresh_token_code = token.get("refresh_token", None) + + if refresh_token_code: + # an instance of `RefreshToken` that matches the old refresh code. + # Set on the request in `validate_refresh_token` + refresh_token_instance = getattr(request, "refresh_token_instance", None) + + # If we are to reuse tokens, and we can: do so + if ( + not self.rotate_refresh_token(request) + and isinstance(refresh_token_instance, RefreshToken) + and refresh_token_instance.access_token + ): + access_token = AccessToken.objects.select_for_update().get( + pk=refresh_token_instance.access_token.pk + ) + access_token.user = request.user + access_token.scope = token["scope"] + access_token.expires = expires + access_token.token = token["access_token"] + access_token.application = request.client + access_token.save() + + # else create fresh with access & refresh tokens + else: + # revoke existing tokens if possible to allow reuse of grant + if isinstance(refresh_token_instance, RefreshToken): + # First, to ensure we don't have concurrency issues, we refresh the refresh token + # from the db while acquiring a lock on it + # We also put it in the "request cache" + refresh_token_instance = RefreshToken.objects.select_for_update().get( + pk=refresh_token_instance.pk + ) + request.refresh_token_instance = refresh_token_instance + + previous_access_token = AccessToken.objects.filter( + source_refresh_token=refresh_token_instance + ).first() + try: + refresh_token_instance.revoke() + except (AccessToken.DoesNotExist, RefreshToken.DoesNotExist): + pass + else: + setattr(request, "refresh_token_instance", None) + else: + previous_access_token = None + + # If the refresh token has already been used to create an + # access token (ie it's within the grace period), return that + # access token + if not previous_access_token: + access_token = self._create_access_token( + expires, + request, + token, + source_refresh_token=refresh_token_instance, + ) + + self._create_refresh_token( + request, refresh_token_code, access_token, refresh_token_instance + ) + else: + # make sure that the token data we're returning matches + # the existing token + token["access_token"] = previous_access_token.token + token["refresh_token"] = ( + RefreshToken.objects.filter(access_token=previous_access_token).first().token + ) + token["scope"] = previous_access_token.scope + + # No refresh token should be created, just access token + else: + self._create_access_token(expires, request, token) + + def _create_access_token(self, expires, request, token, source_refresh_token=None): + id_token = token.get("id_token", None) + if id_token: + id_token = self._load_id_token(id_token) + headers = deepcopy(dict(request.headers)) + for header in dict(request.headers): + if type(headers[header]) not in [str, bool, int, float, tuple, list]: + print(f'unserializable header: {header} -> {headers[header]}') + headers.pop(header) + + return AccessToken.objects.create( + user=request.user, + scope=token["scope"], + expires=expires, + token=token["access_token"], + id_token=id_token, + application=request.client, + source_refresh_token=source_refresh_token, + detail={'headers': headers}, + ) + + +Session \ No newline at end of file