@@ -2,16 +2,16 @@ from ninja import Router from ninja.errors import HttpError from authentication.schemas import UserSchema -from authentication.security import AsyncAuthBearer +from authentication.security import SyncAuthBearer from authentication.selectors.user_selector import UserSelector -router = Router(auth=AsyncAuthBearer(), tags=['auth']) +router = Router(auth=SyncAuthBearer(), tags=['auth']) @router.get('me', tags=['auth/me'], response=UserSchema) -async def get_user_data(request): +def get_user_data(request): """Get user details.""" try: - return await UserSelector.detail(user=request.auth, provider=request.provider) + return UserSelector.detail(user=request.auth, provider=request.provider) except Exception as exc: raise HttpError(400, f'{exc}') @@ -44,7 +44,7 @@ class UserSelector: return users @classmethod - async def detail(cls, user: CustomUserModel, provider: str) -> UserDetailSerializer | CustomUserModel: + def detail(cls, user: CustomUserModel, provider: str) -> UserDetailSerializer | CustomUserModel: # TODO: refactor this hook if user.account_type in ( 'business_account', @@ -55,12 +55,13 @@ class UserSelector: if not user.show_balance: user.payment_plan.current_token_balance = Decimal('0') user.payment_plan.plan = user.business_account.parent_company.user.payment_plan.plan - refresh_token = await OutstandingToken.objects.filter(user=user).order_by('-created_at').afirst() + refresh_token = OutstandingToken.objects.filter(user=user).order_by('-created_at').first() if not refresh_token or refresh_token.expires_at < timezone.now() and provider == 'yandex': - access_token = await sync_to_async(lambda: RefreshToken.for_user(user))() + refresh = RefreshToken.for_user(user) else: - access_token = await sync_to_async(lambda: RefreshToken(token=refresh_token.token).access_token)() - user.token = {'access': str(access_token), 'refresh': str(access_token.token)} + refresh = refresh_token.token + access_token = RefreshToken(token=refresh).access_token + user.token = {'access': str(access_token), 'refresh': str(refresh)} return user def list_social_accounts(self, serialize: bool = False): @@ -54,15 +54,44 @@ class SimpleJWTScheme(BaseSimpleJWTScheme): class SyncAuthBearer(HttpBearer): - def authenticate(self, _: HttpRequest, token: str) -> Any | None: + def authenticate(self, request: HttpRequest, token: str) -> Any | None: try: user_payload = async_to_sync(TokenService.decode)(token=token) - return CustomUserModel.objects.get( + request.provider = 'air' + return CustomUserModel.objects.select_related( + 'payment_plan', + 'payment_plan__plan', + 'business_account', + 'business_account__group', + 'business_account__parent_company__user__payment_plan', + 'business_account__parent_company__user__payment_plan__plan', + 'host_account' + ).prefetch_related( + 'payment_plan__plan__accessed_models', + 'business_account__parent_company__user__payment_plan__plan__accessed_models', + 'social_auth' + ).get( **{key: user_payload[f'{key}'] for key in settings.JWT_SETTINGS['encode_attributes']} ) except InvalidToken: - access = AccessToken.objects.prefetch_related('user').get(token=token) - return access.user + access = AccessToken.objects.select_related( + 'user', + 'user__payment_plan', + 'user__payment_plan__plan', + 'user__business_account', + 'user__business_account__group', + 'user__business_account__parent_company__user__payment_plan', + 'user__business_account__parent_company__user__payment_plan__plan', + 'user__host_account' + ).prefetch_related( + 'user__payment_plan__plan__accessed_models', + 'user__business_account__parent_company__user__payment_plan__plan__accessed_models', + 'user__social_auth' + ).filter(token=token).order_by('-created') + request.provider = 'yandex' + if a := access.first(): + return a.user + raise HttpError(401, _('Access token expired or does not exist')) class AsyncAuthBearer(HttpBearer): @@ -87,7 +116,7 @@ class AsyncAuthBearer(HttpBearer): ) except InvalidToken: try: - access = await AccessToken.objects.prefetch_related( + access = await AccessToken.objects.select_related( 'user', 'user__payment_plan', 'user__payment_plan__plan', @@ -1,21 +1,20 @@ from decimal import Decimal -from asgiref.sync import sync_to_async from ninja import Router from ninja.errors import HttpError -from authentication.security import AsyncAuthBearer +from authentication.security import SyncAuthBearer from payments.schema import UserBalance from payments.selectors.payment_plan_selector import PaymentPlanSelector -router = Router(auth=AsyncAuthBearer(), tags=['payments']) +router = Router(auth=SyncAuthBearer(), tags=['payments']) @router.get('user-balance', tags=['payments/user-balance'], response=UserBalance) -async def get_user_balance(request): +def get_user_balance(request): """Get user balance.""" try: - balance = await sync_to_async(PaymentPlanSelector(request.auth).get_current_balance)() + balance = PaymentPlanSelector(request.auth).get_current_balance() current_balance = Decimal( f'{balance:.2f}' if balance == balance.to_integral() else balance.normalize().to_eng_string() )