@@ -1,6 +1,7 @@ import re import subprocess import zipfile +from collections.abc import Iterable from functools import wraps from uuid import UUID @@ -144,25 +145,15 @@ class ImageFileProcessingService: def get_bytes(self, size: int | None = None) -> bytes: return self.image.read(size) - def get_kind(self, file_bytes: bytes): - kind = filetype.guess(file_bytes) - if not kind: - raise CorruptedFileError - if kind.extension.upper() not in self.ALLOWED_EXTENSIONS: - raise FileExtensionNotSupported(self.ALLOWED_EXTENSIONS) - return kind + def get_kind(self, file_bytes: bytes) -> filetype.Type | None: + return filetype.guess(file_bytes) @__reset_image - def get_dimensions(self, max_pixels: int) -> tuple[int, int]: + def get_dimensions(self) -> tuple[int | None, int | None]: try: - w, h = get_image_dimensions(self.image) - if not (w and h): - raise CorruptedFileError - if w * h > max_pixels: - raise ImageTooLargeError(max_pixels) - return w, h - except Image.DecompressionBombError: - raise ImageTooLargeError(max_pixels) + return get_image_dimensions(self.image) + except Image.DecompressionBombError as exc: + raise ImageTooLargeError(Image.MAX_IMAGE_PIXELS or 0) from exc def get_normalized_image(self, file_bytes: bytes) -> BytesIO: normalized_image = BytesIO(file_bytes) @@ -173,3 +164,23 @@ class ImageFileProcessingService: img.close() normalized_image.seek(0) return normalized_image + + def validate_kind(self, kind: filetype.Type | None) -> filetype.Type: + if not kind: + raise CorruptedFileError + return kind + + def validate_extension( + self, + kind: filetype.Type, + allowed_extensions: Iterable[str] | None = None, + ) -> None: + extensions = allowed_extensions or self.ALLOWED_EXTENSIONS + if kind.extension.upper() not in extensions: + raise FileExtensionNotSupported(extensions) + + def validate_dimensions(self, width: int | None, height: int | None, max_pixels: int) -> None: + if not (width and height): + raise CorruptedFileError + if width * height > max_pixels: + raise ImageTooLargeError(max_pixels) @@ -95,9 +95,11 @@ class Seedream(SimpleService): } if image := input_message.file: image_processor = ImageFileProcessingService(image) - file_bytes = image_processor.get_bytes(20) - image_processor.get_kind(file_bytes) - image_processor.get_dimensions(max_pixels=36000000) + kind = image_processor.get_kind(image_processor.get_bytes(20)) + kind = image_processor.validate_kind(kind) + image_processor.validate_extension(kind) + width, height = image_processor.get_dimensions() + image_processor.validate_dimensions(width, height, max_pixels=36000000) image_processor.image.close() callback_data.update({'image': image.url}) @@ -5,12 +5,12 @@ 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 FileNotProvided +from ml_model.services.FileService import ImageFileProcessingService from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run @@ -50,11 +50,12 @@ class Wan(SimpleService): raise InsufficientBalance(balance, cost) if not input_message.file: raise FileNotProvided('Image') - kind = filetype.guess(input_message.file.read(20)) - mime = kind.mime if kind else 'application/octet-stream' - input_message.file.seek(0) - image = f'data:{mime};base64,{base64.b64encode(input_message.file.read()).decode("utf-8")}' - input_message.file.close() + image_processor = ImageFileProcessingService(input_message.file) + kind = image_processor.get_kind(image_processor.get_bytes(50)) + kind = image_processor.validate_kind(kind) + image_processor.validate_extension(kind) + image = f'data:{kind.mime};base64,{base64.b64encode(image_processor.get_bytes()).decode("utf-8")}' + image_processor.image.close() callback_data = dict( { 'prompt': input_message.content,