diff --git a/fastapi_users/authentication/__init__.py b/fastapi_users/authentication/__init__.py index c7874ccd..eba64491 100644 --- a/fastapi_users/authentication/__init__.py +++ b/fastapi_users/authentication/__init__.py @@ -47,6 +47,38 @@ class Authenticator: self.backends = backends self.user_db = user_db + self.get_current_user = self.get_auth_dependency(required=True) + self.get_current_active_user = self.get_auth_dependency( + required=True, active=True + ) + self.get_current_verified_user = self.get_auth_dependency( + required=True, active=True, verified=True + ) + self.get_current_superuser = self.get_auth_dependency( + required=True, active=True, superuser=True + ) + self.get_current_verified_superuser = self.get_auth_dependency( + required=True, active=True, verified=True, superuser=True + ) + self.get_optional_current_user = self.get_auth_dependency() + self.get_optional_current_active_user = self.get_auth_dependency(active=True) + self.get_optional_current_verified_user = self.get_auth_dependency( + active=True, verified=True + ) + self.get_optional_current_superuser = self.get_auth_dependency( + active=True, superuser=True + ) + self.get_optional_current_verified_superuser = self.get_auth_dependency( + active=True, verified=True, superuser=True + ) + + def get_auth_dependency( + self, + required: bool = False, + active: bool = False, + verified: bool = False, + superuser: bool = False, + ): # Here comes some blood magic 🧙‍♂️ # Thank to "makefun", we are able to generate callable # with a dynamic number of dependencies at runtime. @@ -64,100 +96,46 @@ class Authenticator: except ValueError: raise DuplicateBackendNamesError() - @with_signature(signature, func_name="get_optional_current_user") - async def get_optional_current_user(*args, **kwargs): - return await self._authenticate(*args, **kwargs) + @with_signature(signature) + async def auth_dependency(*args, **kwargs): + return await self._authenticate( + *args, + required=required, + active=active, + verified=verified, + superuser=superuser, + **kwargs + ) - @with_signature(signature, func_name="get_optional_current_active_user") - async def get_optional_current_active_user(*args, **kwargs): - user = await get_optional_current_user(*args, **kwargs) - if not user or not user.is_active: - return None - return user + return auth_dependency - @with_signature(signature, func_name="get_optional_current_verified_user") - async def get_optional_current_verified_user(*args, **kwargs): - user = await get_optional_current_active_user(*args, **kwargs) - if not user or not user.is_verified: - return None - return user - - @with_signature(signature, func_name="get_optional_current_superuser") - async def get_optional_current_superuser(*args, **kwargs): - user = await get_optional_current_active_user(*args, **kwargs) - if not user or not user.is_superuser: - return None - return user - - @with_signature(signature, func_name="get_optional_current_verified_superuser") - async def get_optional_current_verified_superuser(*args, **kwargs): - user = await get_optional_current_verified_user(*args, **kwargs) - if not user or not user.is_superuser: - return None - return user - - @with_signature(signature, func_name="get_current_user") - async def get_current_user(*args, **kwargs): - user = await get_optional_current_user(*args, **kwargs) - if user is None: - raise self._get_credentials_exception() - return user - - @with_signature(signature, func_name="get_current_active_user") - async def get_current_active_user(*args, **kwargs): - user = await get_optional_current_active_user(*args, **kwargs) - if user is None: - raise self._get_credentials_exception() - return user - - @with_signature(signature, func_name="get_current_verified_user") - async def get_current_verified_user(*args, **kwargs): - user = await get_optional_current_verified_user(*args, **kwargs) - if user is None: - raise self._get_credentials_exception() - return user - - @with_signature(signature, func_name="get_current_superuser") - async def get_current_superuser(*args, **kwargs): - user = await get_optional_current_active_user(*args, **kwargs) - if user is None: - raise self._get_credentials_exception() - if not user.is_superuser: - raise self._get_credentials_exception(status.HTTP_403_FORBIDDEN) - return user - - @with_signature(signature, func_name="get_current_verified_superuser") - async def get_current_verified_superuser(*args, **kwargs): - user = await get_optional_current_verified_user(*args, **kwargs) - if user is None: - raise self._get_credentials_exception() - if not user.is_superuser: - raise self._get_credentials_exception(status.HTTP_403_FORBIDDEN) - return user - - self.get_current_user = get_current_user - self.get_current_active_user = get_current_active_user - self.get_current_verified_user = get_current_verified_user - self.get_current_superuser = get_current_superuser - self.get_current_verified_superuser = get_current_verified_superuser - self.get_optional_current_user = get_optional_current_user - self.get_optional_current_active_user = get_optional_current_active_user - self.get_optional_current_verified_user = get_optional_current_verified_user - self.get_optional_current_superuser = get_optional_current_superuser - self.get_optional_current_verified_superuser = ( - get_optional_current_verified_superuser - ) - - async def _authenticate(self, *args, **kwargs) -> Optional[BaseUserDB]: + async def _authenticate( + self, + *args, + required: bool = False, + active: bool = False, + verified: bool = False, + superuser: bool = False, + **kwargs + ) -> Optional[BaseUserDB]: + user: Optional[BaseUserDB] = None for backend in self.backends: token: str = kwargs[name_to_variable_name(backend.name)] if token: user = await backend(token, self.user_db) - if user is not None: - return user - return None + if user: + break - def _get_credentials_exception( - self, status_code: int = status.HTTP_401_UNAUTHORIZED - ) -> HTTPException: - return HTTPException(status_code=status_code) + status_code = status.HTTP_401_UNAUTHORIZED + if user: + if active and not user.is_active: + user = None + elif verified and not user.is_verified: + user = None + elif superuser and not user.is_superuser: + status_code = status.HTTP_403_FORBIDDEN + user = None + + if not user and required: + raise HTTPException(status_code=status_code) + return user