@@ -80,3 +80,4 @@ from ml_model.services.vicuna import Vicuna from ml_model.services.wan import Wan from ml_model.services.wan_lite import Wan_Lite from ml_model.services.whisper import Whisper +from ml_model.services.audio_test_model import Audio_Test_Model \ No newline at end of file @@ -0,0 +1,97 @@ +import time +from datetime import timedelta +from decimal import Decimal +from io import BytesIO +from typing import Any +import filetype + +import requests +from django.core.files import File + +from messages.models import Message +from ml_model.exceptions import CorruptedFileError, InvalidParameterError +from ml_model.services.base import SimpleService +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + +import random + +class Audio_Test_Model(SimpleService): + + TOKENS_COST = Decimal('3') + + PLACEHOLDER_URL=[ + 'https://www.myinstants.com/media/sounds/saliut-eblany-batia-doma-billy-butcher-i-the-boys.mp3', + 'https://www.myinstants.com/media/sounds/zdravstvuite-nichtozhnye-nishchie-smertnye.mp3', + 'https://www.myinstants.com/media/sounds/okh-zria-ia-tuda-polez.mp3'] # Позже убрать + + def calculate_price(self, num_audios: int = 1) -> Decimal: + return self.TOKENS_COST * num_audios + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + num_audios = info.get('num_audios', 1) + + return cls.TOKENS_COST * num_audios + + def save_results( + self, + content: str, + t: timedelta, + audios: list[bytes], + save: bool = True, + ) -> list[Message]: + messages: list[Message] = [] + for audio in audios: + messages.append( + Message( + content_object=self.store, + elapsed_time=t, + content=content, + file=File(BytesIO(audio), '.mp3'), + ) + ) + if save: + return Message.objects.bulk_create(messages) + return messages + + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + cau = input_message.info.get('cau') or self.PLACEHOLDER_URL[random.randint(0, 2)] # Позже убрать + num_audios = input_message.info.get('num_audios', 1) + + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < (cost := self.calculate_price(num_audios)): + raise InsufficientBalance(balance, cost) + + start_time = time.time() + + audio_bytes = self._fetch_audio(cau) + + audios = [audio_bytes] * num_audios + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, num_audios) + + msgs = self.save_results(input_message.content, process_time, audios, save) + return msgs + + + def _fetch_audio(self, url: str): + try: + response = requests.get( + url, + timeout=600 + ) + response.raise_for_status() + except requests.RequestException as exc: + raise InvalidParameterError(f'Invalid audio URL: {exc}') + + kind = filetype.guess(response.content[:120]) + + if not kind: + raise CorruptedFileError + + if not kind.mime.startswith('audio/'): + raise InvalidParameterError('Audio format not supported') + + return response.content \ No newline at end of file