diff --git a/esteid/mixins.py b/esteid/mixins.py index eafd62f..39f5243 100644 --- a/esteid/mixins.py +++ b/esteid/mixins.py @@ -1,7 +1,8 @@ import json import logging +from contextvars import ContextVar from http import HTTPStatus -from typing import Callable, TYPE_CHECKING +from typing import Callable, Optional, TYPE_CHECKING from django.core.exceptions import ValidationError as DjangoValidationError from django.http import Http404, HttpRequest, JsonResponse, QueryDict @@ -30,6 +31,14 @@ class RequestType(HttpRequest): logger = logging.getLogger(__name__) +_current_request: ContextVar[Optional[HttpRequest]] = ContextVar("_current_request", default=None) + + +def get_current_request() -> Optional[HttpRequest]: + """Public getter for the current request stored in context variable.""" + return _current_request.get() + + class SessionViewMixin: """ Provides POST and PATCH method handlers for auth/signing session management. @@ -81,19 +90,25 @@ def post(self, request, *args, **kwargs): """ Handles session start requests """ + token = _current_request.set(request) try: return self.start_session(request, *args, **kwargs) except Exception as e: return self.handle_errors(e, stage="start") + finally: + _current_request.reset(token) def patch(self, request, *args, **kwargs): """ Handles session finish requests """ + token = _current_request.set(request) try: return self.finish_session(request, *args, **kwargs) except Exception as e: return self.handle_errors(e, stage="finish") + finally: + _current_request.reset(token) class DjangoRestCompatibilityMixin(SessionViewMixin):