import asyncio from io import BytesIO import os import logging import uuid from PIL import Image, ImageOps from functools import partial import pika import json import numpy as np from lxml import etree import io import requests from dotenv import load_dotenv load_dotenv() amqp_addr = os.getenv('AMQP_ADDR') or 'amqp://api:gacdownatravy9@51.8.120.154:5672/dev' logging.getLogger().setLevel(logging.INFO) from enum import Enum class QueueProgressKind(Enum): ImageGenerated = "image_generated" ImageGenerating = "image_generating" SamePrompt = "same_prompt" FaceswapGenerated = "faceswap_generated" FaceswapGenerating = "faceswap_generating" Failed = "failed" class MemedeckWorker: class BinaryEventTypes: PREVIEW_IMAGE = 1 UNENCODED_PREVIEW_IMAGE = 2 class JsonEventTypes(Enum): PROGRESS = "progress" EXECUTING = "executing" EXECUTED = "executed" ERROR = "error" STATUS = "status" """ MemedeckWorker is a class that is responsible for relaying messages between comfy and the memedeck backend api it is used to send images to the memedeck backend api and to receive prompts from the memedeck backend api """ def __init__(self, loop): MemedeckWorker.instance = self logging.getLogger().setLevel(logging.INFO) self.loop = loop self.messages = asyncio.Queue() self.ws_id = None self.http_client = None self.prompt_queue = None self.validate_prompt = None self.last_prompt_id = None self.amqp_url = amqp_addr self.queue_name = os.getenv('QUEUE_NAME') or 'generic-queue' self.api_url = os.getenv('API_ADDRESS') or 'http://0.0.0.0:8079/v2' self.api_key = os.getenv('API_KEY') or 'fb46e20e-cc25-4ed4-a39b-d47ca8cc3383' self.is_dev = os.getenv('IS_DEV') or False self.training_only = os.getenv('TRAINING_ONLY') or False self.video_gen_only = False self.azure_storage = MemedeckAzureStorage() if self.queue_name == 'video-gen-queue': print(f"[memedeck]: video gen only mode enabled") self.video_gen_only = True if self.training_only: self.queue_name = 'training-queue' print(f"[memedeck]: training only mode enabled") # Internal job queue self.internal_job_queue = asyncio.Queue() # Dictionary to keep track of tasks by ws_id self.tasks_by_ws_id = {} # Configure logging logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') self.logger = logging.getLogger(__name__) self.logger.info(f"\n[memedeck]: initialized with API URL: {self.api_url} and API Key: {self.api_key}\n") def start(self, prompt_queue, validate_prompt): self.prompt_queue = prompt_queue self.validate_prompt = validate_prompt # Start the process_job_queue task **after** prompt_queue is set self.loop.create_task(self.process_job_queue()) parameters = pika.URLParameters(self.amqp_url) logging.getLogger('pika').setLevel(logging.WARNING) # suppress all logs from pika self.connection = pika.SelectConnection(parameters, on_open_callback=self.on_connection_open) try: self.connection.ioloop.start() except KeyboardInterrupt: self.stop() def stop(self): self.connection.close() self.connection.ioloop.stop() # -------------------------------------------------- # AMQP # -------------------------------------------------- def on_connection_open(self, connection): self.connection = connection # Open the first channel self.connection.channel(on_open_callback=self.on_channel_open) if not self.training_only: # Open the second channel self.connection.channel(on_open_callback=self.on_faceswap_channel_open) def on_channel_open(self, channel): self.channel = channel # Only consume one message at a time self.channel.basic_qos(prefetch_size=0, prefetch_count=1) # Declare the queue and set the callback self.channel.queue_declare( queue=self.queue_name, durable=True, callback=self.on_queue_declared ) def on_queue_declared(self, frame): self.channel.basic_consume(queue=self.queue_name, on_message_callback=self.on_message_received) def on_faceswap_channel_open(self, channel): self.faceswap_channel = channel self.faceswap_channel.basic_qos(prefetch_size=0, prefetch_count=1) # Declare the faceswap queue and set the callback self.faceswap_channel.queue_declare(queue='faceswap-queue', durable=True, callback=self.on_faceswap_queue_declared) def on_faceswap_queue_declared(self, frame): self.faceswap_channel.basic_consume(queue='faceswap-queue', on_message_callback=self.on_faceswap_message_received) def on_faceswap_message_received(self, channel, method, properties, body): self.on_message_received(channel, method, properties, body) def on_message_received(self, channel, method, properties, body): decoded_string = body.decode('utf-8') payload = json.loads(decoded_string) # Execute the task prompt = payload["nodes"] valid = self.validate_prompt(prompt) routing_key = method.routing_key workflow = 'faceswap' if routing_key == 'faceswap-queue' else 'generation' user_id = payload["user_id"] if 'user_id' in payload else None self.logger.info(f"[memedeck]: workflow {workflow} user_id: {user_id}") if self.video_gen_only: workflow = 'video_gen' user_id = payload["user_id"] if self.training_only: workflow = 'training' end_node_id = None video_source_image_id = None output_asset_id = None # output asset id # Find the end_node_id if not self.video_gen_only and not self.training_only: for node in prompt: if isinstance(prompt[node], dict) and prompt[node].get("class_type") == "SaveImageWebsocket": end_node_id = node break elif self.video_gen_only: end_node_id = payload['end_node_id'] video_source_image_id = payload['image_id'] output_asset_id = payload['output_asset_id'] elif self.training_only: end_node_id = "130" self.logger.info(f"[memedeck]: end_node_id: {end_node_id}") # Prepare task_info prompt_id = str(uuid.uuid4()).replace("-", "_") outputs_to_execute = valid[2] task_info = { "workflow": workflow, "prompt_id": prompt_id, "prompt": prompt, "outputs_to_execute": outputs_to_execute, "client_id": "memedeck-1", "is_memedeck": True, "end_node_id": end_node_id, "ws_id": payload["source_ws_id"], "context": payload["req_ctx"] if "req_ctx" in payload else {}, "current_node": None, "current_progress": 0, "delivery_tag": method.delivery_tag, "task_status": "waiting", # based on workflow:flux_training_v1 "training_loop": { "4": 0.1, # loop 1 start "9": 0.25, # loop 1 end "44": 0.25, # loop 2 start "46": 0.5, # loop 2 end "59": 0.5, # loop 3 start "61": 0.75, # loop 3 end "60": 0.75, # loop 4 start "62": 1.0, # loop 4 end }, "training_status": { "4": "loop-1", "44": "loop-2", "59": "loop-3", "60": "loop-4", }, "training_filename": payload['filename'] if 'filename' in payload else None, # video data "image_id": video_source_image_id, "user_id": user_id, "output_asset_id": output_asset_id, } if valid[0]: # Enqueue the task into the internal job queue self.loop.call_soon_threadsafe(self.internal_job_queue.put_nowait, (prompt_id, prompt, task_info)) else: channel.basic_nack(delivery_tag=method.delivery_tag, requeue=False) # Unack the message # -------------------------------------------------- # Internal job queue # -------------------------------------------------- async def process_job_queue(self): while True: prompt_id, prompt, task_info = await self.internal_job_queue.get() # Start a new coroutine for each task self.loop.create_task(self.process_task(prompt_id, prompt, task_info)) async def process_task(self, prompt_id, prompt, task_info): ws_id = task_info['ws_id'] # Add the task to tasks_by_ws_id self.tasks_by_ws_id[ws_id] = task_info # Put the prompt into the prompt_queue self.prompt_queue.put((0, prompt_id, prompt, { "client_id": task_info["client_id"], 'is_memedeck': task_info['is_memedeck'], 'end_node_id': task_info['end_node_id'], 'ws_id': task_info['ws_id'], 'context': task_info['context'] }, task_info['outputs_to_execute'])) if task_info['workflow'] != 'training' and task_info['workflow'] != 'video_gen': if 'faceswap_strength' not in task_info['context']['prompt_config']: # pretty print the prompt config self.logger.info(f"[memedeck]: prompt: {task_info['context']['prompt_config']['character']['id']} {task_info['context']['prompt_config']['positive_prompt']}") # Wait until the current task is completed await self.wait_for_task_completion(ws_id) # Task is done self.internal_job_queue.task_done() async def wait_for_task_completion(self, ws_id): """ Wait until the task with the given ws_id is completed. """ task = self.tasks_by_ws_id[ws_id] delivery_tag = task['delivery_tag'] while ws_id in self.tasks_by_ws_id: await asyncio.sleep(0.25) # Acknowledge the message when the task is completed if task['workflow'] == 'faceswap': self.faceswap_channel.basic_ack(delivery_tag=delivery_tag) else: self.channel.basic_ack(delivery_tag=delivery_tag) # -------------------------------------------------- # allbacks for the prompt queue # -------------------------------------------------- def queue_updated(self): info = self.get_queue_info() # self.send_sync("status", { "status": self.get_queue_info() }) def get_queue_info(self): prompt_info = {} exec_info = {} exec_info['queue_remaining'] = self.prompt_queue.get_tasks_remaining() prompt_info['exec_info'] = exec_info return prompt_info def send_sync(self, event, data, sid=None): self.loop.call_soon_threadsafe( self.messages.put_nowait, (event, data, sid)) async def publish_loop(self): while True: msg = await self.messages.get() await self.send(*msg) async def send(self, event, data, sid=None): if sid is None: self.logger.warning("Received event without sid") return # Retrieve the task based on sid task = self.tasks_by_ws_id.get(sid) if not task: self.logger.warning(f"Received event {event} for unknown sid: {sid}") return if task['workflow'] == 'training': return await self.handle_training_send(event=event, task=task, sid=sid, data=data) if task['workflow'] == 'video_gen': return await self.handle_video_gen_send(event=event, task=task, sid=sid, data=data) # this logic is for generation and faceswap (todo move to separate function) if event == MemedeckWorker.BinaryEventTypes.UNENCODED_PREVIEW_IMAGE: await self.send_preview( data, sid=sid, progress=task['current_progress'], context=task['context'], workflow=task['workflow'] ) else: # Send JSON data / text data if event == "executing": task['current_node'] = data['node'] if task['workflow'] == 'faceswap' and task["task_status"] == "waiting": start_data = { "ws_id": task['ws_id'], "status": "started", "info": None, } await self.send_to_api(start_data) # self.logger.info(f"[memedeck]: faceswap executing: {data}") task["task_status"] = "executing" elif event == "progress": if task['current_node'] == task['end_node_id']: # If the node is the websocket node, then set the progress to 100 task['current_progress'] = 100 else: # If the node is not the websocket node, then set the progress based on the node's progress task['current_progress'] = (data['value'] / data['max']) * 100 if task['current_progress'] == 100 and task['current_node'] != task['end_node_id']: # In case the progress is 100 but the node is not the websocket node, set progress to 95 task['current_progress'] = 95 # Allows the full resolution image to be sent on the 100 progress event if data['value'] == 1 and task['workflow'] != 'training': # If the value is 1, send started to API start_data = { "ws_id": task['ws_id'], "status": "started", "info": None, } task["task_status"] = "executing" await self.send_to_api(start_data) elif event == "status": self.logger.info(f"[memedeck]: sending status event: {data}") elif event == "executed": # self.logger.info(f"[memedeck]: sending executed event: {data}") if task['workflow'] == 'training' and task['current_node'] == task['end_node_id']: self.logger.info(f"[memedeck]: training completed for {sid}") del self.tasks_by_ws_id[sid] # Update the task in tasks_by_ws_id self.tasks_by_ws_id[sid] = task async def handle_training_send(self, event, task, sid, data): # if event == "progress": # self.logger.info(f"[memedeck]: training progress: {data}") if event == "executing": training_progress = task['training_loop'][data['node']] if data['node'] in task['training_loop'] else None status = task['training_status'][data['node']] if data['node'] in task['training_status'] else None if training_progress is not None: await self.send_to_api({ "ws_id": sid, "prompt_id": data['prompt_id'], "progress": training_progress, "status": status }) if event == "executed": if data['node'] == task['end_node_id']: await self.send_to_api({ "ws_id": sid, "progress": 1.0, "status": "completed", "filename": task['training_filename'] }) # training task is done del self.tasks_by_ws_id[sid] async def handle_video_gen_send(self, event, task, sid, data): if event == "progress": if data['max'] > 1: progress = data['value'] max_progress = data['max'] # calculate the percentage percentage = (progress / max_progress) await self.send_to_api({ "ws_id": sid, "image_id": task['image_id'], "user_id": task['user_id'], "status": "generating", "output_asset_id": task['output_asset_id'], "progress": percentage * 0.9 # 90% of the progress is the gen step, 10% is the video encode step }) if event == "executed": if data['node'] == task['end_node_id']: # self.logger.info(f"[memedeck]: video gen completed {data}") file_path = data['output']['gifs'][0]['filename'] metadata = json.loads(data['output']['metadata'][0]) self.logger.info(f"[memedeck]: video gen completed {metadata}") # current_dir = os.path.dirname(os.path.abspath(__file__)) # file_path = os.path.join(current_dir, "output", filename) blob_name = f"{task['user_id']}/video_gen/video_{task['image_id'].replace('image:', '')}_{task['prompt_id']}.mp4" # TODO: take the file path and upload to azure blob storage # load image bytes with open(file_path, "rb") as video_file: video_bytes = video_file.read() self.logger.info(f"[memedeck]: video gen completed for {sid}, file={file_path}, blob={blob_name}") url = await self.azure_storage.save_image(blob_name, "video/mp4", video_bytes) self.logger.info(f"[memedeck]: video gen completed for {sid}, {url}") await self.send_to_api({ "ws_id": sid, "progress": 1.0, "image_id": task['image_id'], "user_id": task['user_id'], "status": "completed", "output_video_url": url, "output_asset_id": task['output_asset_id'], "metadata": metadata }) # video gen task is done del self.tasks_by_ws_id[sid] async def send_preview(self, image_data, sid=None, progress=None, context=None, workflow=None): # self.logger.info(f"[memedeck]: send_preview: {sid}") if sid is None: self.logger.warning("Received preview without sid") return task = self.tasks_by_ws_id.get(sid) if not task: self.logger.warning(f"Received preview for unknown sid: {sid}") return if progress is None: progress = task['current_progress'] # if progress is odd, then don't send the preview if int(progress) % 2 == 1: return image_type = image_data[0] image = image_data[1] max_size = image_data[2] if max_size is not None: if hasattr(Image, 'Resampling'): resampling = Image.Resampling.BILINEAR else: resampling = Image.ANTIALIAS image = ImageOps.contain(image, (max_size, max_size), resampling) bytesIO = BytesIO() image.save(bytesIO, format=image_type, quality=100 if progress == 95 else 75, compress_level=1) preview_bytes = bytesIO.getvalue() kind = "image_generating" if progress < 100 else "image_generated" image_id = None url = None watermarked_url = None if kind == "image_generated" and task['workflow'] != 'faceswap': image_uuid = str(uuid.uuid4()).replace("-", "_") # create uuid for the image blob_name = f"{task['user_id']}/{image_uuid}" # upload to azure blob storage url = await self.azure_storage.save_image(blob_name + ".jpeg", "image/jpeg", preview_bytes) watermarked_url = await self.azure_storage.save_image_watermarked(blob_name + "_watermarked.jpeg", "image/jpeg", preview_bytes) image_id = f"image:{image_uuid}" ai_queue_progress = { "ws_id": sid, "kind": kind, "data": list(preview_bytes) if kind == "image_generating" else None, "progress": int(progress), "context": context, "user_id": task['user_id'], "image_id": image_id, "url": url, "url_watermarked": watermarked_url } # self.logger.info(f"[memedeck]: progress kind: {kind}") # self.logger.info(f"[memedeck]: progress: {progress}") # set the kind to faceswap_generated if workflow is faceswap if workflow == 'faceswap': ai_queue_progress['kind'] = "faceswap_generated" del ai_queue_progress['progress'] await self.send_to_api(ai_queue_progress) if progress == 100 or workflow == 'faceswap': del self.tasks_by_ws_id[sid] # Remove the task from tasks_by_ws_id # self.logger.info(f"[memedeck]: Task {sid} completed") async def send_to_api(self, data): ws_id = data.get('ws_id') if not ws_id: self.logger.error("[memedeck]: Missing ws_id in data") return task = self.tasks_by_ws_id.get(ws_id) if not task: self.logger.error(f"[memedeck]: No task found for ws_id {ws_id}") return if task['end_node_id'] is None: self.logger.error(f"[memedeck]: end_node_id is None for {ws_id}") return api_endpoint = '/generation/update' if task['workflow'] == 'training': api_endpoint = '/training/update' if task['workflow'] == 'video_gen': api_endpoint = '/generation/video/update' try: # this request is not sending properly for faceswap post_func = partial(requests.post, f"{self.api_url}{api_endpoint}", json=data) # run a second time on another port await self.loop.run_in_executor(None, post_func) except Exception as e: if not self.is_dev: self.logger.info(f"[memedeck]: error sending to api: {e}") if self.is_dev: try: # this request is not sending properly for faceswap post_func_2 = partial(requests.post, f"http://0.0.0.0:9091/v2{api_endpoint}", json=data) # run a second time on another port await self.loop.run_in_executor(None, post_func_2) except Exception as e: if not self.is_dev: self.logger.info(f"[memedeck]: error sending to api: {e}") try: # this request is not sending properly for faceswap post_func_3 = partial(requests.post, f"http://0.0.0.0:9092/v2{api_endpoint}", json=data) # run a second time on another port await self.loop.run_in_executor(None, post_func_3) except Exception as e: if not self.is_dev: self.logger.info(f"[memedeck]: error sending to api: {e}") # -------------------------------------------------------------------------- # MemedeckAzureStorage # -------------------------------------------------------------------------- from azure.storage.blob.aio import BlobServiceClient from azure.storage.blob import ContentSettings from typing import Optional, Tuple import cairosvg WATERMARK = """ """ WATERMARK_SIZE = 40 class MemedeckAzureStorage: def __init__(self): # get environment variables self.account = os.getenv('STORAGE_ACCOUNT') self.access_key = os.getenv('STORAGE_ACCESS_KEY') self.container = os.getenv('STORAGE_CONTAINER') logging.getLogger('azure.core.pipeline.policies.http_logging_policy').setLevel(logging.WARNING) logging.getLogger("azure.storage.common.storageclient").setLevel(logging.WARNING) self.logger = logging.getLogger('azure.storage.common') if not all([self.account, self.access_key, self.container]): raise EnvironmentError("Missing STORAGE_ACCOUNT, STORAGE_ACCESS_KEY, or STORAGE_CONTAINER environment variables") # Initialize BlobServiceClient self.blob_service_client = BlobServiceClient( account_url=f"https://{self.account}.blob.core.windows.net", credential=self.access_key ) self.logger.info(f"[memedeck]: Azure Storage connected.") async def save_image( self, blob_name: str, content_type: str, bytes_data: bytes ) -> str: """ Saves image bytes to Azure Blob Storage. Args: blob_name (str): Name of the blob in Azure Storage. content_type (str): MIME type of the content. bytes_data (bytes): Image data in bytes. Returns: str: URL of the uploaded blob. """ blob_client = self.blob_service_client.get_blob_client(container=self.container, blob=blob_name) # Upload the blob try: # prevent logging the request await blob_client.upload_blob( bytes_data, overwrite=True, content_settings=ContentSettings(content_type=content_type) ) except Exception as e: raise Exception(f"Failed to upload blob: {e}") # close the blob client blob_client.close() # Construct and return the blob URL blob_url = f"https://media.memedeck.xyz/{self.container}/{blob_name}" return blob_url async def save_image_watermarked( self, blob_name: str, content_type: str, bytes_data: bytes ) -> str: image = Image.open(BytesIO(bytes_data)) watermarked_image = self.add_watermark_to_image(image) # convert pil to bytes img_byte_arr = BytesIO() # Create an in-memory byte stream watermarked_image.save(img_byte_arr, format=image.format, quality=100, compress_level=1) # Save the image to the in-memory stream watermarked_image_bytes = img_byte_arr.getvalue() return await self.save_image(blob_name, content_type, watermarked_image_bytes) def add_watermark_to_image(self, img, background_brightness=None): """ Adds a watermark to a single PIL Image. Args: img: A PIL Image object. Returns: A PIL Image object with the watermark added. """ padding = 12 x = img.width - WATERMARK_SIZE - padding y = img.height - WATERMARK_SIZE - padding if background_brightness is None: background_brightness = self.analyze_background_brightness(img, x, y, WATERMARK_SIZE) # Generate watermark image (replace this with your actual watermark generation) watermark = self.generate_watermark(WATERMARK_SIZE, background_brightness) # Overlay the watermark img.paste(watermark, (x, y), watermark) return img def analyze_background_brightness(self, img, x, y, size): """ Analyzes the average brightness of a region in the image. Args: img: A PIL Image object. x: The x-coordinate of the top-left corner of the region. y: The y-coordinate of the top-left corner of the region. size: The size of the region (square). Returns: The average brightness of the region as an integer. """ region = img.crop((x, y, x + size, y + size)) pixels = np.array(region) total_brightness = np.sum( 0.299 * pixels[:, :, 0] + 0.587 * pixels[:, :, 1] + 0.114 * pixels[:, :, 2] ) / 1000 print(f"total_brightness: {total_brightness}") return max(0, min(255, total_brightness)) def generate_watermark(self, size, background_brightness): """ Generates a watermark image from an SVG string. Args: size: The size of the watermark (square). background_brightness: The background brightness at the watermark position. Returns: A PIL Image object representing the watermark. """ # Determine watermark color based on background brightness watermark_color = (0, 0, 0, 165) if background_brightness > 128 else (255, 255, 255, 165) # Parse the SVG string svg_tree = etree.fromstring(WATERMARK) # Find the path element and set its fill attribute path_element = svg_tree.find(".//{http://www.w3.org/2000/svg}path") if path_element is not None: r, g, b, a = watermark_color fill_color = f"rgba({r},{g},{b},{a/255})" # Convert to rgba string path_element.set("fill", fill_color) # Convert the modified SVG tree back to a string modified_svg = etree.tostring(svg_tree, encoding="unicode") # Render the modified SVG to a PNG image with a transparent background png_data = cairosvg.svg2png( bytestring=modified_svg, output_width=size, output_height=size, background_color="transparent" ) watermark_img = Image.open(BytesIO(png_data)) # Convert the watermark to RGBA to handle transparency watermark_img = watermark_img.convert("RGBA") return watermark_img