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 base64
|
||||||
import binascii
|
import binascii
|
||||||
import logging
|
import logging
|
||||||
|
from copy import deepcopy
|
||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
from urllib.parse import unquote_plus
|
from urllib.parse import unquote_plus
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
# import service_clients
|
# import service_clients
|
||||||
from django.contrib.auth import get_user_model
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.utils import timezone
|
||||||
from django.utils.timezone import make_aware
|
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 oauth2_provider.oauth2_validators import OAuth2Validator as BaseOAuth2Validator
|
||||||
|
from requests import Session
|
||||||
|
|
||||||
from .settings import oauth2_settings
|
from .settings import oauth2_settings
|
||||||
from django.conf import settings
|
from django.conf import settings
|
||||||
|
from django.db import router, transaction
|
||||||
|
|
||||||
log = logging.getLogger("oauth2_provider")
|
log = logging.getLogger("oauth2_provider")
|
||||||
|
|
||||||
AccessTokenModel = get_access_token_model()
|
|
||||||
UserModel = get_user_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
|
class OAuth2Validator(BaseOAuth2Validator): # pylint: disable=w0223
|
||||||
|
|
@ -77,3 +90,153 @@ class OAuth2Validator(BaseOAuth2Validator): # pylint: disable=w0223
|
||||||
|
|
||||||
def get_oidc_issuer_endpoint(self, request):
|
def get_oidc_issuer_endpoint(self, request):
|
||||||
return oauth2_settings.oidc_issuer(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