@@ -1,4 +1,5 @@ import random +from decimal import Decimal from typing import TYPE_CHECKING, Optional from uuid import uuid4 @@ -17,6 +18,7 @@ from django.utils.translation import gettext_lazy as _ from authentication.models.utm import UTM from core.models import BaseModel +from payments.utils import get_remaining_token_limit if TYPE_CHECKING: from payments.models.referral_account import ( @@ -154,7 +156,27 @@ class CustomUserModel(AbstractBaseUser, PermissionsMixin, BaseModel): @property def balance(self): - return self.payment_plan.current_token_balance + if ( + self.account_type in ('business_account', 'business_admin', 'business_security') + ) and self.business_account.acceptance_status == 'accepted': + balance = ( + self.business_account.token_limit + if self.business_account.token_limit + else self.business_account.group.token_limit + if self.business_account.group and self.business_account.group.token_limit + else None + ) + if balance is None: + balance = self.business_account.parent_company.user.payment_plan.current_token_balance + else: + balance = get_remaining_token_limit(self, balance) + else: + balance = self.payment_plan.current_token_balance + + if balance < 0: + balance = Decimal('0') + + return balance @property def account_type(self): @@ -11,6 +11,7 @@ from authentication.models import ( ) from authentication.models.choices import InvitationStatus from authentication.serializers import ChangePasswordSerializer +from payments.utils import update_remaining_token_limit, delete_remaining_token_limit class BusinessAccountService: @@ -87,6 +88,14 @@ class BusinessAccountService: self.update_status(InvitationStatus.REJECTED) def update_limit(self, new_balance: Decimal | None): + if self.account.token_limit is not None and new_balance is not None: + additional_token_limit = self.account.token_limit - new_balance + update_remaining_token_limit(self.account.user, additional_token_limit) + elif self.account.group and self.account.group.token_limit and new_balance is not None: + additional_token_limit = self.account.group.token_limit - new_balance + update_remaining_token_limit(self.account.user, additional_token_limit) + elif not self.account.group and new_balance is None: + delete_remaining_token_limit(self.account.user) self.account.token_limit = new_balance self.account.save() @@ -161,7 +161,7 @@ class BusinessHostService: user = UserSelector.get_by_email(user_email) if not AccountStatusSelector(user).is_business_account(): raise Exception("Business Account for this user doesn't exist") - if token_limit := serializer.validated_data.get('token_limit', None): + if (token_limit := serializer.validated_data.get('token_limit', False)) is not False: self.update_token_limit( user, token_limit, @@ -293,7 +293,7 @@ class ChangeInvitationStatusSerializer(serializers.Serializer): class AccountDataUpdateSerializer(serializers.Serializer): status = serializers.ChoiceField(choices=InvitationStatus.choices, required=False) - token_limit = serializers.DecimalField(max_digits=50, decimal_places=2, default=None) + token_limit = serializers.DecimalField(max_digits=50, decimal_places=2, allow_null=True, required=False) account_privileges = serializers.ChoiceField( choices=AccountPrivileges.choices, default=AccountPrivileges.REGULAR ) @@ -94,6 +94,7 @@ from authentication.services.email_token_service import EmailTokenService from authentication.services.user_services import EmailService, UserService from messages.models.message import Message from payments.models.invoice import Invoice +from payments.utils import set_remaining_token_limit, delete_remaining_token_limit from tools.chats.models import Chat from tools.copywrite.models import Copywrite from tools.media.models import Image @@ -788,15 +789,22 @@ class BusinessGroupAPIView(APIView): class BusinessGroupAccountsAPIView(APIView): def post(self, request, group_id: UUID, acc_email: UUID, *args, **kwargs): - group = BusinessGroup.objects.get(uid=group_id) business_account = BusinessAccountService.from_user(UserSelector.get_by_email(acc_email)).account + if business_account.group: + return Response({'detail': _('User is already a member of the group')}, status.HTTP_400_BAD_REQUEST) + group = BusinessGroup.objects.get(uid=group_id) business_account.group = group business_account.save() + if business_account.token_limit is None and group.token_limit is not None: + set_remaining_token_limit(business_account.user, group.token_limit) return Response(status=status.HTTP_200_OK) def delete(self, request, group_id: UUID, acc_email: UUID, *args, **kwargs): business_account = BusinessAccountService.from_user(UserSelector.get_by_email(acc_email)).account business_account.group = None + group = BusinessGroup.objects.get(uid=group_id) + if business_account.token_limit is None and group.token_limit is not None: + delete_remaining_token_limit(business_account.user) business_account.save() return Response(status=status.HTTP_204_NO_CONTENT) @@ -292,6 +292,10 @@ CELERY_BEAT_SCHEDULE = { 'task': 'payments.tasks.send_low_balance_message', 'schedule': crontab(0, 8), }, + 'delete_remaining_tokens_cache': { + 'task': 'payments.tasks.delete_remaining_tokens_cache', + 'schedule': crontab(0, 0, 1) + } } CACHES = { @@ -491,7 +491,7 @@ class Chatgpt(SimpleService): model: str = 'gpt-3.5-turbo', embedding_tokens: int = 0 ): - balance = PaymentPlanSelector(self.store.user).get_current_balance() + balance = self.store.user.balance total_tokens = input_tokens if image_size: total_tokens += self.count_image_tokens(image_size) @@ -39,7 +39,7 @@ class Ray(SimpleService): return [msg] def make(self, input_message: Message, save: bool = True) -> list[Message]: - if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < self.TOKENS_COST: + if (balance := self.store.user.balance) < self.TOKENS_COST: raise InsufficientBalance(balance, self.TOKENS_COST) callback_data = dict({'prompt': self.translate_prompt(input_message.content), **input_message.info}) if input_message.file: @@ -38,7 +38,7 @@ class Veo(SimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: version = input_message.info.pop('version', 'veo-3-fast') - if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < self.TOKENS_COST[version]: + if (balance := self.store.user.balance) < self.TOKENS_COST[version]: raise InsufficientBalance(balance, self.TOKENS_COST[version]) callback_data = dict({'prompt': input_message.content, **input_message.info}) if input_message.file: @@ -41,7 +41,7 @@ class Wan(SimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: resolution = input_message.info.pop('resolution', '720p') - if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < self.TOKENS_COST[resolution]: + if (balance := self.store.user.balance) < self.TOKENS_COST[resolution]: raise InsufficientBalance(balance, self.TOKENS_COST[resolution]) callback_data = dict({'prompt': self.translate_prompt(input_message.content), 'resolution': resolution, **input_message.info}) start_time = time.time() @@ -5,7 +5,6 @@ from ninja.errors import HttpError from authentication.security import SyncAuthBearer from payments.schema import UserBalance -from payments.selectors.payment_plan_selector import PaymentPlanSelector router = Router(auth=SyncAuthBearer(), tags=['payments']) @@ -14,7 +13,7 @@ router = Router(auth=SyncAuthBearer(), tags=['payments']) def get_user_balance(request): """Get user balance.""" try: - balance = PaymentPlanSelector(request.auth).get_current_balance() + balance = request.auth.balance current_balance = Decimal( f'{balance:.2f}' if balance == balance.to_integral() else balance.normalize().to_eng_string() ) @@ -58,22 +58,6 @@ class PaymentPlanSelector: return PaymentPlanSerializer(plan) - def get_current_balance(self) -> Decimal: - if ( - self.user.account_type in ('business_account', 'business_admin', 'business_security') - ) and self.user.business_account.acceptance_status == InvitationStatus.ACCEPTED: - balance = ( - self.user.business_account.group.token_limit - if self.user.business_account.group - else self.user.business_account.token_limit - ) - if balance is None: - balance = self.user.business_account.parent_company.user.payment_plan.current_token_balance - else: - balance = self.user.payment_plan.current_token_balance - - return balance - def get_user_balance(self): user_type = UserSelector(self.user).check_account_type() if ( @@ -1,10 +1,9 @@ from decimal import Decimal -from django.utils.translation import gettext_lazy as _ - from authentication.models.user import CustomUserModel -from authentication.selectors.user_selector import UserSelector + from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.utils import update_remaining_token_limit, get_remaining_token_limit class ModelBillingService: @@ -12,36 +11,32 @@ class ModelBillingService: self.user = user def charge(self, amount: Decimal): - user_type = UserSelector(self.user).check_account_type() - if user_type in ['regular', 'business_host']: - plan = self.user.payment_plan - allowance = None - elif user_type in [ + user_type = self.user.account_type + if user_type in ( 'business_account', 'business_admin', 'business_security', - ]: + ) and self.user.business_account.acceptance_status == 'accepted': plan = self.user.business_account.parent_company.user.payment_plan - allowance = ( + balance = ( self.user.business_account.token_limit - if self.user.business_account.group is None + if self.user.business_account.token_limit else self.user.business_account.group.token_limit + if self.user.business_account.group and self.user.business_account.group.token_limit + else None ) + token_limit = get_remaining_token_limit(self.user, balance) if balance else balance else: - raise Exception(_('Unknown account type')) + plan = self.user.payment_plan + token_limit = None if plan.current_token_balance < amount: raise InsufficientBalance(plan.current_token_balance, amount) - if (allowance is not None) and (allowance < amount): - raise InsufficientBalance(allowance, amount) + if (token_limit is not None) and (token_limit < amount): + raise InsufficientBalance(token_limit, amount) + if token_limit is not None: + update_remaining_token_limit(self.user, amount) plan.current_token_balance -= amount - if allowance is not None: - if self.user.business_account.group is not None: - self.user.business_account.group.token_limit -= amount - self.user.business_account.group.save() - else: - self.user.business_account.token_limit -= amount - self.user.business_account.save() plan.save() @@ -3,6 +3,7 @@ from uuid import UUID from celery import shared_task from celery.utils.log import get_task_logger +from django.core.cache import caches from django.db.models import F from authentication.models.business_host import BusinessUserHost @@ -32,3 +33,11 @@ def send_low_balance_message(): def withdraw(user_id: UUID, amount: Decimal): user = CustomUserModel.objects.get(uid=user_id) PaymentPlanService(user).update_per_token_plan_details(amount) + + +@shared_task +def delete_remaining_tokens_cache(): + cache = caches['default'] + keys = cache.keys('remaining_token_limit:*') + if keys: + cache.delete(*keys) \ No newline at end of file @@ -0,0 +1,28 @@ +from decimal import Decimal + +from django.core.cache import caches + +cache = caches['default'] + + +def get_remaining_token_limit(user, default=None) -> Decimal: + remaining_token_limit = cache.get(f'remaining_token_limit:{user.uid}') + if remaining_token_limit is None: + if default: + remaining_token_limit = default + set_remaining_token_limit(user, default) + return remaining_token_limit + + +def set_remaining_token_limit(user, remaining_token_limit) -> None: + cache.set(f'remaining_token_limit:{user.uid}', remaining_token_limit) + + +def update_remaining_token_limit(user, tokens: Decimal) -> None: + remaining_token_limit = get_remaining_token_limit(user) + remaining_token_limit -= tokens + set_remaining_token_limit(user, remaining_token_limit) + + +def delete_remaining_token_limit(user) -> None: + cache.delete(f'remaining_token_limit:{user.uid}') @@ -147,7 +147,7 @@ class UserPlanAPIView(APIView): def get(self, request, *args, **kwargs): """Get user balance""" try: - result = PaymentPlanSelector(self.request.user).get_current_balance() + result = request.user.balance return Response({'current_token_balance': result}, status=status.HTTP_200_OK) except Exception as err: logger.exception(err)