@@ -12,6 +12,9 @@ from messages.models import Message from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + class Veo(SimpleService): TOKENS_COST = { @@ -35,6 +38,8 @@ 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]: + raise InsufficientBalance(balance, self.TOKENS_COST[version]) callback_data = dict({'prompt': input_message.content, **input_message.info}) if input_message.file: kind = filetype.guess(input_message.file.read(20)) @@ -10,6 +10,9 @@ from messages.models import Message from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + class Wan(SimpleService): """ @@ -38,7 +41,9 @@ class Wan(SimpleService): def make(self, input_message: Message, save: bool = True) -> list[Message]: resolution = input_message.info.pop('resolution', '720p') - callback_data = dict({'prompt': input_message.content, 'resolution': resolution, **input_message.info}) + if (balance := PaymentPlanSelector(self.store.user).get_current_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() video = replicate_run('wan-video/wan-2.2-t2v-fast', callback_data) process_time = timedelta(seconds=(time.time() - start_time)) @@ -1,10 +1,12 @@ import logging import sys +from django.utils.translation import gettext_lazy as _ + from drf_spectacular.utils import OpenApiParameter, extend_schema from rest_framework.permissions import IsAuthenticated from rest_framework.response import Response -from rest_framework.status import HTTP_402_PAYMENT_REQUIRED +from rest_framework.status import HTTP_402_PAYMENT_REQUIRED, HTTP_400_BAD_REQUEST from rest_framework.views import APIView from messages.models import Message @@ -123,7 +125,7 @@ class MediaAPIView(APIView): """Create new media content (image, video, audio) message.""" serializer = MessageSerializer(data=request.data) if serializer.is_valid(): - gallery, _ = self.manager.objects.get_or_create( + gallery, created = self.manager.objects.get_or_create( user=request.user, model__slug=model, defaults={ @@ -150,7 +152,14 @@ class MediaAPIView(APIView): input_message.save() if isinstance(exc, InsufficientBalance): return Response({'detail': f'{exc}'}, status=HTTP_402_PAYMENT_REQUIRED) - return Response(f'Error: {exc}', status=400) + return Response( + { + 'detail': _( + 'Error occured when create generation. It may cause NSFW-content not allowed, retry again' + ) + }, + status=HTTP_400_BAD_REQUEST, + ) return Response(MessageSerializer(output_messages, many=True).data, 201) else: return Response(data=serializer.errors, status=400)