@@ -7,12 +7,14 @@ from io import BytesIO import filetype import requests +from PIL import Image from django.core.files import File from django.core.files.images import get_image_dimensions from replicate.exceptions import ModelError from messages.models import Message -from ml_model.exceptions import PredictionInterruptedError, RequestBlocked, GenerationException +from ml_model.exceptions import PredictionInterruptedError, RequestBlocked, GenerationException, \ + FileExtensionNotSupported from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run @@ -79,13 +81,20 @@ class Flux_2(SimpleService): } input_mp = 0 if input_message.file: - file_width, file_height = get_image_dimensions(input_message.file) - input_mp = math.ceil((file_width*file_height) / 1_000_000) - 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")}' + file_bytes = input_message.file.read() input_message.file.close() + kind = filetype.guess(file_bytes[:20]) + extension = kind.extension + if extension.upper() not in (extensions := ['JPG', 'JPEG', 'PNG', 'WEBP']): + raise FileExtensionNotSupported(extensions) + normalized_image = Image.open(BytesIO(file_bytes)).convert('RGB') + buf = BytesIO() + format = 'jpeg' if extension not in ('png', 'jpeg', 'webp') else extension + normalized_image.save(buf, format=format) + file_width, file_height = get_image_dimensions(buf) + input_mp = math.ceil((file_width*file_height) / 1_000_000) + image = f'data:image/{format};base64,{base64.b64encode(buf.getvalue()).decode("utf-8")}' + buf.close() callback_data.update({'input_images': [image]}) try: images = [replicate_run(f'black-forest-labs/{version}', callback_data)] @@ -7,8 +7,10 @@ from io import BytesIO import filetype import requests from django.core.files import File +from replicate.exceptions import ModelError from messages.models import Message +from ml_model.exceptions import RequestBlocked, GenerationException from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run @@ -48,7 +50,12 @@ class Minimaxvideo(SimpleService): input_message.file.close() callback_data.update({'subject_reference': image}) start_time = time.time() - video = replicate_run(f'minimax/{version}', callback_data) + try: + video = replicate_run(f'minimax/{version}', callback_data) + except ModelError as exc: + if any(error in str(exc) for error in ('E005', 'E006', 'sexual')): + raise RequestBlocked + raise GenerationException from exc process_time = timedelta(seconds=(time.time() - start_time)) self.handle_invoice(input_message.content_object.model, version) msgs = self.save_results(input_message.content, process_time, video, save) @@ -18,6 +18,7 @@ from ml_model.exceptions import ( InvalidParameterError, InvalidStyleCombinationError, PromptLengthExceeded, + FileExtensionNotSupported, ) from ml_model.models import NeuronModel from ml_model.services.base import SimpleService @@ -170,6 +171,7 @@ class MediaAPIView(APIView): InvalidStyleCombinationError, InvalidParameterError, PromptLengthExceeded, + FileExtensionNotSupported, ), ): return Response({'detail': f'{exc}'}, status=HTTP_400_BAD_REQUEST)