@@ -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): @@ -43,12 +42,27 @@ class Fluxlorafast(SimpleService): ) -> list[Message]: messages: list[Message] = [] for image in images: + 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}', + ) + + try: + file = File(BytesIO(client.get(image).content), '.png') + except Exception: + continue + + break + messages.append( Message( content_object=self.store, elapsed_time=time, content=prompt, - file=File(BytesIO(httpx.get(image).content), '.png'), + file=file, ) ) if save: @@ -64,20 +78,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 +98,52 @@ 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: - while True: - status = client.get(result['status_url']).json() - if status.get('status') == 'COMPLETED': - break - requests_number += 1 - if requests_number == 271: - raise ModelTimeoutError - time.sleep(1/3) - except Exception as exc: - raise GenerationException from exc + + final_result = self._request(version, input_message, callback_data) process_time = timedelta(seconds=(time.time() - start_time)) - 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 + + def _request(self, version: str, input_message: Message, callback_data: dict) -> dict: + for proxy in Proxy.objects.all(): + requests_number = 0 + 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: + if requests_number == 271 / len(Proxy.objects.all()): + break + + try: + status = client.get(result['status_url']).json() + except Exception: + requests_number += 1 + continue + if status.get('status') == 'COMPLETED': + is_success = True + break + + time.sleep(1 / 3) + + if is_success: + break + + if not is_success: + raise GenerationException from ModelTimeoutError + + final_result = client.get(result['response_url']).json() + + return final_result