store headers in access token detail
This commit is contained in:
parent
4d5f1983c7
commit
014c674961
1 changed files with 165 additions and 2 deletions
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue