Author SHA1 Message Date
Sylvester Meighan 629f87a2c3 fix(comfyui): restore missing custom nodes and worker functionality
- Restore __init__.py with full node registration and SQS worker startup
- Add missing controllers.py with NilorPreset and NilorGroup nodes
- Add missing user_input.py with NilorUserInput_* nodes
- Add missing web/js/media_stream.js for MediaStreamOutput UI extensions
- Add missing web/js/controllers.js for controller node UI extensions
- Add missing worker_consumer.py for SQS job processing
- Add missing .env.example for environment configuration
- Fixes missing custom nodes and SQS worker functionality on working branch
2025-10-01 15:47:20 -07:00
Sylvester Meighan e14ffc2284 feat(storage): implement storage endpoints migration for ComfyUI nodes
- Add BrainApiClient for interacting with Brain API storage endpoints
- Update MediaStreamInput to use storage_id + filename instead of presigned URLs
- Update MediaStreamOutput to use Brain API client for uploads and remove output_object_keys dependency
- Add environment configuration for Brain API connection
- Add test file for Brain API client functionality

This enables ComfyUI nodes to work with the new storage architecture
where Brain API acts as a proxy for MinIO operations.
2025-09-30 12:51:41 -07:00
19 changed files with 841 additions and 3759 deletions
+21 -53
View File
@@ -1,58 +1,26 @@
# NILOR_LOG_LEVEL=INFO # possible log levels: INFO, DEBUG, WARNING, ERROR, CRITICAL
# --- S3 (AWS/MINIO) CREDENTIALS ---
## These are used by both scripts to connect to SQS and S3 (MinIO).
AWS_ACCESS_KEY_ID=minioadmin
AWS_SECRET_ACCESS_KEY=...
AWS_DEFAULT_REGION=us-east-1
# --- 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
# --- SQS SETTINGS ---
## Toggles functionality for the SQS Worker Consumer
SQS_ENABLED=false
## 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
## For media_stream.py (MediaStreamOutput Node)
### Endpoint for the SQS service where completion messages are sent.
SQS_ENDPOINT_URL=http://127.0.0.1:9324
## WebSocket reconnect policy
# NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS=5
# NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS=30.0
### The specific SQS queue the worker should push job status updates to.
SQS_JOB_STATUS_UPDATES_QUEUE_NAME=job_status_updates
# --- 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
## For worker_consumer.py (Job Consumer)
### The specific SQS queue which the worker should poll for new jobs.
SQS_JOBS_TO_PROCESS_QUEUE_NAME=jobs_to_process
## 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
## (Optional) For worker_consumer.py (Job Consumer)
### The local URL of the ComfyUI API server.
# You only need to set this if your ComfyUI server is NOT running on the default port 8188.
# COMFYUI_API_URL=http://127.0.0.1:8188
# COMFYUI_WS_URL=ws://127.0.0.1:8188
-65
View File
@@ -2,10 +2,6 @@
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>
@@ -192,65 +188,4 @@ Uploads video files to a HuggingFace dataset.
| 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>
+6 -15
View File
@@ -1,7 +1,6 @@
import os
import threading
import asyncio
import logging
from dotenv import load_dotenv
# --- Nilor-Nodes Custom Node Registration and Startup ---
@@ -20,12 +19,6 @@ dotenv_path = os.path.join(current_dir, ".env")
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 ---
@@ -36,20 +29,18 @@ def start_consumer_loop():
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:
# Start the SQS Worker Consumer (controlled by SQS_ENABLED)
raw_sqs_enabled = os.getenv("SQS_ENABLED", "false")
env_sqs_enabled = raw_sqs_enabled.strip().lower() == "true"
if env_sqs_enabled:
consumer_thread = threading.Thread(target=start_consumer_loop, daemon=True)
consumer_thread.start()
print(
"✅ Nilor-Nodes: SQS worker consumer thread started (NILOR_SQS_ENABLED=true)."
f"✅ Nilor-Nodes: SQS worker consumer thread started (SQS_ENABLED={raw_sqs_enabled} in .env)."
)
else:
print(
"⚠️ Nilor-Nodes: SQS worker consumer functionality is disabled (NILOR_SQS_ENABLED=false)."
f"⚠️ Nilor-Nodes: SQS worker consumer functionality is disabled (SQS_ENABLED={raw_sqs_enabled} in .env)."
)
+286
View File
@@ -0,0 +1,286 @@
"""
Brain API Client for ComfyUI Nodes
This client provides methods to interact with the Brain API storage endpoints,
replacing the need for pre-signed URLs in the ComfyUI workflow.
"""
import requests
import os
import logging
from typing import Optional, Dict, Any
from dotenv import load_dotenv
# Load environment variables
current_dir = os.path.dirname(os.path.abspath(__file__))
dotenv_path = os.path.join(current_dir, ".env")
load_dotenv(dotenv_path=dotenv_path)
# Setup logging
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
class BrainApiClient:
"""
Client for interacting with Brain API storage endpoints.
This client handles authentication and provides methods for uploading,
downloading, and deleting files through the Brain API storage endpoints.
"""
def __init__(self):
"""Initialize the Brain API client with configuration from environment variables."""
self.base_url = os.getenv("BRANDO_BRAIN_API_BASE_URL", "http://localhost:2024/api")
self.api_key = os.getenv("BRANDO_API_KEY")
if not self.api_key:
raise ValueError(
"BRANDO_API_KEY environment variable is required for Brain API authentication"
)
self.headers = {
"Authorization": f"Bearer {self.api_key}",
"User-Agent": "ComfyUI-NilorNodes/1.0"
}
logging.info(f"Brain API Client initialized with base URL: {self.base_url}")
def upload_file_to_storage(self, file_path: str, filename: str) -> Dict[str, Any]:
"""
Upload a file to Brain API storage and return storage metadata.
Args:
file_path: Local path to the file to upload
filename: Name to use for the uploaded file
Returns:
Dict containing storage_id and filename
Raises:
requests.RequestException: If upload fails
FileNotFoundError: If file_path doesn't exist
"""
if not os.path.exists(file_path):
raise FileNotFoundError(f"File not found: {file_path}")
url = f"{self.base_url}/storage/upload"
try:
with open(file_path, 'rb') as file:
files = {'file': (filename, file, 'application/octet-stream')}
logging.info(f"Uploading file '{filename}' to Brain API storage...")
response = requests.post(
url,
files=files,
headers=self.headers,
timeout=300
)
response.raise_for_status()
result = response.json()
logging.info(f"Upload successful. Storage ID: {result.get('storage_id')}")
return result
except requests.RequestException as e:
logging.error(f"Failed to upload file '{filename}': {e}")
raise
except Exception as e:
logging.error(f"Unexpected error uploading file '{filename}': {e}")
raise
def upload_fileobj_to_storage(self, file_obj, filename: str, content_type: str = 'application/octet-stream') -> Dict[str, Any]:
"""
Upload a file-like object to Brain API storage and return storage metadata.
Args:
file_obj: File-like object to upload
filename: Name to use for the uploaded file
content_type: MIME type of the file
Returns:
Dict containing storage_id and filename
Raises:
requests.RequestException: If upload fails
"""
url = f"{self.base_url}/storage/upload"
try:
files = {'file': (filename, file_obj, content_type)}
logging.info(f"Uploading file object '{filename}' to Brain API storage...")
response = requests.post(
url,
files=files,
headers=self.headers,
timeout=300
)
response.raise_for_status()
result = response.json()
logging.info(f"Upload successful. Storage ID: {result.get('storage_id')}")
return result
except requests.RequestException as e:
logging.error(f"Failed to upload file object '{filename}': {e}")
raise
except Exception as e:
logging.error(f"Unexpected error uploading file object '{filename}': {e}")
raise
def download_file_from_storage(self, storage_id: str, filename: str, dest_path: str) -> str:
"""
Download a file from Brain API storage to a local path.
Args:
storage_id: Storage ID of the file to download
filename: Name of the file to download
dest_path: Local path where the file should be saved
Returns:
Path to the downloaded file
Raises:
requests.RequestException: If download fails
"""
url = f"{self.base_url}/storage/{storage_id}"
params = {'filename': filename}
try:
logging.info(f"Downloading file '{filename}' (storage_id: {storage_id}) from Brain API storage...")
response = requests.get(
url,
params=params,
headers=self.headers,
timeout=300,
stream=True
)
response.raise_for_status()
# Ensure destination directory exists
os.makedirs(os.path.dirname(dest_path), exist_ok=True)
with open(dest_path, 'wb') as f:
for chunk in response.iter_content(chunk_size=8192):
f.write(chunk)
logging.info(f"Download successful. File saved to: {dest_path}")
return dest_path
except requests.RequestException as e:
logging.error(f"Failed to download file '{filename}' (storage_id: {storage_id}): {e}")
raise
except Exception as e:
logging.error(f"Unexpected error downloading file '{filename}': {e}")
raise
def get_file_from_storage(self, storage_id: str, filename: str) -> bytes:
"""
Get file content from Brain API storage as bytes.
Args:
storage_id: Storage ID of the file to download
filename: Name of the file to download
Returns:
File content as bytes
Raises:
requests.RequestException: If download fails
"""
url = f"{self.base_url}/storage/{storage_id}"
params = {'filename': filename}
try:
logging.info(f"Getting file '{filename}' (storage_id: {storage_id}) from Brain API storage...")
response = requests.get(
url,
params=params,
headers=self.headers,
timeout=300
)
response.raise_for_status()
logging.info(f"File retrieval successful. Size: {len(response.content)} bytes")
return response.content
except requests.RequestException as e:
logging.error(f"Failed to get file '{filename}' (storage_id: {storage_id}): {e}")
raise
except Exception as e:
logging.error(f"Unexpected error getting file '{filename}': {e}")
raise
def delete_file_from_storage(self, storage_id: str, filename: str) -> None:
"""
Delete a file from Brain API storage.
Args:
storage_id: Storage ID of the file to delete
filename: Name of the file to delete
Raises:
requests.RequestException: If deletion fails
"""
url = f"{self.base_url}/storage/{storage_id}"
params = {'filename': filename}
try:
logging.info(f"Deleting file '{filename}' (storage_id: {storage_id}) from Brain API storage...")
response = requests.delete(
url,
params=params,
headers=self.headers,
timeout=60
)
response.raise_for_status()
logging.info(f"File deletion successful")
except requests.RequestException as e:
logging.error(f"Failed to delete file '{filename}' (storage_id: {storage_id}): {e}")
raise
except Exception as e:
logging.error(f"Unexpected error deleting file '{filename}': {e}")
raise
def health_check(self) -> bool:
"""
Check if the Brain API is accessible and authentication is working.
Returns:
True if API is accessible, False otherwise
"""
try:
# Try to access a simple endpoint to verify connectivity
url = f"{self.base_url}/health" # Assuming there's a health endpoint
response = requests.get(url, headers=self.headers, timeout=10)
return response.status_code == 200
except:
# If health endpoint doesn't exist, try the storage upload endpoint
# with a HEAD request to check authentication
try:
url = f"{self.base_url}/storage/upload"
response = requests.head(url, headers=self.headers, timeout=10)
return response.status_code in [200, 405] # 405 Method Not Allowed is OK for HEAD
except:
return False
# Global client instance
_brain_api_client = None
def get_brain_api_client() -> BrainApiClient:
"""
Get or create the global Brain API client instance.
Returns:
BrainApiClient instance
"""
global _brain_api_client
if _brain_api_client is None:
_brain_api_client = BrainApiClient()
return _brain_api_client
-910
View File
@@ -1,910 +0,0 @@
"""
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]
-75
View File
@@ -1,75 +0,0 @@
{
// 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: ""
}
-584
View File
@@ -1,584 +0,0 @@
"""
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()
-37
View File
@@ -1,37 +0,0 @@
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"]
+86 -114
View File
@@ -7,13 +7,23 @@ import logging
import imageio.v2 as imageio
import mimetypes
import boto3
import os
import json
from .logger import logger
from .config.config import load_nilor_nodes_config
from dotenv import load_dotenv
from .brain_api_client import get_brain_api_client
# Load shared configuration once
_CFG = load_nilor_nodes_config()
# --- Load Environment Variables ---
# Get the directory of the current script
current_dir = os.path.dirname(os.path.abspath(__file__))
# Construct the path to the .env file
dotenv_path = os.path.join(current_dir, ".env")
# Load the .env file
load_dotenv(dotenv_path=dotenv_path)
# --- Setup Logging ---
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
# --- Node Categories ---
category = "Nilor Nodes 👺"
@@ -40,9 +50,13 @@ class MediaStreamInput:
{"default": "default_input", "multiline": False},
),
"format": (["image", "image_batch", "video"],),
"presigned_download_url": (
"storage_id": (
"STRING",
{"multiline": True, "default": "<auto-filled by system>"},
{"default": "<auto-filled by system>", "multiline": False},
),
"filename": (
"STRING",
{"default": "<auto-filled by system>", "multiline": False},
),
},
"hidden": {},
@@ -55,21 +69,25 @@ class MediaStreamInput:
def download(
self,
presigned_download_url: str,
storage_id: str,
filename: 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}'"
logging.info(
f"ℹ️\u2009 Nilor-Nodes: MediaStreamInput: Downloading file '{filename}' (storage_id: {storage_id}) for input '{input_name}' with format '{format}'"
)
try:
# Get Brain API client
brain_client = get_brain_api_client()
# Two-phase download for batches: manifest first, then assets
if format == "image_batch":
manifest_response = requests.get(presigned_download_url, timeout=60)
manifest_response.raise_for_status()
manifest = manifest_response.json()
# Download manifest file first
manifest_bytes = brain_client.get_file_from_storage(storage_id, filename)
manifest = json.loads(manifest_bytes.decode('utf-8'))
logger.info(
logging.info(
f"ℹ️\u2009 Nilor-Nodes: Processing manifest for '{manifest.get('input_name')}' with {len(manifest.get('files', []))} assets."
)
@@ -78,15 +96,20 @@ class MediaStreamInput:
manifest.get("files", []), key=lambda x: x.get("sequence", 0)
)
# Download all assets in parallel
# Download all assets using Brain API client
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(
# Each file_info should now contain storage_id and filename instead of presigned_url
file_storage_id = file_info.get("storage_id")
file_filename = file_info.get("filename")
if not file_storage_id or not file_filename:
raise ValueError(f"Missing storage_id or filename in manifest file info: {file_info}")
file_bytes = brain_client.get_file_from_storage(file_storage_id, file_filename)
asset_responses.append(file_bytes)
except Exception as e:
logging.error(
f"🛑\u2009 Nilor-Nodes: Failed to download asset {file_info.get('filename')}: {e}"
)
raise # Re-raise to fail the entire process
@@ -94,9 +117,7 @@ class MediaStreamInput:
return self._process_image_batch(asset_responses)
# --- Single-file download ---
response = requests.get(presigned_download_url, timeout=180)
response.raise_for_status()
media_bytes = response.content
media_bytes = brain_client.get_file_from_storage(storage_id, filename)
if format == "video":
return self._process_video(media_bytes)
@@ -105,22 +126,17 @@ class MediaStreamInput:
else:
# Should not happen if UI choices are respected
raise ValueError(
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Unsupported format '{format}' for single media download."
f"[🛑] 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}"
logging.error(
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Failed to download or process media: {e}"
)
return (None,)
def _process_image_batch(self, image_bytes_list):
logger.info(
logging.info(
f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing image batch with {len(image_bytes_list)} images..."
)
output_images = []
@@ -138,13 +154,13 @@ class MediaStreamInput:
# Concatenate along the batch dimension (dim=0)
images_tensor = torch.cat(output_images, dim=0)
logger.info(
logging.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...")
logging.info("ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing as image...")
image_pil = Image.open(io.BytesIO(image_bytes))
# Ensure image is in RGB
@@ -153,11 +169,11 @@ class MediaStreamInput:
np.array(rgb_image_pil).astype(np.float32) / 255.0
).unsqueeze(0)
logger.info("✅ Nilor-Nodes (MediaStreamInput): Image processing successful.")
logging.info("✅ Nilor-Nodes (MediaStreamInput): Image processing successful.")
return (image_tensor,)
def _process_video(self, video_bytes):
logger.info("ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing as video...")
logging.info("ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing as video...")
frames = []
with imageio.get_reader(io.BytesIO(video_bytes), format="mp4") as reader:
for frame in reader:
@@ -169,7 +185,7 @@ class MediaStreamInput:
if not frames:
raise ValueError(
"🛑\u2009 Nilor-Nodes (MediaStreamInput): No frames could be read from the video."
"[🛑] Nilor-Nodes (MediaStreamInput): No frames could be read from the video."
)
# Stack frames into a single tensor (batch of images)
@@ -210,22 +226,10 @@ class MediaStreamOutput:
"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",
@@ -233,8 +237,7 @@ class MediaStreamOutput:
},
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("uploaded_url",)
RETURN_TYPES = ()
FUNCTION = "upload_and_notify"
OUTPUT_NODE = True
CATEGORY = category + subcategories["streaming"]
@@ -247,49 +250,39 @@ class MediaStreamOutput:
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."
"[🛑] 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.
# No longer need to parse output_object_keys since we use storage_ids directly
# 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.
# Upload the media using Brain API client
brain_client = get_brain_api_client()
storage_result = None
if format == "png":
self._upload_image(images[0], presigned_upload_url)
storage_result = self._upload_image(images[0], brain_client, output_name)
elif format == "mp4":
self._upload_video(images, presigned_upload_url, framerate)
storage_result = self._upload_video(images, brain_client, framerate, output_name)
# 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:
# Use the storage_id from the upload result for the SQS message
if not storage_result or not storage_result.get('storage_id'):
logging.error(
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): FATAL -- Could not find object key for output name '{output_name}' in output_object_keys."
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): FATAL -- Upload failed or no storage_id returned."
)
# Send an empty dictionary to signal failure.
final_outputs_for_sqs = {}
else:
final_outputs_for_sqs = {output_name: output_key_for_this_node}
# Use storage_id instead of object key
storage_id = storage_result['storage_id']
final_outputs_for_sqs = {output_name: storage_id}
# After upload, send the filtered dictionary of outputs to the SQS queue.
completion_message = {
@@ -300,38 +293,36 @@ class MediaStreamOutput:
"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,
endpoint_url=os.getenv("SQS_ENDPOINT_URL"),
aws_access_key_id=os.getenv("AWS_ACCESS_KEY_ID", "local"),
aws_secret_access_key=os.getenv("AWS_SECRET_ACCESS_KEY", "local"),
region_name=os.getenv("AWS_DEFAULT_REGION", "us-east-1"),
)
logger.debug(
logging.info(
f"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Sending completion message for content {content_id} to queue: {job_completions_queue_url}"
)
sqs_client.send_message(
QueueUrl=job_completions_queue_url,
MessageBody=json.dumps(completion_message),
)
logger.info(
f"✅ Nilor-Nodes (MediaStreamOutput): Completion message sent successfully for content {content_id} to queue: {job_completions_queue_url}"
logging.info(
"✅ Nilor-Nodes (MediaStreamOutput): Completion message sent successfully."
)
except Exception as e:
logger.error(
logging.error(
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): Failed to send completion message to SQS: {e}"
)
raise # Re-raise to fail the ComfyUI job
return {"ui": {"images": []}, "result": (presigned_upload_url,)}
return {"ui": {"images": []}}
def _upload_image(self, image_tensor, url):
logger.debug(
def _upload_image(self, image_tensor, brain_client, output_name):
logging.info(
"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Uploading as PNG image..."
)
i = 255.0 * image_tensor.cpu().numpy()
@@ -341,10 +332,11 @@ class MediaStreamOutput:
img_pil.save(buffer, format="PNG", compress_level=4)
buffer.seek(0)
self._perform_upload(buffer, url, "image/png")
filename = f"{output_name}.png"
return brain_client.upload_fileobj_to_storage(buffer, filename, "image/png")
def _upload_video(self, image_batch_tensor, url, framerate):
logger.info(
def _upload_video(self, image_batch_tensor, brain_client, framerate, output_name):
logging.info(
f"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Uploading as MP4 video. Frame count: {len(image_batch_tensor)}"
)
frames = []
@@ -357,29 +349,9 @@ class MediaStreamOutput:
imageio.mimwrite(buffer, frames, format="mp4", fps=framerate, quality=8)
buffer.seek(0)
self._perform_upload(buffer, url, "video/mp4")
filename = f"{output_name}.mp4"
return brain_client.upload_fileobj_to_storage(buffer, filename, "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 ---
@@ -389,6 +361,6 @@ NODE_CLASS_MAPPINGS = {
}
NODE_DISPLAY_NAME_MAPPINGS = {
"MediaStreamInput": "👺 Media Stream Input (URL)",
"MediaStreamOutput": "👺 Media Stream Output (URL)",
"MediaStreamInput": "👺 Media Stream Input (Storage)",
"MediaStreamOutput": "👺 Media Stream Output (Storage)",
}
-432
View File
@@ -1,432 +0,0 @@
"""
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",
]
+37 -605
View File
@@ -16,24 +16,7 @@ import torch
import builtins
from pathlib import Path
import cv2
import warnings
from .utils import pil2tensor, tensor2pil
from .logger import logger
from comfy.utils import common_upscale
from comfy import model_management
import sys
from os.path import dirname, join
# Attempt to import ImagePadKJ from comfyui-kjnodes if available
_kj_nodes_path = join(dirname(__file__), "..", "comfyui-kjnodes", "nodes")
if _kj_nodes_path not in sys.path:
sys.path.append(_kj_nodes_path)
try:
from image_nodes import ImagePadKJ # type: ignore
except Exception as _e:
logger.warning(
f"⚠️\u2009 Nilor-Nodes (nilornodes): Could not import ImagePadKJ from comfyui-kjnodes ({_kj_nodes_path}): {_e}"
)
BIGMIN = -(2**53 - 1)
BIGMAX = 2**53 - 1
@@ -183,9 +166,7 @@ class NilorRemapFloatList:
):
# Avoid division by zero
if max_input - min_input == 0:
raise ValueError(
"🛑\u2009 Nilor-Nodes (RemapFloatList): max_input and min_input cannot be the same value."
)
raise ValueError("max_input and min_input cannot be the same value.")
scale = (max_output - min_output) / (max_input - min_input)
return ([min_output + (x - min_input) * scale for x in list_of_floats],)
@@ -240,9 +221,7 @@ class NilorInverseMapFloatList:
def inverse_map_float_list(self, list_of_floats):
if not list_of_floats:
raise ValueError(
"🛑\u2009 Nilor-Nodes (InverseMapFloatList): The input list_of_floats cannot be empty."
)
raise ValueError("The input list_of_floats cannot be empty.")
min_input = min(list_of_floats)
max_input = max(list_of_floats)
@@ -338,9 +317,7 @@ class NilorCountImagesInDirectory:
def count_images_in_directory(self, directory):
if not os.path.isdir(directory):
raise FileNotFoundError(
f"🛑\u2009 Nilor-Nodes (NilorCountImagesInDirectory): Directory '{directory}' cannot be found."
)
raise FileNotFoundError(f"Directory '{directory} cannot be found.")
list_dir = []
list_dir = os.listdir(directory)
@@ -388,9 +365,7 @@ class NilorSelectIndexFromList:
# Ensure the index is within bounds
if index < 0 or index >= len(actual_list):
raise ValueError(
"🛑\u2009 Nilor-Nodes (SelectIndexFromList): Index is outside the bounds of the array."
)
raise ValueError("Index is outside the bounds of the array.")
# Returns the value at the given index
return (actual_list[index],)
@@ -428,9 +403,7 @@ class NilorSaveEXRArbitrary:
self, channels=None, filename_prefix="output", prompt=None, extra_pnginfo=None
):
logger.info(
"ℹ️\u2009 Nilor-Nodes (SaveEXRArbitrary): Running save_exr_arbitrary"
)
print("Running save_exr_arbitrary")
# print(f"channels: {channels}")
# print(f"filename_prefix: {filename_prefix}")
@@ -442,9 +415,7 @@ class NilorSaveEXRArbitrary:
try:
actual_channels[0]
except TypeError:
logger.error(
"🛑\u2009 Nilor-Nodes (SaveEXRArbitrary): actual_channels is not subscriptable"
)
print("actual_channels is not subscriptable")
return
# File path handling
@@ -481,9 +452,7 @@ class NilorSaveEXRArbitrary:
height, width = image_channels[0].shape[-2:]
for tensor in image_channels:
if tensor.shape[-2:] != (height, width):
raise ValueError(
"🛑\u2009 Nilor-Nodes (SaveEXRArbitrary): All input tensors must have the same dimensions"
)
raise ValueError("All input tensors must have the same dimensions")
# Channel naming
default_names = ["R", "G", "B", "A"] + [
@@ -533,13 +502,9 @@ class NilorSaveEXRArbitrary:
exr_file.writePixels(channel_data)
exr_file.close()
logger.info(
f"✅\u2009 Nilor-Nodes (SaveEXRArbitrary): EXR file saved successfully to {writepath}"
)
print(f"EXR file saved successfully to {writepath}")
except Exception as e:
logger.error(
f"🛑\u2009 Nilor-Nodes (SaveEXRArbitrary): Failed to write EXR file: {e}"
)
print(f"Failed to write EXR file: {e}")
class NilorSaveVideoToHFDataset:
@@ -657,14 +622,12 @@ class NilorShuffleImageBatch:
def _check_image_dimensions(self, images):
if images.shape[0] == 0:
raise ValueError(
"🛑\u2009 Nilor-Nodes (ShuffleImageBatch): Input images tensor is empty."
)
raise ValueError("Input images tensor is empty.")
# All images in the batch should have the same dimensions
if len(images.shape) != 4:
raise ValueError(
f"🛑\u2009 Nilor-Nodes (ShuffleImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
f"Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
)
def shuffle_image_batch(self, images: torch.Tensor, seed):
@@ -704,14 +667,12 @@ class NilorRepeatTrimImageBatch:
def _check_image_dimensions(self, images):
if images.shape[0] == 0:
raise ValueError(
"🛑\u2009 Nilor-Nodes (RepeatTrimImageBatch): Input images tensor is empty."
)
raise ValueError("Input images tensor is empty.")
# All images in the batch should have the same dimensions
if len(images.shape) != 4:
raise ValueError(
f"🛑\u2009 Nilor-Nodes (RepeatTrimImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
f"Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
)
def repeat_trim_image_batch(self, images: torch.Tensor, count):
@@ -749,14 +710,12 @@ class NilorRepeatShuffleTrimImageBatch:
def _check_image_dimensions(self, images):
if images.shape[0] == 0:
raise ValueError(
"🛑\u2009 Nilor-Nodes (RepeatShuffleTrimImageBatch): Input images tensor is empty."
)
raise ValueError("Input images tensor is empty.")
# All images in the batch should have the same dimensions
if len(images.shape) != 4:
raise ValueError(
f"🛑\u2009 Nilor-Nodes (RepeatShuffleTrimImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
f"Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
)
def repeat_shuffle_trim_image_batch(self, images: torch.Tensor, seed, count):
@@ -822,16 +781,12 @@ class NilorOutputFilenameString:
if unique_id is not None and extra_pnginfo is not None:
if not isinstance(extra_pnginfo, list):
logger.error(
"🛑\u2009 Nilor-Nodes (OutputFilenameString): extra_pnginfo is not a list"
)
print("Error: extra_pnginfo is not a list")
elif (
not isinstance(extra_pnginfo[0], dict)
or "workflow" not in extra_pnginfo[0]
):
logger.error(
"🛑\u2009 Nilor-Nodes (OutputFilenameString): extra_pnginfo[0] is not a dict or missing 'workflow' key"
)
print("Error: extra_pnginfo[0] is not a dict or missing 'workflow' key")
else:
workflow = extra_pnginfo[0]["workflow"]
node = next(
@@ -882,206 +837,7 @@ class NilorNFractionsOfInt:
elif type == "start + end":
return ([i * numerator // (denominator - 1) for i in range(denominator)],)
else:
raise ValueError(
f"🛑\u2009 Nilor-Nodes (NilorNFractionsOfInt): Unknown type: {type}"
)
class NilorWanTileResolution:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"input_width": (
"INT",
{"default": 1920, "min": 16, "max": BIGMAX, "step": 1},
),
"input_height": (
"INT",
{"default": 1080, "min": 16, "max": BIGMAX, "step": 1},
),
"target_width": (
"INT",
{"default": 3840, "min": 16, "max": BIGMAX, "step": 1},
),
"target_height": (
"INT",
{"default": 2160, "min": 16, "max": BIGMAX, "step": 1},
),
"size_preference": (
["largest", "smallest"],
{"default": "largest"},
),
}
}
RETURN_TYPES = ("INT", "INT")
RETURN_NAMES = ("tile_width", "tile_height")
FUNCTION = "compute_tile_resolution"
CATEGORY = category + subcategories["utilities"]
MIN_TILE_DIM = 384
MAX_TILE_DIM = 1794
MIN_TILE_AREA = 384 * 384
MAX_TILE_AREA = 1024 * 1024
@staticmethod
def _clamp(value, minimum, maximum):
return max(minimum, min(value, maximum))
def compute_tile_resolution(
self,
input_width,
input_height,
target_width,
target_height,
size_preference="largest",
):
"""
Compute (Wt, Ht) tile size (multiples of 16) within
[MIN_TILE_DIM, MAX_TILE_DIM] while keeping area between
[MIN_TILE_AREA, MAX_TILE_AREA]. Emphasise aspect-ratio fidelity to
Wa/Ha while staying within the allowed range.
Among options with comparable aspect error, prefer tiles that do
not hit clamped bounds, then maximise area and width (or minimise both if
size_preference == "smallest").
Assumes Wa, Ha are multiples of 16.
"""
dims = {
"input_width": input_width,
"input_height": input_height,
"target_width": target_width,
"target_height": target_height,
}
for name, value in dims.items():
if value <= 0:
raise ValueError(
f"🛑\u2009 Nilor-Nodes (NilorWanTileResolution): {name} must be a positive integer."
)
if input_width % 16 != 0 or input_height % 16 != 0:
raise ValueError(
"🛑\u2009 Nilor-Nodes (NilorWanTileResolution): input_width and input_height must be multiples of 16."
)
if target_width < self.MIN_TILE_DIM or target_height < self.MIN_TILE_DIM:
raise ValueError(
"🛑\u2009 Nilor-Nodes (NilorWanTileResolution): target_width and target_height must be at least the minimum tile size."
)
min_blocks = self.MIN_TILE_DIM // 16
max_blocks = self.MAX_TILE_DIM // 16
max_width_blocks = min(max_blocks, target_width // 16)
max_height_blocks = min(max_blocks, target_height // 16)
if max_width_blocks < min_blocks or max_height_blocks < min_blocks:
raise ValueError(
"🛑\u2009 Nilor-Nodes (NilorWanTileResolution): Target dimensions do not allow a tile within the supported range."
)
aspect_ratio = input_width / input_height
best_score = None
best_dimensions = None
for height_blocks in range(min_blocks, max_height_blocks + 1):
width_blocks = round(aspect_ratio * height_blocks)
width_blocks = self._clamp(width_blocks, min_blocks, max_width_blocks)
width_px = width_blocks * 16
height_px = height_blocks * 16
area = width_px * height_px
if area < self.MIN_TILE_AREA or area > self.MAX_TILE_AREA:
# Skip tiles that are too small or too large
continue
aspect_error = abs((width_blocks / height_blocks) - aspect_ratio)
width_hits_bound = int(width_blocks in (min_blocks, max_width_blocks))
height_hits_bound = int(height_blocks in (min_blocks, max_height_blocks))
# Penalise tiles that hit the clamped bounds
bound_penalty = width_hits_bound + height_hits_bound
# Score tiles based on size preference
if size_preference == "smallest":
area_score = -area
width_score = -width_px
else:
area_score = area
width_score = width_px
# Combine scores
candidate = (-aspect_error, -bound_penalty, area_score, width_score)
if best_score is None or candidate > best_score:
# Update best score and dimensions if this candidate is better
best_score = candidate
best_dimensions = (width_px, height_px)
if best_dimensions is None:
# If no suitable tile resolution was found, raise an error
raise RuntimeError(
"🛑\u2009 Nilor-Nodes (NilorWanTileResolution): Failed to determine a suitable tile resolution."
)
return best_dimensions
class NilorWanFrameTrim:
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "trim_to_wan_count"
CATEGORY = category + subcategories["utilities"]
def _validate_images(self, images):
if not isinstance(images, torch.Tensor):
raise TypeError(
"🛑\u2009 Nilor-Nodes (WanFrameTrim): images must be a torch.Tensor."
)
if images.dim() != 4:
raise ValueError(
f"🛑\u2009 Nilor-Nodes (WanFrameTrim): Expected 4D tensor (batch, height, width, channels), got shape {tuple(images.shape)}"
)
if images.shape[0] == 0:
raise ValueError(
"🛑\u2009 Nilor-Nodes (WanFrameTrim): Input images tensor is empty."
)
def trim_to_wan_count(self, images: torch.Tensor):
self._validate_images(images)
batch_count = images.shape[0]
# Find the largest m <= batch_count such that m ≡ 1 (mod 4)
wan_count = batch_count - ((batch_count - 1) % 4)
if wan_count <= 0:
raise ValueError(
"🛑\u2009 Nilor-Nodes (WanFrameTrim): Unable to compute a valid 4N+1 frame count from input."
)
trimmed = images[:wan_count]
return (trimmed,)
raise ValueError(f"Unknown type: {type}")
class NilorCategorizeString:
@@ -1203,9 +959,7 @@ class NilorRandomString:
if item.strip()
]
if not options:
raise ValueError(
"🛑\u2009 Nilor-Nodes (NilorRandomString): No valid choices provided."
)
raise ValueError("No valid choices provided.")
# Limit to the first 'max_options' entries if there are more options
if len(options) > max_options:
@@ -1247,9 +1001,7 @@ class NilorLoadImageByIndex:
def load_image_by_index(self, image_directory, seed, sort_mode, reverse_sort):
if not os.path.exists(image_directory):
raise FileNotFoundError(
f"🛑\u2009 Nilor-Nodes (NilorLoadImageByIndex): Image directory {image_directory} does not exist"
)
raise FileNotFoundError(f"Image directory {image_directory} does not exist")
# Get list of image files
files = []
@@ -1261,9 +1013,7 @@ class NilorLoadImageByIndex:
files.append(file_path)
if not files:
raise ValueError(
f"🛑\u2009 Nilor-Nodes (NilorLoadImageByIndex): No image files found in {image_directory}"
)
raise ValueError(f"No image files found in {image_directory}")
# Sort files based on selected mode
if sort_mode == "filename":
@@ -1313,9 +1063,7 @@ class NilorExtractFilenameFromPath:
def extract_filename(self, filepath):
# Ensure the input is a valid path
if not filepath:
raise ValueError(
"🛑\u2009 Nilor-Nodes (ExtractFilenameFromPath): Filepath cannot be empty."
)
raise ValueError("Filepath cannot be empty.")
path = Path(filepath)
@@ -1350,18 +1098,14 @@ class NilorBlurAnalysis:
"""
# Ensure images is a 4D tensor.
if images.dim() != 4:
raise ValueError(
"🛑\u2009 Nilor-Nodes (BlurAnalysis): Input images must be a 4D tensor (batch, channels/height, height/width, width/channels)"
)
raise ValueError("Input images must be a 4D tensor (batch, channels/height, height/width, width/channels)")
# Detect if using NCHW or NHWC.
if images.shape[1] not in (1, 3):
if images.shape[-1] in (1, 3):
images = images.permute(0, 3, 1, 2)
else:
raise ValueError(
"🛑\u2009 Nilor-Nodes (BlurAnalysis): Cannot determine image format (expected channel to be 1 or 3)."
)
raise ValueError("Cannot determine image format (expected channel to be 1 or 3).")
output_images = []
batch_size = images.shape[0]
@@ -1372,7 +1116,9 @@ class NilorBlurAnalysis:
# Convert to grayscale.
if img_np.shape[0] >= 3:
gray = 0.299 * img_np[0] + 0.587 * img_np[1] + 0.114 * img_np[2]
gray = (0.299 * img_np[0] +
0.587 * img_np[1] +
0.114 * img_np[2])
else:
gray = np.squeeze(img_np, axis=0) # shape: (H, W)
@@ -1400,9 +1146,7 @@ class NilorBlurAnalysis:
# Convert the single channel output to a 3-channel image.
# This ensures downstream nodes (like MaskFromRGBCMYBW) that index into channels work properly.
if out_img.ndim == 2:
out_img = np.stack(
[out_img, out_img, out_img], axis=-1
) # shape becomes (H, W, 3)
out_img = np.stack([out_img, out_img, out_img], axis=-1) # shape becomes (H, W, 3)
# Convert from PIL image (or numpy array) to tensor.
# pil2tensor should create a tensor in a format that downstream nodes expect.
@@ -1413,7 +1157,6 @@ class NilorBlurAnalysis:
# If each output has shape, say, (H, W, 3), stacking them gives a tensor of shape (B, H, W, 3).
return (torch.cat(output_images, dim=0),)
class NilorToSparseIndexMethod:
def __init__(self):
pass
@@ -1436,315 +1179,10 @@ class NilorToSparseIndexMethod:
def convert_to_sparse_index_method(self, ints):
indexes_str = ",".join(map(str, ints))
return (indexes_str,)
class NilorImageResizeV2:
upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"]
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"image": ("IMAGE",),
"width": ("INT", {"default": 512, "min": 0, "max": BIGMAX, "step": 1}),
"height": ("INT", {"default": 512, "min": 0, "max": BIGMAX, "step": 1}),
"upscale_method": (s.upscale_methods,),
"keep_proportion": (
[
"stretch",
"resize",
"pad",
"pad_edge",
"pad_edge_pixel",
"crop",
"pillarbox_blur",
],
{"default": False},
),
"pad_color": ("STRING", {"default": "0, 0, 0"}),
"crop_position": (
["center", "top", "bottom", "left", "right"],
{"default": "center"},
),
"divisible_by": (
"INT",
{"default": 2, "min": 0, "max": 512, "step": 1},
),
},
"optional": {
"mask": ("MASK",),
"device": (["cpu", "gpu"],),
"per_batch": (
"INT",
{
"default": 16,
"min": 0,
"max": 4096,
"step": 1,
"tooltip": "Process images in sub-batches. 0 disables.",
},
),
},
"hidden": {"unique_id": "UNIQUE_ID"},
}
RETURN_TYPES = ("IMAGE", "INT", "INT", "MASK")
RETURN_NAMES = ("IMAGE", "width", "height", "mask")
FUNCTION = "resize"
CATEGORY = category + subcategories["utilities"]
DESCRIPTION = """
Resizes images with optional aspect preservation, padding/cropping, and sub-batching to lower peak memory.
"""
def resize(
self,
image,
width,
height,
keep_proportion,
upscale_method,
divisible_by,
pad_color,
crop_position,
unique_id,
device="cpu",
mask=None,
per_batch=16,
):
B, H, W, C = image.shape
if device == "gpu":
if upscale_method == "lanczos":
raise Exception(
"🛑\u2009 Nilor-Nodes (NilorImageResizeV2): Lanczos is not supported on the GPU"
)
device = model_management.get_torch_device()
else:
device = torch.device("cpu")
if width == 0:
width = W
if height == 0:
height = H
pillarbox_blur = keep_proportion == "pillarbox_blur"
if (
keep_proportion == "resize"
or keep_proportion.startswith("pad")
or pillarbox_blur
):
if width == 0 and height != 0:
ratio = height / H
new_width = round(W * ratio)
new_height = height
elif height == 0 and width != 0:
ratio = width / W
new_width = width
new_height = round(H * ratio)
elif width != 0 and height != 0:
ratio = min(width / W, height / H)
new_width = round(W * ratio)
new_height = round(H * ratio)
else:
new_width = width
new_height = height
pad_left = pad_right = pad_top = pad_bottom = 0
if keep_proportion.startswith("pad") or pillarbox_blur:
if crop_position == "center":
pad_left = (width - new_width) // 2
pad_right = width - new_width - pad_left
pad_top = (height - new_height) // 2
pad_bottom = height - new_height - pad_top
elif crop_position == "top":
pad_left = (width - new_width) // 2
pad_right = width - new_width - pad_left
pad_top = 0
pad_bottom = height - new_height
elif crop_position == "bottom":
pad_left = (width - new_width) // 2
pad_right = width - new_width - pad_left
pad_top = height - new_height
pad_bottom = 0
elif crop_position == "left":
pad_left = 0
pad_right = width - new_width
pad_top = (height - new_height) // 2
pad_bottom = height - new_height - pad_top
elif crop_position == "right":
pad_left = width - new_width
pad_right = 0
pad_top = (height - new_height) // 2
pad_bottom = height - new_height - pad_top
width = new_width
height = new_height
if divisible_by > 1:
width = width - (width % divisible_by)
height = height - (height % divisible_by)
if per_batch and B > per_batch:
try:
bytes_per_elem = image.element_size()
est_total_bytes = B * height * width * C * bytes_per_elem
est_mb = est_total_bytes / (1024 * 1024)
logger.info(
f"ℹ️\u2009 Nilor-Nodes (NilorImageResizeV2) Estimated output ~{est_mb:.2f} MB."
)
except:
pass
def _process_subbatch(in_image, in_mask):
out_image = in_image if in_image.device == device else in_image.to(device)
out_mask = (
None
if in_mask is None
else (in_mask if in_mask.device == device else in_mask.to(device))
)
if keep_proportion == "crop":
old_height = out_image.shape[-3]
old_width = out_image.shape[-2]
old_aspect = old_width / old_height
new_aspect = width / height
if old_aspect > new_aspect:
crop_w = round(old_height * new_aspect)
crop_h = old_height
else:
crop_w = old_width
crop_h = round(old_width / new_aspect)
if crop_position == "center":
x = (old_width - crop_w) // 2
y = (old_height - crop_h) // 2
elif crop_position == "top":
x = (old_width - crop_w) // 2
y = 0
elif crop_position == "bottom":
x = (old_width - crop_w) // 2
y = old_height - crop_h
elif crop_position == "left":
x = 0
y = (old_height - crop_h) // 2
elif crop_position == "right":
x = old_width - crop_w
y = (old_height - crop_h) // 2
out_image = out_image.narrow(-2, x, crop_w).narrow(-3, y, crop_h)
if out_mask is not None:
out_mask = out_mask.narrow(-1, x, crop_w).narrow(-2, y, crop_h)
out_image = common_upscale(
out_image.movedim(-1, 1), width, height, upscale_method, crop="disabled"
).movedim(1, -1)
if out_mask is not None:
if upscale_method == "lanczos":
out_mask = common_upscale(
out_mask.unsqueeze(1).repeat(1, 3, 1, 1),
width,
height,
upscale_method,
crop="disabled",
).movedim(1, -1)[:, :, :, 0]
else:
out_mask = common_upscale(
out_mask.unsqueeze(1),
width,
height,
upscale_method,
crop="disabled",
).squeeze(1)
if (keep_proportion.startswith("pad") or pillarbox_blur) and (
pad_left > 0 or pad_right > 0 or pad_top > 0 or pad_bottom > 0
):
padded_width = width + pad_left + pad_right
padded_height = height + pad_top + pad_bottom
if divisible_by > 1:
width_remainder = padded_width % divisible_by
height_remainder = padded_height % divisible_by
if width_remainder > 0:
extra_width = divisible_by - width_remainder
pad_right += extra_width
if height_remainder > 0:
extra_height = divisible_by - height_remainder
pad_bottom += extra_height
pad_mode = (
"pillarbox_blur"
if pillarbox_blur
else (
"edge"
if keep_proportion == "pad_edge"
else (
"edge_pixel"
if keep_proportion == "pad_edge_pixel"
else "color"
)
)
)
out_image, out_mask = ImagePadKJ.pad(
self,
out_image,
pad_left,
pad_right,
pad_top,
pad_bottom,
0,
pad_color,
pad_mode,
mask=out_mask,
)
return out_image, out_mask
if per_batch is None or per_batch == 0 or B <= per_batch:
out_image, out_mask = _process_subbatch(image, mask)
else:
chunks = []
mask_chunks = [] if mask is not None else None
total_batches = (B + per_batch - 1) // per_batch
current_batch = 0
for start_idx in range(0, B, per_batch):
current_batch += 1
end_idx = min(start_idx + per_batch, B)
sub_img = image[start_idx:end_idx]
sub_mask = mask[start_idx:end_idx] if mask is not None else None
sub_out_img, sub_out_mask = _process_subbatch(sub_img, sub_mask)
chunks.append(sub_out_img.cpu())
if mask is not None:
mask_chunks.append(
sub_out_mask.cpu() if sub_out_mask is not None else None
)
try:
logger.info(
f"ℹ️\u2009 Nilor-Nodes (NilorImageResizeV2) Batch {current_batch}/{total_batches} · images {end_idx}/{B}"
)
except:
pass
out_image = torch.cat(chunks, dim=0)
if mask is not None and any(m is not None for m in mask_chunks):
out_mask = torch.cat([m for m in mask_chunks if m is not None], dim=0)
else:
out_mask = None
logger.info(f"✅\u2009 Nilor-Nodes (NilorImageResizeV2) All batches complete.")
return (
out_image.cpu(),
out_image.shape[2],
out_image.shape[1],
(
out_mask.cpu()
if out_mask is not None
else torch.zeros(
64, 64, device=torch.device("cpu"), dtype=torch.float32
)
),
)
# Mapping class names to objects for potential export
NODE_CLASS_MAPPINGS = {
"Nilor Interpolated Float List": NilorInterpolatedFloatList,
@@ -1766,22 +1204,19 @@ NODE_CLASS_MAPPINGS = {
"Nilor n Fractions of Int": NilorNFractionsOfInt,
"Nilor Categorize String": NilorCategorizeString,
"Nilor Random String": NilorRandomString,
"Nilor Wan Tile Resolution": NilorWanTileResolution,
"Nilor Extract Filename from Path": NilorExtractFilenameFromPath,
"Nilor Load Image By Index": NilorLoadImageByIndex,
"Nilor Blur Analysis": NilorBlurAnalysis,
"Nilor To Sparse Index Method": NilorToSparseIndexMethod,
"Nilor Image Resize v2": NilorImageResizeV2,
"Nilor Wan Frame Trim": NilorWanFrameTrim,
}
# Mapping nodes to human-readable names
NODE_DISPLAY_NAME_MAPPINGS = {
"Nilor Interpolated Float List": "👺 Interpolated Float List",
"Nilor One Minus Float List": "👺 One Minus Float List",
"Nilor Remap Float List": "👺 Remap Float List",
"Nilor Remap Float List Auto Input": "👺 Remap Float List Auto Input",
"Nilor Inverse Map Float List": "👺 Inverse Map Float List",
"Nilor Remap Float List": "👺 Nilor Remap Float List",
"Nilor Remap Float List Auto Input": "👺 Nilor Remap Float List Auto Input",
"Nilor Inverse Map Float List": "👺 Nilor Inverse Map Float List",
"Nilor Int To List Of Bools": "👺 Int To List Of Bools",
"Nilor List of Ints": "👺 List of Ints",
"Nilor Count Images In Directory": "👺 Count Images In Directory",
@@ -1789,18 +1224,15 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"Nilor Save Video To HF Dataset": "👺 Save Video To HF Dataset",
"Nilor Select Index From List": "👺 Select Index From List",
"Nilor Save EXR Arbitrary": "👺 Save EXR Arbitrary",
"Nilor Shuffle Image Batch": "👺 Shuffle Image Batch",
"Nilor Repeat & Trim Image Batch": "👺 Repeat & Trim Image Batch",
"Nilor Repeat, Shuffle, & Trim Image Batch": "👺 Repeat, Shuffle, & Trim Image Batch",
"Nilor Output Filename String": "👺 Output Filename String",
"Nilor n Fractions of Int": "👺 n Fractions of Int",
"Nilor Shuffle Image Batch": "👺 Nilor Shuffle Image Batch",
"Nilor Repeat & Trim Image Batch": "👺 Nilor Repeat & Trim Image Batch",
"Nilor Repeat, Shuffle, & Trim Image Batch": "👺 Nilor Repeat, Shuffle, & Trim Image Batch",
"Nilor Output Filename String": "👺 Nilor Output Filename String",
"Nilor n Fractions of Int": "👺 Nilor n Fractions of Int",
"Nilor Categorize String": "👺 Categorize String",
"Nilor Random String": "👺 Random String",
"Nilor Wan Tile Resolution": "👺 Wan Tile Resolution",
"Nilor Extract Filename from Path": "👺 Extract Filename from Path",
"Nilor Load Image By Index": "👺 Load Image By Index",
"Nilor Blur Analysis": "👺 Blur Analysis",
"Nilor To Sparse Index Method": "👺 To Sparse Index Method",
"Nilor Image Resize v2": "👺 Resize Image v2",
"Nilor Wan Frame Trim": "👺 Wan Frame Trim",
}
+2 -20
View File
@@ -1,20 +1,2 @@
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
huggingface_hub
openexr
+151
View File
@@ -0,0 +1,151 @@
#!/usr/bin/env python3
"""
Test script for Brain API Client
This script tests the Brain API client functionality to ensure it can
communicate with the Brain API storage endpoints correctly.
"""
import os
import sys
import tempfile
import logging
from pathlib import Path
# Add the current directory to the Python path
current_dir = Path(__file__).parent
sys.path.insert(0, str(current_dir))
from brain_api_client import get_brain_api_client
# Setup logging
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
def test_brain_api_client():
"""Test the Brain API client functionality."""
print("🧪 Testing Brain API Client...")
try:
# Initialize the client
client = get_brain_api_client()
print("✅ Brain API client initialized successfully")
# Test health check
print("🔍 Testing health check...")
is_healthy = client.health_check()
if is_healthy:
print("✅ Brain API is accessible")
else:
print("⚠️ Brain API health check failed - this might be expected if the API is not running")
# Test file upload
print("📤 Testing file upload...")
test_content = b"Hello, Brain API! This is a test file."
test_filename = "test_file.txt"
# Create a temporary file
with tempfile.NamedTemporaryFile(mode='wb', delete=False, suffix='.txt') as temp_file:
temp_file.write(test_content)
temp_file_path = temp_file.name
try:
# Upload the file
upload_result = client.upload_file_to_storage(temp_file_path, test_filename)
print(f"✅ File uploaded successfully. Storage ID: {upload_result.get('storage_id')}")
storage_id = upload_result.get('storage_id')
if storage_id:
# Test file download
print("📥 Testing file download...")
downloaded_content = client.get_file_from_storage(storage_id, test_filename)
if downloaded_content == test_content:
print("✅ File downloaded successfully and content matches")
else:
print("❌ Downloaded content does not match original")
# Test file deletion
print("🗑️ Testing file deletion...")
client.delete_file_from_storage(storage_id, test_filename)
print("✅ File deleted successfully")
finally:
# Clean up temporary file
os.unlink(temp_file_path)
print("🎉 All tests passed!")
return True
except Exception as e:
print(f"❌ Test failed: {e}")
logging.exception("Test failed with exception:")
return False
def test_fileobj_upload():
"""Test uploading a file-like object."""
print("\n🧪 Testing file object upload...")
try:
client = get_brain_api_client()
# Create a file-like object
import io
test_content = b"Hello from file object!"
file_obj = io.BytesIO(test_content)
# Upload the file object
upload_result = client.upload_fileobj_to_storage(file_obj, "test_fileobj.txt", "text/plain")
print(f"✅ File object uploaded successfully. Storage ID: {upload_result.get('storage_id')}")
storage_id = upload_result.get('storage_id')
if storage_id:
# Test download
downloaded_content = client.get_file_from_storage(storage_id, "test_fileobj.txt")
if downloaded_content == test_content:
print("✅ File object download successful and content matches")
else:
print("❌ Downloaded content does not match original")
# Clean up
client.delete_file_from_storage(storage_id, "test_fileobj.txt")
print("✅ File object deleted successfully")
return True
except Exception as e:
print(f"❌ File object test failed: {e}")
logging.exception("File object test failed with exception:")
return False
if __name__ == "__main__":
print("🚀 Starting Brain API Client Tests")
print("=" * 50)
# Check environment variables
api_key = os.getenv("BRANDO_API_KEY")
base_url = os.getenv("BRANDO_BRAIN_API_BASE_URL", "http://localhost:2024/api")
print(f"API Key: {'✅ Set' if api_key else '❌ Not set'}")
print(f"Base URL: {base_url}")
print()
if not api_key:
print("❌ BRANDO_API_KEY environment variable is not set!")
print("Please set it in your .env file or environment.")
sys.exit(1)
# Run tests
success = True
success &= test_brain_api_client()
success &= test_fileobj_upload()
print("\n" + "=" * 50)
if success:
print("🎉 All tests completed successfully!")
sys.exit(0)
else:
print("❌ Some tests failed!")
sys.exit(1)
-17
View File
@@ -1,17 +0,0 @@
"""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"
+1 -83
View File
@@ -3,9 +3,6 @@ subcategories = {
"io": "/IO",
}
import random
from datetime import datetime
from .controllers import CONTROLLER_HOOK
@@ -53,83 +50,6 @@ class NilorUserInput_Int:
return (value, None)
class NilorUserInput_Seed:
MAX_COMFYUI_SEED = 1125899906842624
SEED_RANDOM_STATE = None
@classmethod
def _ensure_seed_random_state(cls):
if cls.SEED_RANDOM_STATE is not None:
return
initial_random_state = random.getstate()
random.seed(datetime.now().timestamp())
cls.SEED_RANDOM_STATE = random.getstate()
random.setstate(initial_random_state)
@classmethod
def generate_random_seed(cls):
cls._ensure_seed_random_state()
prev_random_state = random.getstate()
random.setstate(cls.SEED_RANDOM_STATE)
seed = random.randint(0, cls.MAX_COMFYUI_SEED)
cls.SEED_RANDOM_STATE = random.getstate()
random.setstate(prev_random_state)
return seed
@classmethod
def resolve_seed(cls, value):
if value in (None, 0, -1):
return cls.generate_random_seed()
try:
return int(value) % (cls.MAX_COMFYUI_SEED + 1)
except (TypeError, ValueError):
return cls.generate_random_seed()
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"input_name": (
"STRING",
{"default": "my_seed_input", "multiline": False},
),
"value": (
"INT",
{
"default": -1,
"min": -1,
"max": cls.MAX_COMFYUI_SEED,
},
),
},
"hidden": {
"prompt": "PROMPT",
"extra_pnginfo": "EXTRA_PNGINFO",
"unique_id": "UNIQUE_ID",
},
}
RETURN_TYPES = ("INT", CONTROLLER_HOOK)
RETURN_NAMES = ("seed", "_controller_hook")
FUNCTION = "get_value"
CATEGORY = category + subcategories["io"]
@classmethod
def IS_CHANGED(
cls, input_name, value, prompt=None, extra_pnginfo=None, unique_id=None
):
# Force node re-execution while using randomize sentinel values.
return cls.resolve_seed(value)
def get_value(
self, input_name, value, prompt=None, extra_pnginfo=None, unique_id=None
):
value = self.resolve_seed(value)
return (value, None)
class NilorUserInput_Float:
@classmethod
def INPUT_TYPES(cls):
@@ -139,7 +59,7 @@ class NilorUserInput_Float:
"STRING",
{"default": "my_float_input", "multiline": False},
),
"value": ("FLOAT", {"default": 0.0, "step": 0.001}),
"value": ("FLOAT", {"default": 0.0}),
}
}
@@ -177,7 +97,6 @@ class NilorUserInput_Boolean:
NODE_CLASS_MAPPINGS = {
"NilorUserInput_String": NilorUserInput_String,
"NilorUserInput_Int": NilorUserInput_Int,
"NilorUserInput_Seed": NilorUserInput_Seed,
"NilorUserInput_Float": NilorUserInput_Float,
"NilorUserInput_Boolean": NilorUserInput_Boolean,
}
@@ -185,7 +104,6 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"NilorUserInput_String": "👺 User Input (String)",
"NilorUserInput_Int": "👺 User Input (Int)",
"NilorUserInput_Seed": "👺 User Input (Seed)",
"NilorUserInput_Float": "👺 User Input (Float)",
"NilorUserInput_Boolean": "👺 User Input (Boolean)",
}
+2 -3
View File
@@ -15,9 +15,8 @@ def numpy2pil(image: np.ndarray, mode=None):
## 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(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)
+37 -88
View File
@@ -19,96 +19,45 @@ function hideWidgets(node, widgetNames) {
});
}
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);
if (node.comfyClass === "MediaStreamOutput") {
// Hide system inputs by default
hideWidgets(node, [
"content_id",
"venue",
"canvas",
"scene",
"presigned_upload_url",
"job_completions_queue_url",
"output_object_keys"
]);
const formatWidget = node.widgets.find((w) => w.name === "format");
// Initial toggle for framerate based on the default format value
toggleFramerateWidget(node, formatWidget.value === "mp4");
// Store original callback to chain it
const originalCallback = formatWidget.callback;
formatWidget.callback = function (value) {
toggleFramerateWidget(node, value === "mp4");
// Recalculate node size after toggling widgets
const size = node.computeSize();
node.onResize?.(size);
if (originalCallback) {
return originalCallback.apply(this, arguments);
}
};
}
if (node.comfyClass === "MediaStreamInput") {
// Hide system inputs by default
hideWidgets(node, ["presigned_download_url"]);
}
},
});
+212 -485
View File
@@ -2,7 +2,7 @@
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,
Its purpose is to poll the `jobs_to_process` SQS queue for new jobs,
submit them to the local ComfyUI server, and manage the message lifecycle.
It also listens to the ComfyUI websocket to send a "running" status update
at the precise moment that job execution begins.
@@ -12,29 +12,50 @@ import os
import json
import logging
import asyncio
import time
import aiohttp
import websockets
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
from dotenv import load_dotenv
from botocore.exceptions import EndpointConnectionError
# --- Load Environment Variables ---
# Load from the .env file in the same directory
current_dir = os.path.dirname(os.path.abspath(__file__))
dotenv_path = os.path.join(current_dir, ".env")
if os.path.exists(dotenv_path):
load_dotenv(dotenv_path=dotenv_path)
logging.info(
f"✅\u2009 Nilor-Nodes: Loaded environment variables from {dotenv_path}"
)
else:
logging.info(
"⚠️\u2009 Nilor-Nodes: No .env file found, relying on shell environment variables."
)
# --- Configuration ---
# Centralized loader provides precedence env > JSON5 and validation
_CFG: NilorNodesConfig = load_nilor_nodes_config()
SQS_ENDPOINT_URL = os.getenv("SQS_ENDPOINT_URL", "http://localhost:9324")
SQS_JOBS_TO_PROCESS_QUEUE_NAME = os.getenv(
"SQS_JOBS_TO_PROCESS_QUEUE_NAME", "jobs_to_process"
)
SQS_JOB_STATUS_UPDATES_QUEUE_NAME = os.getenv(
"SQS_JOB_STATUS_UPDATES_QUEUE_NAME", "job_status_updates"
)
COMFYUI_API_URL = os.getenv("COMFYUI_API_URL", "http://127.0.0.1:8188") + "/prompt"
COMFYUI_WS_URL = os.getenv("COMFYUI_WS_URL", "ws://127.0.0.1:8188") + "/ws"
AWS_ACCESS_KEY_ID = os.getenv("AWS_ACCESS_KEY_ID", "local")
AWS_SECRET_ACCESS_KEY = os.getenv("AWS_SECRET_ACCESS_KEY", "local")
AWS_DEFAULT_REGION = os.getenv("AWS_DEFAULT_REGION", "us-east-1")
POLL_WAIT_TIME_SECONDS = 20 # SQS Long Polling
MAX_MESSAGES = 1
class JobSubmissionError(Exception):
"""Raised when a job cannot be submitted to local ComfyUI."""
# --- Setup Logging ---
logging.basicConfig(
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
)
class WorkerConsumer:
def __init__(self, cfg: NilorNodesConfig):
def __init__(self):
self.session = get_session()
self.prompt_id_to_content_id_map = {}
self.sent_running_status_prompts = set()
@@ -42,45 +63,32 @@ class WorkerConsumer:
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,
region_name=AWS_DEFAULT_REGION,
endpoint_url=SQS_ENDPOINT_URL,
aws_access_key_id=AWS_ACCESS_KEY_ID,
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
) as client:
try:
self.jobs_queue_url = await self._get_queue_url(
client, self.cfg.worker.jobs_queue
client, SQS_JOBS_TO_PROCESS_QUEUE_NAME
)
self.status_updates_queue_url = await self._get_queue_url(
client, self.cfg.worker.status_queue
client, SQS_JOB_STATUS_UPDATES_QUEUE_NAME
)
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}. "
logging.warning(
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS endpoint is unreachable at {SQS_ENDPOINT_URL}: {e}. "
)
return False
except Exception as e:
logger.error(
logging.error(
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Failed to initialize SQS queues: {e}"
)
return False
@@ -91,7 +99,7 @@ class WorkerConsumer:
response = await client.get_queue_url(QueueName=queue_name)
return response["QueueUrl"]
except client.exceptions.QueueDoesNotExist:
logger.error(
logging.error(
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS queue '{queue_name}' does not exist."
)
raise
@@ -99,128 +107,130 @@ class WorkerConsumer:
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."
async with websockets.connect(COMFYUI_WS_URL) as websocket:
logging.info(
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Connected to ComfyUI websocket at {COMFYUI_WS_URL}"
)
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")
while True:
message = await websocket.recv()
if isinstance(message, str):
try:
await self._send_status_update(
content_id,
fail_status,
ctx.get("venue"),
ctx.get("canvas"),
ctx.get("scene"),
ctx.get("job_type"),
event = json.loads(message)
event_type = event.get("type")
data = event.get("data", {})
prompt_id = data.get("prompt_id")
if not prompt_id and "sid" in data:
prompt_id = data["sid"]
if not prompt_id:
continue
# Use the first progress event as a signal that the job is running.
if (
event_type in ["progress", "progress_state"]
and prompt_id in self.prompt_id_to_content_id_map
and prompt_id
not in self.sent_running_status_prompts
):
content_id = self.prompt_id_to_content_id_map[
prompt_id
]
ctx = self.content_context_by_content_id.get(
content_id, {}
)
policy = ctx.get("status_policy") or {}
running_status = policy.get(
"running_status", "running"
)
logging.info(
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Execution started for prompt_id {prompt_id} (content_id: {content_id}) via '{event_type}' event. Sending '{running_status}' status."
)
await self._send_status_update(
content_id,
running_status,
ctx.get("venue"),
ctx.get("canvas"),
ctx.get("scene"),
)
self.sent_running_status_prompts.add(prompt_id)
# Handle execution errors
elif event_type == "execution_error":
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): Received execution error for prompt_id {prompt_id}: {data}"
)
if prompt_id in self.prompt_id_to_content_id_map:
content_id = (
self.prompt_id_to_content_id_map.pop(
prompt_id
)
)
ctx = self.content_context_by_content_id.get(
content_id, {}
)
policy = ctx.get("status_policy") or {}
fail_status = policy.get(
"fail_status", "failed"
)
try:
await self._send_status_update(
content_id,
fail_status,
ctx.get("venue"),
ctx.get("canvas"),
ctx.get("scene"),
)
except Exception:
pass
self.content_context_by_content_id.pop(
content_id, None
)
self.sent_running_status_prompts.discard(prompt_id)
# Log successful execution
elif event_type == "executed":
logging.info(
f"✅ Nilor-Nodes (worker_consumer): Prompt {prompt_id} executed successfully according to websocket event. Final node is responsible for sending completion message."
)
if prompt_id in self.prompt_id_to_content_id_map:
content_id = (
self.prompt_id_to_content_id_map.pop(
prompt_id
)
)
self.content_context_by_content_id.pop(
content_id, None
)
self.sent_running_status_prompts.discard(prompt_id)
elif event_type not in ["progress", "progress_state"]:
logging.info(
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Received ComfyUI websocket event of type '{event_type}': {data}"
)
except json.JSONDecodeError:
logging.debug(
"⚠️\u2009 Nilor-Nodes (worker_consumer): Received non-JSON text message from websocket, ignoring."
)
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",
else:
logging.debug(
"⚠️\u2009 Nilor-Nodes (worker_consumer): Received binary message from websocket, ignoring."
)
# 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,
except (
websockets.exceptions.ConnectionClosedError,
ConnectionRefusedError,
) as e:
logging.warning(
f"🛑\u2009 Nilor-Nodes (worker_consumer): ComfyUI websocket connection failed: {e}. Retrying in 5 seconds..."
)
await asyncio.sleep(5)
except Exception as e:
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): An unexpected error occurred in the websocket listener: {e}",
exc_info=True,
)
await asyncio.sleep(10)
async def consume_loop(self):
"""The main loop to continuously poll for and process messages.
@@ -231,7 +241,6 @@ class WorkerConsumer:
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:
@@ -242,48 +251,30 @@ class WorkerConsumer:
)
await asyncio.sleep(10)
continue
logger.info(
logging.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(
logging.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,
region_name=AWS_DEFAULT_REGION,
endpoint_url=SQS_ENDPOINT_URL,
aws_access_key_id=AWS_ACCESS_KEY_ID,
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
) as client:
response = await client.receive_message(
QueueUrl=self.jobs_queue_url,
MaxNumberOfMessages=self.cfg.worker.max_messages,
WaitTimeSeconds=self.cfg.worker.poll_wait_s,
MaxNumberOfMessages=MAX_MESSAGES,
WaitTimeSeconds=POLL_WAIT_TIME_SECONDS,
)
messages = response.get("Messages", [])
if not messages:
logger.debug(
logging.debug(
"ℹ️\u2009 Nilor-Nodes (worker_consumer): No messages received."
)
continue
@@ -295,80 +286,52 @@ class WorkerConsumer:
# 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,
region_name=AWS_DEFAULT_REGION,
endpoint_url=SQS_ENDPOINT_URL,
aws_access_key_id=AWS_ACCESS_KEY_ID,
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
) as client:
await client.delete_message(
QueueUrl=self.jobs_queue_url,
ReceiptHandle=message["ReceiptHandle"],
)
logger.debug(
logging.info(
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Deleted message {message['MessageId']} from queue."
)
except json.JSONDecodeError:
# This is a poison pill message, log it but don't retry.
# It will be moved to the DLQ after enough failed receives.
logger.error(
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): Message {message['MessageId']} is a poison pill (JSON decode failed) and will be ignored."
)
except Exception as e:
logger.error(
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): Processing failed for message {message['MessageId']}: {e}. It will be returned to the queue for retry."
)
except 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."
logging.warning(
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Lost connection to SQS at {SQS_ENDPOINT_URL}: {e}. Will retry initialization in 10 seconds."
)
self.jobs_queue_url = None
self.status_updates_queue_url = None
await asyncio.sleep(10)
except Exception as e:
logger.error(
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): An error occurred in the consume loop: {e}"
)
await asyncio.sleep(10) # Wait before retrying
finally:
listener_task.cancel()
await asyncio.gather(listener_task, return_exceptions=True)
# Close shared HTTP session if created
if self.http_session is not None:
try:
await self.http_session.close()
except Exception:
pass
logger.info(
logging.info(
"⚠️\u2009 Nilor-Nodes (worker_consumer): Websocket listener stopped."
)
async def process_message(self, message):
"""Processes a single SQS message."""
logger.info(
logging.info(
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Processing message: {message['MessageId']}"
)
@@ -382,38 +345,15 @@ class WorkerConsumer:
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(
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): Invalid message format: missing 'content_id' or 'prompt'. Payload: {job_payload}"
)
return
# Submit to ComfyUI
try:
await self._submit_job_to_comfyui(content_id, job_payload)
except JobSubmissionError as e:
logger.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): Submission failed for content_id {content_id}: {e}. Message will be retried/DLQ'd."
)
await self._emit_failed_status_for_submission_error(
content_id=content_id,
job_payload=job_payload,
error_message=str(e),
)
# Re-raise so consume_loop does not delete the message.
raise
await self._submit_job_to_comfyui(content_id, job_payload)
# Cache context for subsequent status updates
try:
@@ -421,122 +361,50 @@ class WorkerConsumer:
"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(
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): An unexpected error occurred while processing message: {e}. It will be retried."
)
# Re-raise to prevent deletion from queue if we want SQS to handle retry
raise
async def _emit_failed_status_for_submission_error(
self, content_id, job_payload, error_message: str
):
"""Best-effort failed status emission for submit-time errors."""
policy = job_payload.get("status_policy") or {}
fail_status = policy.get("fail_status", "failed")
await self._send_status_update(
content_id,
fail_status,
job_payload.get("venue"),
job_payload.get("canvas"),
job_payload.get("scene"),
job_payload.get("job_type"),
)
logger.info(
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Emitted failed status '%s' for content_id %s after submit error: %s",
fail_status,
content_id,
error_message,
)
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
async with aiohttp.ClientSession() as session:
async with session.post(
COMFYUI_API_URL, json=workflow_data, timeout=30
) as response:
response.raise_for_status()
response_json = await response.json()
prompt_id = response_json.get("prompt_id")
logging.info(
f"✅ Nilor-Nodes (worker_consumer): Successfully submitted job to ComfyUI. Prompt ID: {prompt_id}"
)
self.prompt_id_to_content_id_map[prompt_id] = content_id
# No need to delete here, the consume_loop handles message deletion
except ComfyUIClientError as e:
logger.error(
except aiohttp.ClientError as e:
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to submit job to ComfyUI: {e}. Message will be retried."
)
raise JobSubmissionError(str(e)) from e
except (json.JSONDecodeError, KeyError) as e:
logger.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to parse ComfyUI response: {e}. Message will be retried."
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to parse ComfyUI response: {e}. Discarding malformed response."
)
raise JobSubmissionError(f"Malformed ComfyUI response: {e}") from e
except Exception as e:
logger.error(
logging.error(
f"🛑\u2009 Nilor-Nodes (worker_consumer): An unexpected error occurred while submitting job to ComfyUI: {e}",
exc_info=True,
)
raise JobSubmissionError(str(e)) from e
async def _send_status_update(
self, content_id, status, venue=None, canvas=None, scene=None, job_type=None
self, content_id, status, venue=None, canvas=None, scene=None
):
try:
body = {"content_id": content_id, "status": status}
@@ -546,169 +414,28 @@ class WorkerConsumer:
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,
region_name=AWS_DEFAULT_REGION,
endpoint_url=SQS_ENDPOINT_URL,
aws_access_key_id=AWS_ACCESS_KEY_ID,
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
) as client:
await client.send_message(
QueueUrl=self.status_updates_queue_url, MessageBody=message_body
)
logger.info(
logging.info(
f"✅ Nilor-Nodes (worker_consumer): Sent status update for content {content_id}: {status}"
)
except Exception as e:
logger.error(
logging.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
consumer = WorkerConsumer()
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
-173
View File
@@ -1,173 +0,0 @@
"""
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