@@ -81,3 +81,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 CorruptedFileError, InvalidParameterError +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): + + TOKENS_COST = Decimal('3') + PLACEHOLDER_URL='https://i.pinimg.com/736x/8b/e0/61/8be06158da3986fb4c47497b5660bb29.jpg' + + def calculate_price(self, num_images: int = 1 ) -> Decimal: + return self.TOKENS_COST * num_images + + @classmethod + def predict_price(cls, content: str, file_exists: bool, info: dict[str, Any]) -> Decimal | None: + num_images = info.get('num_images', 1) + + return cls.TOKENS_COST * num_images + + def save_results( + self, + content: str, + t: timedelta, + images: list[bytes], + save: bool = True, + ) -> list[Message]: + messages: list[Message] = [] + for image in images: + messages.append( + Message( + content_object=self.store, + elapsed_time=t, + content=content, + file=File(BytesIO(image), '.png'), + ) + ) + if save: + return Message.objects.bulk_create(messages) + return messages + + + def make(self, input_message: Message, save: bool = True) -> list[Message]: + num_images = input_message.info.get('num_images', 1) + + if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < (cost := self.TOKENS_COST * num_images): + raise InsufficientBalance(balance, cost) + + start_time = time.time() + ciu = input_message.info.get('ciu') or self.PLACEHOLDER_URL + + images = [self._fetch_image(ciu)] * num_images + + process_time = timedelta(seconds=(time.time() - start_time)) + self.handle_invoice(input_message.content_object.model, num_images) + + msgs = self.save_results(input_message.content, process_time, images, 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 not kind: + raise CorruptedFileError + + if not kind.mime.startswith('image/'): + raise InvalidParameterError('Image format not supported') + + return response.content + \ No newline at end of file