diff --git a/api/src/backend/api/base_views.py b/api/src/backend/api/base_views.py index 605afe332d..f8ed8f7c47 100644 --- a/api/src/backend/api/base_views.py +++ b/api/src/backend/api/base_views.py @@ -52,6 +52,8 @@ class BaseRLSViewSet(BaseViewSet): return super().dispatch(request, *args, **kwargs) def initial(self, request, *args, **kwargs): + super().initial(request, *args, **kwargs) + # Ideally, this logic would be in the `.setup()` method but DRF view sets don't call it # https://docs.djangoproject.com/en/5.1/ref/class-based-views/base/#django.views.generic.base.View.setup if request.auth is None: @@ -61,9 +63,25 @@ class BaseRLSViewSet(BaseViewSet): if tenant_id is None: raise NotAuthenticated("Tenant ID is not present in token") +<<<<<<< HEAD with rls_transaction(tenant_id): self.request.tenant_id = tenant_id return super().initial(request, *args, **kwargs) +======= + self.request.tenant_id = tenant_id + + self._rls_cm = rls_transaction(tenant_id) + self._rls_cm.__enter__() + + def finalize_response(self, request, response, *args, **kwargs): + response = super().finalize_response(request, response, *args, **kwargs) + + if hasattr(self, "_rls_cm"): + self._rls_cm.__exit__(None, None, None) + del self._rls_cm + + return response +>>>>>>> 6e7a32cb5 (revert(views): calling order to initial view method (#7921)) def get_serializer_context(self): context = super().get_serializer_context() @@ -109,6 +127,8 @@ class BaseTenantViewset(BaseViewSet): pass # Tenant might not exist, handle gracefully def initial(self, request, *args, **kwargs): + super().initial(request, *args, **kwargs) + if request.auth is None: raise NotAuthenticated @@ -117,8 +137,23 @@ class BaseTenantViewset(BaseViewSet): raise NotAuthenticated("Tenant ID is not present in token") user_id = str(request.user.id) +<<<<<<< HEAD with rls_transaction(value=user_id, parameter=POSTGRES_USER_VAR): return super().initial(request, *args, **kwargs) +======= + + self._rls_cm = rls_transaction(value=user_id, parameter=POSTGRES_USER_VAR) + self._rls_cm.__enter__() + + def finalize_response(self, request, response, *args, **kwargs): + response = super().finalize_response(request, response, *args, **kwargs) + + if hasattr(self, "_rls_cm"): + self._rls_cm.__exit__(None, None, None) + del self._rls_cm + + return response +>>>>>>> 6e7a32cb5 (revert(views): calling order to initial view method (#7921)) class BaseUserViewset(BaseViewSet): @@ -127,9 +162,11 @@ class BaseUserViewset(BaseViewSet): return super().dispatch(request, *args, **kwargs) def initial(self, request, *args, **kwargs): + super().initial(request, *args, **kwargs) + # TODO refactor after improving RLS on users if request.stream is not None and request.stream.method == "POST": - return super().initial(request, *args, **kwargs) + return if request.auth is None: raise NotAuthenticated @@ -137,6 +174,22 @@ class BaseUserViewset(BaseViewSet): if tenant_id is None: raise NotAuthenticated("Tenant ID is not present in token") +<<<<<<< HEAD with rls_transaction(tenant_id): self.request.tenant_id = tenant_id return super().initial(request, *args, **kwargs) +======= + self.request.tenant_id = tenant_id + + self._rls_cm = rls_transaction(tenant_id) + self._rls_cm.__enter__() + + def finalize_response(self, request, response, *args, **kwargs): + response = super().finalize_response(request, response, *args, **kwargs) + + if hasattr(self, "_rls_cm"): + self._rls_cm.__exit__(None, None, None) + del self._rls_cm + + return response +>>>>>>> 6e7a32cb5 (revert(views): calling order to initial view method (#7921))