Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
629f87a2c3 | ||
|
|
e14ffc2284 |
@@ -0,0 +1,26 @@
|
||||
# --- S3 (AWS/MINIO) CREDENTIALS ---
|
||||
## These are used by both scripts to connect to SQS and S3 (MinIO).
|
||||
AWS_ACCESS_KEY_ID=minioadmin
|
||||
AWS_SECRET_ACCESS_KEY=...
|
||||
AWS_DEFAULT_REGION=us-east-1
|
||||
|
||||
# --- SQS SETTINGS ---
|
||||
## Toggles functionality for the SQS Worker Consumer
|
||||
SQS_ENABLED=false
|
||||
|
||||
## For media_stream.py (MediaStreamOutput Node)
|
||||
### Endpoint for the SQS service where completion messages are sent.
|
||||
SQS_ENDPOINT_URL=http://127.0.0.1:9324
|
||||
|
||||
### The specific SQS queue the worker should push job status updates to.
|
||||
SQS_JOB_STATUS_UPDATES_QUEUE_NAME=job_status_updates
|
||||
|
||||
## For worker_consumer.py (Job Consumer)
|
||||
### The specific SQS queue which the worker should poll for new jobs.
|
||||
SQS_JOBS_TO_PROCESS_QUEUE_NAME=jobs_to_process
|
||||
|
||||
## (Optional) For worker_consumer.py (Job Consumer)
|
||||
### The local URL of the ComfyUI API server.
|
||||
# You only need to set this if your ComfyUI server is NOT running on the default port 8188.
|
||||
# COMFYUI_API_URL=http://127.0.0.1:8188
|
||||
# COMFYUI_WS_URL=ws://127.0.0.1:8188
|
||||
+78
-1
@@ -1,3 +1,80 @@
|
||||
from .nilornodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
import os
|
||||
import threading
|
||||
import asyncio
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# --- Nilor-Nodes Custom Node Registration and Startup ---
|
||||
# This file is executed when ComfyUI starts and discovers this custom node directory.
|
||||
# It's responsible for:
|
||||
# 1. Starting background services (like the SQS worker and a FastAPI server).
|
||||
# 2. Registering the custom nodes with ComfyUI so they appear in the menu.
|
||||
|
||||
|
||||
# --- Load Environment Variables ---
|
||||
# Get the directory of the current script
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
# Construct the path to the .env file
|
||||
dotenv_path = os.path.join(current_dir, ".env")
|
||||
# Load the .env file, overriding any pre-existing process env for these keys
|
||||
load_dotenv(dotenv_path=dotenv_path, override=True)
|
||||
|
||||
|
||||
# --- Background Services ---
|
||||
|
||||
|
||||
def start_consumer_loop():
|
||||
"""Synchronous wrapper to run the asyncio event loop for the consumer."""
|
||||
from .worker_consumer import consume_jobs
|
||||
|
||||
asyncio.run(consume_jobs())
|
||||
|
||||
|
||||
# Start the SQS Worker Consumer (controlled by SQS_ENABLED)
|
||||
raw_sqs_enabled = os.getenv("SQS_ENABLED", "false")
|
||||
env_sqs_enabled = raw_sqs_enabled.strip().lower() == "true"
|
||||
if env_sqs_enabled:
|
||||
consumer_thread = threading.Thread(target=start_consumer_loop, daemon=True)
|
||||
consumer_thread.start()
|
||||
print(
|
||||
f"✅ Nilor-Nodes: SQS worker consumer thread started (SQS_ENABLED={raw_sqs_enabled} in .env)."
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f"⚠️ Nilor-Nodes: SQS worker consumer functionality is disabled (SQS_ENABLED={raw_sqs_enabled} in .env)."
|
||||
)
|
||||
|
||||
|
||||
# --- Node Registration ---
|
||||
from .nilornodes import (
|
||||
NODE_CLASS_MAPPINGS as base_NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as base_NODE_DISPLAY_NAME_MAPPINGS,
|
||||
)
|
||||
from .media_stream import (
|
||||
NODE_CLASS_MAPPINGS as ms_NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as ms_NODE_DISPLAY_NAME_MAPPINGS,
|
||||
)
|
||||
from .user_input import (
|
||||
NODE_CLASS_MAPPINGS as ui_NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as ui_NODE_DISPLAY_NAME_MAPPINGS,
|
||||
)
|
||||
from .controllers import (
|
||||
NODE_CLASS_MAPPINGS as ctrl_NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as ctrl_NODE_DISPLAY_NAME_MAPPINGS,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS = dict(base_NODE_CLASS_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS = dict(base_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(ms_NODE_CLASS_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(ms_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(ui_NODE_CLASS_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(ui_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(ctrl_NODE_CLASS_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(ctrl_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
print("✅ Nilor-Nodes: All custom nodes registered.")
|
||||
|
||||
@@ -0,0 +1,286 @@
|
||||
"""
|
||||
Brain API Client for ComfyUI Nodes
|
||||
|
||||
This client provides methods to interact with the Brain API storage endpoints,
|
||||
replacing the need for pre-signed URLs in the ComfyUI workflow.
|
||||
"""
|
||||
|
||||
import requests
|
||||
import os
|
||||
import logging
|
||||
from typing import Optional, Dict, Any
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# Load environment variables
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
dotenv_path = os.path.join(current_dir, ".env")
|
||||
load_dotenv(dotenv_path=dotenv_path)
|
||||
|
||||
# Setup logging
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
|
||||
|
||||
class BrainApiClient:
|
||||
"""
|
||||
Client for interacting with Brain API storage endpoints.
|
||||
|
||||
This client handles authentication and provides methods for uploading,
|
||||
downloading, and deleting files through the Brain API storage endpoints.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize the Brain API client with configuration from environment variables."""
|
||||
self.base_url = os.getenv("BRANDO_BRAIN_API_BASE_URL", "http://localhost:2024/api")
|
||||
self.api_key = os.getenv("BRANDO_API_KEY")
|
||||
|
||||
if not self.api_key:
|
||||
raise ValueError(
|
||||
"BRANDO_API_KEY environment variable is required for Brain API authentication"
|
||||
)
|
||||
|
||||
self.headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"User-Agent": "ComfyUI-NilorNodes/1.0"
|
||||
}
|
||||
|
||||
logging.info(f"Brain API Client initialized with base URL: {self.base_url}")
|
||||
|
||||
def upload_file_to_storage(self, file_path: str, filename: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Upload a file to Brain API storage and return storage metadata.
|
||||
|
||||
Args:
|
||||
file_path: Local path to the file to upload
|
||||
filename: Name to use for the uploaded file
|
||||
|
||||
Returns:
|
||||
Dict containing storage_id and filename
|
||||
|
||||
Raises:
|
||||
requests.RequestException: If upload fails
|
||||
FileNotFoundError: If file_path doesn't exist
|
||||
"""
|
||||
if not os.path.exists(file_path):
|
||||
raise FileNotFoundError(f"File not found: {file_path}")
|
||||
|
||||
url = f"{self.base_url}/storage/upload"
|
||||
|
||||
try:
|
||||
with open(file_path, 'rb') as file:
|
||||
files = {'file': (filename, file, 'application/octet-stream')}
|
||||
|
||||
logging.info(f"Uploading file '{filename}' to Brain API storage...")
|
||||
response = requests.post(
|
||||
url,
|
||||
files=files,
|
||||
headers=self.headers,
|
||||
timeout=300
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
result = response.json()
|
||||
logging.info(f"Upload successful. Storage ID: {result.get('storage_id')}")
|
||||
return result
|
||||
|
||||
except requests.RequestException as e:
|
||||
logging.error(f"Failed to upload file '{filename}': {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
logging.error(f"Unexpected error uploading file '{filename}': {e}")
|
||||
raise
|
||||
|
||||
def upload_fileobj_to_storage(self, file_obj, filename: str, content_type: str = 'application/octet-stream') -> Dict[str, Any]:
|
||||
"""
|
||||
Upload a file-like object to Brain API storage and return storage metadata.
|
||||
|
||||
Args:
|
||||
file_obj: File-like object to upload
|
||||
filename: Name to use for the uploaded file
|
||||
content_type: MIME type of the file
|
||||
|
||||
Returns:
|
||||
Dict containing storage_id and filename
|
||||
|
||||
Raises:
|
||||
requests.RequestException: If upload fails
|
||||
"""
|
||||
url = f"{self.base_url}/storage/upload"
|
||||
|
||||
try:
|
||||
files = {'file': (filename, file_obj, content_type)}
|
||||
|
||||
logging.info(f"Uploading file object '{filename}' to Brain API storage...")
|
||||
response = requests.post(
|
||||
url,
|
||||
files=files,
|
||||
headers=self.headers,
|
||||
timeout=300
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
result = response.json()
|
||||
logging.info(f"Upload successful. Storage ID: {result.get('storage_id')}")
|
||||
return result
|
||||
|
||||
except requests.RequestException as e:
|
||||
logging.error(f"Failed to upload file object '{filename}': {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
logging.error(f"Unexpected error uploading file object '{filename}': {e}")
|
||||
raise
|
||||
|
||||
def download_file_from_storage(self, storage_id: str, filename: str, dest_path: str) -> str:
|
||||
"""
|
||||
Download a file from Brain API storage to a local path.
|
||||
|
||||
Args:
|
||||
storage_id: Storage ID of the file to download
|
||||
filename: Name of the file to download
|
||||
dest_path: Local path where the file should be saved
|
||||
|
||||
Returns:
|
||||
Path to the downloaded file
|
||||
|
||||
Raises:
|
||||
requests.RequestException: If download fails
|
||||
"""
|
||||
url = f"{self.base_url}/storage/{storage_id}"
|
||||
params = {'filename': filename}
|
||||
|
||||
try:
|
||||
logging.info(f"Downloading file '{filename}' (storage_id: {storage_id}) from Brain API storage...")
|
||||
response = requests.get(
|
||||
url,
|
||||
params=params,
|
||||
headers=self.headers,
|
||||
timeout=300,
|
||||
stream=True
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
# Ensure destination directory exists
|
||||
os.makedirs(os.path.dirname(dest_path), exist_ok=True)
|
||||
|
||||
with open(dest_path, 'wb') as f:
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
f.write(chunk)
|
||||
|
||||
logging.info(f"Download successful. File saved to: {dest_path}")
|
||||
return dest_path
|
||||
|
||||
except requests.RequestException as e:
|
||||
logging.error(f"Failed to download file '{filename}' (storage_id: {storage_id}): {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
logging.error(f"Unexpected error downloading file '{filename}': {e}")
|
||||
raise
|
||||
|
||||
def get_file_from_storage(self, storage_id: str, filename: str) -> bytes:
|
||||
"""
|
||||
Get file content from Brain API storage as bytes.
|
||||
|
||||
Args:
|
||||
storage_id: Storage ID of the file to download
|
||||
filename: Name of the file to download
|
||||
|
||||
Returns:
|
||||
File content as bytes
|
||||
|
||||
Raises:
|
||||
requests.RequestException: If download fails
|
||||
"""
|
||||
url = f"{self.base_url}/storage/{storage_id}"
|
||||
params = {'filename': filename}
|
||||
|
||||
try:
|
||||
logging.info(f"Getting file '{filename}' (storage_id: {storage_id}) from Brain API storage...")
|
||||
response = requests.get(
|
||||
url,
|
||||
params=params,
|
||||
headers=self.headers,
|
||||
timeout=300
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
logging.info(f"File retrieval successful. Size: {len(response.content)} bytes")
|
||||
return response.content
|
||||
|
||||
except requests.RequestException as e:
|
||||
logging.error(f"Failed to get file '{filename}' (storage_id: {storage_id}): {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
logging.error(f"Unexpected error getting file '{filename}': {e}")
|
||||
raise
|
||||
|
||||
def delete_file_from_storage(self, storage_id: str, filename: str) -> None:
|
||||
"""
|
||||
Delete a file from Brain API storage.
|
||||
|
||||
Args:
|
||||
storage_id: Storage ID of the file to delete
|
||||
filename: Name of the file to delete
|
||||
|
||||
Raises:
|
||||
requests.RequestException: If deletion fails
|
||||
"""
|
||||
url = f"{self.base_url}/storage/{storage_id}"
|
||||
params = {'filename': filename}
|
||||
|
||||
try:
|
||||
logging.info(f"Deleting file '{filename}' (storage_id: {storage_id}) from Brain API storage...")
|
||||
response = requests.delete(
|
||||
url,
|
||||
params=params,
|
||||
headers=self.headers,
|
||||
timeout=60
|
||||
)
|
||||
response.raise_for_status()
|
||||
|
||||
logging.info(f"File deletion successful")
|
||||
|
||||
except requests.RequestException as e:
|
||||
logging.error(f"Failed to delete file '{filename}' (storage_id: {storage_id}): {e}")
|
||||
raise
|
||||
except Exception as e:
|
||||
logging.error(f"Unexpected error deleting file '{filename}': {e}")
|
||||
raise
|
||||
|
||||
def health_check(self) -> bool:
|
||||
"""
|
||||
Check if the Brain API is accessible and authentication is working.
|
||||
|
||||
Returns:
|
||||
True if API is accessible, False otherwise
|
||||
"""
|
||||
try:
|
||||
# Try to access a simple endpoint to verify connectivity
|
||||
url = f"{self.base_url}/health" # Assuming there's a health endpoint
|
||||
response = requests.get(url, headers=self.headers, timeout=10)
|
||||
return response.status_code == 200
|
||||
except:
|
||||
# If health endpoint doesn't exist, try the storage upload endpoint
|
||||
# with a HEAD request to check authentication
|
||||
try:
|
||||
url = f"{self.base_url}/storage/upload"
|
||||
response = requests.head(url, headers=self.headers, timeout=10)
|
||||
return response.status_code in [200, 405] # 405 Method Not Allowed is OK for HEAD
|
||||
except:
|
||||
return False
|
||||
|
||||
|
||||
# Global client instance
|
||||
_brain_api_client = None
|
||||
|
||||
def get_brain_api_client() -> BrainApiClient:
|
||||
"""
|
||||
Get or create the global Brain API client instance.
|
||||
|
||||
Returns:
|
||||
BrainApiClient instance
|
||||
"""
|
||||
global _brain_api_client
|
||||
if _brain_api_client is None:
|
||||
_brain_api_client = BrainApiClient()
|
||||
return _brain_api_client
|
||||
@@ -0,0 +1,88 @@
|
||||
category = "Nilor Nodes 👺"
|
||||
subcategories = {
|
||||
"io": "/IO",
|
||||
}
|
||||
|
||||
# Unique hook type for controller wiring (used by both Preset and Group controllers)
|
||||
CONTROLLER_HOOK = "CONTROLLER_HOOK"
|
||||
|
||||
|
||||
class NilorPreset:
|
||||
"""
|
||||
Declarative controller that binds a Brando preset group to a set of connected inputs.
|
||||
|
||||
- preset_group_name: Semantic key used to look up choices and values in
|
||||
presets_config.json5 via PresetsService (drives dropdown + value application).
|
||||
- _preset_hook_*: Dynamic inputs that accept CONTROLLER_HOOK from NilorUserInput_* `_controller_hook` outputs.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
# Start with a single hook; dynamic inputs handled by companion JS
|
||||
optional_inputs = {"_preset_hook_1": (CONTROLLER_HOOK,)}
|
||||
return {
|
||||
"required": {
|
||||
# Lookup key in presets_config.json5 (NOT a UI label)
|
||||
"preset_group_name": (
|
||||
"STRING",
|
||||
{"default": "my_preset", "multiline": False},
|
||||
),
|
||||
},
|
||||
"optional": optional_inputs,
|
||||
}
|
||||
|
||||
# No outputs; declarative controller only
|
||||
RETURN_TYPES = tuple()
|
||||
RETURN_NAMES = tuple()
|
||||
FUNCTION = "do_nothing"
|
||||
CATEGORY = category + subcategories["io"]
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def do_nothing(self, **kwargs):
|
||||
# This node performs no computation; it exists for declarative wiring only
|
||||
return tuple()
|
||||
|
||||
|
||||
class NilorGroup:
|
||||
"""
|
||||
Declarative UI-grouper that clusters connected inputs together in the Brando UI.
|
||||
|
||||
- group_label: Purely a visual label for the gr.Group that will contain the inputs.
|
||||
It does NOT look up presets or apply values.
|
||||
- _group_hook_*: Dynamic inputs that accept CONTROLLER_HOOK from NilorUserInput_* `_controller_hook` outputs.
|
||||
Reuses a shared controller hook so no additional output types are required on input nodes.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
# Start with a single hook; dynamic inputs handled by companion JS
|
||||
optional_inputs = {"_group_hook_1": (CONTROLLER_HOOK,)}
|
||||
return {
|
||||
"required": {
|
||||
# UI label only (NOT used to look up presets)
|
||||
"group_label": ("STRING", {"default": "my_group", "multiline": False}),
|
||||
},
|
||||
"optional": optional_inputs,
|
||||
}
|
||||
|
||||
# No outputs; declarative controller only
|
||||
RETURN_TYPES = tuple()
|
||||
RETURN_NAMES = tuple()
|
||||
FUNCTION = "do_nothing"
|
||||
CATEGORY = category + subcategories["io"]
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def do_nothing(self, **kwargs):
|
||||
# Declarative only
|
||||
return tuple()
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"NilorPreset": NilorPreset,
|
||||
"NilorGroup": NilorGroup,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"NilorPreset": "👺 User Input Preset Controller",
|
||||
"NilorGroup": "👺 User Input Group Controller",
|
||||
}
|
||||
+366
@@ -0,0 +1,366 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
import requests
|
||||
import io
|
||||
import logging
|
||||
import imageio.v2 as imageio
|
||||
import mimetypes
|
||||
import boto3
|
||||
import os
|
||||
import json
|
||||
from dotenv import load_dotenv
|
||||
from .brain_api_client import get_brain_api_client
|
||||
|
||||
# --- Load Environment Variables ---
|
||||
# Get the directory of the current script
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
# Construct the path to the .env file
|
||||
dotenv_path = os.path.join(current_dir, ".env")
|
||||
# Load the .env file
|
||||
load_dotenv(dotenv_path=dotenv_path)
|
||||
|
||||
# --- Setup Logging ---
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
|
||||
# --- Node Categories ---
|
||||
category = "Nilor Nodes 👺"
|
||||
subcategories = {
|
||||
"streaming": "/Streaming",
|
||||
}
|
||||
|
||||
|
||||
# --- MediaStreamInput: Universal Media Downloader ---
|
||||
class MediaStreamInput:
|
||||
"""
|
||||
A custom node to download an image/video from a pre-signed URL and provide it as a tensor.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_name": (
|
||||
"STRING",
|
||||
{"default": "default_input", "multiline": False},
|
||||
),
|
||||
"format": (["image", "image_batch", "video"],),
|
||||
"storage_id": (
|
||||
"STRING",
|
||||
{"default": "<auto-filled by system>", "multiline": False},
|
||||
),
|
||||
"filename": (
|
||||
"STRING",
|
||||
{"default": "<auto-filled by system>", "multiline": False},
|
||||
),
|
||||
},
|
||||
"hidden": {},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
FUNCTION = "download"
|
||||
CATEGORY = category + subcategories["streaming"]
|
||||
|
||||
def download(
|
||||
self,
|
||||
storage_id: str,
|
||||
filename: str,
|
||||
format: str,
|
||||
input_name: str = "default_input",
|
||||
):
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes: MediaStreamInput: Downloading file '{filename}' (storage_id: {storage_id}) for input '{input_name}' with format '{format}'"
|
||||
)
|
||||
try:
|
||||
# Get Brain API client
|
||||
brain_client = get_brain_api_client()
|
||||
|
||||
# Two-phase download for batches: manifest first, then assets
|
||||
if format == "image_batch":
|
||||
# Download manifest file first
|
||||
manifest_bytes = brain_client.get_file_from_storage(storage_id, filename)
|
||||
manifest = json.loads(manifest_bytes.decode('utf-8'))
|
||||
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes: Processing manifest for '{manifest.get('input_name')}' with {len(manifest.get('files', []))} assets."
|
||||
)
|
||||
|
||||
# Sort files by sequence number to ensure correct order
|
||||
sorted_files = sorted(
|
||||
manifest.get("files", []), key=lambda x: x.get("sequence", 0)
|
||||
)
|
||||
|
||||
# Download all assets using Brain API client
|
||||
asset_responses = []
|
||||
for file_info in sorted_files:
|
||||
try:
|
||||
# Each file_info should now contain storage_id and filename instead of presigned_url
|
||||
file_storage_id = file_info.get("storage_id")
|
||||
file_filename = file_info.get("filename")
|
||||
if not file_storage_id or not file_filename:
|
||||
raise ValueError(f"Missing storage_id or filename in manifest file info: {file_info}")
|
||||
|
||||
file_bytes = brain_client.get_file_from_storage(file_storage_id, file_filename)
|
||||
asset_responses.append(file_bytes)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes: Failed to download asset {file_info.get('filename')}: {e}"
|
||||
)
|
||||
raise # Re-raise to fail the entire process
|
||||
|
||||
return self._process_image_batch(asset_responses)
|
||||
|
||||
# --- Single-file download ---
|
||||
media_bytes = brain_client.get_file_from_storage(storage_id, filename)
|
||||
|
||||
if format == "video":
|
||||
return self._process_video(media_bytes)
|
||||
elif format == "image":
|
||||
return self._process_image(media_bytes)
|
||||
else:
|
||||
# Should not happen if UI choices are respected
|
||||
raise ValueError(
|
||||
f"[🛑] Nilor-Nodes (MediaStreamInput): Unsupported format '{format}' for single media download."
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Failed to download or process media: {e}"
|
||||
)
|
||||
return (None,)
|
||||
|
||||
def _process_image_batch(self, image_bytes_list):
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing image batch with {len(image_bytes_list)} images..."
|
||||
)
|
||||
output_images = []
|
||||
|
||||
for image_bytes in image_bytes_list:
|
||||
image_pil = Image.open(io.BytesIO(image_bytes))
|
||||
|
||||
rgb_image_pil = image_pil.convert("RGB")
|
||||
image_tensor = torch.from_numpy(
|
||||
np.array(rgb_image_pil).astype(np.float32) / 255.0
|
||||
).unsqueeze(0)
|
||||
|
||||
output_images.append(image_tensor)
|
||||
|
||||
# Concatenate along the batch dimension (dim=0)
|
||||
images_tensor = torch.cat(output_images, dim=0)
|
||||
|
||||
logging.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamInput): Image batch processing successful. Batch shape: {images_tensor.shape}"
|
||||
)
|
||||
return (images_tensor,)
|
||||
|
||||
def _process_image(self, image_bytes):
|
||||
logging.info("ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing as image...")
|
||||
image_pil = Image.open(io.BytesIO(image_bytes))
|
||||
|
||||
# Ensure image is in RGB
|
||||
rgb_image_pil = image_pil.convert("RGB")
|
||||
image_tensor = torch.from_numpy(
|
||||
np.array(rgb_image_pil).astype(np.float32) / 255.0
|
||||
).unsqueeze(0)
|
||||
|
||||
logging.info("✅ Nilor-Nodes (MediaStreamInput): Image processing successful.")
|
||||
return (image_tensor,)
|
||||
|
||||
def _process_video(self, video_bytes):
|
||||
logging.info("ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing as video...")
|
||||
frames = []
|
||||
with imageio.get_reader(io.BytesIO(video_bytes), format="mp4") as reader:
|
||||
for frame in reader:
|
||||
# Convert frame to RGB PIL Image and then to tensor
|
||||
pil_image = Image.fromarray(frame).convert("RGB")
|
||||
numpy_image = np.array(pil_image).astype(np.float32) / 255.0
|
||||
tensor_frame = torch.from_numpy(numpy_image)
|
||||
frames.append(tensor_frame)
|
||||
|
||||
if not frames:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (MediaStreamInput): No frames could be read from the video."
|
||||
)
|
||||
|
||||
# Stack frames into a single tensor (batch of images)
|
||||
video_tensor = torch.stack(frames)
|
||||
|
||||
logging.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamInput): Video processing successful. Image Shape: {video_tensor.shape}"
|
||||
)
|
||||
return (video_tensor,)
|
||||
|
||||
|
||||
# --- MediaStreamOutput: Universal Media Uploader & SQS Notifier ---
|
||||
class MediaStreamOutput:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"output_name": (
|
||||
"STRING",
|
||||
{"default": "default_output", "multiline": False},
|
||||
),
|
||||
"images": ("IMAGE",),
|
||||
"format": (["png", "mp4"],),
|
||||
"framerate": ("INT", {"default": 24, "min": 1, "max": 240, "step": 1}),
|
||||
"content_id": (
|
||||
"STRING",
|
||||
{"default": "<auto-filled by system>", "multiline": False},
|
||||
),
|
||||
"venue": (
|
||||
"STRING",
|
||||
{"default": "<auto-filled by system>", "multiline": False},
|
||||
),
|
||||
"canvas": (
|
||||
"STRING",
|
||||
{"default": "<auto-filled by system>", "multiline": False},
|
||||
),
|
||||
"scene": (
|
||||
"STRING",
|
||||
{"default": "<auto-filled by system>", "multiline": False},
|
||||
),
|
||||
"job_completions_queue_url": (
|
||||
"STRING",
|
||||
{"multiline": True, "default": "<auto-filled by system>"},
|
||||
),
|
||||
},
|
||||
"hidden": {
|
||||
"prompt": "PROMPT",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO",
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "upload_and_notify"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = category + subcategories["streaming"]
|
||||
|
||||
def upload_and_notify(
|
||||
self,
|
||||
images,
|
||||
format,
|
||||
content_id,
|
||||
venue,
|
||||
canvas,
|
||||
scene,
|
||||
job_completions_queue_url,
|
||||
framerate,
|
||||
output_name: str = "default_output",
|
||||
prompt=None,
|
||||
extra_pnginfo=None,
|
||||
):
|
||||
if not content_id:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (MediaStreamOutput): content_id is a required input for MediaStreamOutput."
|
||||
)
|
||||
|
||||
# No longer need to parse output_object_keys since we use storage_ids directly
|
||||
|
||||
# Upload the media using Brain API client
|
||||
brain_client = get_brain_api_client()
|
||||
storage_result = None
|
||||
|
||||
if format == "png":
|
||||
storage_result = self._upload_image(images[0], brain_client, output_name)
|
||||
elif format == "mp4":
|
||||
storage_result = self._upload_video(images, brain_client, framerate, output_name)
|
||||
|
||||
# Use the storage_id from the upload result for the SQS message
|
||||
if not storage_result or not storage_result.get('storage_id'):
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): FATAL -- Upload failed or no storage_id returned."
|
||||
)
|
||||
# Send an empty dictionary to signal failure.
|
||||
final_outputs_for_sqs = {}
|
||||
else:
|
||||
# Use storage_id instead of object key
|
||||
storage_id = storage_result['storage_id']
|
||||
final_outputs_for_sqs = {output_name: storage_id}
|
||||
|
||||
# After upload, send the filtered dictionary of outputs to the SQS queue.
|
||||
completion_message = {
|
||||
"content_id": content_id,
|
||||
"status": "completed",
|
||||
"venue": venue,
|
||||
"canvas": canvas,
|
||||
"scene": scene,
|
||||
"outputs": final_outputs_for_sqs,
|
||||
}
|
||||
|
||||
try:
|
||||
# Re-initialize the client inside the execution to ensure it picks up env vars correctly.
|
||||
sqs_client = boto3.client(
|
||||
"sqs",
|
||||
endpoint_url=os.getenv("SQS_ENDPOINT_URL"),
|
||||
aws_access_key_id=os.getenv("AWS_ACCESS_KEY_ID", "local"),
|
||||
aws_secret_access_key=os.getenv("AWS_SECRET_ACCESS_KEY", "local"),
|
||||
region_name=os.getenv("AWS_DEFAULT_REGION", "us-east-1"),
|
||||
)
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Sending completion message for content {content_id} to queue: {job_completions_queue_url}"
|
||||
)
|
||||
sqs_client.send_message(
|
||||
QueueUrl=job_completions_queue_url,
|
||||
MessageBody=json.dumps(completion_message),
|
||||
)
|
||||
logging.info(
|
||||
"✅ Nilor-Nodes (MediaStreamOutput): Completion message sent successfully."
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): Failed to send completion message to SQS: {e}"
|
||||
)
|
||||
raise # Re-raise to fail the ComfyUI job
|
||||
|
||||
return {"ui": {"images": []}}
|
||||
|
||||
def _upload_image(self, image_tensor, brain_client, output_name):
|
||||
logging.info(
|
||||
"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Uploading as PNG image..."
|
||||
)
|
||||
i = 255.0 * image_tensor.cpu().numpy()
|
||||
img_pil = Image.fromarray(np.clip(i, 0, 255).astype(np.uint8))
|
||||
|
||||
buffer = io.BytesIO()
|
||||
img_pil.save(buffer, format="PNG", compress_level=4)
|
||||
buffer.seek(0)
|
||||
|
||||
filename = f"{output_name}.png"
|
||||
return brain_client.upload_fileobj_to_storage(buffer, filename, "image/png")
|
||||
|
||||
def _upload_video(self, image_batch_tensor, brain_client, framerate, output_name):
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Uploading as MP4 video. Frame count: {len(image_batch_tensor)}"
|
||||
)
|
||||
frames = []
|
||||
for image_tensor in image_batch_tensor:
|
||||
i = 255.0 * image_tensor.cpu().numpy()
|
||||
frame = np.clip(i, 0, 255).astype(np.uint8)
|
||||
frames.append(frame)
|
||||
|
||||
buffer = io.BytesIO()
|
||||
imageio.mimwrite(buffer, frames, format="mp4", fps=framerate, quality=8)
|
||||
buffer.seek(0)
|
||||
|
||||
filename = f"{output_name}.mp4"
|
||||
return brain_client.upload_fileobj_to_storage(buffer, filename, "video/mp4")
|
||||
|
||||
|
||||
|
||||
# --- Node Mappings ---
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"MediaStreamInput": MediaStreamInput,
|
||||
"MediaStreamOutput": MediaStreamOutput,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"MediaStreamInput": "👺 Media Stream Input (Storage)",
|
||||
"MediaStreamOutput": "👺 Media Stream Output (Storage)",
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Test script for Brain API Client
|
||||
|
||||
This script tests the Brain API client functionality to ensure it can
|
||||
communicate with the Brain API storage endpoints correctly.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
# Add the current directory to the Python path
|
||||
current_dir = Path(__file__).parent
|
||||
sys.path.insert(0, str(current_dir))
|
||||
|
||||
from brain_api_client import get_brain_api_client
|
||||
|
||||
# Setup logging
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
|
||||
def test_brain_api_client():
|
||||
"""Test the Brain API client functionality."""
|
||||
print("🧪 Testing Brain API Client...")
|
||||
|
||||
try:
|
||||
# Initialize the client
|
||||
client = get_brain_api_client()
|
||||
print("✅ Brain API client initialized successfully")
|
||||
|
||||
# Test health check
|
||||
print("🔍 Testing health check...")
|
||||
is_healthy = client.health_check()
|
||||
if is_healthy:
|
||||
print("✅ Brain API is accessible")
|
||||
else:
|
||||
print("⚠️ Brain API health check failed - this might be expected if the API is not running")
|
||||
|
||||
# Test file upload
|
||||
print("📤 Testing file upload...")
|
||||
test_content = b"Hello, Brain API! This is a test file."
|
||||
test_filename = "test_file.txt"
|
||||
|
||||
# Create a temporary file
|
||||
with tempfile.NamedTemporaryFile(mode='wb', delete=False, suffix='.txt') as temp_file:
|
||||
temp_file.write(test_content)
|
||||
temp_file_path = temp_file.name
|
||||
|
||||
try:
|
||||
# Upload the file
|
||||
upload_result = client.upload_file_to_storage(temp_file_path, test_filename)
|
||||
print(f"✅ File uploaded successfully. Storage ID: {upload_result.get('storage_id')}")
|
||||
|
||||
storage_id = upload_result.get('storage_id')
|
||||
if storage_id:
|
||||
# Test file download
|
||||
print("📥 Testing file download...")
|
||||
downloaded_content = client.get_file_from_storage(storage_id, test_filename)
|
||||
|
||||
if downloaded_content == test_content:
|
||||
print("✅ File downloaded successfully and content matches")
|
||||
else:
|
||||
print("❌ Downloaded content does not match original")
|
||||
|
||||
# Test file deletion
|
||||
print("🗑️ Testing file deletion...")
|
||||
client.delete_file_from_storage(storage_id, test_filename)
|
||||
print("✅ File deleted successfully")
|
||||
|
||||
finally:
|
||||
# Clean up temporary file
|
||||
os.unlink(temp_file_path)
|
||||
|
||||
print("🎉 All tests passed!")
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ Test failed: {e}")
|
||||
logging.exception("Test failed with exception:")
|
||||
return False
|
||||
|
||||
def test_fileobj_upload():
|
||||
"""Test uploading a file-like object."""
|
||||
print("\n🧪 Testing file object upload...")
|
||||
|
||||
try:
|
||||
client = get_brain_api_client()
|
||||
|
||||
# Create a file-like object
|
||||
import io
|
||||
test_content = b"Hello from file object!"
|
||||
file_obj = io.BytesIO(test_content)
|
||||
|
||||
# Upload the file object
|
||||
upload_result = client.upload_fileobj_to_storage(file_obj, "test_fileobj.txt", "text/plain")
|
||||
print(f"✅ File object uploaded successfully. Storage ID: {upload_result.get('storage_id')}")
|
||||
|
||||
storage_id = upload_result.get('storage_id')
|
||||
if storage_id:
|
||||
# Test download
|
||||
downloaded_content = client.get_file_from_storage(storage_id, "test_fileobj.txt")
|
||||
|
||||
if downloaded_content == test_content:
|
||||
print("✅ File object download successful and content matches")
|
||||
else:
|
||||
print("❌ Downloaded content does not match original")
|
||||
|
||||
# Clean up
|
||||
client.delete_file_from_storage(storage_id, "test_fileobj.txt")
|
||||
print("✅ File object deleted successfully")
|
||||
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ File object test failed: {e}")
|
||||
logging.exception("File object test failed with exception:")
|
||||
return False
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("🚀 Starting Brain API Client Tests")
|
||||
print("=" * 50)
|
||||
|
||||
# Check environment variables
|
||||
api_key = os.getenv("BRANDO_API_KEY")
|
||||
base_url = os.getenv("BRANDO_BRAIN_API_BASE_URL", "http://localhost:2024/api")
|
||||
|
||||
print(f"API Key: {'✅ Set' if api_key else '❌ Not set'}")
|
||||
print(f"Base URL: {base_url}")
|
||||
print()
|
||||
|
||||
if not api_key:
|
||||
print("❌ BRANDO_API_KEY environment variable is not set!")
|
||||
print("Please set it in your .env file or environment.")
|
||||
sys.exit(1)
|
||||
|
||||
# Run tests
|
||||
success = True
|
||||
success &= test_brain_api_client()
|
||||
success &= test_fileobj_upload()
|
||||
|
||||
print("\n" + "=" * 50)
|
||||
if success:
|
||||
print("🎉 All tests completed successfully!")
|
||||
sys.exit(0)
|
||||
else:
|
||||
print("❌ Some tests failed!")
|
||||
sys.exit(1)
|
||||
+109
@@ -0,0 +1,109 @@
|
||||
category = "Nilor Nodes 👺"
|
||||
subcategories = {
|
||||
"io": "/IO",
|
||||
}
|
||||
|
||||
from .controllers import CONTROLLER_HOOK
|
||||
|
||||
|
||||
class NilorUserInput_String:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"input_name": (
|
||||
"STRING",
|
||||
{"default": "my_string_input", "multiline": False},
|
||||
),
|
||||
"value": ("STRING", {"default": "", "multiline": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING", CONTROLLER_HOOK)
|
||||
RETURN_NAMES = ("string", "_controller_hook")
|
||||
FUNCTION = "get_value"
|
||||
CATEGORY = category + subcategories["io"]
|
||||
|
||||
def get_value(self, input_name, value):
|
||||
return (value, None)
|
||||
|
||||
|
||||
class NilorUserInput_Int:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"input_name": (
|
||||
"STRING",
|
||||
{"default": "my_int_input", "multiline": False},
|
||||
),
|
||||
"value": ("INT", {"default": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT", CONTROLLER_HOOK)
|
||||
RETURN_NAMES = ("int", "_controller_hook")
|
||||
FUNCTION = "get_value"
|
||||
CATEGORY = category + subcategories["io"]
|
||||
|
||||
def get_value(self, input_name, value):
|
||||
return (value, None)
|
||||
|
||||
|
||||
class NilorUserInput_Float:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"input_name": (
|
||||
"STRING",
|
||||
{"default": "my_float_input", "multiline": False},
|
||||
),
|
||||
"value": ("FLOAT", {"default": 0.0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT", CONTROLLER_HOOK)
|
||||
RETURN_NAMES = ("float", "_controller_hook")
|
||||
FUNCTION = "get_value"
|
||||
CATEGORY = category + subcategories["io"]
|
||||
|
||||
def get_value(self, input_name, value):
|
||||
return (value, None)
|
||||
|
||||
|
||||
class NilorUserInput_Boolean:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"input_name": (
|
||||
"STRING",
|
||||
{"default": "my_bool_input", "multiline": False},
|
||||
),
|
||||
"value": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BOOLEAN", CONTROLLER_HOOK)
|
||||
RETURN_NAMES = ("boolean", "_controller_hook")
|
||||
FUNCTION = "get_value"
|
||||
CATEGORY = category + subcategories["io"]
|
||||
|
||||
def get_value(self, input_name, value):
|
||||
return (value, None)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"NilorUserInput_String": NilorUserInput_String,
|
||||
"NilorUserInput_Int": NilorUserInput_Int,
|
||||
"NilorUserInput_Float": NilorUserInput_Float,
|
||||
"NilorUserInput_Boolean": NilorUserInput_Boolean,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"NilorUserInput_String": "👺 User Input (String)",
|
||||
"NilorUserInput_Int": "👺 User Input (Int)",
|
||||
"NilorUserInput_Float": "👺 User Input (Float)",
|
||||
"NilorUserInput_Boolean": "👺 User Input (Boolean)",
|
||||
}
|
||||
@@ -0,0 +1,272 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
|
||||
// NilorPreset dynamic inputs extension
|
||||
// Adds a new empty _input_hook_N slot whenever the last slot gets connected, up to a hard cap
|
||||
|
||||
const MAX_INPUTS = 32;
|
||||
// Preset controller constants
|
||||
const CLASS_TYPE = "NilorPreset";
|
||||
const INPUT_PREFIX = "_preset_hook_";
|
||||
|
||||
function isTargetNode(node) {
|
||||
return node && (node.comfyClass === CLASS_TYPE || node.type === CLASS_TYPE);
|
||||
}
|
||||
|
||||
function countHookInputs(node) {
|
||||
return (node.inputs || []).filter((i) => i && i.name?.startsWith(INPUT_PREFIX)).length;
|
||||
}
|
||||
|
||||
function nextInputName(node) {
|
||||
let index = 1;
|
||||
while (index <= MAX_INPUTS) {
|
||||
const key = `${INPUT_PREFIX}${index}`;
|
||||
if (!node.inputs || !node.inputs.find((i) => i.name === key)) {
|
||||
return key;
|
||||
}
|
||||
index++;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function resizeNode(node) {
|
||||
try {
|
||||
const size = node.computeSize();
|
||||
node.onResize?.(size);
|
||||
app.graph?.setDirtyCanvas(true, true);
|
||||
} catch (_) {}
|
||||
}
|
||||
|
||||
function ensureAtLeastOneSlot(node) {
|
||||
if (!isTargetNode(node)) return;
|
||||
if (countHookInputs(node) === 0) {
|
||||
const name = `${INPUT_PREFIX}1`;
|
||||
node.addInput(name, "CONTROLLER_HOOK");
|
||||
resizeNode(node);
|
||||
}
|
||||
}
|
||||
|
||||
function growIfLastLinked(node) {
|
||||
if (!isTargetNode(node)) return;
|
||||
const inputs = (node.inputs || []).filter((i) => i && i.name?.startsWith(INPUT_PREFIX));
|
||||
if (inputs.length === 0) return;
|
||||
const last = inputs[inputs.length - 1];
|
||||
const lastIsLinked = !!last.link;
|
||||
if (lastIsLinked && inputs.length < MAX_INPUTS) {
|
||||
const name = nextInputName(node);
|
||||
if (name) {
|
||||
node.addInput(name, "CONTROLLER_HOOK");
|
||||
resizeNode(node);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function shrinkTrailingUnlinked(node) {
|
||||
if (!isTargetNode(node)) return;
|
||||
const allInputs = node.inputs || [];
|
||||
// Collect indices of hook inputs
|
||||
const hookIndices = [];
|
||||
for (let i = 0; i < allInputs.length; i++) {
|
||||
const inp = allInputs[i];
|
||||
if (inp && inp.name && inp.name.startsWith(INPUT_PREFIX)) {
|
||||
hookIndices.push(i);
|
||||
}
|
||||
}
|
||||
if (hookIndices.length <= 1) return; // always keep at least one
|
||||
|
||||
// Find last linked among hook inputs (by position in hookIndices)
|
||||
let lastLinkedPos = -1;
|
||||
for (let pos = 0; pos < hookIndices.length; pos++) {
|
||||
const idx = hookIndices[pos];
|
||||
if (allInputs[idx]?.link) lastLinkedPos = pos;
|
||||
}
|
||||
|
||||
const targetHookCount = lastLinkedPos >= 0 ? lastLinkedPos + 1 : 1;
|
||||
|
||||
// Remove trailing unlinked beyond targetHookCount
|
||||
for (let pos = hookIndices.length - 1; pos >= targetHookCount; pos--) {
|
||||
const idx = hookIndices[pos];
|
||||
const input = node.inputs[idx];
|
||||
if (input && !input.link) {
|
||||
try {
|
||||
node.removeInput(idx);
|
||||
} catch (e) {
|
||||
console.warn("nilor-preset-dynamic-inputs removeInput error", e);
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
resizeNode(node);
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "comfy.nilor-nodes.userinputPreset",
|
||||
|
||||
// Ensure compatibility with saved/loaded graphs
|
||||
afterConfigureGraph(graph) {
|
||||
try {
|
||||
(graph?._nodes || graph?.nodes || []).forEach((n) => {
|
||||
if (isTargetNode(n)) {
|
||||
ensureAtLeastOneSlot(n);
|
||||
shrinkTrailingUnlinked(n);
|
||||
growIfLastLinked(n);
|
||||
}
|
||||
});
|
||||
} catch (e) {
|
||||
console.warn("nilor-preset-dynamic-inputs afterConfigureGraph error", e);
|
||||
}
|
||||
},
|
||||
|
||||
// Patch the prototype so we always react to connection changes
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, appInstance) {
|
||||
if (nodeData?.name !== CLASS_TYPE) return;
|
||||
const original = nodeType.prototype.onConnectionsChange;
|
||||
nodeType.prototype.onConnectionsChange = function (type, index, connected, link_info) {
|
||||
if (typeof original === "function") {
|
||||
original.apply(this, arguments);
|
||||
}
|
||||
try {
|
||||
shrinkTrailingUnlinked(this);
|
||||
growIfLastLinked(this);
|
||||
} catch (e) {
|
||||
console.warn("nilor-preset-dynamic-inputs onConnectionsChange error", e);
|
||||
}
|
||||
};
|
||||
},
|
||||
|
||||
nodeCreated(node) {
|
||||
if (!isTargetNode(node)) return;
|
||||
ensureAtLeastOneSlot(node);
|
||||
shrinkTrailingUnlinked(node);
|
||||
growIfLastLinked(node);
|
||||
},
|
||||
});
|
||||
|
||||
|
||||
// NilorGroup dynamic inputs extension (mirrors preset behavior)
|
||||
const GROUP_CLASS_TYPE = "NilorGroup";
|
||||
const GROUP_INPUT_PREFIX = "_group_hook_";
|
||||
|
||||
function isGroupNode(node) {
|
||||
return node && (node.comfyClass === GROUP_CLASS_TYPE || node.type === GROUP_CLASS_TYPE);
|
||||
}
|
||||
|
||||
function countGroupHookInputs(node) {
|
||||
return (node.inputs || []).filter((i) => i && i.name?.startsWith(GROUP_INPUT_PREFIX)).length;
|
||||
}
|
||||
|
||||
function nextGroupInputName(node) {
|
||||
let index = 1;
|
||||
while (index <= MAX_INPUTS) {
|
||||
const key = `${GROUP_INPUT_PREFIX}${index}`;
|
||||
if (!node.inputs || !node.inputs.find((i) => i.name === key)) {
|
||||
return key;
|
||||
}
|
||||
index++;
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
function ensureAtLeastOneGroupSlot(node) {
|
||||
if (!isGroupNode(node)) return;
|
||||
if (countGroupHookInputs(node) === 0) {
|
||||
const name = `${GROUP_INPUT_PREFIX}1`;
|
||||
node.addInput(name, "CONTROLLER_HOOK");
|
||||
resizeNode(node);
|
||||
}
|
||||
}
|
||||
|
||||
function growGroupIfLastLinked(node) {
|
||||
if (!isGroupNode(node)) return;
|
||||
const inputs = (node.inputs || []).filter((i) => i && i.name?.startsWith(GROUP_INPUT_PREFIX));
|
||||
if (inputs.length === 0) return;
|
||||
const last = inputs[inputs.length - 1];
|
||||
const lastIsLinked = !!last.link;
|
||||
if (lastIsLinked && inputs.length < MAX_INPUTS) {
|
||||
const name = nextGroupInputName(node);
|
||||
if (name) {
|
||||
node.addInput(name, "CONTROLLER_HOOK");
|
||||
resizeNode(node);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function shrinkGroupTrailingUnlinked(node) {
|
||||
if (!isGroupNode(node)) return;
|
||||
const allInputs = node.inputs || [];
|
||||
const hookIndices = [];
|
||||
for (let i = 0; i < allInputs.length; i++) {
|
||||
const inp = allInputs[i];
|
||||
if (inp && inp.name && inp.name.startsWith(GROUP_INPUT_PREFIX)) {
|
||||
hookIndices.push(i);
|
||||
}
|
||||
}
|
||||
if (hookIndices.length <= 1) return;
|
||||
|
||||
let lastLinkedPos = -1;
|
||||
for (let pos = 0; pos < hookIndices.length; pos++) {
|
||||
const idx = hookIndices[pos];
|
||||
if (allInputs[idx]?.link) lastLinkedPos = pos;
|
||||
}
|
||||
|
||||
const targetHookCount = lastLinkedPos >= 0 ? lastLinkedPos + 1 : 1;
|
||||
|
||||
for (let pos = hookIndices.length - 1; pos >= targetHookCount; pos--) {
|
||||
const idx = hookIndices[pos];
|
||||
const input = node.inputs[idx];
|
||||
if (input && !input.link) {
|
||||
try {
|
||||
node.removeInput(idx);
|
||||
} catch (e) {
|
||||
console.warn("nilor-group-dynamic-inputs removeInput error", e);
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
resizeNode(node);
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "comfy.nilor-nodes.userinputGroup",
|
||||
|
||||
afterConfigureGraph(graph) {
|
||||
try {
|
||||
(graph?._nodes || graph?.nodes || []).forEach((n) => {
|
||||
if (isGroupNode(n)) {
|
||||
ensureAtLeastOneGroupSlot(n);
|
||||
shrinkGroupTrailingUnlinked(n);
|
||||
growGroupIfLastLinked(n);
|
||||
}
|
||||
});
|
||||
} catch (e) {
|
||||
console.warn("nilor-group-dynamic-inputs afterConfigureGraph error", e);
|
||||
}
|
||||
},
|
||||
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, appInstance) {
|
||||
if (nodeData?.name !== GROUP_CLASS_TYPE) return;
|
||||
const original = nodeType.prototype.onConnectionsChange;
|
||||
nodeType.prototype.onConnectionsChange = function (type, index, connected, link_info) {
|
||||
if (typeof original === "function") {
|
||||
original.apply(this, arguments);
|
||||
}
|
||||
try {
|
||||
shrinkGroupTrailingUnlinked(this);
|
||||
growGroupIfLastLinked(this);
|
||||
} catch (e) {
|
||||
console.warn("nilor-group-dynamic-inputs onConnectionsChange error", e);
|
||||
}
|
||||
};
|
||||
},
|
||||
|
||||
nodeCreated(node) {
|
||||
if (!isGroupNode(node)) return;
|
||||
ensureAtLeastOneGroupSlot(node);
|
||||
shrinkGroupTrailingUnlinked(node);
|
||||
growGroupIfLastLinked(node);
|
||||
},
|
||||
});
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
import { app } from "/scripts/app.js";
|
||||
|
||||
function toggleFramerateWidget(node, show) {
|
||||
const framerateWidget = node.widgets.find((w) => w.name === "framerate");
|
||||
if (framerateWidget) {
|
||||
framerateWidget.hidden = !show;
|
||||
// This is a hack to force the node to redraw.
|
||||
//const size = node.computeSize();
|
||||
//node.onResize?.(size);
|
||||
}
|
||||
}
|
||||
|
||||
function hideWidgets(node, widgetNames) {
|
||||
widgetNames.forEach(name => {
|
||||
const widget = node.widgets.find((w) => w.name === name);
|
||||
if (widget) {
|
||||
widget.hidden = true;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "comfy.nilor-nodes.mediaStream",
|
||||
nodeCreated(node) {
|
||||
if (node.comfyClass === "MediaStreamOutput") {
|
||||
// Hide system inputs by default
|
||||
hideWidgets(node, [
|
||||
"content_id",
|
||||
"venue",
|
||||
"canvas",
|
||||
"scene",
|
||||
"presigned_upload_url",
|
||||
"job_completions_queue_url",
|
||||
"output_object_keys"
|
||||
]);
|
||||
|
||||
const formatWidget = node.widgets.find((w) => w.name === "format");
|
||||
|
||||
// Initial toggle for framerate based on the default format value
|
||||
toggleFramerateWidget(node, formatWidget.value === "mp4");
|
||||
|
||||
// Store original callback to chain it
|
||||
const originalCallback = formatWidget.callback;
|
||||
|
||||
formatWidget.callback = function (value) {
|
||||
toggleFramerateWidget(node, value === "mp4");
|
||||
|
||||
// Recalculate node size after toggling widgets
|
||||
const size = node.computeSize();
|
||||
node.onResize?.(size);
|
||||
|
||||
if (originalCallback) {
|
||||
return originalCallback.apply(this, arguments);
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
if (node.comfyClass === "MediaStreamInput") {
|
||||
// Hide system inputs by default
|
||||
hideWidgets(node, ["presigned_download_url"]);
|
||||
}
|
||||
},
|
||||
});
|
||||
@@ -0,0 +1,441 @@
|
||||
"""
|
||||
Worker Consumer Service for ComfyUI
|
||||
|
||||
This script runs as a continuous background service on each ComfyUI worker.
|
||||
Its purpose is to poll the `jobs_to_process` SQS queue for new jobs,
|
||||
submit them to the local ComfyUI server, and manage the message lifecycle.
|
||||
It also listens to the ComfyUI websocket to send a "running" status update
|
||||
at the precise moment that job execution begins.
|
||||
"""
|
||||
|
||||
import os
|
||||
import json
|
||||
import logging
|
||||
import asyncio
|
||||
import aiohttp
|
||||
import websockets
|
||||
from aiobotocore.session import get_session
|
||||
from dotenv import load_dotenv
|
||||
from botocore.exceptions import EndpointConnectionError
|
||||
|
||||
# --- Load Environment Variables ---
|
||||
# Load from the .env file in the same directory
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
dotenv_path = os.path.join(current_dir, ".env")
|
||||
if os.path.exists(dotenv_path):
|
||||
load_dotenv(dotenv_path=dotenv_path)
|
||||
logging.info(
|
||||
f"✅\u2009 Nilor-Nodes: Loaded environment variables from {dotenv_path}"
|
||||
)
|
||||
else:
|
||||
logging.info(
|
||||
"⚠️\u2009 Nilor-Nodes: No .env file found, relying on shell environment variables."
|
||||
)
|
||||
|
||||
# --- Configuration ---
|
||||
SQS_ENDPOINT_URL = os.getenv("SQS_ENDPOINT_URL", "http://localhost:9324")
|
||||
SQS_JOBS_TO_PROCESS_QUEUE_NAME = os.getenv(
|
||||
"SQS_JOBS_TO_PROCESS_QUEUE_NAME", "jobs_to_process"
|
||||
)
|
||||
SQS_JOB_STATUS_UPDATES_QUEUE_NAME = os.getenv(
|
||||
"SQS_JOB_STATUS_UPDATES_QUEUE_NAME", "job_status_updates"
|
||||
)
|
||||
COMFYUI_API_URL = os.getenv("COMFYUI_API_URL", "http://127.0.0.1:8188") + "/prompt"
|
||||
COMFYUI_WS_URL = os.getenv("COMFYUI_WS_URL", "ws://127.0.0.1:8188") + "/ws"
|
||||
AWS_ACCESS_KEY_ID = os.getenv("AWS_ACCESS_KEY_ID", "local")
|
||||
AWS_SECRET_ACCESS_KEY = os.getenv("AWS_SECRET_ACCESS_KEY", "local")
|
||||
AWS_DEFAULT_REGION = os.getenv("AWS_DEFAULT_REGION", "us-east-1")
|
||||
POLL_WAIT_TIME_SECONDS = 20 # SQS Long Polling
|
||||
MAX_MESSAGES = 1
|
||||
|
||||
# --- Setup Logging ---
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
|
||||
|
||||
class WorkerConsumer:
|
||||
def __init__(self):
|
||||
self.session = get_session()
|
||||
self.prompt_id_to_content_id_map = {}
|
||||
self.sent_running_status_prompts = set()
|
||||
self.content_context_by_content_id = {}
|
||||
self.jobs_queue_url = None
|
||||
self.status_updates_queue_url = None
|
||||
self.http_session = None
|
||||
|
||||
async def _initialize_sqs(self):
|
||||
"""Initializes SQS queue URLs. Returns True on success, False on failure."""
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
) as client:
|
||||
try:
|
||||
self.jobs_queue_url = await self._get_queue_url(
|
||||
client, SQS_JOBS_TO_PROCESS_QUEUE_NAME
|
||||
)
|
||||
self.status_updates_queue_url = await self._get_queue_url(
|
||||
client, SQS_JOB_STATUS_UPDATES_QUEUE_NAME
|
||||
)
|
||||
return True
|
||||
except EndpointConnectionError as e:
|
||||
# Quiet the noisy traceback by logging a concise warning instead
|
||||
logging.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS endpoint is unreachable at {SQS_ENDPOINT_URL}: {e}. "
|
||||
)
|
||||
return False
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Failed to initialize SQS queues: {e}"
|
||||
)
|
||||
return False
|
||||
|
||||
async def _get_queue_url(self, client, queue_name):
|
||||
"""Retrieves the SQS queue URL."""
|
||||
try:
|
||||
response = await client.get_queue_url(QueueName=queue_name)
|
||||
return response["QueueUrl"]
|
||||
except client.exceptions.QueueDoesNotExist:
|
||||
logging.error(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS queue '{queue_name}' does not exist."
|
||||
)
|
||||
raise
|
||||
|
||||
async def listen_for_comfy_events(self):
|
||||
while True:
|
||||
try:
|
||||
async with websockets.connect(COMFYUI_WS_URL) as websocket:
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Connected to ComfyUI websocket at {COMFYUI_WS_URL}"
|
||||
)
|
||||
while True:
|
||||
message = await websocket.recv()
|
||||
if isinstance(message, str):
|
||||
try:
|
||||
event = json.loads(message)
|
||||
event_type = event.get("type")
|
||||
data = event.get("data", {})
|
||||
prompt_id = data.get("prompt_id")
|
||||
|
||||
if not prompt_id and "sid" in data:
|
||||
prompt_id = data["sid"]
|
||||
|
||||
if not prompt_id:
|
||||
continue
|
||||
|
||||
# Use the first progress event as a signal that the job is running.
|
||||
if (
|
||||
event_type in ["progress", "progress_state"]
|
||||
and prompt_id in self.prompt_id_to_content_id_map
|
||||
and prompt_id
|
||||
not in self.sent_running_status_prompts
|
||||
):
|
||||
content_id = self.prompt_id_to_content_id_map[
|
||||
prompt_id
|
||||
]
|
||||
ctx = self.content_context_by_content_id.get(
|
||||
content_id, {}
|
||||
)
|
||||
policy = ctx.get("status_policy") or {}
|
||||
running_status = policy.get(
|
||||
"running_status", "running"
|
||||
)
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Execution started for prompt_id {prompt_id} (content_id: {content_id}) via '{event_type}' event. Sending '{running_status}' status."
|
||||
)
|
||||
await self._send_status_update(
|
||||
content_id,
|
||||
running_status,
|
||||
ctx.get("venue"),
|
||||
ctx.get("canvas"),
|
||||
ctx.get("scene"),
|
||||
)
|
||||
self.sent_running_status_prompts.add(prompt_id)
|
||||
|
||||
# Handle execution errors
|
||||
elif event_type == "execution_error":
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Received execution error for prompt_id {prompt_id}: {data}"
|
||||
)
|
||||
if prompt_id in self.prompt_id_to_content_id_map:
|
||||
content_id = (
|
||||
self.prompt_id_to_content_id_map.pop(
|
||||
prompt_id
|
||||
)
|
||||
)
|
||||
ctx = self.content_context_by_content_id.get(
|
||||
content_id, {}
|
||||
)
|
||||
policy = ctx.get("status_policy") or {}
|
||||
fail_status = policy.get(
|
||||
"fail_status", "failed"
|
||||
)
|
||||
try:
|
||||
await self._send_status_update(
|
||||
content_id,
|
||||
fail_status,
|
||||
ctx.get("venue"),
|
||||
ctx.get("canvas"),
|
||||
ctx.get("scene"),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
self.content_context_by_content_id.pop(
|
||||
content_id, None
|
||||
)
|
||||
self.sent_running_status_prompts.discard(prompt_id)
|
||||
|
||||
# Log successful execution
|
||||
elif event_type == "executed":
|
||||
logging.info(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Prompt {prompt_id} executed successfully according to websocket event. Final node is responsible for sending completion message."
|
||||
)
|
||||
if prompt_id in self.prompt_id_to_content_id_map:
|
||||
content_id = (
|
||||
self.prompt_id_to_content_id_map.pop(
|
||||
prompt_id
|
||||
)
|
||||
)
|
||||
self.content_context_by_content_id.pop(
|
||||
content_id, None
|
||||
)
|
||||
self.sent_running_status_prompts.discard(prompt_id)
|
||||
|
||||
elif event_type not in ["progress", "progress_state"]:
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Received ComfyUI websocket event of type '{event_type}': {data}"
|
||||
)
|
||||
|
||||
except json.JSONDecodeError:
|
||||
logging.debug(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Received non-JSON text message from websocket, ignoring."
|
||||
)
|
||||
else:
|
||||
logging.debug(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Received binary message from websocket, ignoring."
|
||||
)
|
||||
except (
|
||||
websockets.exceptions.ConnectionClosedError,
|
||||
ConnectionRefusedError,
|
||||
) as e:
|
||||
logging.warning(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): ComfyUI websocket connection failed: {e}. Retrying in 5 seconds..."
|
||||
)
|
||||
await asyncio.sleep(5)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): An unexpected error occurred in the websocket listener: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
await asyncio.sleep(10)
|
||||
|
||||
async def consume_loop(self):
|
||||
"""The main loop to continuously poll for and process messages.
|
||||
Keeps retrying SQS initialization and polling if the endpoint is down.
|
||||
"""
|
||||
|
||||
# Start the websocket listener in the background immediately
|
||||
listener_task = asyncio.create_task(self.listen_for_comfy_events())
|
||||
|
||||
try:
|
||||
while True:
|
||||
# Ensure SQS is initialized; if not, keep attempting to initialize
|
||||
if self.jobs_queue_url is None or self.status_updates_queue_url is None:
|
||||
initialized = await self._initialize_sqs()
|
||||
if not initialized:
|
||||
logging.warning(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS initialization failed. Retrying in 10 seconds..."
|
||||
)
|
||||
await asyncio.sleep(10)
|
||||
continue
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Starting worker consumer. Polling queue: {self.jobs_queue_url}"
|
||||
)
|
||||
|
||||
logging.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Polling for messages..."
|
||||
)
|
||||
try:
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
) as client:
|
||||
response = await client.receive_message(
|
||||
QueueUrl=self.jobs_queue_url,
|
||||
MaxNumberOfMessages=MAX_MESSAGES,
|
||||
WaitTimeSeconds=POLL_WAIT_TIME_SECONDS,
|
||||
)
|
||||
|
||||
messages = response.get("Messages", [])
|
||||
if not messages:
|
||||
logging.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): No messages received."
|
||||
)
|
||||
continue
|
||||
|
||||
for message in messages:
|
||||
try:
|
||||
await self.process_message(message)
|
||||
|
||||
# On successful processing, delete the message
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
) as client:
|
||||
await client.delete_message(
|
||||
QueueUrl=self.jobs_queue_url,
|
||||
ReceiptHandle=message["ReceiptHandle"],
|
||||
)
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Deleted message {message['MessageId']} from queue."
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
# This is a poison pill message, log it but don't retry.
|
||||
# It will be moved to the DLQ after enough failed receives.
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Message {message['MessageId']} is a poison pill (JSON decode failed) and will be ignored."
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Processing failed for message {message['MessageId']}: {e}. It will be returned to the queue for retry."
|
||||
)
|
||||
|
||||
except EndpointConnectionError as e:
|
||||
# Lost connection to SQS; reset and re-initialize on next loop
|
||||
logging.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Lost connection to SQS at {SQS_ENDPOINT_URL}: {e}. Will retry initialization in 10 seconds."
|
||||
)
|
||||
self.jobs_queue_url = None
|
||||
self.status_updates_queue_url = None
|
||||
await asyncio.sleep(10)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): An error occurred in the consume loop: {e}"
|
||||
)
|
||||
await asyncio.sleep(10) # Wait before retrying
|
||||
finally:
|
||||
listener_task.cancel()
|
||||
await asyncio.gather(listener_task, return_exceptions=True)
|
||||
logging.info(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Websocket listener stopped."
|
||||
)
|
||||
|
||||
async def process_message(self, message):
|
||||
"""Processes a single SQS message."""
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Processing message: {message['MessageId']}"
|
||||
)
|
||||
|
||||
try:
|
||||
body = json.loads(message["Body"])
|
||||
# SQS messages are often double-encoded, with the actual payload inside a 'Message' key.
|
||||
if "Message" in body:
|
||||
job_payload = json.loads(body["Message"])
|
||||
else:
|
||||
job_payload = body
|
||||
|
||||
content_id = job_payload.get("content_id")
|
||||
|
||||
# Validate that the payload has the required keys before submitting.
|
||||
if not content_id or "prompt" not in job_payload:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Invalid message format: missing 'content_id' or 'prompt'. Payload: {job_payload}"
|
||||
)
|
||||
return
|
||||
|
||||
# Submit to ComfyUI
|
||||
await self._submit_job_to_comfyui(content_id, job_payload)
|
||||
|
||||
# Cache context for subsequent status updates
|
||||
try:
|
||||
self.content_context_by_content_id[content_id] = {
|
||||
"venue": job_payload.get("venue"),
|
||||
"canvas": job_payload.get("canvas"),
|
||||
"scene": job_payload.get("scene"),
|
||||
"status_policy": job_payload.get("status_policy") or {},
|
||||
}
|
||||
except Exception:
|
||||
self.content_context_by_content_id[content_id] = {}
|
||||
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): An unexpected error occurred while processing message: {e}. It will be retried."
|
||||
)
|
||||
# Re-raise to prevent deletion from queue if we want SQS to handle retry
|
||||
raise
|
||||
|
||||
async def _submit_job_to_comfyui(self, content_id, workflow_data):
|
||||
"""Submits a single job to the ComfyUI API."""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
COMFYUI_API_URL, json=workflow_data, timeout=30
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
response_json = await response.json()
|
||||
prompt_id = response_json.get("prompt_id")
|
||||
logging.info(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Successfully submitted job to ComfyUI. Prompt ID: {prompt_id}"
|
||||
)
|
||||
self.prompt_id_to_content_id_map[prompt_id] = content_id
|
||||
|
||||
# No need to delete here, the consume_loop handles message deletion
|
||||
except aiohttp.ClientError as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to submit job to ComfyUI: {e}. Message will be retried."
|
||||
)
|
||||
except (json.JSONDecodeError, KeyError) as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to parse ComfyUI response: {e}. Discarding malformed response."
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): An unexpected error occurred while submitting job to ComfyUI: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
async def _send_status_update(
|
||||
self, content_id, status, venue=None, canvas=None, scene=None
|
||||
):
|
||||
try:
|
||||
body = {"content_id": content_id, "status": status}
|
||||
if venue is not None:
|
||||
body["venue"] = venue
|
||||
if canvas is not None:
|
||||
body["canvas"] = canvas
|
||||
if scene is not None:
|
||||
body["scene"] = scene
|
||||
message_body = json.dumps(body)
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
) as client:
|
||||
await client.send_message(
|
||||
QueueUrl=self.status_updates_queue_url, MessageBody=message_body
|
||||
)
|
||||
logging.info(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Sent status update for content {content_id}: {status}"
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to send status update for content {content_id}: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
async def consume_jobs():
|
||||
"""Entry point function to be called in a background thread."""
|
||||
consumer = WorkerConsumer()
|
||||
await consumer.consume_loop()
|
||||
Reference in New Issue
Block a user