Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8e811b11bd | ||
|
|
5ed354b3a1 | ||
|
|
62b71b2978 | ||
|
|
d266978e49 | ||
|
|
552252b84c | ||
|
|
66c66f045f | ||
|
|
266fe598c6 | ||
|
|
ef20c72aff | ||
|
|
527e8dc4ab | ||
|
|
5d05aeaf68 | ||
|
|
5668e09a79 | ||
|
|
a94c1a9d60 | ||
|
|
675fce6a16 | ||
|
|
c1f5eb352c | ||
|
|
0f22deb027 | ||
|
|
58ff7a0bce | ||
|
|
c33b06a092 | ||
|
|
d5109ae0b3 | ||
|
|
63b44ebff4 | ||
|
|
3476fbfbf9 | ||
|
|
aafaf87d18 | ||
|
|
790a3ce38e | ||
|
|
114a3d661f | ||
|
|
b6bf6c1b6c | ||
|
|
3c0bf2222c | ||
|
|
ff0c7f1209 | ||
|
|
75a632b2e5 | ||
|
|
2f707ebc48 | ||
|
|
14862c1712 | ||
|
|
7320a15d08 | ||
|
|
4a56d4dff5 | ||
|
|
5e579b1aed | ||
|
|
9883728f55 | ||
|
|
bfba7b0265 | ||
|
|
a093a51291 | ||
|
|
848899298c | ||
|
|
2333a0d45d | ||
|
|
87b846069e | ||
|
|
6d17e664bb | ||
|
|
c541bc18dd | ||
|
|
7beda857c8 | ||
|
|
4fa9710aea | ||
|
|
e8c0d451d9 | ||
|
|
a491e6463c | ||
|
|
5b8184cc33 | ||
|
|
6b2b482b89 | ||
|
|
ea13b61f21 | ||
|
|
663d2dd739 | ||
|
|
b940ad79ea | ||
|
|
baba20c3c4 | ||
|
|
f3dbe8a3ed | ||
|
|
cb9a98edf1 | ||
|
|
e8ee057dbe | ||
|
|
5f7d00560b | ||
|
|
90e812b0c3 | ||
|
|
51777e3462 | ||
|
|
7806a92772 | ||
|
|
a42221565d | ||
|
|
75a45c6c34 | ||
|
|
414679c676 | ||
|
|
5af42118fe | ||
|
|
7bc6116a12 | ||
|
|
08e2d39d76 | ||
|
|
803fca81f7 | ||
|
|
b44a66bf03 | ||
|
|
35428a8287 | ||
|
|
41372155c1 | ||
|
|
04ac2b655a | ||
|
|
cc9056e11c | ||
|
|
19c693168b | ||
|
|
66d6318b9a | ||
|
|
a94814e632 | ||
|
|
c2edf563f9 | ||
|
|
ccb2b14dfa | ||
|
|
d2ae46d3c2 | ||
|
|
5f723ae82d | ||
|
|
feedc4a002 | ||
|
|
0d6ea9f00c | ||
|
|
55d83abed0 | ||
|
|
85e28d01d5 | ||
|
|
b091d3057b | ||
|
|
f731a55292 | ||
|
|
5e2aa12434 | ||
|
|
a76949e628 | ||
|
|
afc68ccded | ||
|
|
778ed5272f | ||
|
|
e5b165605a | ||
|
|
2387d0e0ad | ||
|
|
67c5159cfc | ||
|
|
dd9e13e148 | ||
|
|
9e5ccb8bb7 | ||
|
|
95aaea78c6 | ||
|
|
8f355e360c | ||
|
|
e8c5c449ac | ||
|
|
31ae8624d3 | ||
|
|
2c7edfc535 | ||
|
|
c2caab535a | ||
|
|
97cae869d4 | ||
|
|
3171063c50 | ||
|
|
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,58 @@
|
||||
# NILOR_LOG_LEVEL=INFO # possible log levels: INFO, DEBUG, WARNING, ERROR, CRITICAL
|
||||
|
||||
# --- Comfy API client ---
|
||||
# NILOR_COMFYUI_API_URL=http://127.0.0.1:8188
|
||||
# NILOR_COMFYUI_WS_URL=ws://127.0.0.1:8188
|
||||
# NILOR_COMFY_API_TIMEOUT_SECONDS=30
|
||||
|
||||
## HTTP idempotent retry policy (GET /system_stats, POST /free only)
|
||||
# NILOR_COMFY_RETRY_BASE_SECONDS=0.25
|
||||
# NILOR_COMFY_RETRY_MULTIPLIER=2.0
|
||||
# NILOR_COMFY_RETRY_JITTER_SECONDS=0.25
|
||||
# NILOR_COMFY_RETRY_MAX_SLEEP_SECONDS=4.0
|
||||
# NILOR_COMFY_RETRY_MAX_ATTEMPTS=3
|
||||
|
||||
## WebSocket reconnect policy
|
||||
# NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS=5
|
||||
# NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS=30.0
|
||||
|
||||
# --- Worker / Queue settings ---
|
||||
# NILOR_SQS_ENABLED=true
|
||||
NILOR_SQS_ENDPOINT_URL=http://127.0.0.1:9324
|
||||
# NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME=jobs_to_process-comfyui
|
||||
# NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME=job_status_updates
|
||||
# NILOR_SQS_POLL_WAIT_TIME=10
|
||||
# NILOR_SQS_MAX_MESSAGES=1
|
||||
|
||||
## Enable/disable workflow normalization based on operating system
|
||||
# NILOR_WORKFLOW_OS_NORMALIZATION_ENABLED=false
|
||||
|
||||
## Non-secret local defaults; provide real values locally
|
||||
# NILOR_AWS_ACCESS_KEY_ID=minioadmin
|
||||
NILOR_AWS_SECRET_ACCESS_KEY=...
|
||||
# NILOR_AWS_DEFAULT_REGION=us-east-1
|
||||
|
||||
## If empty or unset, a stable id will be generated by the loader
|
||||
# NILOR_WORKER_CLIENT_ID=
|
||||
|
||||
# --- Memory Hygiene (Guardian) ---
|
||||
## Enable/disable hygiene between jobs
|
||||
# NILOR_MEMORY_HYGIENE_ENABLED=true
|
||||
|
||||
## How often to check when idle (seconds)
|
||||
# NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS=5
|
||||
|
||||
## Thresholds (either percent usage or absolute free MB can trigger)
|
||||
# NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX=88
|
||||
# NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX=90
|
||||
# NILOR_MEMORY_HYGIENE_VRAM_MIN_FREE_MB=2048
|
||||
# NILOR_MEMORY_HYGIENE_RAM_MIN_FREE_MB=4096
|
||||
|
||||
## Action policy: free|unload|both|auto (auto: free first, escalate to unload once if needed)
|
||||
# NILOR_MEMORY_HYGIENE_ACTION_POLICY=auto
|
||||
|
||||
## Retry/backoff/cycle caps
|
||||
# NILOR_MEMORY_HYGIENE_MAX_RETRIES=2
|
||||
# NILOR_MEMORY_HYGIENE_COOLDOWN_SECONDS=60
|
||||
# NILOR_MEMORY_HYGIENE_SLEEP_BETWEEN_ATTEMPTS_SECONDS=4
|
||||
# NILOR_MEMORY_HYGIENE_MAX_CYCLE_DURATION_SECONDS=15
|
||||
@@ -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,256 @@
|
||||
# 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.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- `comfyui-kjnodes` custom_nodes repo
|
||||
|
||||
## 🏭 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>
|
||||
|
||||
## 📡 Core Nilor Services
|
||||
|
||||
<details>
|
||||
<summary><b>Worker Consumer Service</b></summary>
|
||||
|
||||
The `worker_consumer.py` script is a background service that runs on each ComfyUI worker. It is responsible for pulling jobs from the central ElasticMQ `jobs_to_process` queue and submitting them to its local ComfyUI instance for processing. This service is essential for the distributed architecture of the system.
|
||||
|
||||
**Key Responsibilities:**
|
||||
- Continuously polls the `jobs_to_process` queue for new jobs using long polling.
|
||||
- When a job is received, it extracts the workflow data and submits it to the local ComfyUI server.
|
||||
- Normalizes OS-sensitive path formatting in the ComfyUI `prompt` graph (e.g. converts `Flux1/ae.safetensors` ↔ `Flux1\ae.safetensors`) so workflows authored on Linux/Windows run on the current worker OS.
|
||||
- Deletes the job message from the queue upon successful submission to prevent reprocessing.
|
||||
- If submission fails, the message remains on the queue to be picked up by another worker.
|
||||
|
||||
**Workflow normalization toggle:**
|
||||
|
||||
- Set `NILOR_WORKFLOW_OS_NORMALIZATION_ENABLED=false` (via environment or `config/config.json5`) to disable this behavior (default: enabled).
|
||||
|
||||
</details>
|
||||
|
||||
## ⚙️ Configuration and Runtime Model
|
||||
|
||||
The sidecar uses a small, typed configuration loader with JSON5 defaults and optional environment overrides.
|
||||
|
||||
- **Precedence**: environment variables > `config/config.json5` (controlled by `allow_env_override: true`).
|
||||
- **No hot‑reload**: configuration is loaded once at process start and passed to components.
|
||||
- **Paths/keys**: JSON5 at `ComfyUI/custom_nodes/nilor-nodes/config/config.json5` with `NILOR_*` keys (e.g., `NILOR_COMFYUI_API_URL`, `NILOR_SQS_ENDPOINT_URL`). Secrets (AWS secret) must be set via `.env`.
|
||||
- **Typed object**: loader returns a `NilorNodesConfig` with `comfy`, `worker`, and `hygiene` sections.
|
||||
|
||||
Pseudocode usage:
|
||||
|
||||
```pseudo
|
||||
cfg = load_nilor_nodes_config()
|
||||
# Comfy endpoints
|
||||
http_url = cfg.comfy.api_url + "/prompt"
|
||||
ws_url = cfg.comfy.ws_url + "/ws"
|
||||
# SQS client params
|
||||
endpoint = cfg.worker.sqs_endpoint_url
|
||||
region = cfg.worker.aws_region
|
||||
access_key = cfg.worker.aws_access_key_id
|
||||
secret_key = cfg.worker.aws_secret_access_key
|
||||
client_id = cfg.worker.worker_client_id
|
||||
```
|
||||
|
||||
Current integrations:
|
||||
|
||||
- `worker_consumer.py`: loads config at startup, reuses a single HTTP session, and uses `cfg.comfy`/`cfg.worker` exclusively.
|
||||
- `media_stream.py`: uses `cfg.worker` for SQS completion notifications.
|
||||
|
||||
<details>
|
||||
<summary><b>Environment Variables</b></summary>
|
||||
|
||||
The `nilor-nodes` require a `.env` file to be present in the `ComfyUI` directory to configure the connection to the core services (MinIO, ElasticMQ, and the Brain API). To set it up, create a file named `.env` in the root of your `ComfyUI` directory by copying the `.env.example` template.
|
||||
|
||||
**Instructions:**
|
||||
1. Create a new file named `.env` in the `ComfyUI` directory.
|
||||
2. Copy the contents of the `.env.example` file into your new `.env` file.
|
||||
3. Replace the placeholder values with your actual credentials and endpoint URLs for your local or production environment.
|
||||
|
||||
</details>
|
||||
+87
-1
@@ -1,3 +1,89 @@
|
||||
from .nilornodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
import os
|
||||
import threading
|
||||
import asyncio
|
||||
import logging
|
||||
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)
|
||||
|
||||
|
||||
# --- Package Logger (applied early for all nilor-nodes modules) ---
|
||||
from .logger import configure_from_env, logger
|
||||
|
||||
configure_from_env()
|
||||
|
||||
|
||||
# --- 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())
|
||||
|
||||
|
||||
from .config.config import load_nilor_nodes_config
|
||||
|
||||
cfg = load_nilor_nodes_config()
|
||||
|
||||
# Start the SQS Worker Consumer (controlled by NILOR_SQS_ENABLED)
|
||||
if cfg.sqs_enabled:
|
||||
consumer_thread = threading.Thread(target=start_consumer_loop, daemon=True)
|
||||
consumer_thread.start()
|
||||
print(
|
||||
"✅ Nilor-Nodes: SQS worker consumer thread started (NILOR_SQS_ENABLED=true)."
|
||||
)
|
||||
else:
|
||||
print(
|
||||
"⚠️ Nilor-Nodes: SQS worker consumer functionality is disabled (NILOR_SQS_ENABLED=false)."
|
||||
)
|
||||
|
||||
|
||||
# --- 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,910 @@
|
||||
"""
|
||||
Thin, typed client surface for accessing ComfyUI HTTP endpoints and the websocket.
|
||||
|
||||
This module defines the public protocol, DTOs, and exceptions that callers and
|
||||
tests depend on. Implementations are intentionally minimal at this stage; network
|
||||
behavior will be added in subsequent commits.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncIterator, Dict, Optional, Protocol, TypedDict, Tuple
|
||||
|
||||
from urllib.parse import quote, urlparse
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ComfyUIClientProtocol",
|
||||
"ComfyUILocalClient",
|
||||
"SystemStats",
|
||||
"WsEvent",
|
||||
"ComfyUIClientError",
|
||||
"ComfyUIClientTimeout",
|
||||
"ComfyUIClientWsClosed",
|
||||
]
|
||||
|
||||
|
||||
class WsEvent(TypedDict, total=False):
|
||||
"""Typed view of websocket events emitted by ComfyUI.
|
||||
|
||||
Fields:
|
||||
- type: Event type string (e.g., "status", "progress", "executed").
|
||||
- data: Opaque payload; commonly includes keys like "prompt_id", "node", etc.
|
||||
"""
|
||||
|
||||
type: str
|
||||
data: Dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class SystemStats:
|
||||
"""Subset of system statistics reported by ComfyUI `/system_stats`.
|
||||
|
||||
Known fields are optional; unknown fields should be ignored by parsers. When
|
||||
present, RAM-related metrics may also appear (e.g., `ram_total`, `ram_free`).
|
||||
|
||||
Attributes:
|
||||
vram_total: Total VRAM (bytes) when reported.
|
||||
vram_free: Free VRAM reported by the backend, if available.
|
||||
torch_vram_free: Free VRAM according to torch, if available.
|
||||
ram_total: Total system RAM (bytes) when reported.
|
||||
ram_free: Free system RAM (bytes) when reported.
|
||||
"""
|
||||
|
||||
vram_total: Optional[float] = None
|
||||
vram_free: Optional[float] = None
|
||||
torch_vram_free: Optional[float] = None
|
||||
ram_total: Optional[float] = None
|
||||
ram_free: Optional[float] = None
|
||||
|
||||
|
||||
class ComfyUIClientError(Exception):
|
||||
"""Base error for all ComfyUI client failures.
|
||||
|
||||
Args:
|
||||
message: Human-friendly error message.
|
||||
route: Route path (e.g., "/prompt").
|
||||
method: HTTP method (e.g., "GET", "POST").
|
||||
status: Optional HTTP status code or websocket close code.
|
||||
code: Optional machine-readable error code (e.g., "timeout").
|
||||
body_snippet: Optional diagnostic snippet from a response payload.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
route: Optional[str] = None,
|
||||
method: Optional[str] = None,
|
||||
status: Optional[int] = None,
|
||||
code: Optional[str] = None,
|
||||
body_snippet: Optional[str] = None,
|
||||
) -> None:
|
||||
super().__init__(message)
|
||||
self.route: Optional[str] = route
|
||||
self.method: Optional[str] = method
|
||||
self.status: Optional[int] = status
|
||||
self.code: Optional[str] = code
|
||||
self.body_snippet: Optional[str] = body_snippet
|
||||
|
||||
|
||||
class ComfyUIClientTimeout(ComfyUIClientError):
|
||||
"""Raised when an operation exceeds its allowed timeout."""
|
||||
|
||||
|
||||
class ComfyUIClientWsClosed(ComfyUIClientError):
|
||||
"""Raised when the websocket is closed or cannot be maintained."""
|
||||
|
||||
|
||||
class ComfyUIClientProtocol(Protocol):
|
||||
"""Protocol for ComfyUI clients.
|
||||
|
||||
Callers and tests should depend on this interface rather than a concrete
|
||||
implementation. Methods are asynchronous and may raise subclasses of
|
||||
`ComfyUIClientError`.
|
||||
"""
|
||||
|
||||
async def submit_prompt(self, payload: Dict[str, Any]) -> str:
|
||||
"""Submit a prompt to ComfyUI and return the resulting `prompt_id`.
|
||||
|
||||
Args:
|
||||
payload: JSON-serializable payload for the `/prompt` endpoint.
|
||||
|
||||
Returns:
|
||||
The non-empty `prompt_id` string returned by the server.
|
||||
"""
|
||||
|
||||
async def get_system_stats(self) -> SystemStats:
|
||||
"""Fetch system statistics from `/system_stats`."""
|
||||
|
||||
async def free(
|
||||
self, *, free_memory: bool = False, unload_models: bool = False
|
||||
) -> None:
|
||||
"""Invoke `/free` with the provided flags."""
|
||||
|
||||
async def ws_connect(self, client_id: str) -> AsyncIterator[WsEvent]:
|
||||
"""Connect to the websocket (`/ws?clientId=...`) and yield parsed events."""
|
||||
|
||||
async def probe(self) -> None:
|
||||
"""Lightweight health probe for connectivity/parseability.
|
||||
|
||||
Executes a GET `/system_stats` with a short timeout and no retries. Raises
|
||||
`ComfyUIClientError` subclasses on failure; returns `None` on success.
|
||||
"""
|
||||
|
||||
async def supports_hygiene(self) -> bool:
|
||||
"""Return True if both `/system_stats` and `/free` are supported.
|
||||
|
||||
Performs a one-time capability probe; caches results for the session.
|
||||
"""
|
||||
|
||||
|
||||
class ComfyUILocalClient(ComfyUIClientProtocol):
|
||||
"""Local HTTP/WebSocket client for a running ComfyUI instance.
|
||||
|
||||
This class provides the concrete implementation for the protocol. At this
|
||||
stage it only declares the interface and stores constructor parameters; the
|
||||
network behavior will be implemented in subsequent commits.
|
||||
|
||||
Args:
|
||||
base_url: Base HTTP URL for ComfyUI endpoints (e.g., `/prompt`).
|
||||
ws_url: Base WebSocket URL (e.g., `/ws`).
|
||||
session: Optional externally-managed aiohttp session for reuse.
|
||||
logger: Optional logger compatible with the worker's logging API.
|
||||
timeout: Default timeout in seconds for HTTP operations.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
ws_url: str,
|
||||
session: Optional["aiohttp.ClientSession"] = None,
|
||||
logger: Optional[Any] = None,
|
||||
timeout: float = 30.0,
|
||||
) -> None:
|
||||
self._base_url: str = base_url
|
||||
self._ws_url: str = ws_url
|
||||
self._session = session # type: ignore[assignment]
|
||||
self._logger = logger
|
||||
self._timeout: float = float(timeout)
|
||||
self._owned_session: Optional[aiohttp.ClientSession] = None
|
||||
# Backoff defaults per plan (commit 3)
|
||||
self._retry_base_seconds: float = 0.25
|
||||
self._retry_multiplier: float = 2.0
|
||||
self._retry_jitter_seconds: float = 0.25
|
||||
self._retry_max_sleep_seconds: float = 4.0
|
||||
self._retry_max_attempts: int = 3
|
||||
# WebSocket reconnect policy (commit 4)
|
||||
self._ws_max_reconnect_attempts: int = 5
|
||||
self._ws_max_total_backoff_seconds: float = 30.0
|
||||
# Capability probe cache (None = unknown, True/False = probed)
|
||||
self._supports_system_stats: Optional[bool] = None
|
||||
self._supports_free: Optional[bool] = None
|
||||
self._capability_warning_emitted: bool = False
|
||||
|
||||
# Lifecycle methods may be implemented later; for now they act as no-ops.
|
||||
async def __aenter__(self) -> "ComfyUILocalClient":
|
||||
"""Enter async context; create an internal session when none provided."""
|
||||
if self._session is None and self._owned_session is None:
|
||||
self._owned_session = aiohttp.ClientSession()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None: # type: ignore[override]
|
||||
"""Exit async context; close internal session if owned by this client."""
|
||||
if self._owned_session is not None:
|
||||
try:
|
||||
await self._owned_session.close()
|
||||
finally:
|
||||
self._owned_session = None
|
||||
return None
|
||||
|
||||
# Protocol methods — to be implemented in subsequent commits.
|
||||
async def submit_prompt(self, payload: Dict[str, Any]) -> str: # type: ignore[override]
|
||||
route = "/prompt"
|
||||
method = "POST"
|
||||
url = self._join_http(route)
|
||||
try:
|
||||
data = await self._http_request_json(
|
||||
method,
|
||||
url,
|
||||
json_payload=payload,
|
||||
timeout_s=self._timeout,
|
||||
retry_idempotent=False,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ComfyUIClientTimeout(
|
||||
f"Timeout while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code="timeout",
|
||||
) from e
|
||||
except _MappedHttpError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
except _MappedConnError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
|
||||
prompt_id = data.get("prompt_id") if isinstance(data, dict) else None
|
||||
if not isinstance(prompt_id, str) or not prompt_id.strip():
|
||||
snippet = _safe_preview(data)
|
||||
raise ComfyUIClientError(
|
||||
"Invalid response payload: missing non-empty prompt_id",
|
||||
route=route,
|
||||
method=method,
|
||||
code="bad_json",
|
||||
body_snippet=snippet,
|
||||
)
|
||||
return prompt_id
|
||||
|
||||
async def get_system_stats(self) -> SystemStats: # type: ignore[override]
|
||||
route = "/system_stats"
|
||||
method = "GET"
|
||||
url = self._join_http(route)
|
||||
try:
|
||||
data = await self._http_request_json(
|
||||
method,
|
||||
url,
|
||||
json_payload=None,
|
||||
timeout_s=self._timeout,
|
||||
retry_idempotent=True,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ComfyUIClientTimeout(
|
||||
f"Timeout while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code="timeout",
|
||||
) from e
|
||||
except _MappedHttpError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
except _MappedConnError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
|
||||
# Parse known fields, tolerate missing/unknown; try alternate shapes; fallback to torch
|
||||
vram_total = None
|
||||
vram_free = None
|
||||
torch_vram_free = None
|
||||
ram_total = None
|
||||
ram_free = None
|
||||
|
||||
if isinstance(data, dict):
|
||||
vram_total = _coerce_optional_float(data.get("vram_total"))
|
||||
vram_free = _coerce_optional_float(data.get("vram_free"))
|
||||
torch_vram_free = _coerce_optional_float(data.get("torch_vram_free"))
|
||||
ram_total = _coerce_optional_float(data.get("ram_total"))
|
||||
ram_free = _coerce_optional_float(data.get("ram_free"))
|
||||
|
||||
# Common alternates
|
||||
if vram_total is None:
|
||||
vram_total = _coerce_optional_float(data.get("total_vram"))
|
||||
if vram_free is None:
|
||||
vram_free = _coerce_optional_float(data.get("free_vram"))
|
||||
|
||||
vram_obj = data.get("vram") if isinstance(data.get("vram"), dict) else None
|
||||
if vram_obj:
|
||||
if vram_total is None:
|
||||
vram_total = _coerce_optional_float(vram_obj.get("total"))
|
||||
if vram_free is None:
|
||||
vram_free = _coerce_optional_float(vram_obj.get("free"))
|
||||
|
||||
ram_obj = data.get("ram") if isinstance(data.get("ram"), dict) else None
|
||||
if ram_obj:
|
||||
if ram_total is None:
|
||||
ram_total = _coerce_optional_float(ram_obj.get("total"))
|
||||
if ram_free is None:
|
||||
ram_free = _coerce_optional_float(ram_obj.get("free"))
|
||||
|
||||
# devices[0] fallback (common in mock/alt servers)
|
||||
devices = (
|
||||
data.get("devices") if isinstance(data.get("devices"), list) else None
|
||||
)
|
||||
if devices and len(devices) > 0 and isinstance(devices[0], dict):
|
||||
dev0 = devices[0]
|
||||
if vram_total is None:
|
||||
vram_total = _coerce_optional_float(dev0.get("vram_total"))
|
||||
if vram_free is None:
|
||||
vram_free = _coerce_optional_float(dev0.get("vram_free"))
|
||||
# Note: we rely solely on server-reported values; no local torch fallback
|
||||
|
||||
# Optional debug logging of reported stats
|
||||
if self._logger:
|
||||
try:
|
||||
self._logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (comfyui_client): /system_stats: vram_total=%s vram_free=%s torch_vram_free=%s ram_total=%s ram_free=%s",
|
||||
vram_total,
|
||||
vram_free,
|
||||
torch_vram_free,
|
||||
ram_total,
|
||||
ram_free,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return SystemStats(
|
||||
vram_total=vram_total,
|
||||
vram_free=vram_free,
|
||||
torch_vram_free=torch_vram_free,
|
||||
ram_total=ram_total,
|
||||
ram_free=ram_free,
|
||||
)
|
||||
|
||||
async def free(self, *, free_memory: bool = False, unload_models: bool = False) -> None: # type: ignore[override]
|
||||
route = "/free"
|
||||
method = "POST"
|
||||
url = self._join_http(route)
|
||||
body = {"free_memory": bool(free_memory), "unload_models": bool(unload_models)}
|
||||
try:
|
||||
await self._http_request_json(
|
||||
method,
|
||||
url,
|
||||
json_payload=body,
|
||||
timeout_s=self._timeout,
|
||||
retry_idempotent=True, # idempotent when flags identical
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ComfyUIClientTimeout(
|
||||
f"Timeout while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code="timeout",
|
||||
) from e
|
||||
except _MappedHttpError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
except _MappedConnError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
|
||||
async def ws_connect(self, client_id: str) -> AsyncIterator[WsEvent]: # type: ignore[override]
|
||||
url = _ws_url(self._ws_url, client_id)
|
||||
attempts = 0
|
||||
total_backoff = 0.0
|
||||
base = self._retry_base_seconds
|
||||
multiplier = self._retry_multiplier
|
||||
jitter = self._retry_jitter_seconds
|
||||
max_sleep = self._retry_max_sleep_seconds
|
||||
|
||||
while True:
|
||||
try:
|
||||
# Allow large frames and set ping/pong defaults
|
||||
async with websockets.connect(
|
||||
url,
|
||||
max_size=None,
|
||||
max_queue=4,
|
||||
ping_interval=20,
|
||||
ping_timeout=20,
|
||||
) as websocket:
|
||||
# On successful connect, reset counters
|
||||
attempts = 0
|
||||
total_backoff = 0.0
|
||||
if self._logger:
|
||||
try:
|
||||
self._logger.debug(
|
||||
f"✅ Nilor-Nodes (comfyui_client): connected to websocket {url}"
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
while True:
|
||||
message = await websocket.recv()
|
||||
if isinstance(message, bytes):
|
||||
try:
|
||||
message = message.decode("utf-8", errors="replace")
|
||||
except Exception:
|
||||
yield WsEvent(type="binary", data={"length": len(message)}) # type: ignore[call-arg]
|
||||
continue
|
||||
try:
|
||||
payload = json.loads(message)
|
||||
except Exception:
|
||||
yield WsEvent(type="text", data={"message": message}) # type: ignore[call-arg]
|
||||
continue
|
||||
|
||||
if isinstance(payload, dict):
|
||||
event_type = str(payload.get("type", "event"))
|
||||
event_data = payload.get("data")
|
||||
if not isinstance(event_data, dict):
|
||||
event_data = {"raw": payload}
|
||||
yield WsEvent(type=event_type, data=event_data) # type: ignore[call-arg]
|
||||
else:
|
||||
yield WsEvent(type="event", data={"raw": payload}) # type: ignore[call-arg]
|
||||
|
||||
except asyncio.CancelledError:
|
||||
# Allow clean shutdown by propagating cancellation
|
||||
raise
|
||||
except ws_exc.ConnectionClosedOK as e:
|
||||
raise ComfyUIClientWsClosed(
|
||||
"WebSocket closed normally",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
status=getattr(e, "code", 1000),
|
||||
code="ws_closed",
|
||||
body_snippet=str(getattr(e, "reason", ""))[:256],
|
||||
) from e
|
||||
except ws_exc.ConnectionClosedError as e:
|
||||
# Abnormal close: attempt bounded reconnect
|
||||
attempts += 1
|
||||
if attempts > self._ws_max_reconnect_attempts:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket reconnect attempts exhausted",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="ws_closed",
|
||||
body_snippet=str(getattr(e, "reason", ""))[:256],
|
||||
) from e
|
||||
delay = _next_backoff(attempts - 1, base, multiplier, jitter, max_sleep)
|
||||
total_backoff += delay
|
||||
if total_backoff > self._ws_max_total_backoff_seconds:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket reconnect backoff budget exhausted",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="ws_closed",
|
||||
body_snippet=str(getattr(e, "reason", ""))[:256],
|
||||
) from e
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
except ws_exc.InvalidStatus as e:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket handshake failed",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="ws_handshake",
|
||||
) from e
|
||||
except ws_exc.InvalidURI as e:
|
||||
raise ComfyUIClientError(
|
||||
"Invalid WebSocket URI",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="invalid_ws_uri",
|
||||
) from e
|
||||
except Exception as e:
|
||||
# Treat as connection error; bounded reconnect
|
||||
attempts += 1
|
||||
if attempts > self._ws_max_reconnect_attempts:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket reconnect attempts exhausted",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="connection_error",
|
||||
) from e
|
||||
delay = _next_backoff(attempts - 1, base, multiplier, jitter, max_sleep)
|
||||
total_backoff += delay
|
||||
if total_backoff > self._ws_max_total_backoff_seconds:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket reconnect backoff budget exhausted",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="connection_error",
|
||||
) from e
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
|
||||
async def probe(self) -> None: # type: ignore[override]
|
||||
route = "/system_stats"
|
||||
method = "GET"
|
||||
url = self._join_http(route)
|
||||
short_timeout = min(self._timeout, 3.0)
|
||||
try:
|
||||
# No retries: retry_idempotent=False
|
||||
await self._http_request_json(
|
||||
method,
|
||||
url,
|
||||
json_payload=None,
|
||||
timeout_s=short_timeout,
|
||||
retry_idempotent=False,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ComfyUIClientTimeout(
|
||||
f"Timeout while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code="timeout",
|
||||
) from e
|
||||
except _MappedHttpError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
except _MappedConnError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
|
||||
async def supports_hygiene(self) -> bool: # type: ignore[override]
|
||||
# If both probed, return cached decision
|
||||
if self._supports_system_stats is not None and self._supports_free is not None:
|
||||
return bool(self._supports_system_stats and self._supports_free)
|
||||
|
||||
await self._probe_capabilities_once()
|
||||
supported = bool(
|
||||
(self._supports_system_stats is True) and (self._supports_free is True)
|
||||
)
|
||||
if not supported and not self._capability_warning_emitted and self._logger:
|
||||
try:
|
||||
self._logger.warning(
|
||||
"⚠️\u2009 Nilor-Nodes (comfyui_client): hygiene disabled — missing /system_stats or /free support"
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
self._capability_warning_emitted = True
|
||||
return supported
|
||||
|
||||
async def _probe_capabilities_once(self) -> None:
|
||||
"""Probe `/system_stats` and `/free` capabilities once and cache results.
|
||||
|
||||
Only marks capabilities as False on definitive 404/405 responses. Transient
|
||||
failures leave the capability as None so a future call may retry.
|
||||
"""
|
||||
short_timeout = min(self._timeout, 3.0)
|
||||
|
||||
# Probe /system_stats support
|
||||
if self._supports_system_stats is None:
|
||||
route = "/system_stats"
|
||||
url = self._join_http(route)
|
||||
try:
|
||||
await self._http_request_json(
|
||||
"GET",
|
||||
url,
|
||||
json_payload=None,
|
||||
timeout_s=short_timeout,
|
||||
retry_idempotent=False,
|
||||
)
|
||||
self._supports_system_stats = True
|
||||
except ComfyUIClientError as e:
|
||||
if getattr(e, "status", None) in (404, 405):
|
||||
self._supports_system_stats = False
|
||||
|
||||
# Probe /free support (no-op body)
|
||||
if self._supports_free is None:
|
||||
route = "/free"
|
||||
url = self._join_http(route)
|
||||
try:
|
||||
await self._http_request_json(
|
||||
"POST",
|
||||
url,
|
||||
json_payload={"free_memory": False, "unload_models": False},
|
||||
timeout_s=short_timeout,
|
||||
retry_idempotent=False,
|
||||
)
|
||||
self._supports_free = True
|
||||
except ComfyUIClientError as e:
|
||||
if getattr(e, "status", None) in (404, 405):
|
||||
self._supports_free = False
|
||||
|
||||
|
||||
# Runtime dependency; imported here to avoid issues if module is scanned without execution
|
||||
import aiohttp # type: ignore
|
||||
import websockets # type: ignore
|
||||
from websockets import exceptions as ws_exc # type: ignore
|
||||
|
||||
|
||||
# ---- Internal helpers (HTTP) ----
|
||||
|
||||
|
||||
def _coerce_optional_float(value: Any) -> Optional[float]:
|
||||
try:
|
||||
if value is None:
|
||||
return None
|
||||
return float(value)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _is_transient_status(status: int) -> bool:
|
||||
return status == 429 or 500 <= status <= 599
|
||||
|
||||
|
||||
def _safe_preview(data: Any, limit: int = 512) -> str:
|
||||
try:
|
||||
text = json.dumps(data, ensure_ascii=False)
|
||||
except Exception:
|
||||
text = str(data)
|
||||
if len(text) > limit:
|
||||
return text[:limit] + "…"
|
||||
return text
|
||||
|
||||
|
||||
class _MappedHttpError(Exception):
|
||||
def __init__(self, *, status: Optional[int], body_snippet: Optional[str]) -> None:
|
||||
self.status = status
|
||||
self.body_snippet = body_snippet
|
||||
|
||||
def to_public_error(self, *, route: str, method: str) -> ComfyUIClientError:
|
||||
return ComfyUIClientError(
|
||||
f"🛑\u2009 Nilor-Nodes (comfyui_client): HTTP error while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
status=self.status,
|
||||
code="http_error",
|
||||
body_snippet=self.body_snippet,
|
||||
)
|
||||
|
||||
|
||||
class _MappedConnError(Exception):
|
||||
def __init__(self, *, code: str = "connection_error") -> None:
|
||||
self.code = code
|
||||
|
||||
def to_public_error(self, *, route: str, method: str) -> ComfyUIClientError:
|
||||
return ComfyUIClientError(
|
||||
f"🛑\u2009 Nilor-Nodes (comfyui_client): Connection error while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code=self.code,
|
||||
)
|
||||
|
||||
|
||||
async def _read_limited_text(resp: aiohttp.ClientResponse, limit: int = 512) -> str:
|
||||
try:
|
||||
raw = await resp.read()
|
||||
# Truncate at byte level then decode safely
|
||||
raw = raw[:limit]
|
||||
return raw.decode("utf-8", errors="replace")
|
||||
except Exception:
|
||||
return ""
|
||||
|
||||
|
||||
def _with_jitter(seconds: float, jitter: float) -> float:
|
||||
if jitter <= 0:
|
||||
return seconds
|
||||
return max(0.0, seconds + random.uniform(-jitter, jitter))
|
||||
|
||||
|
||||
def _next_backoff(
|
||||
attempt_index: int,
|
||||
base: float,
|
||||
multiplier: float,
|
||||
jitter: float,
|
||||
max_sleep: float,
|
||||
) -> float:
|
||||
# attempt_index is 0-based
|
||||
delay = base * (multiplier**attempt_index)
|
||||
delay = min(delay, max_sleep)
|
||||
return _with_jitter(delay, jitter)
|
||||
|
||||
|
||||
def _should_retry(
|
||||
*,
|
||||
retry_idempotent: bool,
|
||||
exc: Optional[BaseException] = None,
|
||||
status: Optional[int] = None,
|
||||
) -> bool:
|
||||
if not retry_idempotent:
|
||||
return False
|
||||
if isinstance(exc, (asyncio.TimeoutError, aiohttp.ClientConnectionError)):
|
||||
return True
|
||||
if status is not None and _is_transient_status(status):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _finalize_attempts(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
last_exc: Optional[BaseException],
|
||||
last_status: Optional[int],
|
||||
last_body_snippet: Optional[str],
|
||||
) -> BaseException:
|
||||
if isinstance(last_exc, asyncio.TimeoutError):
|
||||
return ComfyUIClientTimeout(
|
||||
f"🛑\u2009 Nilor-Nodes (comfyui_client): Timeout while calling {method} {url}",
|
||||
route=_route_from_url(url),
|
||||
method=method,
|
||||
code="timeout",
|
||||
)
|
||||
if isinstance(last_exc, aiohttp.ClientConnectionError):
|
||||
return _MappedConnError().to_public_error(
|
||||
route=_route_from_url(url), method=method
|
||||
)
|
||||
# Otherwise treat as HTTP error
|
||||
return _MappedHttpError(
|
||||
status=last_status, body_snippet=last_body_snippet
|
||||
).to_public_error(route=_route_from_url(url), method=method)
|
||||
|
||||
|
||||
def _route_from_url(url: str) -> str:
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
return parsed.path or "/"
|
||||
except Exception:
|
||||
return url
|
||||
|
||||
|
||||
class _TempSession:
|
||||
"""Context manager that yields an aiohttp session, reusing if provided."""
|
||||
|
||||
def __init__(self, session: Optional[aiohttp.ClientSession]):
|
||||
self._provided = session
|
||||
self._owned: Optional[aiohttp.ClientSession] = None
|
||||
|
||||
async def __aenter__(self) -> aiohttp.ClientSession:
|
||||
if self._provided is not None:
|
||||
return self._provided
|
||||
self._owned = aiohttp.ClientSession()
|
||||
return self._owned
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None: # type: ignore[override]
|
||||
if self._owned is not None:
|
||||
await self._owned.close()
|
||||
|
||||
|
||||
async def _json_or_text(resp: aiohttp.ClientResponse) -> Any:
|
||||
ctype = resp.headers.get("Content-Type", "").lower()
|
||||
text = await _read_limited_text(resp) # limited read to use for both cases
|
||||
if "json" in ctype:
|
||||
try:
|
||||
return json.loads(text)
|
||||
except Exception:
|
||||
# Fallthrough to treat as bad JSON
|
||||
raise _MappedHttpError(status=resp.status, body_snippet=text)
|
||||
# Not JSON; return raw text
|
||||
return text
|
||||
|
||||
|
||||
async def _raise_for_status_with_snippet(resp: aiohttp.ClientResponse) -> None:
|
||||
if 200 <= resp.status <= 299:
|
||||
return
|
||||
snippet = await _read_limited_text(resp)
|
||||
raise _MappedHttpError(status=resp.status, body_snippet=snippet)
|
||||
|
||||
|
||||
async def _request_once(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
session: aiohttp.ClientSession,
|
||||
json_payload: Optional[Dict[str, Any]],
|
||||
timeout_s: float,
|
||||
) -> Tuple[Optional[int], Optional[str], Any]:
|
||||
timeout = aiohttp.ClientTimeout(total=timeout_s)
|
||||
try:
|
||||
async with session.request(
|
||||
method, url, json=json_payload, timeout=timeout
|
||||
) as resp:
|
||||
status = resp.status
|
||||
await _raise_for_status_with_snippet(resp)
|
||||
# Success path: attempt to parse JSON body; if not JSON, return text
|
||||
try:
|
||||
data = await resp.json(content_type=None)
|
||||
except aiohttp.ContentTypeError:
|
||||
# Not JSON; use limited text
|
||||
data = await _read_limited_text(resp)
|
||||
return status, None, data
|
||||
except asyncio.TimeoutError:
|
||||
raise
|
||||
except aiohttp.ClientConnectionError as e:
|
||||
raise e
|
||||
except aiohttp.ClientPayloadError as e:
|
||||
# Map as payload error
|
||||
raise _MappedHttpError(status=None, body_snippet=str(e))
|
||||
|
||||
|
||||
async def _http_request_core(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
session: aiohttp.ClientSession,
|
||||
json_payload: Optional[Dict[str, Any]],
|
||||
timeout_s: float,
|
||||
retry_idempotent: bool,
|
||||
base: float,
|
||||
multiplier: float,
|
||||
jitter: float,
|
||||
max_sleep: float,
|
||||
max_attempts: int,
|
||||
) -> Any:
|
||||
last_exc: Optional[BaseException] = None
|
||||
last_status: Optional[int] = None
|
||||
last_body: Optional[str] = None
|
||||
|
||||
attempts = max(1, int(max_attempts))
|
||||
for attempt in range(attempts):
|
||||
try:
|
||||
status, body_snippet, data = await _request_once(
|
||||
method,
|
||||
url,
|
||||
session=session,
|
||||
json_payload=json_payload,
|
||||
timeout_s=timeout_s,
|
||||
)
|
||||
return data
|
||||
except asyncio.TimeoutError as e:
|
||||
last_exc = e
|
||||
if attempt < attempts - 1 and _should_retry(
|
||||
retry_idempotent=retry_idempotent, exc=e
|
||||
):
|
||||
await asyncio.sleep(
|
||||
_next_backoff(attempt, base, multiplier, jitter, max_sleep)
|
||||
)
|
||||
continue
|
||||
break
|
||||
except aiohttp.ClientConnectionError as e:
|
||||
last_exc = e
|
||||
if attempt < attempts - 1 and _should_retry(
|
||||
retry_idempotent=retry_idempotent, exc=e
|
||||
):
|
||||
await asyncio.sleep(
|
||||
_next_backoff(attempt, base, multiplier, jitter, max_sleep)
|
||||
)
|
||||
continue
|
||||
break
|
||||
except _MappedHttpError as e:
|
||||
last_exc = None
|
||||
last_status = e.status
|
||||
last_body = e.body_snippet
|
||||
if attempt < attempts - 1 and _should_retry(
|
||||
retry_idempotent=retry_idempotent, status=e.status
|
||||
):
|
||||
await asyncio.sleep(
|
||||
_next_backoff(attempt, base, multiplier, jitter, max_sleep)
|
||||
)
|
||||
continue
|
||||
break
|
||||
|
||||
raise _finalize_attempts(
|
||||
method,
|
||||
url,
|
||||
last_exc=last_exc,
|
||||
last_status=last_status,
|
||||
last_body_snippet=last_body,
|
||||
)
|
||||
|
||||
|
||||
async def _http_request_json(
|
||||
self: "ComfyUILocalClient",
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
json_payload: Optional[Dict[str, Any]],
|
||||
timeout_s: float,
|
||||
retry_idempotent: bool,
|
||||
) -> Any:
|
||||
async with _TempSession(self._session or self._owned_session) as session:
|
||||
return await _http_request_core(
|
||||
method,
|
||||
url,
|
||||
session=session,
|
||||
json_payload=json_payload,
|
||||
timeout_s=timeout_s,
|
||||
retry_idempotent=retry_idempotent,
|
||||
base=self._retry_base_seconds,
|
||||
multiplier=self._retry_multiplier,
|
||||
jitter=self._retry_jitter_seconds,
|
||||
max_sleep=self._retry_max_sleep_seconds,
|
||||
max_attempts=self._retry_max_attempts,
|
||||
)
|
||||
|
||||
|
||||
def _join_http_base(base: str, route: str) -> str:
|
||||
if not route:
|
||||
return base
|
||||
return f"{base.rstrip('/')}{route}"
|
||||
|
||||
|
||||
def _join_ws_base(base: str, route: str) -> str:
|
||||
if not route:
|
||||
return base
|
||||
return f"{base.rstrip('/')}{route}"
|
||||
|
||||
|
||||
def _ensure_scheme(base: str, allowed: Tuple[str, ...]) -> None:
|
||||
parsed = urlparse(base)
|
||||
if not parsed.scheme or parsed.scheme.lower() not in allowed:
|
||||
allowed_str = ", ".join(allowed)
|
||||
raise ValueError(
|
||||
f"🛑\u2009 Nilor-Nodes (comfyui_client): Base URL must start with one of [{allowed_str}]; got: {base!r}"
|
||||
)
|
||||
|
||||
|
||||
def _validate_bases(http_base: str, ws_base: str) -> None:
|
||||
_ensure_scheme(http_base, ("http", "https"))
|
||||
_ensure_scheme(ws_base, ("ws", "wss"))
|
||||
|
||||
|
||||
def _quote_client_id(client_id: str) -> str:
|
||||
return quote(client_id, safe="")
|
||||
|
||||
|
||||
def _ws_url(base_ws: str, client_id: str) -> str:
|
||||
return f"{base_ws.rstrip('/')}/ws?clientId={_quote_client_id(client_id)}"
|
||||
|
||||
|
||||
# Bind helper methods to class namespace (private) without exposing publicly
|
||||
ComfyUILocalClient._http_request_json = _http_request_json # type: ignore[attr-defined]
|
||||
ComfyUILocalClient._join_http = lambda self, route: _join_http_base(self._base_url, route) # type: ignore[attr-defined]
|
||||
ComfyUILocalClient._join_ws = lambda self, route: _join_ws_base(self._ws_url, route) # type: ignore[attr-defined]
|
||||
@@ -0,0 +1,75 @@
|
||||
{
|
||||
// Nilor-Nodes sidecar configuration defaults (non-secret) for local development
|
||||
// Precedence model (enforced by the loader): environment variables > this file
|
||||
// Secrets MUST NOT be committed here; use .env to override sensitive values
|
||||
|
||||
allow_env_override: true,
|
||||
|
||||
// Log level for nilor-nodes components (DEBUG, INFO, WARNING, ERROR, CRITICAL)
|
||||
NILOR_LOG_LEVEL: "INFO",
|
||||
// Enable/disable SQS job consumption from Nilor brain_rnd
|
||||
NILOR_SQS_ENABLED: false,
|
||||
|
||||
// ---- Comfy API client (consumed by worker_consumer.py) ----
|
||||
// Base HTTP URL for ComfyUI REST API (worker submits to `${api_url}/prompt`)
|
||||
NILOR_COMFYUI_API_URL: "http://127.0.0.1:8188",
|
||||
// Base WS URL for ComfyUI websocket events (worker listens at `${ws_url}/ws`)
|
||||
NILOR_COMFYUI_WS_URL: "ws://127.0.0.1:8188",
|
||||
// Request timeout in seconds for ComfyUI HTTP calls
|
||||
NILOR_COMFY_API_TIMEOUT_SECONDS: 30,
|
||||
|
||||
// HTTP idempotent retry policy (applies only to GET /system_stats and POST /free)
|
||||
NILOR_COMFY_RETRY_BASE_SECONDS: 0.25,
|
||||
NILOR_COMFY_RETRY_MULTIPLIER: 2.0,
|
||||
NILOR_COMFY_RETRY_JITTER_SECONDS: 0.25,
|
||||
NILOR_COMFY_RETRY_MAX_SLEEP_SECONDS: 4.0,
|
||||
NILOR_COMFY_RETRY_MAX_ATTEMPTS: 3,
|
||||
|
||||
// WebSocket reconnect policy
|
||||
NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS: 5,
|
||||
NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS: 30.0,
|
||||
|
||||
// ---- Memory Hygiene (Guardian) defaults ----
|
||||
// Enable/disable hygiene between jobs
|
||||
NILOR_MEMORY_HYGIENE_ENABLED: false,
|
||||
// How often to check when idle (seconds)
|
||||
NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS: 5,
|
||||
// Thresholds (either percent usage or absolute free MB can trigger)
|
||||
NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX: 88,
|
||||
NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX: 90,
|
||||
NILOR_MEMORY_HYGIENE_VRAM_MIN_FREE_MB: 2048,
|
||||
NILOR_MEMORY_HYGIENE_RAM_MIN_FREE_MB: 4096,
|
||||
// Action policy: free|unload|both|auto (auto: stage free then unload once if needed)
|
||||
NILOR_MEMORY_HYGIENE_ACTION_POLICY: "auto",
|
||||
// Retry/backoff/cycle caps
|
||||
NILOR_MEMORY_HYGIENE_MAX_RETRIES: 2,
|
||||
NILOR_MEMORY_HYGIENE_COOLDOWN_SECONDS: 60,
|
||||
NILOR_MEMORY_HYGIENE_SLEEP_BETWEEN_ATTEMPTS_SECONDS: 4,
|
||||
NILOR_MEMORY_HYGIENE_MAX_CYCLE_DURATION_SECONDS: 15,
|
||||
|
||||
// ---- Worker / Queue settings ----
|
||||
// Workflow normalization (applies in worker_consumer before submission)
|
||||
// When enabled, nilor-nodes will normalize OS-specific path formatting inside
|
||||
// incoming workflows (e.g. Windows backslashes vs POSIX slashes).
|
||||
NILOR_WORKFLOW_OS_NORMALIZATION_ENABLED: true,
|
||||
|
||||
// ElasticMQ/SQS endpoint URL (used by consumer and status updates)
|
||||
NILOR_SQS_ENDPOINT_URL: "http://localhost:9324",
|
||||
// Queue to pull new jobs from (polled by worker_consumer)
|
||||
NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME: "jobs_to_process-comfyui",
|
||||
// Queue to publish job status updates to (sent by worker_consumer)
|
||||
NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME: "job_status_updates",
|
||||
// Long poll wait time (seconds); must be within [0, 20]
|
||||
NILOR_SQS_POLL_WAIT_TIME: 10,
|
||||
// Max number of messages to pull per poll
|
||||
NILOR_SQS_MAX_MESSAGES: 1,
|
||||
|
||||
// Non-secret local defaults; override via .env for real deployments
|
||||
NILOR_AWS_ACCESS_KEY_ID: "local",
|
||||
NILOR_AWS_SECRET_ACCESS_KEY: "local", // secret must be provided via .env
|
||||
NILOR_AWS_DEFAULT_REGION: "us-east-1",
|
||||
|
||||
// If empty or unset, a stable id will be generated by the loader
|
||||
NILOR_WORKER_CLIENT_ID: ""
|
||||
}
|
||||
|
||||
@@ -0,0 +1,584 @@
|
||||
"""
|
||||
Typed configuration scaffolding for the Nilor-Nodes ComfyUI sidecar.
|
||||
|
||||
This module defines the dataclasses and public loader API contract. The actual
|
||||
implementation of precedence, parsing, and validation is added in a subsequent
|
||||
commit. For now, only type definitions and the public `Config.load` signature
|
||||
are provided to enable incremental integration without behavior changes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import socket
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, Optional
|
||||
from urllib.parse import urlparse
|
||||
from pathlib import Path
|
||||
from ..logger import logger
|
||||
|
||||
# Ensure we cache config globally
|
||||
_CONFIG: Optional["NilorNodesConfig"] = None
|
||||
|
||||
try: # json5 is declared in nilor-nodes/requirements.txt
|
||||
import json5 # type: ignore
|
||||
except Exception as _e: # pragma: no cover
|
||||
json5 = None # lazy failure in loader
|
||||
|
||||
|
||||
class BaseConfig: # type: ignore
|
||||
@classmethod
|
||||
def get_instance(cls):
|
||||
# Minimal fallback: load JSON5 directly when BaseConfig is unavailable
|
||||
path = os.path.join(os.path.dirname(__file__), "config.json5")
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json5.load(f) if json5 is not None else {}
|
||||
return cls.from_dict(data) # type: ignore[attr-defined]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ComfyApiConfig:
|
||||
"""Configuration for the ComfyUI client.
|
||||
|
||||
Args:
|
||||
api_url: Base HTTP URL for the ComfyUI REST API (e.g., "http://127.0.0.1:8188").
|
||||
ws_url: Base WebSocket URL for ComfyUI events (e.g., "ws://127.0.0.1:8188").
|
||||
timeout_s: Request timeout in seconds for ComfyUI HTTP calls.
|
||||
retry_base_seconds: Base backoff seconds for idempotent retries.
|
||||
retry_multiplier: Exponential backoff multiplier.
|
||||
retry_jitter_seconds: Jitter range (±seconds) added to backoff.
|
||||
retry_max_sleep_seconds: Maximum sleep per backoff step.
|
||||
retry_max_attempts: Maximum retry attempts for idempotent routes.
|
||||
ws_max_reconnect_attempts: Maximum websocket reconnect attempts.
|
||||
ws_max_total_backoff_seconds: Cap on total backoff time during WS reconnects.
|
||||
"""
|
||||
|
||||
api_url: str
|
||||
ws_url: str
|
||||
timeout_s: int
|
||||
retry_base_seconds: float
|
||||
retry_multiplier: float
|
||||
retry_jitter_seconds: float
|
||||
retry_max_sleep_seconds: float
|
||||
retry_max_attempts: int
|
||||
ws_max_reconnect_attempts: int
|
||||
ws_max_total_backoff_seconds: float
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkerConfig:
|
||||
"""Configuration for the worker and its SQS integration.
|
||||
|
||||
Args:
|
||||
sqs_endpoint_url: URL of the SQS-compatible endpoint (e.g., ElasticMQ).
|
||||
jobs_queue: Name of the queue from which to pull new jobs.
|
||||
status_queue: Name of the queue to which job status updates are published.
|
||||
poll_wait_s: Long poll wait time in seconds (expected to be within [0, 20]).
|
||||
max_messages: Max number of messages pulled per poll.
|
||||
aws_access_key_id: Access key id for the SQS client (non-secret default acceptable for local dev).
|
||||
aws_secret_access_key: Secret access key for the SQS client (must be overridden via environment for real deployments).
|
||||
aws_region: AWS region name used by the SQS client.
|
||||
worker_client_id: Stable identifier for routing websocket events to this worker.
|
||||
workflow_os_normalization_enabled: When true, normalize OS-specific path formatting
|
||||
inside ComfyUI prompt graphs before submission.
|
||||
"""
|
||||
|
||||
sqs_endpoint_url: str
|
||||
jobs_queue: str
|
||||
status_queue: str
|
||||
poll_wait_s: int
|
||||
max_messages: int
|
||||
aws_access_key_id: str
|
||||
aws_secret_access_key: str
|
||||
aws_region: str
|
||||
worker_client_id: str
|
||||
workflow_os_normalization_enabled: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MemoryHygieneConfig:
|
||||
"""Configuration for Memory Guardian hygiene between jobs.
|
||||
|
||||
Args:
|
||||
enabled: Feature flag to enable/disable hygiene.
|
||||
idle_poll_seconds: How often to check while idle.
|
||||
vram_usage_pct_max: If used VRAM exceeds this percent, trigger remediation.
|
||||
ram_usage_pct_max: If used RAM exceeds this percent, trigger remediation.
|
||||
vram_min_free_mb: Absolute minimum free VRAM (MB) threshold.
|
||||
ram_min_free_mb: Absolute minimum free RAM (MB) threshold.
|
||||
action_policy: One of {free, unload, both, auto}.
|
||||
max_retries: Max remediation attempts per cycle.
|
||||
cooldown_seconds: Cooldown between remediation cycles.
|
||||
sleep_between_attempts_seconds: Wait time after /free before re-checking.
|
||||
max_cycle_duration_seconds: Hard cap for a single remediation cycle.
|
||||
"""
|
||||
|
||||
enabled: bool
|
||||
idle_poll_seconds: int
|
||||
vram_usage_pct_max: int
|
||||
ram_usage_pct_max: int
|
||||
vram_min_free_mb: int
|
||||
ram_min_free_mb: int
|
||||
action_policy: str
|
||||
max_retries: int
|
||||
cooldown_seconds: int
|
||||
sleep_between_attempts_seconds: int
|
||||
max_cycle_duration_seconds: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class NilorNodesConfig(BaseConfig):
|
||||
"""Aggregate configuration for the Nilor-Nodes sidecar.
|
||||
|
||||
Args:
|
||||
comfy: Configuration for the ComfyUI HTTP/WS client.
|
||||
worker: Configuration for SQS and worker identity.
|
||||
allow_env_override: When true, environment variables may override file values.
|
||||
"""
|
||||
|
||||
comfy: ComfyApiConfig
|
||||
worker: WorkerConfig
|
||||
allow_env_override: bool
|
||||
sqs_enabled: bool
|
||||
hygiene: MemoryHygieneConfig
|
||||
|
||||
@classmethod
|
||||
def _get_config_path(cls) -> str:
|
||||
return os.path.join(os.path.dirname(__file__), "config.json5")
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config_dict: Dict[str, object]) -> "NilorNodesConfig":
|
||||
allow_env_override = bool(config_dict.get("allow_env_override", True))
|
||||
sqs_enabled = _coerce_bool(config_dict.get("NILOR_SQS_ENABLED", False))
|
||||
|
||||
# Build nested from flat NILOR_* keys present in JSON5
|
||||
comfy_cfg = ComfyApiConfig(
|
||||
api_url=str(config_dict.get("NILOR_COMFYUI_API_URL", "")).strip(),
|
||||
ws_url=str(config_dict.get("NILOR_COMFYUI_WS_URL", "")).strip(),
|
||||
timeout_s=int(config_dict.get("NILOR_COMFY_API_TIMEOUT_SECONDS", 30)),
|
||||
retry_base_seconds=float(
|
||||
config_dict.get("NILOR_COMFY_RETRY_BASE_SECONDS", 0.25)
|
||||
),
|
||||
retry_multiplier=float(
|
||||
config_dict.get("NILOR_COMFY_RETRY_MULTIPLIER", 2.0)
|
||||
),
|
||||
retry_jitter_seconds=float(
|
||||
config_dict.get("NILOR_COMFY_RETRY_JITTER_SECONDS", 0.25)
|
||||
),
|
||||
retry_max_sleep_seconds=float(
|
||||
config_dict.get("NILOR_COMFY_RETRY_MAX_SLEEP_SECONDS", 4.0)
|
||||
),
|
||||
retry_max_attempts=int(
|
||||
config_dict.get("NILOR_COMFY_RETRY_MAX_ATTEMPTS", 3)
|
||||
),
|
||||
ws_max_reconnect_attempts=int(
|
||||
config_dict.get("NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS", 5)
|
||||
),
|
||||
ws_max_total_backoff_seconds=float(
|
||||
config_dict.get("NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS", 30.0)
|
||||
),
|
||||
)
|
||||
|
||||
worker_client_id = (
|
||||
str(config_dict.get("NILOR_WORKER_CLIENT_ID", "")).strip()
|
||||
or _generate_worker_client_id()
|
||||
)
|
||||
|
||||
worker_cfg = WorkerConfig(
|
||||
sqs_endpoint_url=str(config_dict.get("NILOR_SQS_ENDPOINT_URL", "")).strip(),
|
||||
jobs_queue=str(
|
||||
config_dict.get("NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME", "")
|
||||
).strip(),
|
||||
status_queue=str(
|
||||
config_dict.get("NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME", "")
|
||||
).strip(),
|
||||
poll_wait_s=int(config_dict.get("NILOR_SQS_POLL_WAIT_TIME", 10)),
|
||||
max_messages=int(config_dict.get("NILOR_SQS_MAX_MESSAGES", 1)),
|
||||
aws_access_key_id=str(
|
||||
config_dict.get("NILOR_AWS_ACCESS_KEY_ID", "")
|
||||
).strip(),
|
||||
aws_secret_access_key=str(
|
||||
config_dict.get("NILOR_AWS_SECRET_ACCESS_KEY", "")
|
||||
).strip(),
|
||||
aws_region=str(config_dict.get("NILOR_AWS_DEFAULT_REGION", "")).strip(),
|
||||
worker_client_id=worker_client_id,
|
||||
workflow_os_normalization_enabled=_coerce_bool(
|
||||
config_dict.get("NILOR_WORKFLOW_OS_NORMALIZATION_ENABLED", True)
|
||||
),
|
||||
)
|
||||
|
||||
# Memory hygiene
|
||||
hygiene_cfg = MemoryHygieneConfig(
|
||||
enabled=_coerce_bool(config_dict.get("NILOR_MEMORY_HYGIENE_ENABLED", True)),
|
||||
idle_poll_seconds=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS", 5)
|
||||
),
|
||||
vram_usage_pct_max=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX", 88)
|
||||
),
|
||||
ram_usage_pct_max=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX", 90)
|
||||
),
|
||||
vram_min_free_mb=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_VRAM_MIN_FREE_MB", 2048)
|
||||
),
|
||||
ram_min_free_mb=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_RAM_MIN_FREE_MB", 4096)
|
||||
),
|
||||
action_policy=str(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_ACTION_POLICY", "auto")
|
||||
).strip(),
|
||||
max_retries=int(config_dict.get("NILOR_MEMORY_HYGIENE_MAX_RETRIES", 2)),
|
||||
cooldown_seconds=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_COOLDOWN_SECONDS", 60)
|
||||
),
|
||||
sleep_between_attempts_seconds=int(
|
||||
config_dict.get(
|
||||
"NILOR_MEMORY_HYGIENE_SLEEP_BETWEEN_ATTEMPTS_SECONDS", 4
|
||||
)
|
||||
),
|
||||
max_cycle_duration_seconds=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_MAX_CYCLE_DURATION_SECONDS", 15)
|
||||
),
|
||||
)
|
||||
_validate_hygiene_config(hygiene_cfg)
|
||||
|
||||
cfg = cls(
|
||||
comfy=comfy_cfg,
|
||||
worker=worker_cfg,
|
||||
allow_env_override=allow_env_override,
|
||||
sqs_enabled=sqs_enabled,
|
||||
hygiene=hygiene_cfg,
|
||||
)
|
||||
_validate_comfy_config(cfg.comfy)
|
||||
_validate_worker_config(cfg.worker)
|
||||
return cfg
|
||||
|
||||
|
||||
# ---- Internal helpers ----
|
||||
|
||||
|
||||
def _apply_env_overrides(cfg: NilorNodesConfig) -> None:
|
||||
if not cfg.allow_env_override:
|
||||
return
|
||||
|
||||
# Feature flags
|
||||
sqs_enabled = os.getenv("NILOR_SQS_ENABLED", cfg.sqs_enabled)
|
||||
cfg.sqs_enabled = _coerce_bool(sqs_enabled)
|
||||
|
||||
# Comfy (rebuild frozen dataclass)
|
||||
comfy_api_url = os.getenv("NILOR_COMFYUI_API_URL", cfg.comfy.api_url)
|
||||
comfy_ws_url = os.getenv("NILOR_COMFYUI_WS_URL", cfg.comfy.ws_url)
|
||||
comfy_timeout_s = int(
|
||||
os.getenv("NILOR_COMFY_API_TIMEOUT_SECONDS", cfg.comfy.timeout_s)
|
||||
)
|
||||
comfy_retry_base = float(
|
||||
os.getenv("NILOR_COMFY_RETRY_BASE_SECONDS", cfg.comfy.retry_base_seconds)
|
||||
)
|
||||
comfy_retry_multiplier = float(
|
||||
os.getenv("NILOR_COMFY_RETRY_MULTIPLIER", cfg.comfy.retry_multiplier)
|
||||
)
|
||||
comfy_retry_jitter = float(
|
||||
os.getenv("NILOR_COMFY_RETRY_JITTER_SECONDS", cfg.comfy.retry_jitter_seconds)
|
||||
)
|
||||
comfy_retry_max_sleep = float(
|
||||
os.getenv(
|
||||
"NILOR_COMFY_RETRY_MAX_SLEEP_SECONDS", cfg.comfy.retry_max_sleep_seconds
|
||||
)
|
||||
)
|
||||
comfy_retry_max_attempts = int(
|
||||
os.getenv("NILOR_COMFY_RETRY_MAX_ATTEMPTS", cfg.comfy.retry_max_attempts)
|
||||
)
|
||||
comfy_ws_max_reconnect = int(
|
||||
os.getenv(
|
||||
"NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS",
|
||||
cfg.comfy.ws_max_reconnect_attempts,
|
||||
)
|
||||
)
|
||||
comfy_ws_max_total_backoff = float(
|
||||
os.getenv(
|
||||
"NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS",
|
||||
cfg.comfy.ws_max_total_backoff_seconds,
|
||||
)
|
||||
)
|
||||
cfg.comfy = ComfyApiConfig(
|
||||
api_url=str(comfy_api_url),
|
||||
ws_url=str(comfy_ws_url),
|
||||
timeout_s=comfy_timeout_s,
|
||||
retry_base_seconds=comfy_retry_base,
|
||||
retry_multiplier=comfy_retry_multiplier,
|
||||
retry_jitter_seconds=comfy_retry_jitter,
|
||||
retry_max_sleep_seconds=comfy_retry_max_sleep,
|
||||
retry_max_attempts=comfy_retry_max_attempts,
|
||||
ws_max_reconnect_attempts=comfy_ws_max_reconnect,
|
||||
ws_max_total_backoff_seconds=comfy_ws_max_total_backoff,
|
||||
)
|
||||
|
||||
# Worker (rebuild frozen dataclass)
|
||||
worker_sqs_endpoint_url = os.getenv(
|
||||
"NILOR_SQS_ENDPOINT_URL", cfg.worker.sqs_endpoint_url
|
||||
)
|
||||
worker_jobs_queue = os.getenv(
|
||||
"NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME", cfg.worker.jobs_queue
|
||||
)
|
||||
worker_status_queue = os.getenv(
|
||||
"NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME", cfg.worker.status_queue
|
||||
)
|
||||
worker_poll_wait_s = int(
|
||||
os.getenv("NILOR_SQS_POLL_WAIT_TIME", cfg.worker.poll_wait_s)
|
||||
)
|
||||
worker_max_messages = int(
|
||||
os.getenv("NILOR_SQS_MAX_MESSAGES", cfg.worker.max_messages)
|
||||
)
|
||||
worker_access_key = os.getenv(
|
||||
"NILOR_AWS_ACCESS_KEY_ID", cfg.worker.aws_access_key_id
|
||||
)
|
||||
worker_secret_key = os.getenv(
|
||||
"NILOR_AWS_SECRET_ACCESS_KEY", cfg.worker.aws_secret_access_key
|
||||
)
|
||||
worker_region = os.getenv("NILOR_AWS_DEFAULT_REGION", cfg.worker.aws_region)
|
||||
worker_client_id = os.getenv("NILOR_WORKER_CLIENT_ID", cfg.worker.worker_client_id)
|
||||
worker_workflow_os_norm = os.getenv(
|
||||
"NILOR_WORKFLOW_OS_NORMALIZATION_ENABLED",
|
||||
cfg.worker.workflow_os_normalization_enabled,
|
||||
)
|
||||
cfg.worker = WorkerConfig(
|
||||
sqs_endpoint_url=str(worker_sqs_endpoint_url),
|
||||
jobs_queue=str(worker_jobs_queue),
|
||||
status_queue=str(worker_status_queue),
|
||||
poll_wait_s=worker_poll_wait_s,
|
||||
max_messages=worker_max_messages,
|
||||
aws_access_key_id=str(worker_access_key),
|
||||
aws_secret_access_key=str(worker_secret_key),
|
||||
aws_region=str(worker_region),
|
||||
worker_client_id=str(worker_client_id),
|
||||
workflow_os_normalization_enabled=_coerce_bool(worker_workflow_os_norm),
|
||||
)
|
||||
|
||||
# Re-validate after overrides
|
||||
_validate_comfy_config(cfg.comfy)
|
||||
_validate_worker_config(cfg.worker)
|
||||
# Hygiene (rebuild frozen dataclass)
|
||||
hygiene_enabled = _coerce_bool(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_ENABLED", cfg.hygiene.enabled)
|
||||
)
|
||||
hygiene_idle_poll = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS", cfg.hygiene.idle_poll_seconds
|
||||
)
|
||||
)
|
||||
hygiene_vram_pct = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX", cfg.hygiene.vram_usage_pct_max
|
||||
)
|
||||
)
|
||||
hygiene_ram_pct = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX", cfg.hygiene.ram_usage_pct_max
|
||||
)
|
||||
)
|
||||
hygiene_vram_min = int(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_VRAM_MIN_FREE_MB", cfg.hygiene.vram_min_free_mb)
|
||||
)
|
||||
hygiene_ram_min = int(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_RAM_MIN_FREE_MB", cfg.hygiene.ram_min_free_mb)
|
||||
)
|
||||
hygiene_policy = str(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_ACTION_POLICY", cfg.hygiene.action_policy)
|
||||
).strip()
|
||||
hygiene_max_retries = int(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_MAX_RETRIES", cfg.hygiene.max_retries)
|
||||
)
|
||||
hygiene_cooldown = int(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_COOLDOWN_SECONDS", cfg.hygiene.cooldown_seconds)
|
||||
)
|
||||
hygiene_sleep_between = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_SLEEP_BETWEEN_ATTEMPTS_SECONDS",
|
||||
cfg.hygiene.sleep_between_attempts_seconds,
|
||||
)
|
||||
)
|
||||
hygiene_max_cycle = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_MAX_CYCLE_DURATION_SECONDS",
|
||||
cfg.hygiene.max_cycle_duration_seconds,
|
||||
)
|
||||
)
|
||||
cfg.hygiene = MemoryHygieneConfig(
|
||||
enabled=hygiene_enabled,
|
||||
idle_poll_seconds=hygiene_idle_poll,
|
||||
vram_usage_pct_max=hygiene_vram_pct,
|
||||
ram_usage_pct_max=hygiene_ram_pct,
|
||||
vram_min_free_mb=hygiene_vram_min,
|
||||
ram_min_free_mb=hygiene_ram_min,
|
||||
action_policy=hygiene_policy,
|
||||
max_retries=hygiene_max_retries,
|
||||
cooldown_seconds=hygiene_cooldown,
|
||||
sleep_between_attempts_seconds=hygiene_sleep_between,
|
||||
max_cycle_duration_seconds=hygiene_max_cycle,
|
||||
)
|
||||
_validate_hygiene_config(cfg.hygiene)
|
||||
|
||||
|
||||
def _validate_comfy_config(cfg: ComfyApiConfig) -> None:
|
||||
_require_url_scheme(cfg.api_url, {"http", "https"}, "NILOR_COMFYUI_API_URL")
|
||||
_require_url_scheme(cfg.ws_url, {"ws", "wss"}, "NILOR_COMFYUI_WS_URL")
|
||||
if cfg.timeout_s <= 0:
|
||||
raise ValueError(
|
||||
f"NILOR_COMFY_API_TIMEOUT_SECONDS must be a positive integer; got {cfg.timeout_s}"
|
||||
)
|
||||
if (
|
||||
cfg.retry_base_seconds < 0
|
||||
or cfg.retry_multiplier <= 0
|
||||
or cfg.retry_max_sleep_seconds <= 0
|
||||
):
|
||||
raise ValueError("Invalid retry backoff parameters in Comfy client config")
|
||||
if cfg.retry_max_attempts <= 0:
|
||||
raise ValueError("NILOR_COMFY_RETRY_MAX_ATTEMPTS must be a positive integer")
|
||||
if cfg.ws_max_reconnect_attempts < 0 or cfg.ws_max_total_backoff_seconds < 0:
|
||||
raise ValueError(
|
||||
"Invalid websocket reconnect parameters in Comfy client config"
|
||||
)
|
||||
|
||||
|
||||
def _validate_worker_config(cfg: WorkerConfig) -> None:
|
||||
if cfg.poll_wait_s < 0 or cfg.poll_wait_s > 20:
|
||||
raise ValueError(
|
||||
f"NILOR_SQS_POLL_WAIT_TIME must be within [0, 20]; got {cfg.poll_wait_s}"
|
||||
)
|
||||
if cfg.max_messages <= 0:
|
||||
raise ValueError(
|
||||
f"NILOR_SQS_MAX_MESSAGES must be a positive integer; got {cfg.max_messages}"
|
||||
)
|
||||
|
||||
|
||||
def _validate_hygiene_config(cfg: MemoryHygieneConfig) -> None:
|
||||
if cfg.idle_poll_seconds < 0:
|
||||
raise ValueError("NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS must be >= 0")
|
||||
if not 0 <= cfg.vram_usage_pct_max <= 100:
|
||||
raise ValueError(
|
||||
"NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX must be within [0, 100]"
|
||||
)
|
||||
if not 0 <= cfg.ram_usage_pct_max <= 100:
|
||||
raise ValueError(
|
||||
"NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX must be within [0, 100]"
|
||||
)
|
||||
if cfg.vram_min_free_mb < 0 or cfg.ram_min_free_mb < 0:
|
||||
raise ValueError("Memory hygiene min free MB must be >= 0")
|
||||
if cfg.max_retries < 0:
|
||||
raise ValueError("NILOR_MEMORY_HYGIENE_MAX_RETRIES must be >= 0")
|
||||
if cfg.cooldown_seconds < 0 or cfg.sleep_between_attempts_seconds < 0:
|
||||
raise ValueError("Memory hygiene cooldown/sleep must be >= 0")
|
||||
if cfg.max_cycle_duration_seconds < 0:
|
||||
raise ValueError("Memory hygiene max cycle duration must be >= 0")
|
||||
allowed = {"free", "unload", "both", "auto"}
|
||||
if cfg.action_policy not in allowed:
|
||||
allowed_str = ", ".join(sorted(allowed))
|
||||
raise ValueError(
|
||||
f"NILOR_MEMORY_HYGIENE_ACTION_POLICY must be one of {{{allowed_str}}}"
|
||||
)
|
||||
|
||||
|
||||
def _coerce_bool(value: object) -> bool:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if value is None:
|
||||
return False
|
||||
text = str(value).strip().lower()
|
||||
if text in {"1", "true", "yes", "on"}:
|
||||
return True
|
||||
if text in {"0", "false", "no", "off"}:
|
||||
return False
|
||||
return bool(text)
|
||||
|
||||
|
||||
def _require_url_scheme(url: str, allowed: set[str], key_name: str) -> None:
|
||||
parsed = urlparse(url)
|
||||
if not parsed.scheme or parsed.scheme.lower() not in allowed:
|
||||
allowed_str = ", ".join(sorted(allowed))
|
||||
raise ValueError(
|
||||
f"{key_name} must start with one of [{allowed_str}]; got: {url!r}"
|
||||
)
|
||||
|
||||
|
||||
def _generate_worker_client_id() -> str:
|
||||
host = socket.gethostname().strip() or "worker"
|
||||
suffix = _random_base36_suffix(5)
|
||||
return f"nilor-worker-{host}-{suffix}"
|
||||
|
||||
|
||||
def _random_base36_suffix(length: int = 5) -> str:
|
||||
import random
|
||||
|
||||
n = random.getrandbits(32)
|
||||
base36 = _to_base36(n)
|
||||
return base36[-length:]
|
||||
|
||||
|
||||
def _to_base36(n: int) -> str:
|
||||
if n == 0:
|
||||
return "0"
|
||||
digits = "0123456789abcdefghijklmnopqrstuvwxyz"
|
||||
sign = "-" if n < 0 else ""
|
||||
n = abs(n)
|
||||
res = []
|
||||
while n:
|
||||
n, r = divmod(n, 36)
|
||||
res.append(digits[r])
|
||||
return sign + "".join(reversed(res))
|
||||
|
||||
|
||||
def load_nilor_nodes_config() -> NilorNodesConfig:
|
||||
"""Convenience loader aligning with brain_rnd config pattern.
|
||||
|
||||
- Loads environment variables via python-dotenv if available
|
||||
- Reads JSON5 defaults and applies env overrides when enabled
|
||||
- Returns a typed `NilorNodesConfig`
|
||||
"""
|
||||
global _CONFIG
|
||||
if _CONFIG is not None:
|
||||
return _CONFIG
|
||||
try:
|
||||
from dotenv import load_dotenv # optional dependency present in sidecar
|
||||
|
||||
try:
|
||||
# Load .env next to nilor-nodes (ComfyUI/custom_nodes/nilor-nodes/.env) first
|
||||
here = Path(__file__).resolve().parent.parent # .../nilor-nodes/
|
||||
dotenv_path = here / ".env"
|
||||
loaded = load_dotenv(dotenv_path=dotenv_path)
|
||||
try:
|
||||
if loaded:
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes: Loaded environment variables from {dotenv_path}"
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"⚠️\u2009 Nilor-Nodes: No .env file found, relying on shell environment variables."
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
# Also allow default search (repo root / current working dir)
|
||||
load_dotenv()
|
||||
except Exception:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
# Use BaseConfig-backed singleton load
|
||||
if json5 is None:
|
||||
raise RuntimeError(
|
||||
"json5 module is required to load configuration from JSON5 file"
|
||||
)
|
||||
cfg = NilorNodesConfig.get_instance() # type: ignore[attr-defined]
|
||||
_apply_env_overrides(cfg)
|
||||
_CONFIG = cfg
|
||||
return _CONFIG
|
||||
|
||||
|
||||
def refresh_nilor_nodes_config() -> NilorNodesConfig:
|
||||
"""Clear the cached config and reload it. Intended for explicit hot-reload."""
|
||||
global _CONFIG
|
||||
_CONFIG = None
|
||||
return load_nilor_nodes_config()
|
||||
@@ -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",
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
import logging
|
||||
import os
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def configure_from_env(
|
||||
primary_env_var: str = "NILOR_LOG_LEVEL", fallback_env_var: str = "LOG_LEVEL"
|
||||
) -> None:
|
||||
# Prefer NILOR_LOG_LEVEL; fall back to LOG_LEVEL for backward compatibility
|
||||
chosen_var = primary_env_var if os.getenv(primary_env_var) else fallback_env_var
|
||||
value = os.getenv(chosen_var)
|
||||
|
||||
default_level = logging.INFO
|
||||
level = getattr(logging, value.upper(), None) if value else default_level
|
||||
if not isinstance(level, int):
|
||||
level = default_level
|
||||
logger.setLevel(level)
|
||||
|
||||
# Announce the effective level to the terminal via the global handlers
|
||||
effective_name = logging.getLevelName(level)
|
||||
if value:
|
||||
if getattr(logging, value.upper(), None) is None:
|
||||
logging.warning(
|
||||
f"⚠️ Nilor-Nodes: {chosen_var}='{value}' is invalid; defaulting to {effective_name}"
|
||||
)
|
||||
else:
|
||||
logging.info(
|
||||
f"ℹ️ Nilor-Nodes: {chosen_var}='{value}' → level set to {effective_name}"
|
||||
)
|
||||
else:
|
||||
logging.info(
|
||||
f"ℹ️ Nilor-Nodes: {primary_env_var} or {fallback_env_var} not set; defaulting to {effective_name}"
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["logger", "configure_from_env"]
|
||||
+430
@@ -0,0 +1,430 @@
|
||||
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 json
|
||||
|
||||
import tempfile
|
||||
import os
|
||||
from .logger import logger
|
||||
from .config.config import load_nilor_nodes_config
|
||||
|
||||
# Load shared configuration once
|
||||
_CFG = load_nilor_nodes_config()
|
||||
|
||||
# --- 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"],),
|
||||
"presigned_download_url": (
|
||||
"STRING",
|
||||
{"multiline": True, "default": "<auto-filled by system>"},
|
||||
),
|
||||
},
|
||||
"hidden": {},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
FUNCTION = "download"
|
||||
CATEGORY = category + subcategories["streaming"]
|
||||
|
||||
def download(
|
||||
self,
|
||||
presigned_download_url: str,
|
||||
format: str,
|
||||
input_name: str = "default_input",
|
||||
):
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes: MediaStreamInput: Downloading from {presigned_download_url} for input '{input_name}' with format '{format}'"
|
||||
)
|
||||
try:
|
||||
# Two-phase download for batches: manifest first, then assets
|
||||
if format == "image_batch":
|
||||
manifest_response = requests.get(presigned_download_url, timeout=60)
|
||||
manifest_response.raise_for_status()
|
||||
manifest = manifest_response.json()
|
||||
|
||||
logger.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 in parallel
|
||||
asset_responses = []
|
||||
for file_info in sorted_files:
|
||||
try:
|
||||
resp = requests.get(file_info["presigned_url"], timeout=180)
|
||||
resp.raise_for_status()
|
||||
asset_responses.append(resp.content)
|
||||
except requests.RequestException as e:
|
||||
logger.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 ---
|
||||
if format == "video":
|
||||
# Stream video to temp file to avoid loading entire video into RAM
|
||||
temp_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
|
||||
try:
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Streaming video to temp file: {temp_file.name}")
|
||||
with requests.get(presigned_download_url, timeout=180, stream=True) as response:
|
||||
response.raise_for_status()
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
temp_file.write(chunk)
|
||||
temp_file.close()
|
||||
return self._process_video(temp_file.name)
|
||||
finally:
|
||||
# Clean up temp file
|
||||
if os.path.exists(temp_file.name):
|
||||
os.unlink(temp_file.name)
|
||||
else:
|
||||
# For images, load into memory (they're small)
|
||||
response = requests.get(presigned_download_url, timeout=180)
|
||||
response.raise_for_status()
|
||||
media_bytes = response.content
|
||||
|
||||
if format == "image":
|
||||
return self._process_image(media_bytes)
|
||||
else:
|
||||
# Should not happen if UI choices are respected
|
||||
raise ValueError(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Unsupported format '{format}' for single media download."
|
||||
)
|
||||
|
||||
except requests.RequestException as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Failed to download file: {e}"
|
||||
)
|
||||
return (None,)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Failed to process media: {e}"
|
||||
)
|
||||
return (None,)
|
||||
|
||||
def _process_image_batch(self, image_bytes_list):
|
||||
logger.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)
|
||||
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamInput): Image batch processing successful. Batch shape: {images_tensor.shape}"
|
||||
)
|
||||
return (images_tensor,)
|
||||
|
||||
def _process_image(self, image_bytes):
|
||||
logger.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)
|
||||
|
||||
logger.info("✅ Nilor-Nodes (MediaStreamInput): Image processing successful.")
|
||||
return (image_tensor,)
|
||||
|
||||
def _process_video(self, video_path):
|
||||
logger.info(f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing video from {video_path}...")
|
||||
|
||||
# Open video to get metadata first
|
||||
with imageio.get_reader(video_path, format="mp4") as reader:
|
||||
# Get video metadata
|
||||
metadata = reader.get_meta_data()
|
||||
num_frames = reader.count_frames()
|
||||
|
||||
if num_frames == 0:
|
||||
raise ValueError(
|
||||
"🛑\u2009 Nilor-Nodes (MediaStreamInput): No frames could be read from the video."
|
||||
)
|
||||
|
||||
# Read first frame to get dimensions
|
||||
first_frame = reader.get_data(0)
|
||||
height, width = first_frame.shape[:2]
|
||||
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Video has {num_frames} frames at {width}x{height}"
|
||||
)
|
||||
|
||||
# Pre-allocate tensor for all frames (N, H, W, 3)
|
||||
video_tensor = torch.empty((num_frames, height, width, 3), dtype=torch.float32)
|
||||
|
||||
# Process first frame (already read for dimensions)
|
||||
pil_image = Image.fromarray(first_frame).convert("RGB")
|
||||
numpy_image = np.array(pil_image).astype(np.float32) / 255.0
|
||||
video_tensor[0] = torch.from_numpy(numpy_image)
|
||||
|
||||
# Read remaining frames by explicit index to avoid iterator position ambiguity
|
||||
for i in range(1, num_frames):
|
||||
frame = reader.get_data(i)
|
||||
pil_image = Image.fromarray(frame).convert("RGB")
|
||||
numpy_image = np.array(pil_image).astype(np.float32) / 255.0
|
||||
video_tensor[i] = torch.from_numpy(numpy_image)
|
||||
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamInput): Video processing successful. Tensor 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},
|
||||
),
|
||||
"presigned_upload_url": (
|
||||
"STRING",
|
||||
{"multiline": True, "default": "<auto-filled by system>"},
|
||||
),
|
||||
"job_completions_queue_url": (
|
||||
"STRING",
|
||||
{"multiline": True, "default": "<auto-filled by system>"},
|
||||
),
|
||||
"output_object_keys": (
|
||||
"STRING",
|
||||
{"multiline": False, "default": "<auto-filled by system>"},
|
||||
),
|
||||
"job_type": (
|
||||
"STRING",
|
||||
{"default": "<auto-filled by system>", "multiline": False},
|
||||
),
|
||||
},
|
||||
"hidden": {
|
||||
"prompt": "PROMPT",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO",
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("uploaded_url",)
|
||||
FUNCTION = "upload_and_notify"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = category + subcategories["streaming"]
|
||||
|
||||
def upload_and_notify(
|
||||
self,
|
||||
images,
|
||||
format,
|
||||
content_id,
|
||||
venue,
|
||||
canvas,
|
||||
scene,
|
||||
presigned_upload_url,
|
||||
job_completions_queue_url,
|
||||
output_object_keys,
|
||||
framerate,
|
||||
output_name: str = "default_output",
|
||||
prompt=None,
|
||||
extra_pnginfo=None,
|
||||
job_type: str | None = None,
|
||||
):
|
||||
if not content_id:
|
||||
raise ValueError(
|
||||
"🛑\u2009 Nilor-Nodes (MediaStreamOutput): content_id is a required input for MediaStreamOutput."
|
||||
)
|
||||
|
||||
# The `output_object_keys` is received as a string representation of a dictionary.
|
||||
# We must parse it back into a dictionary.
|
||||
final_outputs_dict = {}
|
||||
try:
|
||||
# The string may use single quotes, so we replace them for valid JSON.
|
||||
final_outputs_dict = json.loads(output_object_keys.replace("'", '"'))
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): FATAL -- Could not parse output_object_keys from string: {output_object_keys}. Error: {e}"
|
||||
)
|
||||
final_outputs_dict = {} # Send empty dict on failure.
|
||||
|
||||
# The presigned_upload_url provided to this node is specific to its output_name.
|
||||
# We don't need to re-select it. We just need to perform the upload.
|
||||
if format == "png":
|
||||
self._upload_image(images[0], presigned_upload_url)
|
||||
elif format == "mp4":
|
||||
self._upload_video(images, presigned_upload_url, framerate)
|
||||
|
||||
# This node is responsible for a single output. We find its corresponding object key.
|
||||
output_key_for_this_node = final_outputs_dict.get(output_name)
|
||||
if not output_key_for_this_node:
|
||||
logging.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): FATAL -- Could not find object key for output name '{output_name}' in output_object_keys."
|
||||
)
|
||||
# Send an empty dictionary to signal failure.
|
||||
final_outputs_for_sqs = {}
|
||||
else:
|
||||
final_outputs_for_sqs = {output_name: output_key_for_this_node}
|
||||
|
||||
# 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,
|
||||
}
|
||||
if job_type:
|
||||
completion_message["job_type"] = job_type
|
||||
|
||||
try:
|
||||
# Re-initialize the client inside the execution to ensure it picks up env vars correctly.
|
||||
sqs_client = boto3.client(
|
||||
"sqs",
|
||||
endpoint_url=_CFG.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=_CFG.worker.aws_access_key_id,
|
||||
aws_secret_access_key=_CFG.worker.aws_secret_access_key,
|
||||
region_name=_CFG.worker.aws_region,
|
||||
)
|
||||
logger.debug(
|
||||
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),
|
||||
)
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamOutput): Completion message sent successfully for content {content_id} to queue: {job_completions_queue_url}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.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": []}, "result": (presigned_upload_url,)}
|
||||
|
||||
def _upload_image(self, image_tensor, url):
|
||||
logger.debug(
|
||||
"ℹ️\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)
|
||||
|
||||
self._perform_upload(buffer, url, "image/png")
|
||||
|
||||
def _upload_video(self, image_batch_tensor, url, framerate):
|
||||
logger.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)
|
||||
|
||||
self._perform_upload(buffer, url, "video/mp4")
|
||||
|
||||
def _perform_upload(self, buffer, url, content_type):
|
||||
try:
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Uploading to {url} with Content-Type: {content_type}"
|
||||
)
|
||||
headers = {"Content-Type": content_type}
|
||||
response = requests.put(
|
||||
url, data=buffer.read(), headers=headers, timeout=300
|
||||
)
|
||||
response.raise_for_status()
|
||||
logger.info("✅ Nilor-Nodes (MediaStreamOutput): Upload successful.")
|
||||
except requests.RequestException as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): Failed to upload media: {e}"
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): Failed to process and upload media: {e}"
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
# --- Node Mappings ---
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"MediaStreamInput": MediaStreamInput,
|
||||
"MediaStreamOutput": MediaStreamOutput,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"MediaStreamInput": "👺 Media Stream Input (URL)",
|
||||
"MediaStreamOutput": "👺 Media Stream Output (URL)",
|
||||
}
|
||||
@@ -0,0 +1,432 @@
|
||||
"""
|
||||
Memory Hygiene (Guardian) — typed skeleton and public API.
|
||||
|
||||
This module provides a lightweight, dependency-injected component that can be
|
||||
invoked between jobs to assess memory pressure and (optionally) remediate via
|
||||
the ComfyUI server's `/free` endpoint. This commit introduces the types and
|
||||
public API only; detailed policy and remediation logic will be implemented in
|
||||
subsequent commits.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import asyncio
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Literal, Any
|
||||
|
||||
from .config.config import MemoryHygieneConfig
|
||||
from .comfyui_client import ComfyUIClientProtocol, SystemStats
|
||||
|
||||
|
||||
RemediationAction = Literal["none", "free", "unload", "both", "auto"]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RemediationResult:
|
||||
"""Outcome of a hygiene check/remediation cycle.
|
||||
|
||||
Attributes:
|
||||
before: Stats captured before any remediation attempt.
|
||||
after: Stats captured after remediation (if attempted), or None.
|
||||
action: Action that was selected/executed for this cycle.
|
||||
attempts: Number of remediation attempts performed.
|
||||
elapsed_seconds: Wall-clock duration of the cycle in seconds.
|
||||
reason: Optional decision rationale (e.g., threshold that triggered).
|
||||
success: True when targets were met or no action was needed.
|
||||
"""
|
||||
|
||||
before: SystemStats
|
||||
after: Optional[SystemStats]
|
||||
action: RemediationAction
|
||||
attempts: int
|
||||
elapsed_seconds: float
|
||||
reason: Optional[str]
|
||||
success: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DerivedStats:
|
||||
"""Computed metrics used by the policy engine.
|
||||
|
||||
Attributes:
|
||||
vram_total: Total VRAM (bytes) if known.
|
||||
vram_free: Free VRAM (bytes) if known.
|
||||
vram_used_pct: Percent VRAM used in [0, 100] when computable.
|
||||
ram_total: Total RAM (bytes) if known.
|
||||
ram_free: Free RAM (bytes) if known.
|
||||
ram_used_pct: Percent RAM used in [0, 100] when computable.
|
||||
"""
|
||||
|
||||
vram_total: Optional[float]
|
||||
vram_free: Optional[float]
|
||||
vram_used_pct: Optional[float]
|
||||
ram_total: Optional[float]
|
||||
ram_free: Optional[float]
|
||||
ram_used_pct: Optional[float]
|
||||
|
||||
|
||||
class MemoryHygiene:
|
||||
"""Memory Guardian component orchestrating checks and remediation between jobs."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
client: ComfyUIClientProtocol,
|
||||
cfg: MemoryHygieneConfig,
|
||||
logger: Optional[Any] = None,
|
||||
) -> None:
|
||||
self._client = client
|
||||
self._cfg = cfg
|
||||
self._logger = logger
|
||||
self._cooldown_until: float = 0.0
|
||||
# Hardened capability state: disable guardian for session after unsupported
|
||||
self._capability_disabled: bool = False
|
||||
self._unsupported_warned: bool = False
|
||||
|
||||
async def check_and_remediate(self) -> RemediationResult:
|
||||
"""Run a single hygiene cycle.
|
||||
|
||||
This skeleton performs capability and enablement checks, then captures a
|
||||
baseline stats snapshot and returns without modification. Subsequent
|
||||
commits implement policy evaluation and remediation.
|
||||
"""
|
||||
start_ts = time.monotonic()
|
||||
|
||||
if not self._cfg.enabled:
|
||||
before, _ = await self._collect_metrics()
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.debug(
|
||||
"⚠️\u2009 Nilor-Nodes (memory_hygiene): disabled; skipping remediation. vram_free=%s ram_free=%s",
|
||||
getattr(before, "vram_free", None),
|
||||
getattr(before, "ram_free", None),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason="disabled",
|
||||
success=True,
|
||||
)
|
||||
|
||||
# Session-wide disable if previously detected unsupported endpoints
|
||||
if self._capability_disabled:
|
||||
return RemediationResult(
|
||||
before=await self._get_stats_safe(),
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason="capability_unsupported",
|
||||
success=True,
|
||||
)
|
||||
|
||||
try:
|
||||
supported = await self._client.supports_hygiene()
|
||||
except Exception:
|
||||
supported = False
|
||||
|
||||
before, derived = await self._collect_metrics()
|
||||
|
||||
if not supported:
|
||||
# Disable for the rest of the session and emit a single warning
|
||||
self._capability_disabled = True
|
||||
if not self._unsupported_warned:
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.warning(
|
||||
"⚠️\u2009 Nilor-Nodes (memory_hygiene): unsupported endpoints; disabling for this session."
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
self._unsupported_warned = True
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason="capability_unsupported",
|
||||
success=True,
|
||||
)
|
||||
|
||||
now = time.monotonic()
|
||||
if now < self._cooldown_until:
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (memory_hygiene): cooldown active for %.2fs; skipping.",
|
||||
max(0.0, self._cooldown_until - now),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, now - start_ts),
|
||||
reason="cooldown_active",
|
||||
success=True,
|
||||
)
|
||||
|
||||
pressure_reason = self._pressure_reason(derived)
|
||||
if pressure_reason is None:
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (memory_hygiene): no pressure; nothing to do."
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason="no_pressure",
|
||||
success=True,
|
||||
)
|
||||
|
||||
action = self._choose_initial_action(self._cfg.action_policy)
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.info(
|
||||
"✅ Nilor-Nodes (memory_hygiene): start cycle action=%s reason=%s vram_free=%s ram_free=%s",
|
||||
action,
|
||||
pressure_reason,
|
||||
getattr(before, "vram_free", None),
|
||||
getattr(before, "ram_free", None),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
after, attempts, final_action, outcome_reason, success = (
|
||||
await self._remediate_cycle(
|
||||
initial_action=action,
|
||||
initial_reason=pressure_reason,
|
||||
cycle_start=start_ts,
|
||||
)
|
||||
)
|
||||
# Set cooldown after any attempted cycle (success or not)
|
||||
self._cooldown_until = time.monotonic() + float(self._cfg.cooldown_seconds)
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.info(
|
||||
"✅ Nilor-Nodes (memory_hygiene): end cycle action=%s attempts=%s success=%s reason=%s vram_free_before=%s vram_free_after=%s ram_free_before=%s ram_free_after=%s elapsed=%.2fs",
|
||||
final_action,
|
||||
attempts,
|
||||
success,
|
||||
outcome_reason,
|
||||
getattr(before, "vram_free", None),
|
||||
getattr(after, "vram_free", None) if after else None,
|
||||
getattr(before, "ram_free", None),
|
||||
getattr(after, "ram_free", None) if after else None,
|
||||
max(0.0, time.monotonic() - start_ts),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=after,
|
||||
action=final_action,
|
||||
attempts=attempts,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason=outcome_reason,
|
||||
success=success,
|
||||
)
|
||||
|
||||
async def _get_stats_safe(self) -> SystemStats:
|
||||
try:
|
||||
return await self._client.get_system_stats()
|
||||
except Exception:
|
||||
# Return an empty struct; callers tolerate partial data
|
||||
return SystemStats()
|
||||
|
||||
async def _collect_metrics(self) -> tuple[SystemStats, DerivedStats]:
|
||||
base = await self._get_stats_safe()
|
||||
vram_used_pct = _compute_used_pct(base.vram_total, base.vram_free)
|
||||
ram_used_pct = _compute_used_pct(base.ram_total, base.ram_free)
|
||||
derived = DerivedStats(
|
||||
vram_total=base.vram_total,
|
||||
vram_free=base.vram_free,
|
||||
vram_used_pct=vram_used_pct,
|
||||
ram_total=base.ram_total,
|
||||
ram_free=base.ram_free,
|
||||
ram_used_pct=ram_used_pct,
|
||||
)
|
||||
return base, derived
|
||||
|
||||
def _pressure_reason(self, d: DerivedStats) -> Optional[str]:
|
||||
if _gt_pct(d.vram_used_pct, self._cfg.vram_usage_pct_max):
|
||||
return f"vram_used_pct {d.vram_used_pct}% > {self._cfg.vram_usage_pct_max}%"
|
||||
if _lt_bytes(d.vram_free, self._cfg.vram_min_free_mb):
|
||||
return f"vram_free below {self._cfg.vram_min_free_mb}MB"
|
||||
if _gt_pct(d.ram_used_pct, self._cfg.ram_usage_pct_max):
|
||||
return f"ram_used_pct {d.ram_used_pct}% > {self._cfg.ram_usage_pct_max}%"
|
||||
if _lt_bytes(d.ram_free, self._cfg.ram_min_free_mb):
|
||||
return f"ram_free below {self._cfg.ram_min_free_mb}MB"
|
||||
return None
|
||||
|
||||
def _choose_initial_action(self, policy: str) -> RemediationAction:
|
||||
return _initial_action_for(policy)
|
||||
|
||||
async def _remediate_cycle(
|
||||
self,
|
||||
*,
|
||||
initial_action: RemediationAction,
|
||||
initial_reason: str,
|
||||
cycle_start: float,
|
||||
) -> tuple[Optional[SystemStats], int, RemediationAction, str, bool]:
|
||||
attempts = 0
|
||||
action = initial_action
|
||||
escalated = False
|
||||
last_stats: Optional[SystemStats] = None
|
||||
|
||||
max_retries = max(0, int(self._cfg.max_retries))
|
||||
sleep_between = max(0, int(self._cfg.sleep_between_attempts_seconds))
|
||||
max_cycle_s = max(0, int(self._cfg.max_cycle_duration_seconds))
|
||||
|
||||
def time_budget_exhausted() -> bool:
|
||||
if max_cycle_s <= 0:
|
||||
return False
|
||||
return (time.monotonic() - cycle_start) >= max_cycle_s
|
||||
|
||||
outcome_reason = initial_reason
|
||||
|
||||
while True:
|
||||
# Execute remediation step
|
||||
free_flag, unload_flag = _flags_for_action(action)
|
||||
try:
|
||||
if self._logger:
|
||||
try:
|
||||
self._logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (memory_hygiene): calling /free free_memory=%s unload_models=%s",
|
||||
free_flag,
|
||||
unload_flag,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
await self._client.free(
|
||||
free_memory=free_flag, unload_models=unload_flag
|
||||
)
|
||||
except Exception:
|
||||
# Continue even on errors; treat as unsuccessful attempt
|
||||
pass
|
||||
|
||||
# Wait and re-measure
|
||||
if sleep_between > 0:
|
||||
try:
|
||||
await asyncio.sleep(sleep_between)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
last, derived = await self._collect_metrics()
|
||||
last_stats = last
|
||||
if self._pressure_reason(derived) is None:
|
||||
outcome_reason = "targets_met"
|
||||
return last_stats, attempts + 1, action, outcome_reason, True
|
||||
|
||||
attempts += 1
|
||||
if attempts > max_retries:
|
||||
outcome_reason = "max_retries_exhausted"
|
||||
return last_stats, attempts, action, outcome_reason, False
|
||||
|
||||
if time_budget_exhausted():
|
||||
outcome_reason = "max_duration_reached"
|
||||
return last_stats, attempts, action, outcome_reason, False
|
||||
|
||||
# Escalation logic for auto: free -> unload (once)
|
||||
if initial_action == "auto" and not escalated:
|
||||
action = "unload"
|
||||
escalated = True
|
||||
# For explicit free/unload/both, keep same action for subsequent attempts
|
||||
|
||||
|
||||
def _compute_used_pct(total: Optional[float], free: Optional[float]) -> Optional[float]:
|
||||
try:
|
||||
if total is None or free is None:
|
||||
return None
|
||||
total_f = float(total)
|
||||
free_f = float(free)
|
||||
if total_f <= 0:
|
||||
return None
|
||||
used = max(0.0, min(1.0, (total_f - max(0.0, free_f)) / total_f))
|
||||
return round(used * 100.0, 2)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _mb_to_bytes(mb: Optional[int]) -> Optional[float]:
|
||||
if mb is None:
|
||||
return None
|
||||
try:
|
||||
return float(max(0, int(mb))) * 1024.0 * 1024.0
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _gt_pct(value: Optional[float], threshold_pct: Optional[int]) -> bool:
|
||||
if value is None or threshold_pct is None:
|
||||
return False
|
||||
try:
|
||||
return float(value) > float(threshold_pct)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _lt_bytes(value_bytes: Optional[float], threshold_mb: Optional[int]) -> bool:
|
||||
if value_bytes is None or threshold_mb is None:
|
||||
return False
|
||||
thr = _mb_to_bytes(threshold_mb)
|
||||
if thr is None:
|
||||
return False
|
||||
try:
|
||||
return float(value_bytes) < float(thr)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _normalize_policy(policy: str) -> str:
|
||||
try:
|
||||
return str(policy).strip().lower()
|
||||
except Exception:
|
||||
return "auto"
|
||||
|
||||
|
||||
def _action_literal(policy: str) -> RemediationAction:
|
||||
p = _normalize_policy(policy)
|
||||
if p in ("free", "unload", "both"):
|
||||
return p # type: ignore[return-value]
|
||||
return "auto"
|
||||
|
||||
|
||||
def _initial_action_for(policy: str) -> RemediationAction:
|
||||
lit = _action_literal(policy)
|
||||
if lit == "auto":
|
||||
return "free"
|
||||
return lit
|
||||
|
||||
|
||||
def _flags_for_action(action: RemediationAction) -> tuple[bool, bool]:
|
||||
if action == "free":
|
||||
return True, False
|
||||
if action == "unload":
|
||||
return False, True
|
||||
if action == "both":
|
||||
return True, True
|
||||
# auto is staged; when executing a step, treat like free unless escalated update chooses unload
|
||||
return True, False
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RemediationResult",
|
||||
"RemediationAction",
|
||||
"MemoryHygiene",
|
||||
"DerivedStats",
|
||||
]
|
||||
+1378
-31
File diff suppressed because it is too large
Load Diff
@@ -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 = ""
|
||||
+20
-1
@@ -1 +1,20 @@
|
||||
huggingface_hub
|
||||
|
||||
aiobotocore==2.24.2
|
||||
aiofiles>=23.2.1
|
||||
aiohttp==3.12.14
|
||||
boto3==1.40.15
|
||||
fastapi==0.110.0
|
||||
huggingface_hub==0.34.0
|
||||
imageio==2.37.0
|
||||
imageio-ffmpeg==0.6.0
|
||||
numpy>=1.26.4
|
||||
opencv-python>=4.6.0.66
|
||||
openexr==3.3.4
|
||||
Pillow==10.4.0
|
||||
python-dotenv==1.0.1
|
||||
python-multipart==0.0.9
|
||||
requests==2.31.0
|
||||
uvicorn==0.27.1
|
||||
websockets==11.0.3
|
||||
json5>=0.9.0
|
||||
--prefer-binary
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
"""Shared types and enums for the Nilor-Nodes sidecar.
|
||||
|
||||
This module intentionally contains minimal placeholders to support the
|
||||
configuration loader and future extensions without introducing unnecessary
|
||||
complexity at this stage.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ConfigSource(Enum):
|
||||
"""Represents the origin of a configuration value."""
|
||||
|
||||
ENV = "env"
|
||||
JSON5 = "json5"
|
||||
+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, "step": 0.001}),
|
||||
}
|
||||
}
|
||||
|
||||
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,38 @@
|
||||
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,114 @@
|
||||
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;
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function setupMediaStreamOutput(node) {
|
||||
// Hide system inputs by default
|
||||
hideWidgets(node, [
|
||||
"content_id",
|
||||
"venue",
|
||||
"canvas",
|
||||
"scene",
|
||||
"job_type",
|
||||
"presigned_upload_url",
|
||||
"job_completions_queue_url",
|
||||
"output_object_keys",
|
||||
]);
|
||||
|
||||
const formatWidget = node.widgets.find((w) => w.name === "format");
|
||||
if (!formatWidget) return;
|
||||
|
||||
// Apply current value
|
||||
toggleFramerateWidget(node, formatWidget.value === "mp4");
|
||||
try {
|
||||
const size = node.computeSize();
|
||||
node.onResize?.(size);
|
||||
app.graph?.setDirtyCanvas(true, true);
|
||||
} catch (_) {}
|
||||
|
||||
// Chain the widget callback once
|
||||
if (!formatWidget.__nilorPatched) {
|
||||
const originalCallback = formatWidget.callback;
|
||||
formatWidget.callback = function (value) {
|
||||
toggleFramerateWidget(node, value === "mp4");
|
||||
try {
|
||||
const size = node.computeSize();
|
||||
node.onResize?.(size);
|
||||
app.graph?.setDirtyCanvas(true, true);
|
||||
} catch (_) {}
|
||||
if (originalCallback) return originalCallback.apply(this, arguments);
|
||||
};
|
||||
formatWidget.__nilorPatched = true;
|
||||
}
|
||||
}
|
||||
|
||||
function setupMediaStreamInput(node) {
|
||||
hideWidgets(node, ["presigned_download_url"]);
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "comfy.nilor-nodes.mediaStream",
|
||||
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, appInstance) {
|
||||
if (nodeData?.name === "MediaStreamOutput") {
|
||||
const origAdded = nodeType.prototype.onAdded;
|
||||
nodeType.prototype.onAdded = function () {
|
||||
if (typeof origAdded === "function") origAdded.apply(this, arguments);
|
||||
try { setTimeout(() => setupMediaStreamOutput(this), 0); } catch (_) {}
|
||||
};
|
||||
|
||||
const origConfigure = nodeType.prototype.onConfigure;
|
||||
nodeType.prototype.onConfigure = function () {
|
||||
if (typeof origConfigure === "function") origConfigure.apply(this, arguments);
|
||||
try { setTimeout(() => setupMediaStreamOutput(this), 0); } catch (_) {}
|
||||
};
|
||||
}
|
||||
|
||||
if (nodeData?.name === "MediaStreamInput") {
|
||||
const origAddedIn = nodeType.prototype.onAdded;
|
||||
nodeType.prototype.onAdded = function () {
|
||||
if (typeof origAddedIn === "function") origAddedIn.apply(this, arguments);
|
||||
try { setTimeout(() => setupMediaStreamInput(this), 0); } catch (_) {}
|
||||
};
|
||||
|
||||
const origConfigureIn = nodeType.prototype.onConfigure;
|
||||
nodeType.prototype.onConfigure = function () {
|
||||
if (typeof origConfigureIn === "function") origConfigureIn.apply(this, arguments);
|
||||
try { setTimeout(() => setupMediaStreamInput(this), 0); } catch (_) {}
|
||||
};
|
||||
}
|
||||
},
|
||||
|
||||
afterConfigureGraph(graph) {
|
||||
try {
|
||||
(graph?._nodes || graph?.nodes || []).forEach((n) => {
|
||||
if (n?.comfyClass === "MediaStreamOutput") setupMediaStreamOutput(n);
|
||||
if (n?.comfyClass === "MediaStreamInput") setupMediaStreamInput(n);
|
||||
});
|
||||
} catch (e) {
|
||||
console.warn("nilor-media-stream afterConfigureGraph error", e);
|
||||
}
|
||||
},
|
||||
|
||||
nodeCreated(node) {
|
||||
if (node.comfyClass === "MediaStreamOutput") setupMediaStreamOutput(node);
|
||||
if (node.comfyClass === "MediaStreamInput") setupMediaStreamInput(node);
|
||||
},
|
||||
});
|
||||
@@ -0,0 +1,673 @@
|
||||
"""
|
||||
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-comfyui` 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 time
|
||||
import aiohttp
|
||||
|
||||
from aiobotocore.session import get_session
|
||||
from botocore.exceptions import EndpointConnectionError, ClientError
|
||||
from .logger import logger
|
||||
from .comfyui_client import ComfyUILocalClient, ComfyUIClientError
|
||||
from .memory_hygiene import MemoryHygiene
|
||||
from .workflow_normalizer import normalize_comfyui_prompt_for_current_os
|
||||
from .config.config import load_nilor_nodes_config, NilorNodesConfig
|
||||
|
||||
|
||||
# --- Configuration ---
|
||||
# Centralized loader provides precedence env > JSON5 and validation
|
||||
_CFG: NilorNodesConfig = load_nilor_nodes_config()
|
||||
|
||||
|
||||
class WorkerConsumer:
|
||||
def __init__(self, cfg: NilorNodesConfig):
|
||||
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
|
||||
self.comfy_client = None
|
||||
self.is_busy = False
|
||||
self.websocket_sid = None
|
||||
# Memory hygiene component (initialized once a client is available)
|
||||
self.hygiene = None
|
||||
|
||||
# Stable client_id for routing events to this worker
|
||||
self.cfg = cfg
|
||||
self.worker_client_id = cfg.worker.worker_client_id
|
||||
|
||||
self.current_prompt_id = None
|
||||
# Hygiene cadence tracking
|
||||
self._last_hygiene_check_ts = 0.0
|
||||
|
||||
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=self.cfg.worker.aws_region,
|
||||
endpoint_url=self.cfg.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=self.cfg.worker.aws_access_key_id,
|
||||
aws_secret_access_key=self.cfg.worker.aws_secret_access_key,
|
||||
) as client:
|
||||
try:
|
||||
self.jobs_queue_url = await self._get_queue_url(
|
||||
client, self.cfg.worker.jobs_queue
|
||||
)
|
||||
self.status_updates_queue_url = await self._get_queue_url(
|
||||
client, self.cfg.worker.status_queue
|
||||
)
|
||||
return True
|
||||
except EndpointConnectionError as e:
|
||||
# Quiet the noisy traceback by logging a concise warning instead
|
||||
logger.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS endpoint is unreachable at {self.cfg.worker.sqs_endpoint_url}: {e}. "
|
||||
)
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.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:
|
||||
logger.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:
|
||||
# Wait until a websocket-capable client is available
|
||||
if self.comfy_client is None:
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Comfy client not constructed; skipping WS listen this cycle."
|
||||
)
|
||||
await asyncio.sleep(5)
|
||||
continue
|
||||
|
||||
# Consume events from client iterator (handles reconnects internally)
|
||||
async for evt in self.comfy_client.ws_connect(self.worker_client_id):
|
||||
event_type = evt.get("type")
|
||||
data = (
|
||||
evt.get("data", {}) if isinstance(evt.get("data"), dict) else {}
|
||||
)
|
||||
prompt_id = data.get("prompt_id")
|
||||
|
||||
if not prompt_id and "sid" in data:
|
||||
prompt_id = data["sid"]
|
||||
|
||||
# Allow 'executing' events through even if missing prompt_id
|
||||
if not prompt_id and event_type != "executing":
|
||||
continue
|
||||
|
||||
# Capture our websocket client id from initial status message
|
||||
if event_type == "status":
|
||||
sid = data.get("sid")
|
||||
if sid:
|
||||
self.websocket_sid = sid
|
||||
logger.debug(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Captured websocket SID: {sid}"
|
||||
)
|
||||
|
||||
# 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")
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Execution started for prompt_id {prompt_id} (content_id: {content_id}) via '{event_type}' event."
|
||||
)
|
||||
# Mark worker busy as soon as execution starts
|
||||
self.is_busy = True
|
||||
|
||||
await self._send_status_update(
|
||||
content_id,
|
||||
running_status,
|
||||
ctx.get("venue"),
|
||||
ctx.get("canvas"),
|
||||
ctx.get("scene"),
|
||||
ctx.get("job_type"),
|
||||
)
|
||||
self.sent_running_status_prompts.add(prompt_id)
|
||||
|
||||
# Handle execution errors
|
||||
elif event_type == "execution_error":
|
||||
logger.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"),
|
||||
ctx.get("job_type"),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
self.content_context_by_content_id.pop(content_id, None)
|
||||
effective_prompt_id = prompt_id or self.current_prompt_id
|
||||
if effective_prompt_id:
|
||||
self._finalize_prompt(
|
||||
effective_prompt_id,
|
||||
reason="execution_error",
|
||||
)
|
||||
|
||||
# Node-level executed event (many per prompt) — ignore for busy/reset
|
||||
elif event_type == "executed":
|
||||
logger.debug(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Received node executed event for prompt_id {prompt_id}: {data}"
|
||||
)
|
||||
|
||||
# Prompt-level completion signal: executing with node None
|
||||
elif event_type == "executing":
|
||||
node_id = data.get("node")
|
||||
logger.debug(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Received 'executing' event. prompt_id={prompt_id}, node_id={node_id}"
|
||||
)
|
||||
|
||||
norm_node_id = self._normalize_node_id(node_id)
|
||||
if norm_node_id is None:
|
||||
effective_prompt_id = prompt_id or self.current_prompt_id
|
||||
if effective_prompt_id:
|
||||
self._finalize_prompt(
|
||||
effective_prompt_id,
|
||||
reason="executing node=None",
|
||||
)
|
||||
|
||||
elif event_type == "execution_success":
|
||||
effective_prompt_id = prompt_id or self.current_prompt_id
|
||||
if effective_prompt_id:
|
||||
self._finalize_prompt(
|
||||
effective_prompt_id,
|
||||
reason="execution_success",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Websocket listener error: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
await asyncio.sleep(5)
|
||||
|
||||
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
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Starting worker consumer. Polling queue: {self.jobs_queue_url}"
|
||||
)
|
||||
|
||||
# Capacity gate: avoid pulling a new job while local ComfyUI is busy
|
||||
if self.is_busy:
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Skipping poll; worker is busy executing a job."
|
||||
)
|
||||
await asyncio.sleep(5)
|
||||
continue
|
||||
|
||||
# Run memory hygiene between jobs (no-op if disabled/unsupported)
|
||||
try:
|
||||
now = time.monotonic()
|
||||
min_interval = max(0, int(self.cfg.hygiene.idle_poll_seconds))
|
||||
if now - self._last_hygiene_check_ts >= min_interval:
|
||||
await self._run_memory_hygiene(debounce_seconds=0)
|
||||
self._last_hygiene_check_ts = now
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Polling for messages..."
|
||||
)
|
||||
try:
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=self.cfg.worker.aws_region,
|
||||
endpoint_url=self.cfg.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=self.cfg.worker.aws_access_key_id,
|
||||
aws_secret_access_key=self.cfg.worker.aws_secret_access_key,
|
||||
) as client:
|
||||
response = await client.receive_message(
|
||||
QueueUrl=self.jobs_queue_url,
|
||||
MaxNumberOfMessages=self.cfg.worker.max_messages,
|
||||
WaitTimeSeconds=self.cfg.worker.poll_wait_s,
|
||||
)
|
||||
|
||||
messages = response.get("Messages", [])
|
||||
if not messages:
|
||||
logger.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=self.cfg.worker.aws_region,
|
||||
endpoint_url=self.cfg.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=self.cfg.worker.aws_access_key_id,
|
||||
aws_secret_access_key=self.cfg.worker.aws_secret_access_key,
|
||||
) as client:
|
||||
await client.delete_message(
|
||||
QueueUrl=self.jobs_queue_url,
|
||||
ReceiptHandle=message["ReceiptHandle"],
|
||||
)
|
||||
logger.debug(
|
||||
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.
|
||||
logger.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:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Processing failed for message {message['MessageId']}: {e}. It will be returned to the queue for retry."
|
||||
)
|
||||
except ClientError as e:
|
||||
# Some SQS providers/endpoints may sporadically return 503 for ReceiveMessage when the queue is idle
|
||||
error_code = None
|
||||
try:
|
||||
error_code = e.response.get("Error", {}).get("Code")
|
||||
except Exception:
|
||||
pass
|
||||
operation_name = getattr(e, "operation_name", "")
|
||||
if operation_name == "ReceiveMessage" and str(error_code) in (
|
||||
"503",
|
||||
"ServiceUnavailable",
|
||||
):
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Queue is empty or endpoint timed out (ReceiveMessage 503). Polling again shortly..."
|
||||
)
|
||||
# await asyncio.sleep(2)
|
||||
continue
|
||||
|
||||
# Unhandled ClientError; fall back to generic handling
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): SQS client error during ReceiveMessage: {e}"
|
||||
)
|
||||
await asyncio.sleep(10)
|
||||
except EndpointConnectionError as e:
|
||||
# Lost connection to SQS; reset and re-initialize on next loop
|
||||
logger.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Lost connection to SQS at {self.cfg.worker.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:
|
||||
logger.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)
|
||||
# Close shared HTTP session if created
|
||||
if self.http_session is not None:
|
||||
try:
|
||||
await self.http_session.close()
|
||||
except Exception:
|
||||
pass
|
||||
logger.info(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Websocket listener stopped."
|
||||
)
|
||||
|
||||
async def process_message(self, message):
|
||||
"""Processes a single SQS message."""
|
||||
logger.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")
|
||||
|
||||
# Prefer execution_spec for other engines; ComfyUI strictly requires 'prompt'
|
||||
if "execution_spec" in job_payload and "prompt" not in job_payload:
|
||||
logger.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Received 'execution_spec' without 'prompt'. ComfyUI path requires 'prompt'; skipping message {message['MessageId']}."
|
||||
)
|
||||
return
|
||||
if "execution_spec" in job_payload and "prompt" in job_payload:
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): 'execution_spec' present alongside 'prompt'; ignoring 'execution_spec' for ComfyUI."
|
||||
)
|
||||
|
||||
# Validate that the payload has the required keys before submitting.
|
||||
if not content_id or "prompt" not in job_payload:
|
||||
logger.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"),
|
||||
"job_type": job_payload.get("job_type"),
|
||||
"status_policy": job_payload.get("status_policy") or {},
|
||||
}
|
||||
except Exception:
|
||||
self.content_context_by_content_id[content_id] = {}
|
||||
|
||||
except Exception as e:
|
||||
logger.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:
|
||||
# Attach/override websocket client_id so server targets events to this worker
|
||||
payload = (
|
||||
dict(workflow_data)
|
||||
if isinstance(workflow_data, dict)
|
||||
else workflow_data
|
||||
)
|
||||
|
||||
if isinstance(payload, dict):
|
||||
# Normalize OS-sensitive path formatting inside the ComfyUI prompt graph.
|
||||
# This lets us accept workflows authored on a different OS (e.g. Windows
|
||||
# backslashes) and run them on the current worker OS.
|
||||
enabled = self.cfg.worker.workflow_os_normalization_enabled
|
||||
if enabled and "prompt" in payload:
|
||||
try:
|
||||
normalized_prompt, rewritten = (
|
||||
normalize_comfyui_prompt_for_current_os(payload["prompt"])
|
||||
)
|
||||
if rewritten:
|
||||
payload["prompt"] = normalized_prompt
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Normalized %s path-like values in ComfyUI prompt for os=%s.",
|
||||
rewritten,
|
||||
os.name,
|
||||
)
|
||||
except Exception as e:
|
||||
# Don't fail the job if normalization fails; submit as-is.
|
||||
logger.warning(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Workflow normalization failed; submitting original prompt. Error: %s",
|
||||
e,
|
||||
)
|
||||
|
||||
# Force top-level client_id to this worker's stable ID
|
||||
payload["client_id"] = self.worker_client_id
|
||||
# Ensure extra_data exists and force its client_id too
|
||||
extra = payload.get("extra_data") or {}
|
||||
if isinstance(extra, dict):
|
||||
extra["client_id"] = self.worker_client_id
|
||||
payload["extra_data"] = extra
|
||||
|
||||
if self.comfy_client is None:
|
||||
# Construct client using shared session when available; fallback to creating a temp one
|
||||
self.comfy_client = ComfyUILocalClient(
|
||||
base_url=self.cfg.comfy.api_url,
|
||||
ws_url=self.cfg.comfy.ws_url,
|
||||
session=self.http_session,
|
||||
logger=logger,
|
||||
timeout=float(self.cfg.comfy.timeout_s),
|
||||
)
|
||||
|
||||
prompt_id = await self.comfy_client.submit_prompt(payload)
|
||||
logger.debug(
|
||||
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
|
||||
# Mark worker busy after successful submission to avoid over-queuing on this machine
|
||||
self.is_busy = True
|
||||
self.current_prompt_id = prompt_id
|
||||
|
||||
# No need to delete here, the consume_loop handles message deletion
|
||||
except ComfyUIClientError as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to submit job to ComfyUI: {e}. Message will be retried."
|
||||
)
|
||||
except (json.JSONDecodeError, KeyError) as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to parse ComfyUI response: {e}. Discarding malformed response."
|
||||
)
|
||||
except Exception as e:
|
||||
logger.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, job_type=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
|
||||
if job_type is not None:
|
||||
body["job_type"] = job_type
|
||||
message_body = json.dumps(body)
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=self.cfg.worker.aws_region,
|
||||
endpoint_url=self.cfg.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=self.cfg.worker.aws_access_key_id,
|
||||
aws_secret_access_key=self.cfg.worker.aws_secret_access_key,
|
||||
) as client:
|
||||
await client.send_message(
|
||||
QueueUrl=self.status_updates_queue_url, MessageBody=message_body
|
||||
)
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Sent status update for content {content_id}: {status}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to send status update for content {content_id}: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
def _finalize_prompt(self, prompt_id, reason: str | None = None):
|
||||
# No-op if we're not currently busy; avoids redundant work on duplicate signals
|
||||
if not self.is_busy:
|
||||
return
|
||||
try:
|
||||
if reason:
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Finalizing prompt via '{reason}'. Using prompt_id={prompt_id}"
|
||||
)
|
||||
content_id = self.prompt_id_to_content_id_map.pop(prompt_id, None)
|
||||
if content_id is not None:
|
||||
self.content_context_by_content_id.pop(content_id, None)
|
||||
self.sent_running_status_prompts.discard(prompt_id)
|
||||
finally:
|
||||
# Only clear busy/state if this finalize corresponds to the current in-flight prompt
|
||||
if self.current_prompt_id == prompt_id:
|
||||
self.current_prompt_id = None
|
||||
self.is_busy = False
|
||||
# Debounced hygiene after job completion
|
||||
try:
|
||||
asyncio.create_task(self._run_memory_hygiene(debounce_seconds=1))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _normalize_node_id(node_id):
|
||||
if node_id is None:
|
||||
return None
|
||||
if isinstance(node_id, str) and node_id.strip() in ("", "None"):
|
||||
return None
|
||||
return node_id
|
||||
|
||||
async def _run_memory_hygiene(self, debounce_seconds: int = 0) -> None:
|
||||
"""Run memory hygiene with optional debounce, guarded by busy state.
|
||||
|
||||
Delegates to the shared hygiene component and ensures we do not
|
||||
pull new work while remediation is running.
|
||||
"""
|
||||
await _run_hygiene_guarded(self, debounce_seconds)
|
||||
|
||||
|
||||
async def consume_jobs():
|
||||
"""Entry point function to be called in a background thread."""
|
||||
# Create shared HTTP session once and reuse throughout lifecycle
|
||||
consumer = WorkerConsumer(cfg=_CFG)
|
||||
# Initialize shared HTTP session at startup
|
||||
consumer.http_session = aiohttp.ClientSession()
|
||||
|
||||
# Construct ComfyUI client using shared session (WS is always client-driven)
|
||||
consumer.comfy_client = ComfyUILocalClient(
|
||||
base_url=_CFG.comfy.api_url,
|
||||
ws_url=_CFG.comfy.ws_url,
|
||||
session=consumer.http_session,
|
||||
logger=logger,
|
||||
timeout=float(_CFG.comfy.timeout_s),
|
||||
)
|
||||
|
||||
# Initialize Memory Hygiene with the constructed client and loaded config
|
||||
try:
|
||||
consumer.hygiene = MemoryHygiene(
|
||||
client=consumer.comfy_client,
|
||||
cfg=_CFG.hygiene,
|
||||
logger=logger,
|
||||
)
|
||||
except Exception:
|
||||
consumer.hygiene = None
|
||||
|
||||
# Emit concise startup configuration summary (no secrets)
|
||||
try:
|
||||
logger.info(
|
||||
(
|
||||
"ℹ️\u2009 Nilor-Nodes: startup config — Comfy API=%s, WS=%s, timeout_s=%s, "
|
||||
"SQS endpoint=%s, jobs_queue=%s, status_queue=%s"
|
||||
),
|
||||
_CFG.comfy.api_url,
|
||||
_CFG.comfy.ws_url,
|
||||
str(_CFG.comfy.timeout_s),
|
||||
_CFG.worker.sqs_endpoint_url,
|
||||
_CFG.worker.jobs_queue,
|
||||
_CFG.worker.status_queue,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Emit concise hygiene summary (effective values)
|
||||
try:
|
||||
h = _CFG.hygiene
|
||||
logger.info(
|
||||
(
|
||||
"ℹ️\u2009 Nilor-Nodes: startup hygiene — enabled=%s, idle_poll_s=%s, "
|
||||
"vram_used_pct_max=%s, ram_used_pct_max=%s, vram_min_free_mb=%s, ram_min_free_mb=%s, "
|
||||
"policy=%s, max_retries=%s, cooldown_s=%s, sleep_between_s=%s, max_cycle_s=%s"
|
||||
),
|
||||
str(h.enabled),
|
||||
str(h.idle_poll_seconds),
|
||||
str(h.vram_usage_pct_max),
|
||||
str(h.ram_usage_pct_max),
|
||||
str(h.vram_min_free_mb),
|
||||
str(h.ram_min_free_mb),
|
||||
h.action_policy,
|
||||
str(h.max_retries),
|
||||
str(h.cooldown_seconds),
|
||||
str(h.sleep_between_attempts_seconds),
|
||||
str(h.max_cycle_duration_seconds),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
await consumer.consume_loop()
|
||||
|
||||
|
||||
async def _sleep_seconds(seconds: int) -> None:
|
||||
try:
|
||||
await asyncio.sleep(max(0, int(seconds)))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def _maybe_bool(value) -> bool:
|
||||
try:
|
||||
return bool(value)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
async def _run_hygiene_guarded(
|
||||
self_ref: "WorkerConsumer", debounce_seconds: int
|
||||
) -> None:
|
||||
# No-op if not available
|
||||
if not getattr(self_ref, "hygiene", None):
|
||||
return
|
||||
# Optional debounce
|
||||
if debounce_seconds > 0:
|
||||
await _sleep_seconds(debounce_seconds)
|
||||
# Set maintenance busy gate
|
||||
if self_ref.is_busy:
|
||||
return
|
||||
self_ref.is_busy = True
|
||||
try:
|
||||
await self_ref.hygiene.check_and_remediate()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
self_ref.is_busy = False
|
||||
@@ -0,0 +1,173 @@
|
||||
"""
|
||||
Workflow normalization helpers.
|
||||
|
||||
Goal: accept ComfyUI "prompt" (API workflow graph) authored on a different OS and
|
||||
rewrite OS-sensitive path formatting (primarily path separators) so that the
|
||||
local ComfyUI instance can resolve model filenames correctly.
|
||||
|
||||
This is intentionally conservative: we only rewrite strings that look like file
|
||||
paths / model names (e.g. end in ".safetensors") and we avoid touching URLs and
|
||||
free-form prompt text.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Iterable, List, Tuple
|
||||
|
||||
__all__ = [
|
||||
"PathRemap",
|
||||
"normalize_comfyui_prompt_for_current_os",
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PathRemap:
|
||||
"""Prefix remap for absolute paths across OSes.
|
||||
|
||||
Example:
|
||||
PathRemap(from_prefix="D:\\ComfyUI\\models", to_prefix="/mnt/models")
|
||||
|
||||
Matching is performed in a canonicalized form (both prefixes and candidate
|
||||
paths have backslashes converted to forward slashes). After remapping, path
|
||||
separators are normalized for the current OS.
|
||||
"""
|
||||
|
||||
from_prefix: str
|
||||
to_prefix: str
|
||||
|
||||
|
||||
_URL_RE = re.compile(r"^[a-zA-Z][a-zA-Z0-9+.-]*://")
|
||||
_WIN_DRIVE_RE = re.compile(r"^[A-Za-z]:[\\/]")
|
||||
|
||||
# Common file extensions encountered in ComfyUI prompts (models + media).
|
||||
_PATH_EXTS = {
|
||||
".safetensors",
|
||||
".pt",
|
||||
".pth",
|
||||
".ckpt",
|
||||
".bin",
|
||||
".onnx",
|
||||
".json",
|
||||
".json5",
|
||||
".yaml",
|
||||
".yml",
|
||||
".txt",
|
||||
".png",
|
||||
".jpg",
|
||||
".jpeg",
|
||||
".webp",
|
||||
".gif",
|
||||
".bmp",
|
||||
".tif",
|
||||
".tiff",
|
||||
".exr",
|
||||
".mp4",
|
||||
".mov",
|
||||
".mkv",
|
||||
".webm",
|
||||
".wav",
|
||||
".mp3",
|
||||
".flac",
|
||||
}
|
||||
|
||||
|
||||
def _canonicalize_for_prefix_match(path: str) -> str:
|
||||
# Use forward slashes for prefix matching regardless of host OS.
|
||||
return path.replace("\\", "/")
|
||||
|
||||
|
||||
def _normalize_separators_for_current_os(path: str) -> str:
|
||||
# ComfyUI often uses OS-native separators in its model name registry.
|
||||
# We normalize to the current OS so lookup keys match.
|
||||
if os.name == "nt":
|
||||
return path.replace("/", "\\")
|
||||
return path.replace("\\", "/")
|
||||
|
||||
|
||||
def _looks_like_path_value(s: str) -> bool:
|
||||
s_stripped = s.strip()
|
||||
if not s_stripped:
|
||||
return False
|
||||
|
||||
# Don't touch URLs (presigned uploads, http inputs, etc.)
|
||||
if _URL_RE.match(s_stripped):
|
||||
return False
|
||||
|
||||
# Don't touch placeholder tokens used by the system.
|
||||
if s_stripped.startswith("<") and s_stripped.endswith(">"):
|
||||
return False
|
||||
|
||||
# Avoid common sentinel.
|
||||
if s_stripped == "None":
|
||||
return False
|
||||
|
||||
lower = s_stripped.lower()
|
||||
if any(lower.endswith(ext) for ext in _PATH_EXTS):
|
||||
return True
|
||||
|
||||
# Absolute Windows paths even without extensions.
|
||||
if _WIN_DRIVE_RE.match(s_stripped):
|
||||
return True
|
||||
|
||||
# Relative paths with explicit prefixes.
|
||||
if s_stripped.startswith(("./", "../", "~/", "~\\")):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def _apply_prefix_remaps(path: str, remaps: Iterable[PathRemap]) -> str:
|
||||
if not remaps:
|
||||
return path
|
||||
|
||||
cand = _canonicalize_for_prefix_match(path)
|
||||
for r in remaps:
|
||||
frm = _canonicalize_for_prefix_match(str(r.from_prefix))
|
||||
if cand.startswith(frm):
|
||||
to = str(r.to_prefix)
|
||||
# Replace using canonical representation then return in that form;
|
||||
# the caller will normalize separators for current OS afterwards.
|
||||
replaced = to + cand[len(frm) :]
|
||||
return replaced
|
||||
return path
|
||||
|
||||
|
||||
def normalize_comfyui_prompt_for_current_os(
|
||||
prompt: Any, *, path_remaps: Iterable[PathRemap] | None = None
|
||||
) -> Tuple[Any, int]:
|
||||
"""Normalize a ComfyUI API `prompt` graph for the current OS.
|
||||
|
||||
Args:
|
||||
prompt: The value of the `/prompt` payload's `prompt` key. Typically a
|
||||
dict mapping node ids to `{class_type, inputs, ...}`.
|
||||
path_remaps: Optional prefix remaps applied before separator
|
||||
normalization (useful for absolute paths).
|
||||
|
||||
Returns:
|
||||
(normalized_prompt, num_rewritten_strings)
|
||||
"""
|
||||
|
||||
remaps: List[PathRemap] = list(path_remaps or [])
|
||||
rewritten = 0
|
||||
|
||||
def walk(x: Any) -> Any:
|
||||
nonlocal rewritten
|
||||
if isinstance(x, dict):
|
||||
return {k: walk(v) for k, v in x.items()}
|
||||
if isinstance(x, list):
|
||||
return [walk(v) for v in x]
|
||||
if isinstance(x, tuple):
|
||||
return tuple(walk(v) for v in x)
|
||||
if isinstance(x, str) and _looks_like_path_value(x):
|
||||
y = _apply_prefix_remaps(x, remaps)
|
||||
y = _normalize_separators_for_current_os(y)
|
||||
if y != x:
|
||||
rewritten += 1
|
||||
return y
|
||||
return x
|
||||
|
||||
return walk(prompt), rewritten
|
||||
|
||||
Reference in New Issue
Block a user