Compare commits
38
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
629f87a2c3 | ||
|
|
e14ffc2284 | ||
|
|
fcecb2d772 | ||
|
|
6ae6487666 | ||
|
|
377e4a166a | ||
|
|
d1cdbeb7ad | ||
|
|
a139ea1801 | ||
|
|
6786d94d44 | ||
|
|
4f45e0130e | ||
|
|
eed9044703 | ||
|
|
c7db4fae18 | ||
|
|
6e564d2356 | ||
|
|
7af35295c6 | ||
|
|
761fa26f0a | ||
|
|
ea320def65 | ||
|
|
c817292af9 | ||
|
|
39008556a2 | ||
|
|
42c9b3589e | ||
|
|
ac780e1d55 | ||
|
|
ebe728b286 | ||
|
|
6ca1a7b201 | ||
|
|
156b87ec59 | ||
|
|
5dbbb68b81 | ||
|
|
120ee1f183 | ||
|
|
9ab44563df | ||
|
|
7d12d43613 | ||
|
|
1b2af4e2cc | ||
|
|
f51c647010 | ||
|
|
ae16552617 | ||
|
|
6f98c20d96 | ||
|
|
16a8bb19e2 | ||
|
|
cbd87cf960 | ||
|
|
1ba5ddeb03 | ||
|
|
bb1583bfc8 | ||
|
|
be04cf39bf | ||
|
|
a116a42062 | ||
|
|
336217df89 | ||
|
|
aeaabd483d |
@@ -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
|
||||
@@ -0,0 +1,25 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'nilor-corp' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
@@ -1,2 +1,191 @@
|
||||
# nilor-nodes
|
||||
Custom utility nodes for ComfyUI
|
||||
# Nilor Nodes Documentation 👺
|
||||
|
||||
A collection of utility nodes for ComfyUI focusing on list manipulation, batch operations, and advanced I/O functionality.
|
||||
|
||||
## 🏭 Generators
|
||||
|
||||
<details>
|
||||
<summary><b>Interpolated Float List</b></summary>
|
||||
|
||||
Generates a list of interpolated float values based on sections.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| number_of_floats | INT | Total number of float values to generate |
|
||||
| number_of_sections | INT | Number of sections to divide into |
|
||||
| section_number | INT | Current section being processed |
|
||||
| interpolation_type | ["slinear", "quadratic", "cubic"] | Type of interpolation |
|
||||
|
||||
| Output | Type | Description |
|
||||
|--------|------|-------------|
|
||||
| floats | FLOAT | List of interpolated float values |
|
||||
|
||||
**Notes**: Creates smooth transitions between values using scipy's interpolation.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>One Minus Float List</b></summary>
|
||||
|
||||
Creates an inverted list of float values (1 - x).
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| list_of_floats | FLOAT | Input float list |
|
||||
|
||||
| Output | Type | Description |
|
||||
|--------|------|-------------|
|
||||
| floats | FLOAT | Inverted float values |
|
||||
|
||||
**Notes**: Simple inversion operation, useful for creating complementary values.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Remap Float List</b></summary>
|
||||
|
||||
Remaps a list of float values from one range to another.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| list_of_floats | FLOAT | Input float list |
|
||||
| min_input | FLOAT | Minimum input value (default: 0.0) |
|
||||
| max_input | FLOAT | Maximum input value (default: 1.0) |
|
||||
| min_output | FLOAT | Minimum output value (default: 0.0) |
|
||||
| max_output | FLOAT | Maximum output value (default: 1.0) |
|
||||
|
||||
| Output | Type | Description |
|
||||
|--------|------|-------------|
|
||||
| remapped_floats | FLOAT | Remapped float values |
|
||||
|
||||
**Notes**: Useful for scaling values between different ranges while preserving relationships.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Inverse Map Float List</b></summary>
|
||||
|
||||
Creates a mirror mapping of float values around their midpoint.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| list_of_floats | FLOAT | Input float list |
|
||||
|
||||
| Output | Type | Description |
|
||||
|--------|------|-------------|
|
||||
| floats | FLOAT | Inverse mapped values |
|
||||
|
||||
**Notes**: Automatically determines min/max from input list.
|
||||
</details>
|
||||
|
||||
## 🛠️ Utilities
|
||||
|
||||
<details>
|
||||
<summary><b>Int To List Of Bools</b></summary>
|
||||
|
||||
Converts an integer into a list of boolean values.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| number_of_images | INT | Number to convert |
|
||||
|
||||
| Output | Type | Description |
|
||||
|--------|------|-------------|
|
||||
| booleans | BOOLEAN | List of boolean values |
|
||||
|
||||
**Notes**: Creates a list where first N values are True, rest are False.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>List of Ints</b></summary>
|
||||
|
||||
Generates a sequential or shuffled list of integers.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| min | INT | Starting integer (default: 0) |
|
||||
| max | INT | Ending integer (default: 9) |
|
||||
| shuffle | BOOLEAN | Whether to randomize order |
|
||||
|
||||
| Output | Type | Description |
|
||||
|--------|------|-------------|
|
||||
| ints | INT | List of integers |
|
||||
|
||||
**Notes**: Output is always a list, even for single values.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Select Index From List</b></summary>
|
||||
|
||||
Extracts a single item from a list at the specified index.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| list_of_any | any | Input list of any type |
|
||||
| index | INT | Index to select (default: 0) |
|
||||
|
||||
| Output | Type | Description |
|
||||
|--------|------|-------------|
|
||||
| any | any | Selected item |
|
||||
|
||||
**Notes**: Uses custom AnyType to accept any input type. Handles tensor unpacking automatically.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Shuffle Image Batch</b></summary>
|
||||
|
||||
Randomly reorders images in a batch.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| images | IMAGE | Batch of images |
|
||||
| seed | INT | Random seed for shuffling |
|
||||
|
||||
| Output | Type | Description |
|
||||
|--------|------|-------------|
|
||||
| images | IMAGE | Shuffled image batch |
|
||||
|
||||
**Notes**: Maintains batch dimensions while randomizing order.
|
||||
</details>
|
||||
|
||||
## 💾 I/O Operations
|
||||
|
||||
<details>
|
||||
<summary><b>Save Image To HF Dataset</b></summary>
|
||||
|
||||
Uploads images to a HuggingFace dataset.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| image | IMAGE | Image to upload |
|
||||
| repository_id | STRING | HuggingFace dataset repository |
|
||||
| hf_auth_token | STRING | HuggingFace authentication token |
|
||||
| filename_prefix | STRING | Prefix for saved files |
|
||||
|
||||
**Notes**: Requires HuggingFace authentication token and repository access.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Save EXR Arbitrary</b></summary>
|
||||
|
||||
Saves multi-channel data as an OpenEXR file.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| channels | any | List of tensor channels |
|
||||
| filename_prefix | STRING | Output filename prefix |
|
||||
|
||||
**Notes**: Supports arbitrary number of channels. Each channel must have same dimensions.
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><b>Save Video To HF Dataset</b></summary>
|
||||
|
||||
Uploads video files to a HuggingFace dataset.
|
||||
|
||||
| Input | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| filenames | VHS_FILENAMES | List of video files |
|
||||
| repository_id | STRING | HuggingFace dataset repository |
|
||||
| hf_auth_token | STRING | HuggingFace authentication token |
|
||||
| filename_prefix | STRING | Prefix for saved files |
|
||||
|
||||
**Notes**: Handles batch upload of multiple video files.
|
||||
</details>
|
||||
+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)",
|
||||
}
|
||||
+803
-24
@@ -13,6 +13,13 @@ import OpenEXR
|
||||
import Imath
|
||||
import folder_paths
|
||||
import torch
|
||||
import builtins
|
||||
from pathlib import Path
|
||||
import cv2
|
||||
from .utils import pil2tensor, tensor2pil
|
||||
|
||||
BIGMIN = -(2**53 - 1)
|
||||
BIGMAX = 2**53 - 1
|
||||
|
||||
category = "Nilor Nodes 👺"
|
||||
subcategories = {
|
||||
@@ -20,6 +27,8 @@ subcategories = {
|
||||
"utilities": "/Utilities",
|
||||
"io": "/IO",
|
||||
}
|
||||
|
||||
|
||||
class AnyType(str):
|
||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||
|
||||
@@ -33,9 +42,7 @@ class AnyType(str):
|
||||
any = AnyType("*")
|
||||
|
||||
|
||||
|
||||
|
||||
class NilorInterpolatedFloatList: # Generate interpolated float values based on a number of sections
|
||||
class NilorInterpolatedFloatList: # Generate interpolated float values based on a number of sections
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@@ -47,7 +54,10 @@ class NilorInterpolatedFloatList: # Generate interpolated float values based on
|
||||
"number_of_floats": ("INT", {"forceInput": False}),
|
||||
"number_of_sections": ("INT", {"forceInput": False}),
|
||||
"section_number": ("INT", {"forceInput": False}),
|
||||
"interpolation_type": (["slinear","quadratic", "cubic"], {}), # Type of interpolation to use
|
||||
"interpolation_type": (
|
||||
["slinear", "quadratic", "cubic"],
|
||||
{},
|
||||
), # Type of interpolation to use
|
||||
},
|
||||
}
|
||||
|
||||
@@ -67,7 +77,9 @@ class NilorInterpolatedFloatList: # Generate interpolated float values based on
|
||||
f = interp1d(x, y, kind=interp_type)
|
||||
return f(x)
|
||||
|
||||
def generate_float_list(self, number_of_floats, number_of_sections, section_number, interpolation_type):
|
||||
def generate_float_list(
|
||||
self, number_of_floats, number_of_sections, section_number, interpolation_type
|
||||
):
|
||||
# Initializes the array with zeros
|
||||
my_floats = [0.0] * number_of_floats
|
||||
# Calculate the length of each portion based on total frames and number of images
|
||||
@@ -75,11 +87,15 @@ class NilorInterpolatedFloatList: # Generate interpolated float values based on
|
||||
|
||||
# Handling the first image (special case for the first segment)
|
||||
if section_number == 1:
|
||||
portion_values = self.interpolate_values(1, 0, portion_length, interpolation_type)
|
||||
portion_values = self.interpolate_values(
|
||||
1, 0, portion_length, interpolation_type
|
||||
)
|
||||
my_floats[0:portion_length] = portion_values
|
||||
# Handling the last image (special case for the last segment)
|
||||
elif section_number == number_of_sections:
|
||||
portion_values = self.interpolate_values(0, 1, portion_length, interpolation_type)
|
||||
portion_values = self.interpolate_values(
|
||||
0, 1, portion_length, interpolation_type
|
||||
)
|
||||
start_index = int((number_of_sections - 2) * portion_length)
|
||||
my_floats[start_index:] = portion_values
|
||||
# Handling middle images (general case for dual segments)
|
||||
@@ -96,6 +112,123 @@ class NilorInterpolatedFloatList: # Generate interpolated float values based on
|
||||
# Returns the modified list of float values
|
||||
return (my_floats,)
|
||||
|
||||
|
||||
class NilorOneMinusFloatList:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
# Dictionary that defines input types for each field
|
||||
return {
|
||||
"required": {
|
||||
"list_of_floats": ("FLOAT", {"input_is_list": True}),
|
||||
},
|
||||
}
|
||||
|
||||
# Define return types and names for outputs of the node
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
RETURN_NAMES = ("floats",)
|
||||
|
||||
FUNCTION = "one_minus_float_list"
|
||||
CATEGORY = category + subcategories["generators"]
|
||||
|
||||
def one_minus_float_list(self, list_of_floats):
|
||||
return ([1 - x for x in list_of_floats],)
|
||||
|
||||
|
||||
class NilorRemapFloatList:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
# Dictionary that defines input types for each field
|
||||
return {
|
||||
"required": {
|
||||
"list_of_floats": ("FLOAT", {"input_is_list": True}),
|
||||
"min_input": ("FLOAT", {"default": 0.0}),
|
||||
"max_input": ("FLOAT", {"default": 1.0}),
|
||||
"min_output": ("FLOAT", {"default": 0.0}),
|
||||
"max_output": ("FLOAT", {"default": 1.0}),
|
||||
},
|
||||
}
|
||||
|
||||
# Define return types and names for outputs of the node
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
RETURN_NAMES = ("remapped_floats",)
|
||||
|
||||
FUNCTION = "remap_float_list"
|
||||
CATEGORY = category + subcategories["generators"]
|
||||
|
||||
def remap_float_list(
|
||||
self, list_of_floats, min_input, max_input, min_output, max_output
|
||||
):
|
||||
# Avoid division by zero
|
||||
if max_input - min_input == 0:
|
||||
raise ValueError("max_input and min_input cannot be the same value.")
|
||||
|
||||
scale = (max_output - min_output) / (max_input - min_input)
|
||||
return ([min_output + (x - min_input) * scale for x in list_of_floats],)
|
||||
|
||||
|
||||
class NilorRemapFloatListAutoInput:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"list_of_floats": ("FLOAT", {"input_is_list": True}),
|
||||
"min_output": ("FLOAT", {"default": 0.0}),
|
||||
"max_output": ("FLOAT", {"default": 1.0}),
|
||||
},
|
||||
}
|
||||
|
||||
# Define return types and names for outputs of the node
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
RETURN_NAMES = ("remapped_list",)
|
||||
|
||||
FUNCTION = "remap_float_list_auto_input"
|
||||
CATEGORY = category + subcategories["generators"]
|
||||
|
||||
def remap_float_list_auto_input(self, list_of_floats, min_output, max_output):
|
||||
min_input = min(list_of_floats)
|
||||
max_input = max(list_of_floats)
|
||||
|
||||
scale = (max_output - min_output) / (max_input - min_input)
|
||||
return ([min_output + (x - min_input) * scale for x in list_of_floats],)
|
||||
|
||||
|
||||
class NilorInverseMapFloatList:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"list_of_floats": ("FLOAT", {"input_is_list": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
RETURN_NAMES = ("floats",)
|
||||
|
||||
FUNCTION = "inverse_map_float_list"
|
||||
CATEGORY = category + subcategories["generators"]
|
||||
|
||||
def inverse_map_float_list(self, list_of_floats):
|
||||
if not list_of_floats:
|
||||
raise ValueError("The input list_of_floats cannot be empty.")
|
||||
|
||||
min_input = min(list_of_floats)
|
||||
max_input = max(list_of_floats)
|
||||
|
||||
return ([min_input + max_input - x for x in list_of_floats],)
|
||||
|
||||
|
||||
class NilorIntToListOfBools:
|
||||
def __init__(self):
|
||||
pass
|
||||
@@ -128,6 +261,7 @@ class NilorIntToListOfBools:
|
||||
|
||||
return (my_bools,)
|
||||
|
||||
|
||||
class NilorListOfInts:
|
||||
def __init__(self):
|
||||
pass
|
||||
@@ -146,7 +280,9 @@ class NilorListOfInts:
|
||||
RETURN_NAMES = ("ints",)
|
||||
FUNCTION = "int_list"
|
||||
CATEGORY = category + subcategories["generators"]
|
||||
OUTPUT_IS_LIST = (True,) # Indicates that the output should be processed as a list of individual elements
|
||||
OUTPUT_IS_LIST = (
|
||||
True,
|
||||
) # Indicates that the output should be processed as a list of individual elements
|
||||
|
||||
def int_list(self, min=1, max=10, shuffle=False):
|
||||
# Generate the list
|
||||
@@ -156,6 +292,7 @@ class NilorListOfInts:
|
||||
|
||||
return (ints_list,)
|
||||
|
||||
|
||||
class NilorCountImagesInDirectory:
|
||||
def __init__(self):
|
||||
pass
|
||||
@@ -191,6 +328,7 @@ class NilorCountImagesInDirectory:
|
||||
|
||||
return [count]
|
||||
|
||||
|
||||
class NilorSelectIndexFromList:
|
||||
def __init__(self):
|
||||
pass
|
||||
@@ -199,7 +337,10 @@ class NilorSelectIndexFromList:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"list_of_any": (any, {"forceInput": False}), # Marking as lazy if processing could be deferred
|
||||
"list_of_any": (
|
||||
any,
|
||||
{"forceInput": False},
|
||||
), # Marking as lazy if processing could be deferred
|
||||
"index": ("INT", {"default": 0}),
|
||||
},
|
||||
}
|
||||
@@ -221,7 +362,7 @@ class NilorSelectIndexFromList:
|
||||
# Handle index access safely
|
||||
if isinstance(index, list):
|
||||
index = index[0]
|
||||
|
||||
|
||||
# Ensure the index is within bounds
|
||||
if index < 0 or index >= len(actual_list):
|
||||
raise ValueError("Index is outside the bounds of the array.")
|
||||
@@ -229,6 +370,7 @@ class NilorSelectIndexFromList:
|
||||
# Returns the value at the given index
|
||||
return (actual_list[index],)
|
||||
|
||||
|
||||
class NilorSaveEXRArbitrary:
|
||||
def __init__(self):
|
||||
self.output_dir = folder_paths.get_output_directory()
|
||||
@@ -238,7 +380,9 @@ class NilorSaveEXRArbitrary:
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"channels": (any,), # This should match the 'any' type list from List of Any
|
||||
"channels": (
|
||||
any,
|
||||
), # This should match the 'any' type list from List of Any
|
||||
"filename_prefix": ("STRING", {"default": "output"}),
|
||||
},
|
||||
"hidden": {
|
||||
@@ -250,12 +394,14 @@ class NilorSaveEXRArbitrary:
|
||||
RETURN_TYPES = ()
|
||||
|
||||
FUNCTION = "save_exr_arbitrary" # The execution function
|
||||
CATEGORY = category + subcategories["io"]
|
||||
|
||||
# INPUT_IS_LIST = True
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = category + subcategories["io"]
|
||||
|
||||
def save_exr_arbitrary(self, channels=None, filename_prefix="output", prompt=None, extra_pnginfo=None):
|
||||
# INPUT_IS_LIST = True
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def save_exr_arbitrary(
|
||||
self, channels=None, filename_prefix="output", prompt=None, extra_pnginfo=None
|
||||
):
|
||||
|
||||
print("Running save_exr_arbitrary")
|
||||
# print(f"channels: {channels}")
|
||||
@@ -275,10 +421,19 @@ class NilorSaveEXRArbitrary:
|
||||
# File path handling
|
||||
useabs = os.path.isabs(filename_prefix)
|
||||
if not useabs:
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix, self.output_dir, actual_channels[0].shape[-1], actual_channels[0].shape[-2])
|
||||
full_output_folder, filename, counter, subfolder, filename_prefix = (
|
||||
folder_paths.get_save_image_path(
|
||||
filename_prefix,
|
||||
self.output_dir,
|
||||
actual_channels[0].shape[-1],
|
||||
actual_channels[0].shape[-2],
|
||||
)
|
||||
)
|
||||
|
||||
# Determine if the input contains a batch
|
||||
is_batch = len(actual_channels[0].shape) == 3 # If batch, shape is [batch_size, height, width]
|
||||
is_batch = (
|
||||
len(actual_channels[0].shape) == 3
|
||||
) # If batch, shape is [batch_size, height, width]
|
||||
if is_batch:
|
||||
batch_size = actual_channels[0].shape[0]
|
||||
else:
|
||||
@@ -287,7 +442,9 @@ class NilorSaveEXRArbitrary:
|
||||
for i in range(batch_size):
|
||||
# Extract each image's channels
|
||||
if is_batch:
|
||||
image_channels = [tensor[i] for tensor in actual_channels] # For batch, select i-th image
|
||||
image_channels = [
|
||||
tensor[i] for tensor in actual_channels
|
||||
] # For batch, select i-th image
|
||||
else:
|
||||
image_channels = actual_channels # For single image, use channels as is
|
||||
|
||||
@@ -298,8 +455,10 @@ class NilorSaveEXRArbitrary:
|
||||
raise ValueError("All input tensors must have the same dimensions")
|
||||
|
||||
# Channel naming
|
||||
default_names = ["R", "G", "B", "A"] + [f"Channel{j}" for j in range(4, len(image_channels))]
|
||||
|
||||
default_names = ["R", "G", "B", "A"] + [
|
||||
f"Channel{j}" for j in range(4, len(image_channels))
|
||||
]
|
||||
|
||||
# Prepare data for EXR writing
|
||||
exr_data = {}
|
||||
for j, tensor in enumerate(image_channels):
|
||||
@@ -325,22 +484,29 @@ class NilorSaveEXRArbitrary:
|
||||
|
||||
# Create the EXR file header with dynamic channel names
|
||||
header = OpenEXR.Header(width, height)
|
||||
header['channels'] = {name: Imath.Channel(Imath.PixelType(Imath.PixelType.FLOAT)) for name in exr_data.keys()}
|
||||
header["channels"] = {
|
||||
name: Imath.Channel(Imath.PixelType(Imath.PixelType.FLOAT))
|
||||
for name in exr_data.keys()
|
||||
}
|
||||
|
||||
# Create the EXR file
|
||||
exr_file = OpenEXR.OutputFile(writepath, header)
|
||||
|
||||
# Prepare the data for each channel
|
||||
channel_data = {name: data.astype(np.float32).tobytes() for name, data in exr_data.items()}
|
||||
channel_data = {
|
||||
name: data.astype(np.float32).tobytes()
|
||||
for name, data in exr_data.items()
|
||||
}
|
||||
|
||||
# Write the channel data to the EXR file
|
||||
exr_file.writePixels(channel_data)
|
||||
exr_file.close()
|
||||
|
||||
|
||||
print(f"EXR file saved successfully to {writepath}")
|
||||
except Exception as e:
|
||||
print(f"Failed to write EXR file: {e}")
|
||||
|
||||
|
||||
class NilorSaveVideoToHFDataset:
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
@@ -379,6 +545,7 @@ class NilorSaveVideoToHFDataset:
|
||||
results.append(name)
|
||||
return {"ui": {"string_field": results}}
|
||||
|
||||
|
||||
class NilorSaveImageToHFDataset:
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
@@ -434,9 +601,595 @@ class NilorSaveImageToHFDataset:
|
||||
return {"ui": {"string_field": results}}
|
||||
|
||||
|
||||
class NilorShuffleImageBatch:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
|
||||
FUNCTION = "shuffle_image_batch"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
|
||||
def _check_image_dimensions(self, images):
|
||||
if images.shape[0] == 0:
|
||||
raise ValueError("Input images tensor is empty.")
|
||||
|
||||
# All images in the batch should have the same dimensions
|
||||
if len(images.shape) != 4:
|
||||
raise ValueError(
|
||||
f"Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
)
|
||||
|
||||
def shuffle_image_batch(self, images: torch.Tensor, seed):
|
||||
self._check_image_dimensions(images)
|
||||
|
||||
# Get the number of images in the batch
|
||||
num_images = images.shape[0]
|
||||
|
||||
# Generate indices and shuffle them
|
||||
torch.manual_seed(seed)
|
||||
indices = torch.randperm(num_images)
|
||||
|
||||
# Shuffle the images using the indices
|
||||
shuffled_images = images[indices]
|
||||
|
||||
return (shuffled_images,)
|
||||
|
||||
|
||||
class NilorRepeatTrimImageBatch:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"count": ("INT", {"default": 1, "min": 1, "max": BIGMAX, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
|
||||
FUNCTION = "repeat_trim_image_batch"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
|
||||
def _check_image_dimensions(self, images):
|
||||
if images.shape[0] == 0:
|
||||
raise ValueError("Input images tensor is empty.")
|
||||
|
||||
# All images in the batch should have the same dimensions
|
||||
if len(images.shape) != 4:
|
||||
raise ValueError(
|
||||
f"Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
)
|
||||
|
||||
def repeat_trim_image_batch(self, images: torch.Tensor, count):
|
||||
self._check_image_dimensions(images)
|
||||
|
||||
batch_count = images.size(0)
|
||||
amount = math.ceil(count / batch_count)
|
||||
|
||||
appended_tensors = (images.repeat(amount, 1, 1, 1),)
|
||||
batched_tensors = torch.cat(appended_tensors, dim=0)
|
||||
trimmed_tensors = batched_tensors[:count]
|
||||
|
||||
return (trimmed_tensors,)
|
||||
|
||||
|
||||
class NilorRepeatShuffleTrimImageBatch:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": BIGMAX, "step": 1}),
|
||||
"count": ("INT", {"default": 1, "min": 1, "max": BIGMAX, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
|
||||
FUNCTION = "repeat_shuffle_trim_image_batch"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
|
||||
def _check_image_dimensions(self, images):
|
||||
if images.shape[0] == 0:
|
||||
raise ValueError("Input images tensor is empty.")
|
||||
|
||||
# All images in the batch should have the same dimensions
|
||||
if len(images.shape) != 4:
|
||||
raise ValueError(
|
||||
f"Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
)
|
||||
|
||||
def repeat_shuffle_trim_image_batch(self, images: torch.Tensor, seed, count):
|
||||
self._check_image_dimensions(images)
|
||||
|
||||
torch.manual_seed(seed)
|
||||
|
||||
batch_count = images.size(0)
|
||||
amount = math.ceil(count / batch_count)
|
||||
|
||||
appended_tensors = []
|
||||
while len(appended_tensors) < count:
|
||||
indices = torch.randperm(batch_count)
|
||||
appended_tensors.append(images[indices])
|
||||
|
||||
batched_tensors = torch.cat(appended_tensors, dim=0)
|
||||
trimmed_tensors = batched_tensors[:count]
|
||||
|
||||
return (trimmed_tensors,)
|
||||
|
||||
|
||||
class NilorOutputFilenameString:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"client": ("STRING", {"default": "nilor"}),
|
||||
"project": ("STRING", {"default": "research"}),
|
||||
"section": ("STRING", {"default": "test-1"}),
|
||||
"name": ("STRING", {"default": "out-1"}),
|
||||
},
|
||||
"hidden": {
|
||||
"unique_id": "UNIQUE_ID",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO",
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("string",)
|
||||
FUNCTION = "notify"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
OUTPUT_NODE = True
|
||||
IS_CHANGED = True
|
||||
|
||||
def get_time(self, format: str):
|
||||
now = datetime.now()
|
||||
return now.strftime(format)
|
||||
|
||||
def notify(
|
||||
self, client, project, section, name, unique_id=None, extra_pnginfo=None
|
||||
):
|
||||
time = self.get_time("%y%m%d-%H%M%S")
|
||||
|
||||
client = client or "nilor"
|
||||
project = project or "research"
|
||||
section = section or "test-1"
|
||||
name = name or "out-1"
|
||||
|
||||
text = f"{client}_{project}/{section}/{time}_{section}/{time}_{client}_{project}_{section}_{name}"
|
||||
|
||||
if unique_id is not None and extra_pnginfo is not None:
|
||||
if not isinstance(extra_pnginfo, list):
|
||||
print("Error: extra_pnginfo is not a list")
|
||||
elif (
|
||||
not isinstance(extra_pnginfo[0], dict)
|
||||
or "workflow" not in extra_pnginfo[0]
|
||||
):
|
||||
print("Error: extra_pnginfo[0] is not a dict or missing 'workflow' key")
|
||||
else:
|
||||
workflow = extra_pnginfo[0]["workflow"]
|
||||
node = next(
|
||||
(x for x in workflow["nodes"] if str(x["id"]) == str(unique_id[0])),
|
||||
None,
|
||||
)
|
||||
if node:
|
||||
node["widgets_values"] = [text]
|
||||
|
||||
# TODO: make this node's text string preview widget work
|
||||
return {"ui": {"text": text}, "result": (text,)}
|
||||
|
||||
|
||||
class NilorNFractionsOfInt:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"numerator": ("INT", {"default": 10}),
|
||||
"denominator": ("INT", {"default": 2}),
|
||||
"type": (["starts", "ends", "centres", "start + end"], {}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
RETURN_NAMES = ("fractions",)
|
||||
|
||||
FUNCTION = "n_fractions_of_int"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
|
||||
def n_fractions_of_int(self, numerator, denominator, type):
|
||||
# the number of fractions to generate is the denominator
|
||||
if type == "starts":
|
||||
return ([i * numerator // denominator for i in range(denominator)],)
|
||||
elif type == "ends":
|
||||
return ([(i + 1) * numerator // denominator for i in range(denominator)],)
|
||||
elif type == "centres":
|
||||
return (
|
||||
[
|
||||
(i * numerator + numerator // 2) // denominator
|
||||
for i in range(denominator)
|
||||
],
|
||||
)
|
||||
elif type == "start + end":
|
||||
return ([i * numerator // (denominator - 1) for i in range(denominator)],)
|
||||
else:
|
||||
raise ValueError(f"Unknown type: {type}")
|
||||
|
||||
|
||||
class NilorCategorizeString:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_string": ("STRING", {"default": ""}),
|
||||
"number_of_categories": ("INT", {"default": 2, "min": 1, "max": 10}),
|
||||
"category_0": ("STRING", {"default": "apple, red fruit"}),
|
||||
"category_1": ("STRING", {"default": "banana, yellow fruit"}),
|
||||
},
|
||||
"optional": {
|
||||
"category_2": ("STRING", {"default": ""}),
|
||||
"category_3": ("STRING", {"default": ""}),
|
||||
"category_4": ("STRING", {"default": ""}),
|
||||
"category_5": ("STRING", {"default": ""}),
|
||||
"category_6": ("STRING", {"default": ""}),
|
||||
"category_7": ("STRING", {"default": ""}),
|
||||
"category_8": ("STRING", {"default": ""}),
|
||||
"category_9": ("STRING", {"default": ""}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
RETURN_NAMES = ("category_index",)
|
||||
FUNCTION = "categorize_string"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
|
||||
def categorize_string(
|
||||
self,
|
||||
input_string,
|
||||
number_of_categories,
|
||||
category_0,
|
||||
category_1,
|
||||
category_2="",
|
||||
category_3="",
|
||||
category_4="",
|
||||
category_5="",
|
||||
category_6="",
|
||||
category_7="",
|
||||
category_8="",
|
||||
category_9="",
|
||||
):
|
||||
# Convert input string to lowercase for case-insensitive matching
|
||||
input_string = input_string.lower()
|
||||
|
||||
# Create categories dictionary from inputs
|
||||
categories = {}
|
||||
all_categories = [
|
||||
category_0,
|
||||
category_1,
|
||||
category_2,
|
||||
category_3,
|
||||
category_4,
|
||||
category_5,
|
||||
category_6,
|
||||
category_7,
|
||||
category_8,
|
||||
category_9,
|
||||
]
|
||||
|
||||
# Only process the number of categories specified
|
||||
for i in range(number_of_categories):
|
||||
if all_categories[i]: # Only add non-empty categories
|
||||
# Split the comma-separated string and clean up whitespace
|
||||
keywords = [k.strip().lower() for k in all_categories[i].split(",")]
|
||||
categories[i] = keywords
|
||||
|
||||
# Check each category's keywords against the input string
|
||||
for index, keywords in categories.items():
|
||||
if builtins.any(keyword in input_string for keyword in keywords):
|
||||
return (index,)
|
||||
|
||||
return (-1,) # Default case if no matches found
|
||||
|
||||
|
||||
class NilorRandomString:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"multiline_text": (
|
||||
"STRING",
|
||||
{"default": "option1, option2, option3", "multiline": True},
|
||||
),
|
||||
"max_options": ("INT", {"default": 3, "min": 1}),
|
||||
"delimiter": ("STRING", {"default": ","}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("chosen_string",)
|
||||
FUNCTION = "choose_random_string"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
|
||||
def choose_random_string(self, multiline_text, max_options, delimiter, seed):
|
||||
import random
|
||||
|
||||
random.seed(seed)
|
||||
|
||||
# If the delimiter is literally "\n", use the actual newline character.
|
||||
if delimiter == r"\n" or delimiter == "\\n":
|
||||
actual_delimiter = "\n"
|
||||
else:
|
||||
actual_delimiter = delimiter
|
||||
|
||||
# Split the input text using the actual delimiter and remove any extra whitespace
|
||||
options = [
|
||||
item.strip()
|
||||
for item in multiline_text.split(actual_delimiter)
|
||||
if item.strip()
|
||||
]
|
||||
if not options:
|
||||
raise ValueError("No valid choices provided.")
|
||||
|
||||
# Limit to the first 'max_options' entries if there are more options
|
||||
if len(options) > max_options:
|
||||
options = options[:max_options]
|
||||
|
||||
chosen = random.choice(options)
|
||||
return (chosen,)
|
||||
|
||||
|
||||
class NilorLoadImageByIndex:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image_directory": (
|
||||
"STRING",
|
||||
{"default": "", "placeholder": "Image Directory"},
|
||||
),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
"sort_mode": (
|
||||
["filename", "creation_time", "modification_time", "size"],
|
||||
{"default": "filename"},
|
||||
),
|
||||
"reverse_sort": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "STRING", "STRING")
|
||||
RETURN_NAMES = ("image", "filename", "filepath")
|
||||
FUNCTION = "load_image_by_index"
|
||||
CATEGORY = category + subcategories["io"]
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, image_directory, seed, sort_mode, reverse_sort):
|
||||
return seed
|
||||
|
||||
def load_image_by_index(self, image_directory, seed, sort_mode, reverse_sort):
|
||||
if not os.path.exists(image_directory):
|
||||
raise FileNotFoundError(f"Image directory {image_directory} does not exist")
|
||||
|
||||
# Get list of image files
|
||||
files = []
|
||||
for f in os.listdir(image_directory):
|
||||
file_path = os.path.join(image_directory, f)
|
||||
if os.path.isfile(file_path) and f.lower().endswith(
|
||||
(".png", ".jpg", ".jpeg", ".webp", ".bmp", ".gif")
|
||||
):
|
||||
files.append(file_path)
|
||||
|
||||
if not files:
|
||||
raise ValueError(f"No image files found in {image_directory}")
|
||||
|
||||
# Sort files based on selected mode
|
||||
if sort_mode == "filename":
|
||||
files.sort()
|
||||
elif sort_mode == "creation_time":
|
||||
files.sort(key=lambda x: os.path.getctime(x))
|
||||
elif sort_mode == "modification_time":
|
||||
files.sort(key=lambda x: os.path.getmtime(x))
|
||||
elif sort_mode == "size":
|
||||
files.sort(key=lambda x: os.path.getsize(x))
|
||||
|
||||
# Apply reverse sort if requested
|
||||
if reverse_sort:
|
||||
files.reverse()
|
||||
|
||||
# Get file at index (with wrapping)
|
||||
file_index = seed % len(files)
|
||||
selected_file = files[file_index]
|
||||
|
||||
# Get filename
|
||||
filename = os.path.basename(selected_file)
|
||||
|
||||
# Load image using PIL and convert to tensor using our helper function
|
||||
img = Image.open(selected_file)
|
||||
img_tensor = pil2tensor(img)
|
||||
|
||||
return (img_tensor, filename, selected_file)
|
||||
|
||||
|
||||
class NilorExtractFilenameFromPath:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"filepath": ("STRING", {"default": ""}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING", "STRING")
|
||||
RETURN_NAMES = ("name", "name_with_extension")
|
||||
FUNCTION = "extract_filename"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
|
||||
def extract_filename(self, filepath):
|
||||
# Ensure the input is a valid path
|
||||
if not filepath:
|
||||
raise ValueError("Filepath cannot be empty.")
|
||||
|
||||
path = Path(filepath)
|
||||
|
||||
# Extract filename with and without extension
|
||||
name = path.stem # Filename without extension
|
||||
name_with_extension = path.name # Filename with extension
|
||||
|
||||
return (name, name_with_extension)
|
||||
|
||||
|
||||
class NilorBlurAnalysis:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",), # Input image batch as a 4D tensor.
|
||||
"block_size": ("INT", {"default": 32, "min": 1, "max": 128, "step": 1}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("blur_analysis",)
|
||||
FUNCTION = "analyze_blur"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
|
||||
def analyze_blur(self, images, block_size):
|
||||
"""
|
||||
Performs blur analysis on each image using OpenCV's Laplacian method.
|
||||
"""
|
||||
# Ensure images is a 4D tensor.
|
||||
if images.dim() != 4:
|
||||
raise ValueError("Input images must be a 4D tensor (batch, channels/height, height/width, width/channels)")
|
||||
|
||||
# Detect if using NCHW or NHWC.
|
||||
if images.shape[1] not in (1, 3):
|
||||
if images.shape[-1] in (1, 3):
|
||||
images = images.permute(0, 3, 1, 2)
|
||||
else:
|
||||
raise ValueError("Cannot determine image format (expected channel to be 1 or 3).")
|
||||
|
||||
output_images = []
|
||||
batch_size = images.shape[0]
|
||||
for i in range(batch_size):
|
||||
# Get the i-th image (in NCHW: [channels, height, width]).
|
||||
img_tensor = images[i].cpu()
|
||||
img_np = img_tensor.numpy() # shape: (C, H, W)
|
||||
|
||||
# Convert to grayscale.
|
||||
if img_np.shape[0] >= 3:
|
||||
gray = (0.299 * img_np[0] +
|
||||
0.587 * img_np[1] +
|
||||
0.114 * img_np[2])
|
||||
else:
|
||||
gray = np.squeeze(img_np, axis=0) # shape: (H, W)
|
||||
|
||||
# Scale from [0, 1] to [0, 255] and convert to uint8.
|
||||
gray = np.clip(gray * 255.0, 0, 255).astype(np.uint8)
|
||||
|
||||
# Compute Laplacian using a 3x3 kernel.
|
||||
lap = cv2.Laplacian(gray, cv2.CV_64F, ksize=3)
|
||||
abs_lap = np.absolute(lap)
|
||||
|
||||
# Apply local averaging using cv2.blur with window size = (block_size, block_size).
|
||||
local_edge = cv2.blur(abs_lap, (block_size, block_size))
|
||||
|
||||
# Normalize and invert the edge response.
|
||||
max_val = local_edge.max()
|
||||
if max_val > 0:
|
||||
norm_edge = local_edge / max_val
|
||||
else:
|
||||
norm_edge = local_edge
|
||||
blur_map = 1.0 - norm_edge
|
||||
|
||||
# Scale back to 0-255 and convert to uint8.
|
||||
out_img = (blur_map * 255.0).astype(np.uint8)
|
||||
|
||||
# Convert the single channel output to a 3-channel image.
|
||||
# This ensures downstream nodes (like MaskFromRGBCMYBW) that index into channels work properly.
|
||||
if out_img.ndim == 2:
|
||||
out_img = np.stack([out_img, out_img, out_img], axis=-1) # shape becomes (H, W, 3)
|
||||
|
||||
# Convert from PIL image (or numpy array) to tensor.
|
||||
# pil2tensor should create a tensor in a format that downstream nodes expect.
|
||||
output_images.append(pil2tensor(out_img))
|
||||
|
||||
# ---
|
||||
# Fix 2: Use torch.stack to preserve the batch dimension.
|
||||
# If each output has shape, say, (H, W, 3), stacking them gives a tensor of shape (B, H, W, 3).
|
||||
return (torch.cat(output_images, dim=0),)
|
||||
|
||||
class NilorToSparseIndexMethod:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"ints": ("INT", {"default": 0}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("sparse_method",)
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
INPUT_IS_LIST = True
|
||||
|
||||
FUNCTION = "convert_to_sparse_index_method"
|
||||
CATEGORY = category + subcategories["utilities"]
|
||||
|
||||
def convert_to_sparse_index_method(self, ints):
|
||||
indexes_str = ",".join(map(str, ints))
|
||||
|
||||
return (indexes_str,)
|
||||
|
||||
|
||||
# Mapping class names to objects for potential export
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Nilor Interpolated Float List": NilorInterpolatedFloatList,
|
||||
"Nilor One Minus Float List": NilorOneMinusFloatList,
|
||||
"Nilor Remap Float List": NilorRemapFloatList,
|
||||
"Nilor Remap Float List Auto Input": NilorRemapFloatListAutoInput,
|
||||
"Nilor Inverse Map Float List": NilorInverseMapFloatList,
|
||||
"Nilor Int To List Of Bools": NilorIntToListOfBools,
|
||||
"Nilor List of Ints": NilorListOfInts,
|
||||
"Nilor Count Images In Directory": NilorCountImagesInDirectory,
|
||||
@@ -444,11 +1197,26 @@ NODE_CLASS_MAPPINGS = {
|
||||
"Nilor Save Video To HF Dataset": NilorSaveVideoToHFDataset,
|
||||
"Nilor Select Index From List": NilorSelectIndexFromList,
|
||||
"Nilor Save EXR Arbitrary": NilorSaveEXRArbitrary,
|
||||
"Nilor Shuffle Image Batch": NilorShuffleImageBatch,
|
||||
"Nilor Repeat & Trim Image Batch": NilorRepeatTrimImageBatch,
|
||||
"Nilor Repeat, Shuffle, & Trim Image Batch": NilorRepeatShuffleTrimImageBatch,
|
||||
"Nilor Output Filename String": NilorOutputFilenameString,
|
||||
"Nilor n Fractions of Int": NilorNFractionsOfInt,
|
||||
"Nilor Categorize String": NilorCategorizeString,
|
||||
"Nilor Random String": NilorRandomString,
|
||||
"Nilor Extract Filename from Path": NilorExtractFilenameFromPath,
|
||||
"Nilor Load Image By Index": NilorLoadImageByIndex,
|
||||
"Nilor Blur Analysis": NilorBlurAnalysis,
|
||||
"Nilor To Sparse Index Method": NilorToSparseIndexMethod,
|
||||
}
|
||||
|
||||
# Mapping nodes to human-readable names
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Nilor Interpolated Float List": "👺 Interpolated Float List",
|
||||
"Nilor One Minus Float List": "👺 One Minus Float List",
|
||||
"Nilor Remap Float List": "👺 Nilor Remap Float List",
|
||||
"Nilor Remap Float List Auto Input": "👺 Nilor Remap Float List Auto Input",
|
||||
"Nilor Inverse Map Float List": "👺 Nilor Inverse Map Float List",
|
||||
"Nilor Int To List Of Bools": "👺 Int To List Of Bools",
|
||||
"Nilor List of Ints": "👺 List of Ints",
|
||||
"Nilor Count Images In Directory": "👺 Count Images In Directory",
|
||||
@@ -456,4 +1224,15 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"Nilor Save Video To HF Dataset": "👺 Save Video To HF Dataset",
|
||||
"Nilor Select Index From List": "👺 Select Index From List",
|
||||
"Nilor Save EXR Arbitrary": "👺 Save EXR Arbitrary",
|
||||
"Nilor Shuffle Image Batch": "👺 Nilor Shuffle Image Batch",
|
||||
"Nilor Repeat & Trim Image Batch": "👺 Nilor Repeat & Trim Image Batch",
|
||||
"Nilor Repeat, Shuffle, & Trim Image Batch": "👺 Nilor Repeat, Shuffle, & Trim Image Batch",
|
||||
"Nilor Output Filename String": "👺 Nilor Output Filename String",
|
||||
"Nilor n Fractions of Int": "👺 Nilor n Fractions of Int",
|
||||
"Nilor Categorize String": "👺 Categorize String",
|
||||
"Nilor Random String": "👺 Random String",
|
||||
"Nilor Extract Filename from Path": "👺 Extract Filename from Path",
|
||||
"Nilor Load Image By Index": "👺 Load Image By Index",
|
||||
"Nilor Blur Analysis": "👺 Blur Analysis",
|
||||
"Nilor To Sparse Index Method": "👺 To Sparse Index Method",
|
||||
}
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
[project]
|
||||
name = "nilor-nodes"
|
||||
description = "Custom utility nodes for ComfyUI by Nilor Corp. Probably not useful for most people, but contains stuff for working with lists, filenames, image batches, etc in a very specifc way."
|
||||
version = "1.0.1"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["huggingface_hub", "openexr"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/nilor-corp/nilor-nodes"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "stephennilor"
|
||||
DisplayName = "nilor-nodes"
|
||||
Icon = ""
|
||||
+2
-1
@@ -1 +1,2 @@
|
||||
huggingface_hub
|
||||
huggingface_hub
|
||||
openexr
|
||||
@@ -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,37 @@
|
||||
import io
|
||||
import torch
|
||||
import base64
|
||||
import numpy as np
|
||||
from pkg_resources import parse_version
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def pil2numpy(image: Image.Image):
|
||||
return np.array(image).astype(np.float32) / 255.0
|
||||
|
||||
|
||||
def numpy2pil(image: np.ndarray, mode=None):
|
||||
return Image.fromarray(np.clip(255.0 * image, 0, 255).astype(np.uint8), mode)
|
||||
|
||||
|
||||
## Helper function equivalent to Mikey's pil2tensor
|
||||
#def pil2tensor(self, image):
|
||||
# return torch.from_numpy(np.array(image).astype(np.float32) / 255.0).unsqueeze(0)
|
||||
|
||||
def pil2tensor(image: Image.Image):
|
||||
return torch.from_numpy(pil2numpy(image)).unsqueeze(0)
|
||||
|
||||
|
||||
def tensor2pil(image: torch.Tensor, mode=None):
|
||||
return numpy2pil(image.cpu().numpy().squeeze(), mode=mode)
|
||||
|
||||
|
||||
def tensor2bytes(image: torch.Tensor) -> bytes:
|
||||
return tensor2pil(image).tobytes()
|
||||
|
||||
|
||||
def pil2base64(image: Image.Image):
|
||||
buffered = io.BytesIO()
|
||||
image.save(buffered, format="PNG")
|
||||
img_str = base64.b64encode(buffered.getvalue()).decode("utf-8")
|
||||
return img_str
|
||||
@@ -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