@@ -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.image_test_model import Image_Test_Model @@ -0,0 +1,90 @@ +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 InvalidParameterError +from ml_model.models import ModelCategory, ModelInput, ModelParameter +from ml_model.services.base import SimpleService +from payments.exceptions.insufficient_balance import InsufficientBalance +from payments.selectors.payment_plan_selector import PaymentPlanSelector + + + +class Image_Test_Model(SimpleService): + title = 'Image Test Model' + description = 'Тестовая нейросеть для генерации изображений' + category = ModelCategory(title='Изображения', slug='images') + inputs = [ModelInput(type=ModelInput.TypeChoices.TEXT, required=True)] + + parameters = [ModelParameter( + name='Кастомный URL изображения', + key='ciu', + type=ModelParameter.TypeChoices.STR, + )] + + TOKENS_COST = Decimal('3') + PLACEHOLDER_URL='https://i.pinimg.com/736x/8b/e0/61/8be06158da3986fb4c47497b5660bb29.jpg' + + def calculate_price(self) -> Decimal: + return self.TOKENS_COST + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + return cls.TOKENS_COST + + def save_results(self, content: str, t: timedelta, image_bytes: bytes, image_ext: str = '.png', save: bool = True) -> list[Message]: + msg = Message( + content=content, + content_object=self.store, + elapsed_time=t, + file=File(BytesIO(image_bytes), image_ext), + ) + if save: + return Message.objects.bulk_create([msg]) + 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: + raise InsufficientBalance(balance, self.TOKENS_COST) + + start_time = time.time() + ciu = input_message.info.get('ciu', None) + + + if ciu: + image_bytes, image_ext = self._fetch_image(ciu) + else: + image_bytes, image_ext = self._fetch_image(self.PLACEHOLDER_URL) + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model) + + msgs = self.save_results(input_message.content, process_time, image_bytes, image_ext, save) + return msgs + + + + def _fetch_image(self, url: str): + try: + response = requests.get(url, timeout=10) + response.raise_for_status() + except requests.RequestException: + raise InvalidParameterError('Invalid image URL') + + kind = filetype.guess(response.content[:20]) + + if kind is None: + raise InvalidParameterError('Invalid image URL') + + if not kind.mime.startswith('image/'): + raise InvalidParameterError('Image format not supported') + + return response.content, f'.{kind.extension}' + \ No newline at end of file