@@ -1,19 +1,13 @@ import time -import requests - from datetime import timedelta from decimal import Decimal from io import BytesIO from typing import Any +import requests from django.core.files import File from messages.models import Message -from ml_model.exceptions import ( - ImageContentNotFound, - GenerationException, - RequestBlocked, -) from ml_model.services.base import SimpleService from ml_model.tasks import replicate_run from payments.exceptions.insufficient_balance import InsufficientBalance @@ -51,12 +45,12 @@ class Grok_Image(SimpleService): raise InsufficientBalance(balance, self.TOKENS_COST) callback_data = dict( { - 'prompt': f"{self.translate_prompt(input_message.content)}\n{self.OPTIMIZATION_PROMPT}", + 'prompt': f'{self.translate_prompt(input_message.content)}\n{self.OPTIMIZATION_PROMPT}', **input_message.info, } ) if image := input_message.file: - callback_data.update({'image': image.url}) + callback_data.update({'image': image}) start_time = time.time() images = replicate_run('xai/grok-imagine-image', callback_data) process_time = timedelta(seconds=(time.time() - start_time)) @@ -68,14 +68,14 @@ class Grok_Image_Ultra(SimpleService): file_bytes = input_message.file.read() kind = filetype.guess(file_bytes[:20]) extension = kind.extension - if extension.upper() not in (extensions := ['JPG', 'JPEG', 'PNG', 'WEBP']): + if extension.upper() not in (extensions := ['JPG', 'JPEG', 'JFIF', 'PNG', 'WEBP']): raise FileExtensionNotSupported(extensions) file_width, file_height = get_image_dimensions(BytesIO(file_bytes)) if file_width and file_height: input_mp = math.ceil((file_width * file_height) / 1_000_000) else: input_mp = 1 - callback_data.update({'image': input_message.file.url}) + callback_data.update({'image': input_message.file}) input_message.file.close() if (balance := PaymentPlanSelector(self.store.user).get_current_balance()) < ( @@ -6,6 +6,7 @@ import time # import uuid from io import BytesIO +from pathlib import PurePosixPath from typing import IO, Any, Dict import deepl @@ -15,6 +16,8 @@ import replicate import requests from celery import shared_task from deepl.translator import TextResult +from django.core.files.storage import Storage +from django.db.models.fields.files import FieldFile from django.utils.translation import gettext as _ from replicate.exceptions import ModelError from requests import Response @@ -28,6 +31,7 @@ from ml_model.exceptions import ( DeploymentDisabled, ExceededContextLengthError, FaceNotFoundError, + FileExtensionNotSupported, GenerationException, ImageAnalysisError, ImageContentNotFound, @@ -36,6 +40,7 @@ from ml_model.exceptions import ( ModelTimeoutError, PredictionInterruptedError, RequestBlocked, + ServiceHighDemandError, ) from poller.models import Proxy @@ -115,10 +120,37 @@ def transcript_audio(payload: dict[str, Any]): ) +def _prepare_replicate_image(image: FieldFile) -> tuple[str, tuple[Storage, str] | None]: + if PurePosixPath(image.name).suffix.lower() != '.jfif': + return image.url, None + + # JFIF already contains JPEG data, so copy it without decoding or re-encoding. + image.open('rb') + image.seek(0) + jpeg_name = str(PurePosixPath(image.name).with_suffix('.jpg')) + try: + saved_name = image.storage.save(jpeg_name, image) + finally: + image.close() + + try: + image_url = image.storage.url(saved_name) + except Exception: + image.storage.delete(saved_name) + + raise + + return image_url, (image.storage, saved_name) + + @shared_task def replicate_run(callback_url: str, payload: dict[str, Any]): replicate_client = replicate.Client(settings.REPLICATE_API_KEY) + temporary_file = None try: + if isinstance(image := payload.get('image'), FieldFile): + payload['image'], temporary_file = _prepare_replicate_image(image) + return replicate_client.run( ref=callback_url, input=payload, @@ -126,6 +158,14 @@ def replicate_run(callback_url: str, payload: dict[str, Any]): except ModelError as exc: prediction_error = getattr(getattr(exc, 'prediction', None), 'error', '') or '' error_text = str(exc) + if 'ModelRateLimitError' in error_text or 'E003' in error_text: + raise ServiceHighDemandError from exc + if ( + 'Music upload failed' in error_text + and 'audio format' in error_text + and 'is not supported' in error_text + ): + raise FileExtensionNotSupported(('MP3', 'WAV')) from exc if any(error in error_text for error in ('E005', 'E006', 'sexual', 'NSFW')): raise RequestBlocked if 'PA' in error_text: @@ -147,6 +187,10 @@ def replicate_run(callback_url: str, payload: dict[str, Any]): if 'PROMPT_TOO_LONG' in error_text: raise ExceededContextLengthError raise GenerationException from exc + finally: + if temporary_file: + storage, file_name = temporary_file + storage.delete(file_name) @shared_task