@@ -1,13 +1,19 @@ import re import subprocess import zipfile +from uuid import UUID + import docx2txt import fitz import openpyxl from io import BytesIO +from django.db.models.fields.files import FieldFile + +from authentication.models import CustomUserModel from ml_model.exceptions import UnrecognizedFileError +from tools.media.models import Voice, Preset class FileProcessingService: @@ -85,3 +91,19 @@ class FileProcessingService: else: return 'Файл пуст или содержит изображения, из которых невозможно извлечь текст.' + @classmethod + def get_voice_file( + cls, + voice_id: int | None, + preset_id: UUID | None, + user: CustomUserModel, + default_voice_slug: str = 'russian_1', + ) -> FieldFile: + if voice_id: + voice = Voice.objects.get(pk=voice_id, user=user) + elif preset_id: + voice = Preset.objects.get(uid=preset_id) + else: + voice = Preset.objects.get(slug=default_voice_slug) + return voice.file + @@ -1,28 +1,38 @@ import time +from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import timedelta from decimal import Decimal from io import BytesIO from typing import Any +import filetype import requests +from backend import settings from django.core.files import File -from replicate.exceptions import ModelError from messages.models import Message -from ml_model.exceptions import GenerationException, PromptLengthExceeded, RequestBlocked +from ml_model.exceptions import ( + GenerationException, + CorruptedFileError, + FileExtensionNotSupported, +) +from ml_model.services.EmbeddingService import EmbeddingService +from ml_model.services.FileService import FileProcessingService 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 +from pydub import AudioSegment class Elevenlabs(SimpleService): TOKENS_PER_1K_CHARS = Decimal('6') - MAX_CONTENT_LENGTH = 2000 @classmethod def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: - chars = min(len(content), cls.MAX_CONTENT_LENGTH) + if file_exists: + return None + chars = len(content) price = cls.TOKENS_PER_1K_CHARS * Decimal(chars) / Decimal(1000) return price.quantize(Decimal('0.1'), rounding='ROUND_UP') @@ -31,38 +41,85 @@ class Elevenlabs(SimpleService): price = self.TOKENS_PER_1K_CHARS * Decimal(chars) / Decimal(1000) return price.quantize(Decimal('0.1'), rounding='ROUND_UP') - def save_results(self, content: str, t: timedelta, audio_url: str, save: bool = True) -> list[Message]: + def save_results( + self, content: str, t: timedelta, audio_file: BytesIO, save: bool = True + ) -> list[Message]: msg = Message( content=content, content_object=self.store, elapsed_time=t, - file=File(BytesIO(requests.get(audio_url).content), '.mp3'), + file=File(audio_file, name='result.mp3'), ) if save: return Message.objects.bulk_create([msg]) return [msg] - def make(self, input_message: Message, save: bool = True) -> list[Message]: - if len(input_message.content) > self.MAX_CONTENT_LENGTH: - raise PromptLengthExceeded(max_length=self.MAX_CONTENT_LENGTH) - cost = self.calculate_price(input_message.content) - if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < cost: - raise InsufficientBalance(balance, cost) - callback_data = { - 'text': input_message.content, - 'mode': 'voice_clone', - 'reference_audio': input_message.file.url, - } - if transcription := input_message.info.get('transcription', ''): - callback_data.update({'reference_text': transcription}) - start_time = time.time() + def _run_one_chunk(self, payload: dict[str, Any]) -> str | None: try: - result = replicate_run('qwen/qwen3-tts', callback_data) - except ModelError as exc: - if any(error in str(exc) for error in ('E005', 'E006', 'sexual')): - raise RequestBlocked - raise GenerationException from exc + return replicate_run('qwen/qwen3-tts', payload) + except Exception: + return None + def make(self, input_message: Message, save: bool = True) -> list[Message]: + raw_text = input_message.content + file_service = FileProcessingService + balance = PaymentPlanSelector(self.store.user).get_current_balance() + if file := input_message.file: + file_bytes = file.read() + kind = filetype.guess(file_bytes[:20]) + if not kind: + raise CorruptedFileError + raw_file_extension = kind.extension + file_extension = file_service.get_file_extension(raw_file_extension, file_bytes) + if file_extension in ('pdf', 'doc', 'docx'): + raw_text = ( + file_service.get_file_data(file_extension, file_bytes) + .replace('\n', ' ') + ) + else: + raise FileExtensionNotSupported(['PDF', 'DOC', 'DOCX']) + max_affordable_chars = max( + int((balance * 1000) / self.TOKENS_PER_1K_CHARS) + - int((Decimal('0.1') * 1000) / self.TOKENS_PER_1K_CHARS), + 0, + ) + text_chunks = EmbeddingService.split_text_to_chunks( + raw_text[:max_affordable_chars], + chunk_size=1000, + overlap=0, + ) + if not text_chunks: + raise InsufficientBalance(balance, Decimal('1')) + reference_audio = file_service.get_voice_file( + voice_id=input_message.info.get('voice_id'), + preset_id=input_message.info.get('preset_id'), + user=self.store.user, + ) + callback_data = {'mode': 'voice_clone', 'reference_audio': reference_audio.url} + total_cost = self.calculate_price(''.join(text_chunks)) + if balance < total_cost: + raise InsufficientBalance(balance, total_cost) + start_time = time.time() + chunk_results: list[tuple[int, str, str | None]] = [] + with ThreadPoolExecutor(max_workers=settings.MAX_THREADS) as executor: + future_to_data = { + executor.submit(self._run_one_chunk, callback_data | {'text': chunk}): (index, chunk) + for index, chunk in enumerate(text_chunks) + } + for future in as_completed(future_to_data): + index, chunk = future_to_data[future] + chunk_results.append((index, chunk, future.result())) + chunk_results.sort(key=lambda item: item[0]) + parts = [url for _, _, url in chunk_results if url] + final_text = ''.join(chunk for _, chunk, url in chunk_results if url) + if not parts: + raise GenerationException + audio = AudioSegment.empty() + for part in parts: + audio += AudioSegment.from_file(BytesIO(requests.get(part).content)) + result = BytesIO() + audio.export(result, format='mp3') + result.seek(0) process_time = timedelta(seconds=(time.time() - start_time)) - self.handle_invoice(input_message.content_object.model, content=input_message.content) + self.handle_invoice(input_message.content_object.model, content=final_text) return self.save_results(input_message.content, process_time, result, save) @@ -271,23 +271,10 @@ class ModelVoiceCloneAPIView(MediaAPIView): }, ) def post(self, request, model: str, *args, **kwargs): - if not (request.FILES.get('file') or request.data.get('file')): - try: - if voice_id := request.data.pop('voice_id', None): - voice = Voice.objects.get(pk=voice_id, user=request.user) - transcription = voice.transcription - elif preset_id := request.data.pop('preset_id', None): - voice = Preset.objects.get(uid=preset_id) - transcription = voice.metadata.get('transcription', '') - else: - voice = Preset.objects.get(slug='russian_1') - transcription = voice.metadata.get('transcription', '') - except (Voice.DoesNotExist, Preset.DoesNotExist): - return Response( - {'detail': _('Voice not found.')}, - status=HTTP_400_BAD_REQUEST, - ) - request.data.update( - {'file': voice.file, 'info': {'transcription': transcription, **request.data['info']}} - ) + info = request.data.get('info', {}) or {} + if voice_id := request.data.get('voice_id'): + info.update({'voice_id': voice_id}) + elif preset_id := request.data.get('preset_id'): + info.update({'preset_id': preset_id}) + request.data.update({'info': info}) return super().post(request, model, *args, **kwargs) @@ -24,8 +24,7 @@ RUN --mount=target=/var/lib/apt/lists,type=cache,sharing=locked \ --mount=target=/var/cache/apt,type=cache,sharing=locked \ rm -f /etc/apt/apt.conf.d/docker-clean \ && apt-get update \ - && apt-get -y --no-install-recommends install -y gettext \ - && apt-get -y install antiword + && apt-get -y --no-install-recommends install gettext antiword ffmpeg RUN --mount=type=cache,target=/root/.cache/pip pip install -r requirements.txt @@ -3852,6 +3852,18 @@ azure-key-vault = ["azure-identity (>=1.16.0)", "azure-keyvault-secrets (>=4.8.0 toml = ["tomli (>=2.0.1)"] yaml = ["pyyaml (>=6.0.1)"] +[[package]] +name = "pydub" +version = "0.25.1" +description = "Manipulate audio with an simple and easy high level interface" +optional = false +python-versions = "*" +groups = ["main"] +files = [ + {file = "pydub-0.25.1-py2.py3-none-any.whl", hash = "sha256:65617e33033874b59d87db603aa1ed450633288aefead953b30bded59cb599a6"}, + {file = "pydub-0.25.1.tar.gz", hash = "sha256:980a33ce9949cab2a569606b65674d748ecbca4f0796887fd6f46173a7b0d30f"}, +] + [[package]] name = "pyjwt" version = "2.10.1" @@ -3956,16 +3968,16 @@ image = ["Pillow"] [[package]] name = "pyroscope-io" -version = "0.8.11" +version = "0.8.16" description = "Pyroscope Python integration" optional = false python-versions = "*" groups = ["main"] files = [ - {file = "pyroscope_io-0.8.11-py2.py3-none-macosx_11_0_arm64.whl", hash = "sha256:644dbd81b162b6d678ef9989649bf936c62373fd7bba16fb6490272f453ff159"}, - {file = "pyroscope_io-0.8.11-py2.py3-none-macosx_11_0_x86_64.whl", hash = "sha256:2df4cd4cbfb451c27cad20f905bf612ffc306a96820e316f7575c84433eadb24"}, - {file = "pyroscope_io-0.8.11-py2.py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:50164c96cf5533ce795a114c6181c5d2aa162c9dae59277b6ac557820c411b7a"}, - {file = "pyroscope_io-0.8.11-py2.py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:a415072dd7e8964d66001fff2446d7426a505cfca1a89f3c912314537622a4cd"}, + {file = "pyroscope_io-0.8.16-py2.py3-none-macosx_11_0_arm64.whl", hash = "sha256:e07edcfd59f5bdce42948b92c9b118c824edbd551730305f095a6b9af401a9e8"}, + {file = "pyroscope_io-0.8.16-py2.py3-none-macosx_11_0_x86_64.whl", hash = "sha256:dc98355e27c0b7b61f27066500fe1045b70e9459bb8b9a3082bc4755cb6392b6"}, + {file = "pyroscope_io-0.8.16-py2.py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:86f0f047554ff62bd92c3e5a26bc2809ccd467d11fbacb9fef898ba299dbda59"}, + {file = "pyroscope_io-0.8.16-py2.py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:6b91ce5b240f8de756c16a17022ca8e25ef8a4eed461c7d074b8a0841cf7b445"}, ] [package.dependencies] @@ -5568,4 +5580,4 @@ testing = ["coverage[toml]", "zope.event", "zope.testing"] [metadata] lock-version = "2.1" python-versions = "^3.12" -content-hash = "fce16784daa634e30c1fc600f28c473937a5b026af04864d27817cb525450713" +content-hash = "be19d70adce9ad107dd4b610d061d2477ea062690bad3e553131bf6d61acbf11" @@ -62,6 +62,7 @@ httptools = "^0.6.4" wsproto = "^1.2.0" sentry-sdk = {extras = ["django"], version = "^2.39.0"} googletrans = "^4.0.2" +pydub = "^0.25.1" [tool.poetry.group.test.dependencies]