@@ -1,18 +1,17 @@ import time -import httpx - -from backend import settings -from decimal import Decimal from datetime import timedelta +from decimal import Decimal from io import BytesIO from typing import Any +import httpx from django.core.files import File +from backend import settings from messages.models import Message - -from ml_model.exceptions import ModelTimeoutError, GenerationException +from ml_model.exceptions import GenerationException, ModelTimeoutError from ml_model.services.base import SimpleService +from poller.models import Proxy class Fluxlorafast(SimpleService): @@ -64,20 +63,19 @@ class Fluxlorafast(SimpleService): '4:3': 'landscape_4_3', '16:9': 'landscape_16_9', } - requests_number = 0 start_time = time.time() version = input_message.info.get('version') translated_prompt = self.translate_prompt(input_message.content) callback_data = dict( { 'prompt': f'in style of raif3_corporate Isometric illustration, ' - f'contemporary vector art style, 3/4 perspective view: {translated_prompt}', + f'contemporary vector art style, 3/4 perspective view: {translated_prompt}', 'model_version': 'fb90c17a-d410-41e7-9961-dc7c687bc627', 'image_size': sizes.get(input_message.info.get('image_size', '1:1')), 'loras': [ { 'path': 'https://v3.fal.media/files/elephant/JthCZoCdAr7' - 'LqnOiNNVCC_pytorch_lora_weights.safetensors' + 'LqnOiNNVCC_pytorch_lora_weights.safetensors' } ], 'guidance_scale': 5, @@ -85,29 +83,47 @@ class Fluxlorafast(SimpleService): 'num_images': input_message.info.get('num_images', 4), } ) - client = httpx.Client( - base_url="https://queue.fal.run", - headers={"Authorization": f"Key {settings.FAL_API_KEY}"}, - timeout=600, - ) - result = client.post( - f'fal-ai/{version}', - json={'prompt': input_message.content, **callback_data}, - ).json() - try: + + final_result = self._request(version, input_message, callback_data) + process_time = timedelta(seconds=(time.time() - start_time)) + + images = [img['url'] for img in final_result['images']] + self.handle_invoice(input_message.content_object.model, input_message=input_message, version=version) + msgs = self.save_results(input_message.content, images, process_time, save) + return msgs + + def _request(self, version: str, input_message: Message, callback_data: dict) -> dict: + requests_number = 0 + for proxy in Proxy.objects.all(): + client = httpx.Client( + base_url='https://queue.fal.run', + headers={'Authorization': f'Key {settings.FAL_API_KEY}'}, + timeout=600, + proxy=f'{proxy.protocol}://{proxy.address}', + ) + result = client.post( + f'fal-ai/{version}', + json={'prompt': input_message.content, **callback_data}, + ).json() + + is_success = False + while True: status = client.get(result['status_url']).json() if status.get('status') == 'COMPLETED': + is_success = True break requests_number += 1 if requests_number == 271: - raise ModelTimeoutError - time.sleep(1/3) - except Exception as exc: - raise GenerationException from exc - process_time = timedelta(seconds=(time.time() - start_time)) + continue + time.sleep(1 / 3) + + if is_success: + break + + if not is_success: + raise GenerationException from ModelTimeoutError + final_result = client.get(result['response_url']).json() - images = [img['url'] for img in final_result['images']] - self.handle_invoice(input_message.content_object.model, input_message=input_message, version=version) - msgs = self.save_results(input_message.content, images, process_time, save) - return msgs + + return final_result