@@ -3,9 +3,10 @@ from datetime import timedelta from decimal import Decimal from io import BytesIO from typing import Any -import filetype +import filetype import requests +from PIL import Image from django.core.files import File from messages.models import Message @@ -15,19 +16,39 @@ 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' + PLACEHOLDER_URL = 'https://i.pinimg.com/736x/8b/e0/61/8be06158da3986fb4c47497b5660bb29.jpg' + SUPPORTED_RATIOS = frozenset( + { + '1:1', + '16:9', + '9:16', + '4:5', + '5:4', + '3:4', + '4:3', + '2:3', + '3:2', + '9:21', + '20:9', + '21:9', + '1:2', + '2:1', + '19.5:9', + '9:20', + '9:19.5', + '7:4', + } + ) - def calculate_price(self, num_images: int = 1 ) -> Decimal: + 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( @@ -51,40 +72,69 @@ class Image_Test_Model(SimpleService): 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): + ratio = input_message.info.get('ratio', '1: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 - + raw_image = self._fetch_image(ciu) + resized_image = self._resize_to_ratio(raw_image, ratio) + images = [resized_image] * 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): + def _fetch_image(self, url: str) -> bytes: 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 + + def _parse_ratio(self, ratio: str) -> tuple[float, float]: + if ratio not in self.SUPPORTED_RATIOS: + raise InvalidParameterError(f'Unsupported ratio: {ratio}') + rw, rh = map(float, ratio.split(':')) + return rw, rh + + def _resize_to_ratio(self, image_bytes: bytes, ratio: str) -> bytes: + rw, rh = self._parse_ratio(ratio) + + with Image.open(BytesIO(image_bytes)) as img: + img = img.convert('RGB') + src_w, src_h = img.size + target_ratio = rw / rh + src_ratio = src_w / src_h + + if src_ratio > target_ratio: + new_w = int(src_h * target_ratio) + left = (src_w - new_w) // 2 + cropped = img.crop((left, 0, left + new_w, src_h)) + else: + new_h = int(src_w / target_ratio) + top = (src_h - new_h) // 2 + cropped = img.crop((0, top, src_w, top + new_h)) + + buf = BytesIO() + cropped.save(buf, format='PNG') + return buf.getvalue()