Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8e811b11bd | ||
|
|
5ed354b3a1 | ||
|
|
62b71b2978 | ||
|
|
d266978e49 | ||
|
|
552252b84c | ||
|
|
66c66f045f | ||
|
|
266fe598c6 | ||
|
|
ef20c72aff | ||
|
|
527e8dc4ab | ||
|
|
5d05aeaf68 | ||
|
|
5668e09a79 | ||
|
|
a94c1a9d60 | ||
|
|
675fce6a16 | ||
|
|
c1f5eb352c | ||
|
|
0f22deb027 | ||
|
|
58ff7a0bce | ||
|
|
c33b06a092 | ||
|
|
d5109ae0b3 | ||
|
|
63b44ebff4 | ||
|
|
3476fbfbf9 | ||
|
|
aafaf87d18 | ||
|
|
790a3ce38e | ||
|
|
114a3d661f | ||
|
|
b6bf6c1b6c | ||
|
|
3c0bf2222c | ||
|
|
ff0c7f1209 | ||
|
|
75a632b2e5 | ||
|
|
2f707ebc48 | ||
|
|
14862c1712 | ||
|
|
7320a15d08 | ||
|
|
4a56d4dff5 | ||
|
|
5e579b1aed | ||
|
|
9883728f55 | ||
|
|
bfba7b0265 | ||
|
|
a093a51291 | ||
|
|
848899298c | ||
|
|
2333a0d45d | ||
|
|
87b846069e | ||
|
|
6d17e664bb | ||
|
|
c541bc18dd | ||
|
|
7beda857c8 | ||
|
|
4fa9710aea | ||
|
|
e8c0d451d9 | ||
|
|
a491e6463c | ||
|
|
5b8184cc33 | ||
|
|
6b2b482b89 | ||
|
|
ea13b61f21 | ||
|
|
663d2dd739 | ||
|
|
b940ad79ea | ||
|
|
baba20c3c4 | ||
|
|
f3dbe8a3ed | ||
|
|
cb9a98edf1 | ||
|
|
e8ee057dbe | ||
|
|
5f7d00560b |
+53
-21
@@ -1,26 +1,58 @@
|
||||
# --- 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
|
||||
# NILOR_LOG_LEVEL=INFO # possible log levels: INFO, DEBUG, WARNING, ERROR, CRITICAL
|
||||
|
||||
# --- SQS SETTINGS ---
|
||||
## Toggles functionality for the SQS Worker Consumer
|
||||
SQS_ENABLED=false
|
||||
# --- 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
|
||||
|
||||
## 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
|
||||
## 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
|
||||
|
||||
### The specific SQS queue the worker should push job status updates to.
|
||||
SQS_JOB_STATUS_UPDATES_QUEUE_NAME=job_status_updates
|
||||
## WebSocket reconnect policy
|
||||
# NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS=5
|
||||
# NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS=30.0
|
||||
|
||||
## 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
|
||||
# --- 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
|
||||
|
||||
## (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
|
||||
## 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
|
||||
|
||||
@@ -2,6 +2,10 @@
|
||||
|
||||
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>
|
||||
@@ -200,11 +204,45 @@ The `worker_consumer.py` script is a background service that runs on each ComfyU
|
||||
**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>
|
||||
|
||||
|
||||
+15
-6
@@ -1,6 +1,7 @@
|
||||
import os
|
||||
import threading
|
||||
import asyncio
|
||||
import logging
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# --- Nilor-Nodes Custom Node Registration and Startup ---
|
||||
@@ -19,6 +20,12 @@ 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 ---
|
||||
|
||||
|
||||
@@ -29,18 +36,20 @@ def start_consumer_loop():
|
||||
asyncio.run(consume_jobs())
|
||||
|
||||
|
||||
# Start the SQS Worker Consumer (controlled by SQS_ENABLED)
|
||||
raw_sqs_enabled = os.getenv("SQS_ENABLED", "false")
|
||||
env_sqs_enabled = raw_sqs_enabled.strip().lower() == "true"
|
||||
if env_sqs_enabled:
|
||||
from .config.config import load_nilor_nodes_config
|
||||
|
||||
cfg = load_nilor_nodes_config()
|
||||
|
||||
# Start the SQS Worker Consumer (controlled by NILOR_SQS_ENABLED)
|
||||
if cfg.sqs_enabled:
|
||||
consumer_thread = threading.Thread(target=start_consumer_loop, daemon=True)
|
||||
consumer_thread.start()
|
||||
print(
|
||||
f"✅ Nilor-Nodes: SQS worker consumer thread started (SQS_ENABLED={raw_sqs_enabled} in .env)."
|
||||
"✅ Nilor-Nodes: SQS worker consumer thread started (NILOR_SQS_ENABLED=true)."
|
||||
)
|
||||
else:
|
||||
print(
|
||||
f"⚠️ Nilor-Nodes: SQS worker consumer functionality is disabled (SQS_ENABLED={raw_sqs_enabled} in .env)."
|
||||
"⚠️ Nilor-Nodes: SQS worker consumer functionality is disabled (NILOR_SQS_ENABLED=false)."
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,910 @@
|
||||
"""
|
||||
Thin, typed client surface for accessing ComfyUI HTTP endpoints and the websocket.
|
||||
|
||||
This module defines the public protocol, DTOs, and exceptions that callers and
|
||||
tests depend on. Implementations are intentionally minimal at this stage; network
|
||||
behavior will be added in subsequent commits.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, AsyncIterator, Dict, Optional, Protocol, TypedDict, Tuple
|
||||
|
||||
from urllib.parse import quote, urlparse
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ComfyUIClientProtocol",
|
||||
"ComfyUILocalClient",
|
||||
"SystemStats",
|
||||
"WsEvent",
|
||||
"ComfyUIClientError",
|
||||
"ComfyUIClientTimeout",
|
||||
"ComfyUIClientWsClosed",
|
||||
]
|
||||
|
||||
|
||||
class WsEvent(TypedDict, total=False):
|
||||
"""Typed view of websocket events emitted by ComfyUI.
|
||||
|
||||
Fields:
|
||||
- type: Event type string (e.g., "status", "progress", "executed").
|
||||
- data: Opaque payload; commonly includes keys like "prompt_id", "node", etc.
|
||||
"""
|
||||
|
||||
type: str
|
||||
data: Dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class SystemStats:
|
||||
"""Subset of system statistics reported by ComfyUI `/system_stats`.
|
||||
|
||||
Known fields are optional; unknown fields should be ignored by parsers. When
|
||||
present, RAM-related metrics may also appear (e.g., `ram_total`, `ram_free`).
|
||||
|
||||
Attributes:
|
||||
vram_total: Total VRAM (bytes) when reported.
|
||||
vram_free: Free VRAM reported by the backend, if available.
|
||||
torch_vram_free: Free VRAM according to torch, if available.
|
||||
ram_total: Total system RAM (bytes) when reported.
|
||||
ram_free: Free system RAM (bytes) when reported.
|
||||
"""
|
||||
|
||||
vram_total: Optional[float] = None
|
||||
vram_free: Optional[float] = None
|
||||
torch_vram_free: Optional[float] = None
|
||||
ram_total: Optional[float] = None
|
||||
ram_free: Optional[float] = None
|
||||
|
||||
|
||||
class ComfyUIClientError(Exception):
|
||||
"""Base error for all ComfyUI client failures.
|
||||
|
||||
Args:
|
||||
message: Human-friendly error message.
|
||||
route: Route path (e.g., "/prompt").
|
||||
method: HTTP method (e.g., "GET", "POST").
|
||||
status: Optional HTTP status code or websocket close code.
|
||||
code: Optional machine-readable error code (e.g., "timeout").
|
||||
body_snippet: Optional diagnostic snippet from a response payload.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
route: Optional[str] = None,
|
||||
method: Optional[str] = None,
|
||||
status: Optional[int] = None,
|
||||
code: Optional[str] = None,
|
||||
body_snippet: Optional[str] = None,
|
||||
) -> None:
|
||||
super().__init__(message)
|
||||
self.route: Optional[str] = route
|
||||
self.method: Optional[str] = method
|
||||
self.status: Optional[int] = status
|
||||
self.code: Optional[str] = code
|
||||
self.body_snippet: Optional[str] = body_snippet
|
||||
|
||||
|
||||
class ComfyUIClientTimeout(ComfyUIClientError):
|
||||
"""Raised when an operation exceeds its allowed timeout."""
|
||||
|
||||
|
||||
class ComfyUIClientWsClosed(ComfyUIClientError):
|
||||
"""Raised when the websocket is closed or cannot be maintained."""
|
||||
|
||||
|
||||
class ComfyUIClientProtocol(Protocol):
|
||||
"""Protocol for ComfyUI clients.
|
||||
|
||||
Callers and tests should depend on this interface rather than a concrete
|
||||
implementation. Methods are asynchronous and may raise subclasses of
|
||||
`ComfyUIClientError`.
|
||||
"""
|
||||
|
||||
async def submit_prompt(self, payload: Dict[str, Any]) -> str:
|
||||
"""Submit a prompt to ComfyUI and return the resulting `prompt_id`.
|
||||
|
||||
Args:
|
||||
payload: JSON-serializable payload for the `/prompt` endpoint.
|
||||
|
||||
Returns:
|
||||
The non-empty `prompt_id` string returned by the server.
|
||||
"""
|
||||
|
||||
async def get_system_stats(self) -> SystemStats:
|
||||
"""Fetch system statistics from `/system_stats`."""
|
||||
|
||||
async def free(
|
||||
self, *, free_memory: bool = False, unload_models: bool = False
|
||||
) -> None:
|
||||
"""Invoke `/free` with the provided flags."""
|
||||
|
||||
async def ws_connect(self, client_id: str) -> AsyncIterator[WsEvent]:
|
||||
"""Connect to the websocket (`/ws?clientId=...`) and yield parsed events."""
|
||||
|
||||
async def probe(self) -> None:
|
||||
"""Lightweight health probe for connectivity/parseability.
|
||||
|
||||
Executes a GET `/system_stats` with a short timeout and no retries. Raises
|
||||
`ComfyUIClientError` subclasses on failure; returns `None` on success.
|
||||
"""
|
||||
|
||||
async def supports_hygiene(self) -> bool:
|
||||
"""Return True if both `/system_stats` and `/free` are supported.
|
||||
|
||||
Performs a one-time capability probe; caches results for the session.
|
||||
"""
|
||||
|
||||
|
||||
class ComfyUILocalClient(ComfyUIClientProtocol):
|
||||
"""Local HTTP/WebSocket client for a running ComfyUI instance.
|
||||
|
||||
This class provides the concrete implementation for the protocol. At this
|
||||
stage it only declares the interface and stores constructor parameters; the
|
||||
network behavior will be implemented in subsequent commits.
|
||||
|
||||
Args:
|
||||
base_url: Base HTTP URL for ComfyUI endpoints (e.g., `/prompt`).
|
||||
ws_url: Base WebSocket URL (e.g., `/ws`).
|
||||
session: Optional externally-managed aiohttp session for reuse.
|
||||
logger: Optional logger compatible with the worker's logging API.
|
||||
timeout: Default timeout in seconds for HTTP operations.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
ws_url: str,
|
||||
session: Optional["aiohttp.ClientSession"] = None,
|
||||
logger: Optional[Any] = None,
|
||||
timeout: float = 30.0,
|
||||
) -> None:
|
||||
self._base_url: str = base_url
|
||||
self._ws_url: str = ws_url
|
||||
self._session = session # type: ignore[assignment]
|
||||
self._logger = logger
|
||||
self._timeout: float = float(timeout)
|
||||
self._owned_session: Optional[aiohttp.ClientSession] = None
|
||||
# Backoff defaults per plan (commit 3)
|
||||
self._retry_base_seconds: float = 0.25
|
||||
self._retry_multiplier: float = 2.0
|
||||
self._retry_jitter_seconds: float = 0.25
|
||||
self._retry_max_sleep_seconds: float = 4.0
|
||||
self._retry_max_attempts: int = 3
|
||||
# WebSocket reconnect policy (commit 4)
|
||||
self._ws_max_reconnect_attempts: int = 5
|
||||
self._ws_max_total_backoff_seconds: float = 30.0
|
||||
# Capability probe cache (None = unknown, True/False = probed)
|
||||
self._supports_system_stats: Optional[bool] = None
|
||||
self._supports_free: Optional[bool] = None
|
||||
self._capability_warning_emitted: bool = False
|
||||
|
||||
# Lifecycle methods may be implemented later; for now they act as no-ops.
|
||||
async def __aenter__(self) -> "ComfyUILocalClient":
|
||||
"""Enter async context; create an internal session when none provided."""
|
||||
if self._session is None and self._owned_session is None:
|
||||
self._owned_session = aiohttp.ClientSession()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None: # type: ignore[override]
|
||||
"""Exit async context; close internal session if owned by this client."""
|
||||
if self._owned_session is not None:
|
||||
try:
|
||||
await self._owned_session.close()
|
||||
finally:
|
||||
self._owned_session = None
|
||||
return None
|
||||
|
||||
# Protocol methods — to be implemented in subsequent commits.
|
||||
async def submit_prompt(self, payload: Dict[str, Any]) -> str: # type: ignore[override]
|
||||
route = "/prompt"
|
||||
method = "POST"
|
||||
url = self._join_http(route)
|
||||
try:
|
||||
data = await self._http_request_json(
|
||||
method,
|
||||
url,
|
||||
json_payload=payload,
|
||||
timeout_s=self._timeout,
|
||||
retry_idempotent=False,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ComfyUIClientTimeout(
|
||||
f"Timeout while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code="timeout",
|
||||
) from e
|
||||
except _MappedHttpError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
except _MappedConnError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
|
||||
prompt_id = data.get("prompt_id") if isinstance(data, dict) else None
|
||||
if not isinstance(prompt_id, str) or not prompt_id.strip():
|
||||
snippet = _safe_preview(data)
|
||||
raise ComfyUIClientError(
|
||||
"Invalid response payload: missing non-empty prompt_id",
|
||||
route=route,
|
||||
method=method,
|
||||
code="bad_json",
|
||||
body_snippet=snippet,
|
||||
)
|
||||
return prompt_id
|
||||
|
||||
async def get_system_stats(self) -> SystemStats: # type: ignore[override]
|
||||
route = "/system_stats"
|
||||
method = "GET"
|
||||
url = self._join_http(route)
|
||||
try:
|
||||
data = await self._http_request_json(
|
||||
method,
|
||||
url,
|
||||
json_payload=None,
|
||||
timeout_s=self._timeout,
|
||||
retry_idempotent=True,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ComfyUIClientTimeout(
|
||||
f"Timeout while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code="timeout",
|
||||
) from e
|
||||
except _MappedHttpError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
except _MappedConnError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
|
||||
# Parse known fields, tolerate missing/unknown; try alternate shapes; fallback to torch
|
||||
vram_total = None
|
||||
vram_free = None
|
||||
torch_vram_free = None
|
||||
ram_total = None
|
||||
ram_free = None
|
||||
|
||||
if isinstance(data, dict):
|
||||
vram_total = _coerce_optional_float(data.get("vram_total"))
|
||||
vram_free = _coerce_optional_float(data.get("vram_free"))
|
||||
torch_vram_free = _coerce_optional_float(data.get("torch_vram_free"))
|
||||
ram_total = _coerce_optional_float(data.get("ram_total"))
|
||||
ram_free = _coerce_optional_float(data.get("ram_free"))
|
||||
|
||||
# Common alternates
|
||||
if vram_total is None:
|
||||
vram_total = _coerce_optional_float(data.get("total_vram"))
|
||||
if vram_free is None:
|
||||
vram_free = _coerce_optional_float(data.get("free_vram"))
|
||||
|
||||
vram_obj = data.get("vram") if isinstance(data.get("vram"), dict) else None
|
||||
if vram_obj:
|
||||
if vram_total is None:
|
||||
vram_total = _coerce_optional_float(vram_obj.get("total"))
|
||||
if vram_free is None:
|
||||
vram_free = _coerce_optional_float(vram_obj.get("free"))
|
||||
|
||||
ram_obj = data.get("ram") if isinstance(data.get("ram"), dict) else None
|
||||
if ram_obj:
|
||||
if ram_total is None:
|
||||
ram_total = _coerce_optional_float(ram_obj.get("total"))
|
||||
if ram_free is None:
|
||||
ram_free = _coerce_optional_float(ram_obj.get("free"))
|
||||
|
||||
# devices[0] fallback (common in mock/alt servers)
|
||||
devices = (
|
||||
data.get("devices") if isinstance(data.get("devices"), list) else None
|
||||
)
|
||||
if devices and len(devices) > 0 and isinstance(devices[0], dict):
|
||||
dev0 = devices[0]
|
||||
if vram_total is None:
|
||||
vram_total = _coerce_optional_float(dev0.get("vram_total"))
|
||||
if vram_free is None:
|
||||
vram_free = _coerce_optional_float(dev0.get("vram_free"))
|
||||
# Note: we rely solely on server-reported values; no local torch fallback
|
||||
|
||||
# Optional debug logging of reported stats
|
||||
if self._logger:
|
||||
try:
|
||||
self._logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (comfyui_client): /system_stats: vram_total=%s vram_free=%s torch_vram_free=%s ram_total=%s ram_free=%s",
|
||||
vram_total,
|
||||
vram_free,
|
||||
torch_vram_free,
|
||||
ram_total,
|
||||
ram_free,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return SystemStats(
|
||||
vram_total=vram_total,
|
||||
vram_free=vram_free,
|
||||
torch_vram_free=torch_vram_free,
|
||||
ram_total=ram_total,
|
||||
ram_free=ram_free,
|
||||
)
|
||||
|
||||
async def free(self, *, free_memory: bool = False, unload_models: bool = False) -> None: # type: ignore[override]
|
||||
route = "/free"
|
||||
method = "POST"
|
||||
url = self._join_http(route)
|
||||
body = {"free_memory": bool(free_memory), "unload_models": bool(unload_models)}
|
||||
try:
|
||||
await self._http_request_json(
|
||||
method,
|
||||
url,
|
||||
json_payload=body,
|
||||
timeout_s=self._timeout,
|
||||
retry_idempotent=True, # idempotent when flags identical
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ComfyUIClientTimeout(
|
||||
f"Timeout while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code="timeout",
|
||||
) from e
|
||||
except _MappedHttpError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
except _MappedConnError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
|
||||
async def ws_connect(self, client_id: str) -> AsyncIterator[WsEvent]: # type: ignore[override]
|
||||
url = _ws_url(self._ws_url, client_id)
|
||||
attempts = 0
|
||||
total_backoff = 0.0
|
||||
base = self._retry_base_seconds
|
||||
multiplier = self._retry_multiplier
|
||||
jitter = self._retry_jitter_seconds
|
||||
max_sleep = self._retry_max_sleep_seconds
|
||||
|
||||
while True:
|
||||
try:
|
||||
# Allow large frames and set ping/pong defaults
|
||||
async with websockets.connect(
|
||||
url,
|
||||
max_size=None,
|
||||
max_queue=4,
|
||||
ping_interval=20,
|
||||
ping_timeout=20,
|
||||
) as websocket:
|
||||
# On successful connect, reset counters
|
||||
attempts = 0
|
||||
total_backoff = 0.0
|
||||
if self._logger:
|
||||
try:
|
||||
self._logger.debug(
|
||||
f"✅ Nilor-Nodes (comfyui_client): connected to websocket {url}"
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
while True:
|
||||
message = await websocket.recv()
|
||||
if isinstance(message, bytes):
|
||||
try:
|
||||
message = message.decode("utf-8", errors="replace")
|
||||
except Exception:
|
||||
yield WsEvent(type="binary", data={"length": len(message)}) # type: ignore[call-arg]
|
||||
continue
|
||||
try:
|
||||
payload = json.loads(message)
|
||||
except Exception:
|
||||
yield WsEvent(type="text", data={"message": message}) # type: ignore[call-arg]
|
||||
continue
|
||||
|
||||
if isinstance(payload, dict):
|
||||
event_type = str(payload.get("type", "event"))
|
||||
event_data = payload.get("data")
|
||||
if not isinstance(event_data, dict):
|
||||
event_data = {"raw": payload}
|
||||
yield WsEvent(type=event_type, data=event_data) # type: ignore[call-arg]
|
||||
else:
|
||||
yield WsEvent(type="event", data={"raw": payload}) # type: ignore[call-arg]
|
||||
|
||||
except asyncio.CancelledError:
|
||||
# Allow clean shutdown by propagating cancellation
|
||||
raise
|
||||
except ws_exc.ConnectionClosedOK as e:
|
||||
raise ComfyUIClientWsClosed(
|
||||
"WebSocket closed normally",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
status=getattr(e, "code", 1000),
|
||||
code="ws_closed",
|
||||
body_snippet=str(getattr(e, "reason", ""))[:256],
|
||||
) from e
|
||||
except ws_exc.ConnectionClosedError as e:
|
||||
# Abnormal close: attempt bounded reconnect
|
||||
attempts += 1
|
||||
if attempts > self._ws_max_reconnect_attempts:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket reconnect attempts exhausted",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="ws_closed",
|
||||
body_snippet=str(getattr(e, "reason", ""))[:256],
|
||||
) from e
|
||||
delay = _next_backoff(attempts - 1, base, multiplier, jitter, max_sleep)
|
||||
total_backoff += delay
|
||||
if total_backoff > self._ws_max_total_backoff_seconds:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket reconnect backoff budget exhausted",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="ws_closed",
|
||||
body_snippet=str(getattr(e, "reason", ""))[:256],
|
||||
) from e
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
except ws_exc.InvalidStatus as e:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket handshake failed",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="ws_handshake",
|
||||
) from e
|
||||
except ws_exc.InvalidURI as e:
|
||||
raise ComfyUIClientError(
|
||||
"Invalid WebSocket URI",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="invalid_ws_uri",
|
||||
) from e
|
||||
except Exception as e:
|
||||
# Treat as connection error; bounded reconnect
|
||||
attempts += 1
|
||||
if attempts > self._ws_max_reconnect_attempts:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket reconnect attempts exhausted",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="connection_error",
|
||||
) from e
|
||||
delay = _next_backoff(attempts - 1, base, multiplier, jitter, max_sleep)
|
||||
total_backoff += delay
|
||||
if total_backoff > self._ws_max_total_backoff_seconds:
|
||||
raise ComfyUIClientError(
|
||||
"WebSocket reconnect backoff budget exhausted",
|
||||
route="/ws",
|
||||
method="GET",
|
||||
code="connection_error",
|
||||
) from e
|
||||
await asyncio.sleep(delay)
|
||||
continue
|
||||
|
||||
async def probe(self) -> None: # type: ignore[override]
|
||||
route = "/system_stats"
|
||||
method = "GET"
|
||||
url = self._join_http(route)
|
||||
short_timeout = min(self._timeout, 3.0)
|
||||
try:
|
||||
# No retries: retry_idempotent=False
|
||||
await self._http_request_json(
|
||||
method,
|
||||
url,
|
||||
json_payload=None,
|
||||
timeout_s=short_timeout,
|
||||
retry_idempotent=False,
|
||||
)
|
||||
except asyncio.TimeoutError as e:
|
||||
raise ComfyUIClientTimeout(
|
||||
f"Timeout while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code="timeout",
|
||||
) from e
|
||||
except _MappedHttpError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
except _MappedConnError as e:
|
||||
raise e.to_public_error(route=route, method=method)
|
||||
|
||||
async def supports_hygiene(self) -> bool: # type: ignore[override]
|
||||
# If both probed, return cached decision
|
||||
if self._supports_system_stats is not None and self._supports_free is not None:
|
||||
return bool(self._supports_system_stats and self._supports_free)
|
||||
|
||||
await self._probe_capabilities_once()
|
||||
supported = bool(
|
||||
(self._supports_system_stats is True) and (self._supports_free is True)
|
||||
)
|
||||
if not supported and not self._capability_warning_emitted and self._logger:
|
||||
try:
|
||||
self._logger.warning(
|
||||
"⚠️\u2009 Nilor-Nodes (comfyui_client): hygiene disabled — missing /system_stats or /free support"
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
self._capability_warning_emitted = True
|
||||
return supported
|
||||
|
||||
async def _probe_capabilities_once(self) -> None:
|
||||
"""Probe `/system_stats` and `/free` capabilities once and cache results.
|
||||
|
||||
Only marks capabilities as False on definitive 404/405 responses. Transient
|
||||
failures leave the capability as None so a future call may retry.
|
||||
"""
|
||||
short_timeout = min(self._timeout, 3.0)
|
||||
|
||||
# Probe /system_stats support
|
||||
if self._supports_system_stats is None:
|
||||
route = "/system_stats"
|
||||
url = self._join_http(route)
|
||||
try:
|
||||
await self._http_request_json(
|
||||
"GET",
|
||||
url,
|
||||
json_payload=None,
|
||||
timeout_s=short_timeout,
|
||||
retry_idempotent=False,
|
||||
)
|
||||
self._supports_system_stats = True
|
||||
except ComfyUIClientError as e:
|
||||
if getattr(e, "status", None) in (404, 405):
|
||||
self._supports_system_stats = False
|
||||
|
||||
# Probe /free support (no-op body)
|
||||
if self._supports_free is None:
|
||||
route = "/free"
|
||||
url = self._join_http(route)
|
||||
try:
|
||||
await self._http_request_json(
|
||||
"POST",
|
||||
url,
|
||||
json_payload={"free_memory": False, "unload_models": False},
|
||||
timeout_s=short_timeout,
|
||||
retry_idempotent=False,
|
||||
)
|
||||
self._supports_free = True
|
||||
except ComfyUIClientError as e:
|
||||
if getattr(e, "status", None) in (404, 405):
|
||||
self._supports_free = False
|
||||
|
||||
|
||||
# Runtime dependency; imported here to avoid issues if module is scanned without execution
|
||||
import aiohttp # type: ignore
|
||||
import websockets # type: ignore
|
||||
from websockets import exceptions as ws_exc # type: ignore
|
||||
|
||||
|
||||
# ---- Internal helpers (HTTP) ----
|
||||
|
||||
|
||||
def _coerce_optional_float(value: Any) -> Optional[float]:
|
||||
try:
|
||||
if value is None:
|
||||
return None
|
||||
return float(value)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _is_transient_status(status: int) -> bool:
|
||||
return status == 429 or 500 <= status <= 599
|
||||
|
||||
|
||||
def _safe_preview(data: Any, limit: int = 512) -> str:
|
||||
try:
|
||||
text = json.dumps(data, ensure_ascii=False)
|
||||
except Exception:
|
||||
text = str(data)
|
||||
if len(text) > limit:
|
||||
return text[:limit] + "…"
|
||||
return text
|
||||
|
||||
|
||||
class _MappedHttpError(Exception):
|
||||
def __init__(self, *, status: Optional[int], body_snippet: Optional[str]) -> None:
|
||||
self.status = status
|
||||
self.body_snippet = body_snippet
|
||||
|
||||
def to_public_error(self, *, route: str, method: str) -> ComfyUIClientError:
|
||||
return ComfyUIClientError(
|
||||
f"🛑\u2009 Nilor-Nodes (comfyui_client): HTTP error while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
status=self.status,
|
||||
code="http_error",
|
||||
body_snippet=self.body_snippet,
|
||||
)
|
||||
|
||||
|
||||
class _MappedConnError(Exception):
|
||||
def __init__(self, *, code: str = "connection_error") -> None:
|
||||
self.code = code
|
||||
|
||||
def to_public_error(self, *, route: str, method: str) -> ComfyUIClientError:
|
||||
return ComfyUIClientError(
|
||||
f"🛑\u2009 Nilor-Nodes (comfyui_client): Connection error while calling {method} {route}",
|
||||
route=route,
|
||||
method=method,
|
||||
code=self.code,
|
||||
)
|
||||
|
||||
|
||||
async def _read_limited_text(resp: aiohttp.ClientResponse, limit: int = 512) -> str:
|
||||
try:
|
||||
raw = await resp.read()
|
||||
# Truncate at byte level then decode safely
|
||||
raw = raw[:limit]
|
||||
return raw.decode("utf-8", errors="replace")
|
||||
except Exception:
|
||||
return ""
|
||||
|
||||
|
||||
def _with_jitter(seconds: float, jitter: float) -> float:
|
||||
if jitter <= 0:
|
||||
return seconds
|
||||
return max(0.0, seconds + random.uniform(-jitter, jitter))
|
||||
|
||||
|
||||
def _next_backoff(
|
||||
attempt_index: int,
|
||||
base: float,
|
||||
multiplier: float,
|
||||
jitter: float,
|
||||
max_sleep: float,
|
||||
) -> float:
|
||||
# attempt_index is 0-based
|
||||
delay = base * (multiplier**attempt_index)
|
||||
delay = min(delay, max_sleep)
|
||||
return _with_jitter(delay, jitter)
|
||||
|
||||
|
||||
def _should_retry(
|
||||
*,
|
||||
retry_idempotent: bool,
|
||||
exc: Optional[BaseException] = None,
|
||||
status: Optional[int] = None,
|
||||
) -> bool:
|
||||
if not retry_idempotent:
|
||||
return False
|
||||
if isinstance(exc, (asyncio.TimeoutError, aiohttp.ClientConnectionError)):
|
||||
return True
|
||||
if status is not None and _is_transient_status(status):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _finalize_attempts(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
last_exc: Optional[BaseException],
|
||||
last_status: Optional[int],
|
||||
last_body_snippet: Optional[str],
|
||||
) -> BaseException:
|
||||
if isinstance(last_exc, asyncio.TimeoutError):
|
||||
return ComfyUIClientTimeout(
|
||||
f"🛑\u2009 Nilor-Nodes (comfyui_client): Timeout while calling {method} {url}",
|
||||
route=_route_from_url(url),
|
||||
method=method,
|
||||
code="timeout",
|
||||
)
|
||||
if isinstance(last_exc, aiohttp.ClientConnectionError):
|
||||
return _MappedConnError().to_public_error(
|
||||
route=_route_from_url(url), method=method
|
||||
)
|
||||
# Otherwise treat as HTTP error
|
||||
return _MappedHttpError(
|
||||
status=last_status, body_snippet=last_body_snippet
|
||||
).to_public_error(route=_route_from_url(url), method=method)
|
||||
|
||||
|
||||
def _route_from_url(url: str) -> str:
|
||||
try:
|
||||
parsed = urlparse(url)
|
||||
return parsed.path or "/"
|
||||
except Exception:
|
||||
return url
|
||||
|
||||
|
||||
class _TempSession:
|
||||
"""Context manager that yields an aiohttp session, reusing if provided."""
|
||||
|
||||
def __init__(self, session: Optional[aiohttp.ClientSession]):
|
||||
self._provided = session
|
||||
self._owned: Optional[aiohttp.ClientSession] = None
|
||||
|
||||
async def __aenter__(self) -> aiohttp.ClientSession:
|
||||
if self._provided is not None:
|
||||
return self._provided
|
||||
self._owned = aiohttp.ClientSession()
|
||||
return self._owned
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> None: # type: ignore[override]
|
||||
if self._owned is not None:
|
||||
await self._owned.close()
|
||||
|
||||
|
||||
async def _json_or_text(resp: aiohttp.ClientResponse) -> Any:
|
||||
ctype = resp.headers.get("Content-Type", "").lower()
|
||||
text = await _read_limited_text(resp) # limited read to use for both cases
|
||||
if "json" in ctype:
|
||||
try:
|
||||
return json.loads(text)
|
||||
except Exception:
|
||||
# Fallthrough to treat as bad JSON
|
||||
raise _MappedHttpError(status=resp.status, body_snippet=text)
|
||||
# Not JSON; return raw text
|
||||
return text
|
||||
|
||||
|
||||
async def _raise_for_status_with_snippet(resp: aiohttp.ClientResponse) -> None:
|
||||
if 200 <= resp.status <= 299:
|
||||
return
|
||||
snippet = await _read_limited_text(resp)
|
||||
raise _MappedHttpError(status=resp.status, body_snippet=snippet)
|
||||
|
||||
|
||||
async def _request_once(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
session: aiohttp.ClientSession,
|
||||
json_payload: Optional[Dict[str, Any]],
|
||||
timeout_s: float,
|
||||
) -> Tuple[Optional[int], Optional[str], Any]:
|
||||
timeout = aiohttp.ClientTimeout(total=timeout_s)
|
||||
try:
|
||||
async with session.request(
|
||||
method, url, json=json_payload, timeout=timeout
|
||||
) as resp:
|
||||
status = resp.status
|
||||
await _raise_for_status_with_snippet(resp)
|
||||
# Success path: attempt to parse JSON body; if not JSON, return text
|
||||
try:
|
||||
data = await resp.json(content_type=None)
|
||||
except aiohttp.ContentTypeError:
|
||||
# Not JSON; use limited text
|
||||
data = await _read_limited_text(resp)
|
||||
return status, None, data
|
||||
except asyncio.TimeoutError:
|
||||
raise
|
||||
except aiohttp.ClientConnectionError as e:
|
||||
raise e
|
||||
except aiohttp.ClientPayloadError as e:
|
||||
# Map as payload error
|
||||
raise _MappedHttpError(status=None, body_snippet=str(e))
|
||||
|
||||
|
||||
async def _http_request_core(
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
session: aiohttp.ClientSession,
|
||||
json_payload: Optional[Dict[str, Any]],
|
||||
timeout_s: float,
|
||||
retry_idempotent: bool,
|
||||
base: float,
|
||||
multiplier: float,
|
||||
jitter: float,
|
||||
max_sleep: float,
|
||||
max_attempts: int,
|
||||
) -> Any:
|
||||
last_exc: Optional[BaseException] = None
|
||||
last_status: Optional[int] = None
|
||||
last_body: Optional[str] = None
|
||||
|
||||
attempts = max(1, int(max_attempts))
|
||||
for attempt in range(attempts):
|
||||
try:
|
||||
status, body_snippet, data = await _request_once(
|
||||
method,
|
||||
url,
|
||||
session=session,
|
||||
json_payload=json_payload,
|
||||
timeout_s=timeout_s,
|
||||
)
|
||||
return data
|
||||
except asyncio.TimeoutError as e:
|
||||
last_exc = e
|
||||
if attempt < attempts - 1 and _should_retry(
|
||||
retry_idempotent=retry_idempotent, exc=e
|
||||
):
|
||||
await asyncio.sleep(
|
||||
_next_backoff(attempt, base, multiplier, jitter, max_sleep)
|
||||
)
|
||||
continue
|
||||
break
|
||||
except aiohttp.ClientConnectionError as e:
|
||||
last_exc = e
|
||||
if attempt < attempts - 1 and _should_retry(
|
||||
retry_idempotent=retry_idempotent, exc=e
|
||||
):
|
||||
await asyncio.sleep(
|
||||
_next_backoff(attempt, base, multiplier, jitter, max_sleep)
|
||||
)
|
||||
continue
|
||||
break
|
||||
except _MappedHttpError as e:
|
||||
last_exc = None
|
||||
last_status = e.status
|
||||
last_body = e.body_snippet
|
||||
if attempt < attempts - 1 and _should_retry(
|
||||
retry_idempotent=retry_idempotent, status=e.status
|
||||
):
|
||||
await asyncio.sleep(
|
||||
_next_backoff(attempt, base, multiplier, jitter, max_sleep)
|
||||
)
|
||||
continue
|
||||
break
|
||||
|
||||
raise _finalize_attempts(
|
||||
method,
|
||||
url,
|
||||
last_exc=last_exc,
|
||||
last_status=last_status,
|
||||
last_body_snippet=last_body,
|
||||
)
|
||||
|
||||
|
||||
async def _http_request_json(
|
||||
self: "ComfyUILocalClient",
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
json_payload: Optional[Dict[str, Any]],
|
||||
timeout_s: float,
|
||||
retry_idempotent: bool,
|
||||
) -> Any:
|
||||
async with _TempSession(self._session or self._owned_session) as session:
|
||||
return await _http_request_core(
|
||||
method,
|
||||
url,
|
||||
session=session,
|
||||
json_payload=json_payload,
|
||||
timeout_s=timeout_s,
|
||||
retry_idempotent=retry_idempotent,
|
||||
base=self._retry_base_seconds,
|
||||
multiplier=self._retry_multiplier,
|
||||
jitter=self._retry_jitter_seconds,
|
||||
max_sleep=self._retry_max_sleep_seconds,
|
||||
max_attempts=self._retry_max_attempts,
|
||||
)
|
||||
|
||||
|
||||
def _join_http_base(base: str, route: str) -> str:
|
||||
if not route:
|
||||
return base
|
||||
return f"{base.rstrip('/')}{route}"
|
||||
|
||||
|
||||
def _join_ws_base(base: str, route: str) -> str:
|
||||
if not route:
|
||||
return base
|
||||
return f"{base.rstrip('/')}{route}"
|
||||
|
||||
|
||||
def _ensure_scheme(base: str, allowed: Tuple[str, ...]) -> None:
|
||||
parsed = urlparse(base)
|
||||
if not parsed.scheme or parsed.scheme.lower() not in allowed:
|
||||
allowed_str = ", ".join(allowed)
|
||||
raise ValueError(
|
||||
f"🛑\u2009 Nilor-Nodes (comfyui_client): Base URL must start with one of [{allowed_str}]; got: {base!r}"
|
||||
)
|
||||
|
||||
|
||||
def _validate_bases(http_base: str, ws_base: str) -> None:
|
||||
_ensure_scheme(http_base, ("http", "https"))
|
||||
_ensure_scheme(ws_base, ("ws", "wss"))
|
||||
|
||||
|
||||
def _quote_client_id(client_id: str) -> str:
|
||||
return quote(client_id, safe="")
|
||||
|
||||
|
||||
def _ws_url(base_ws: str, client_id: str) -> str:
|
||||
return f"{base_ws.rstrip('/')}/ws?clientId={_quote_client_id(client_id)}"
|
||||
|
||||
|
||||
# Bind helper methods to class namespace (private) without exposing publicly
|
||||
ComfyUILocalClient._http_request_json = _http_request_json # type: ignore[attr-defined]
|
||||
ComfyUILocalClient._join_http = lambda self, route: _join_http_base(self._base_url, route) # type: ignore[attr-defined]
|
||||
ComfyUILocalClient._join_ws = lambda self, route: _join_ws_base(self._ws_url, route) # type: ignore[attr-defined]
|
||||
@@ -0,0 +1,75 @@
|
||||
{
|
||||
// Nilor-Nodes sidecar configuration defaults (non-secret) for local development
|
||||
// Precedence model (enforced by the loader): environment variables > this file
|
||||
// Secrets MUST NOT be committed here; use .env to override sensitive values
|
||||
|
||||
allow_env_override: true,
|
||||
|
||||
// Log level for nilor-nodes components (DEBUG, INFO, WARNING, ERROR, CRITICAL)
|
||||
NILOR_LOG_LEVEL: "INFO",
|
||||
// Enable/disable SQS job consumption from Nilor brain_rnd
|
||||
NILOR_SQS_ENABLED: false,
|
||||
|
||||
// ---- Comfy API client (consumed by worker_consumer.py) ----
|
||||
// Base HTTP URL for ComfyUI REST API (worker submits to `${api_url}/prompt`)
|
||||
NILOR_COMFYUI_API_URL: "http://127.0.0.1:8188",
|
||||
// Base WS URL for ComfyUI websocket events (worker listens at `${ws_url}/ws`)
|
||||
NILOR_COMFYUI_WS_URL: "ws://127.0.0.1:8188",
|
||||
// Request timeout in seconds for ComfyUI HTTP calls
|
||||
NILOR_COMFY_API_TIMEOUT_SECONDS: 30,
|
||||
|
||||
// HTTP idempotent retry policy (applies only to GET /system_stats and POST /free)
|
||||
NILOR_COMFY_RETRY_BASE_SECONDS: 0.25,
|
||||
NILOR_COMFY_RETRY_MULTIPLIER: 2.0,
|
||||
NILOR_COMFY_RETRY_JITTER_SECONDS: 0.25,
|
||||
NILOR_COMFY_RETRY_MAX_SLEEP_SECONDS: 4.0,
|
||||
NILOR_COMFY_RETRY_MAX_ATTEMPTS: 3,
|
||||
|
||||
// WebSocket reconnect policy
|
||||
NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS: 5,
|
||||
NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS: 30.0,
|
||||
|
||||
// ---- Memory Hygiene (Guardian) defaults ----
|
||||
// Enable/disable hygiene between jobs
|
||||
NILOR_MEMORY_HYGIENE_ENABLED: false,
|
||||
// How often to check when idle (seconds)
|
||||
NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS: 5,
|
||||
// Thresholds (either percent usage or absolute free MB can trigger)
|
||||
NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX: 88,
|
||||
NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX: 90,
|
||||
NILOR_MEMORY_HYGIENE_VRAM_MIN_FREE_MB: 2048,
|
||||
NILOR_MEMORY_HYGIENE_RAM_MIN_FREE_MB: 4096,
|
||||
// Action policy: free|unload|both|auto (auto: stage free then unload once if needed)
|
||||
NILOR_MEMORY_HYGIENE_ACTION_POLICY: "auto",
|
||||
// Retry/backoff/cycle caps
|
||||
NILOR_MEMORY_HYGIENE_MAX_RETRIES: 2,
|
||||
NILOR_MEMORY_HYGIENE_COOLDOWN_SECONDS: 60,
|
||||
NILOR_MEMORY_HYGIENE_SLEEP_BETWEEN_ATTEMPTS_SECONDS: 4,
|
||||
NILOR_MEMORY_HYGIENE_MAX_CYCLE_DURATION_SECONDS: 15,
|
||||
|
||||
// ---- Worker / Queue settings ----
|
||||
// Workflow normalization (applies in worker_consumer before submission)
|
||||
// When enabled, nilor-nodes will normalize OS-specific path formatting inside
|
||||
// incoming workflows (e.g. Windows backslashes vs POSIX slashes).
|
||||
NILOR_WORKFLOW_OS_NORMALIZATION_ENABLED: true,
|
||||
|
||||
// ElasticMQ/SQS endpoint URL (used by consumer and status updates)
|
||||
NILOR_SQS_ENDPOINT_URL: "http://localhost:9324",
|
||||
// Queue to pull new jobs from (polled by worker_consumer)
|
||||
NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME: "jobs_to_process-comfyui",
|
||||
// Queue to publish job status updates to (sent by worker_consumer)
|
||||
NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME: "job_status_updates",
|
||||
// Long poll wait time (seconds); must be within [0, 20]
|
||||
NILOR_SQS_POLL_WAIT_TIME: 10,
|
||||
// Max number of messages to pull per poll
|
||||
NILOR_SQS_MAX_MESSAGES: 1,
|
||||
|
||||
// Non-secret local defaults; override via .env for real deployments
|
||||
NILOR_AWS_ACCESS_KEY_ID: "local",
|
||||
NILOR_AWS_SECRET_ACCESS_KEY: "local", // secret must be provided via .env
|
||||
NILOR_AWS_DEFAULT_REGION: "us-east-1",
|
||||
|
||||
// If empty or unset, a stable id will be generated by the loader
|
||||
NILOR_WORKER_CLIENT_ID: ""
|
||||
}
|
||||
|
||||
@@ -0,0 +1,584 @@
|
||||
"""
|
||||
Typed configuration scaffolding for the Nilor-Nodes ComfyUI sidecar.
|
||||
|
||||
This module defines the dataclasses and public loader API contract. The actual
|
||||
implementation of precedence, parsing, and validation is added in a subsequent
|
||||
commit. For now, only type definitions and the public `Config.load` signature
|
||||
are provided to enable incremental integration without behavior changes.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import socket
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, Optional
|
||||
from urllib.parse import urlparse
|
||||
from pathlib import Path
|
||||
from ..logger import logger
|
||||
|
||||
# Ensure we cache config globally
|
||||
_CONFIG: Optional["NilorNodesConfig"] = None
|
||||
|
||||
try: # json5 is declared in nilor-nodes/requirements.txt
|
||||
import json5 # type: ignore
|
||||
except Exception as _e: # pragma: no cover
|
||||
json5 = None # lazy failure in loader
|
||||
|
||||
|
||||
class BaseConfig: # type: ignore
|
||||
@classmethod
|
||||
def get_instance(cls):
|
||||
# Minimal fallback: load JSON5 directly when BaseConfig is unavailable
|
||||
path = os.path.join(os.path.dirname(__file__), "config.json5")
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json5.load(f) if json5 is not None else {}
|
||||
return cls.from_dict(data) # type: ignore[attr-defined]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ComfyApiConfig:
|
||||
"""Configuration for the ComfyUI client.
|
||||
|
||||
Args:
|
||||
api_url: Base HTTP URL for the ComfyUI REST API (e.g., "http://127.0.0.1:8188").
|
||||
ws_url: Base WebSocket URL for ComfyUI events (e.g., "ws://127.0.0.1:8188").
|
||||
timeout_s: Request timeout in seconds for ComfyUI HTTP calls.
|
||||
retry_base_seconds: Base backoff seconds for idempotent retries.
|
||||
retry_multiplier: Exponential backoff multiplier.
|
||||
retry_jitter_seconds: Jitter range (±seconds) added to backoff.
|
||||
retry_max_sleep_seconds: Maximum sleep per backoff step.
|
||||
retry_max_attempts: Maximum retry attempts for idempotent routes.
|
||||
ws_max_reconnect_attempts: Maximum websocket reconnect attempts.
|
||||
ws_max_total_backoff_seconds: Cap on total backoff time during WS reconnects.
|
||||
"""
|
||||
|
||||
api_url: str
|
||||
ws_url: str
|
||||
timeout_s: int
|
||||
retry_base_seconds: float
|
||||
retry_multiplier: float
|
||||
retry_jitter_seconds: float
|
||||
retry_max_sleep_seconds: float
|
||||
retry_max_attempts: int
|
||||
ws_max_reconnect_attempts: int
|
||||
ws_max_total_backoff_seconds: float
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkerConfig:
|
||||
"""Configuration for the worker and its SQS integration.
|
||||
|
||||
Args:
|
||||
sqs_endpoint_url: URL of the SQS-compatible endpoint (e.g., ElasticMQ).
|
||||
jobs_queue: Name of the queue from which to pull new jobs.
|
||||
status_queue: Name of the queue to which job status updates are published.
|
||||
poll_wait_s: Long poll wait time in seconds (expected to be within [0, 20]).
|
||||
max_messages: Max number of messages pulled per poll.
|
||||
aws_access_key_id: Access key id for the SQS client (non-secret default acceptable for local dev).
|
||||
aws_secret_access_key: Secret access key for the SQS client (must be overridden via environment for real deployments).
|
||||
aws_region: AWS region name used by the SQS client.
|
||||
worker_client_id: Stable identifier for routing websocket events to this worker.
|
||||
workflow_os_normalization_enabled: When true, normalize OS-specific path formatting
|
||||
inside ComfyUI prompt graphs before submission.
|
||||
"""
|
||||
|
||||
sqs_endpoint_url: str
|
||||
jobs_queue: str
|
||||
status_queue: str
|
||||
poll_wait_s: int
|
||||
max_messages: int
|
||||
aws_access_key_id: str
|
||||
aws_secret_access_key: str
|
||||
aws_region: str
|
||||
worker_client_id: str
|
||||
workflow_os_normalization_enabled: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MemoryHygieneConfig:
|
||||
"""Configuration for Memory Guardian hygiene between jobs.
|
||||
|
||||
Args:
|
||||
enabled: Feature flag to enable/disable hygiene.
|
||||
idle_poll_seconds: How often to check while idle.
|
||||
vram_usage_pct_max: If used VRAM exceeds this percent, trigger remediation.
|
||||
ram_usage_pct_max: If used RAM exceeds this percent, trigger remediation.
|
||||
vram_min_free_mb: Absolute minimum free VRAM (MB) threshold.
|
||||
ram_min_free_mb: Absolute minimum free RAM (MB) threshold.
|
||||
action_policy: One of {free, unload, both, auto}.
|
||||
max_retries: Max remediation attempts per cycle.
|
||||
cooldown_seconds: Cooldown between remediation cycles.
|
||||
sleep_between_attempts_seconds: Wait time after /free before re-checking.
|
||||
max_cycle_duration_seconds: Hard cap for a single remediation cycle.
|
||||
"""
|
||||
|
||||
enabled: bool
|
||||
idle_poll_seconds: int
|
||||
vram_usage_pct_max: int
|
||||
ram_usage_pct_max: int
|
||||
vram_min_free_mb: int
|
||||
ram_min_free_mb: int
|
||||
action_policy: str
|
||||
max_retries: int
|
||||
cooldown_seconds: int
|
||||
sleep_between_attempts_seconds: int
|
||||
max_cycle_duration_seconds: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class NilorNodesConfig(BaseConfig):
|
||||
"""Aggregate configuration for the Nilor-Nodes sidecar.
|
||||
|
||||
Args:
|
||||
comfy: Configuration for the ComfyUI HTTP/WS client.
|
||||
worker: Configuration for SQS and worker identity.
|
||||
allow_env_override: When true, environment variables may override file values.
|
||||
"""
|
||||
|
||||
comfy: ComfyApiConfig
|
||||
worker: WorkerConfig
|
||||
allow_env_override: bool
|
||||
sqs_enabled: bool
|
||||
hygiene: MemoryHygieneConfig
|
||||
|
||||
@classmethod
|
||||
def _get_config_path(cls) -> str:
|
||||
return os.path.join(os.path.dirname(__file__), "config.json5")
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config_dict: Dict[str, object]) -> "NilorNodesConfig":
|
||||
allow_env_override = bool(config_dict.get("allow_env_override", True))
|
||||
sqs_enabled = _coerce_bool(config_dict.get("NILOR_SQS_ENABLED", False))
|
||||
|
||||
# Build nested from flat NILOR_* keys present in JSON5
|
||||
comfy_cfg = ComfyApiConfig(
|
||||
api_url=str(config_dict.get("NILOR_COMFYUI_API_URL", "")).strip(),
|
||||
ws_url=str(config_dict.get("NILOR_COMFYUI_WS_URL", "")).strip(),
|
||||
timeout_s=int(config_dict.get("NILOR_COMFY_API_TIMEOUT_SECONDS", 30)),
|
||||
retry_base_seconds=float(
|
||||
config_dict.get("NILOR_COMFY_RETRY_BASE_SECONDS", 0.25)
|
||||
),
|
||||
retry_multiplier=float(
|
||||
config_dict.get("NILOR_COMFY_RETRY_MULTIPLIER", 2.0)
|
||||
),
|
||||
retry_jitter_seconds=float(
|
||||
config_dict.get("NILOR_COMFY_RETRY_JITTER_SECONDS", 0.25)
|
||||
),
|
||||
retry_max_sleep_seconds=float(
|
||||
config_dict.get("NILOR_COMFY_RETRY_MAX_SLEEP_SECONDS", 4.0)
|
||||
),
|
||||
retry_max_attempts=int(
|
||||
config_dict.get("NILOR_COMFY_RETRY_MAX_ATTEMPTS", 3)
|
||||
),
|
||||
ws_max_reconnect_attempts=int(
|
||||
config_dict.get("NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS", 5)
|
||||
),
|
||||
ws_max_total_backoff_seconds=float(
|
||||
config_dict.get("NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS", 30.0)
|
||||
),
|
||||
)
|
||||
|
||||
worker_client_id = (
|
||||
str(config_dict.get("NILOR_WORKER_CLIENT_ID", "")).strip()
|
||||
or _generate_worker_client_id()
|
||||
)
|
||||
|
||||
worker_cfg = WorkerConfig(
|
||||
sqs_endpoint_url=str(config_dict.get("NILOR_SQS_ENDPOINT_URL", "")).strip(),
|
||||
jobs_queue=str(
|
||||
config_dict.get("NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME", "")
|
||||
).strip(),
|
||||
status_queue=str(
|
||||
config_dict.get("NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME", "")
|
||||
).strip(),
|
||||
poll_wait_s=int(config_dict.get("NILOR_SQS_POLL_WAIT_TIME", 10)),
|
||||
max_messages=int(config_dict.get("NILOR_SQS_MAX_MESSAGES", 1)),
|
||||
aws_access_key_id=str(
|
||||
config_dict.get("NILOR_AWS_ACCESS_KEY_ID", "")
|
||||
).strip(),
|
||||
aws_secret_access_key=str(
|
||||
config_dict.get("NILOR_AWS_SECRET_ACCESS_KEY", "")
|
||||
).strip(),
|
||||
aws_region=str(config_dict.get("NILOR_AWS_DEFAULT_REGION", "")).strip(),
|
||||
worker_client_id=worker_client_id,
|
||||
workflow_os_normalization_enabled=_coerce_bool(
|
||||
config_dict.get("NILOR_WORKFLOW_OS_NORMALIZATION_ENABLED", True)
|
||||
),
|
||||
)
|
||||
|
||||
# Memory hygiene
|
||||
hygiene_cfg = MemoryHygieneConfig(
|
||||
enabled=_coerce_bool(config_dict.get("NILOR_MEMORY_HYGIENE_ENABLED", True)),
|
||||
idle_poll_seconds=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS", 5)
|
||||
),
|
||||
vram_usage_pct_max=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX", 88)
|
||||
),
|
||||
ram_usage_pct_max=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX", 90)
|
||||
),
|
||||
vram_min_free_mb=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_VRAM_MIN_FREE_MB", 2048)
|
||||
),
|
||||
ram_min_free_mb=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_RAM_MIN_FREE_MB", 4096)
|
||||
),
|
||||
action_policy=str(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_ACTION_POLICY", "auto")
|
||||
).strip(),
|
||||
max_retries=int(config_dict.get("NILOR_MEMORY_HYGIENE_MAX_RETRIES", 2)),
|
||||
cooldown_seconds=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_COOLDOWN_SECONDS", 60)
|
||||
),
|
||||
sleep_between_attempts_seconds=int(
|
||||
config_dict.get(
|
||||
"NILOR_MEMORY_HYGIENE_SLEEP_BETWEEN_ATTEMPTS_SECONDS", 4
|
||||
)
|
||||
),
|
||||
max_cycle_duration_seconds=int(
|
||||
config_dict.get("NILOR_MEMORY_HYGIENE_MAX_CYCLE_DURATION_SECONDS", 15)
|
||||
),
|
||||
)
|
||||
_validate_hygiene_config(hygiene_cfg)
|
||||
|
||||
cfg = cls(
|
||||
comfy=comfy_cfg,
|
||||
worker=worker_cfg,
|
||||
allow_env_override=allow_env_override,
|
||||
sqs_enabled=sqs_enabled,
|
||||
hygiene=hygiene_cfg,
|
||||
)
|
||||
_validate_comfy_config(cfg.comfy)
|
||||
_validate_worker_config(cfg.worker)
|
||||
return cfg
|
||||
|
||||
|
||||
# ---- Internal helpers ----
|
||||
|
||||
|
||||
def _apply_env_overrides(cfg: NilorNodesConfig) -> None:
|
||||
if not cfg.allow_env_override:
|
||||
return
|
||||
|
||||
# Feature flags
|
||||
sqs_enabled = os.getenv("NILOR_SQS_ENABLED", cfg.sqs_enabled)
|
||||
cfg.sqs_enabled = _coerce_bool(sqs_enabled)
|
||||
|
||||
# Comfy (rebuild frozen dataclass)
|
||||
comfy_api_url = os.getenv("NILOR_COMFYUI_API_URL", cfg.comfy.api_url)
|
||||
comfy_ws_url = os.getenv("NILOR_COMFYUI_WS_URL", cfg.comfy.ws_url)
|
||||
comfy_timeout_s = int(
|
||||
os.getenv("NILOR_COMFY_API_TIMEOUT_SECONDS", cfg.comfy.timeout_s)
|
||||
)
|
||||
comfy_retry_base = float(
|
||||
os.getenv("NILOR_COMFY_RETRY_BASE_SECONDS", cfg.comfy.retry_base_seconds)
|
||||
)
|
||||
comfy_retry_multiplier = float(
|
||||
os.getenv("NILOR_COMFY_RETRY_MULTIPLIER", cfg.comfy.retry_multiplier)
|
||||
)
|
||||
comfy_retry_jitter = float(
|
||||
os.getenv("NILOR_COMFY_RETRY_JITTER_SECONDS", cfg.comfy.retry_jitter_seconds)
|
||||
)
|
||||
comfy_retry_max_sleep = float(
|
||||
os.getenv(
|
||||
"NILOR_COMFY_RETRY_MAX_SLEEP_SECONDS", cfg.comfy.retry_max_sleep_seconds
|
||||
)
|
||||
)
|
||||
comfy_retry_max_attempts = int(
|
||||
os.getenv("NILOR_COMFY_RETRY_MAX_ATTEMPTS", cfg.comfy.retry_max_attempts)
|
||||
)
|
||||
comfy_ws_max_reconnect = int(
|
||||
os.getenv(
|
||||
"NILOR_COMFY_WS_MAX_RECONNECT_ATTEMPTS",
|
||||
cfg.comfy.ws_max_reconnect_attempts,
|
||||
)
|
||||
)
|
||||
comfy_ws_max_total_backoff = float(
|
||||
os.getenv(
|
||||
"NILOR_COMFY_WS_MAX_TOTAL_BACKOFF_SECONDS",
|
||||
cfg.comfy.ws_max_total_backoff_seconds,
|
||||
)
|
||||
)
|
||||
cfg.comfy = ComfyApiConfig(
|
||||
api_url=str(comfy_api_url),
|
||||
ws_url=str(comfy_ws_url),
|
||||
timeout_s=comfy_timeout_s,
|
||||
retry_base_seconds=comfy_retry_base,
|
||||
retry_multiplier=comfy_retry_multiplier,
|
||||
retry_jitter_seconds=comfy_retry_jitter,
|
||||
retry_max_sleep_seconds=comfy_retry_max_sleep,
|
||||
retry_max_attempts=comfy_retry_max_attempts,
|
||||
ws_max_reconnect_attempts=comfy_ws_max_reconnect,
|
||||
ws_max_total_backoff_seconds=comfy_ws_max_total_backoff,
|
||||
)
|
||||
|
||||
# Worker (rebuild frozen dataclass)
|
||||
worker_sqs_endpoint_url = os.getenv(
|
||||
"NILOR_SQS_ENDPOINT_URL", cfg.worker.sqs_endpoint_url
|
||||
)
|
||||
worker_jobs_queue = os.getenv(
|
||||
"NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME", cfg.worker.jobs_queue
|
||||
)
|
||||
worker_status_queue = os.getenv(
|
||||
"NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME", cfg.worker.status_queue
|
||||
)
|
||||
worker_poll_wait_s = int(
|
||||
os.getenv("NILOR_SQS_POLL_WAIT_TIME", cfg.worker.poll_wait_s)
|
||||
)
|
||||
worker_max_messages = int(
|
||||
os.getenv("NILOR_SQS_MAX_MESSAGES", cfg.worker.max_messages)
|
||||
)
|
||||
worker_access_key = os.getenv(
|
||||
"NILOR_AWS_ACCESS_KEY_ID", cfg.worker.aws_access_key_id
|
||||
)
|
||||
worker_secret_key = os.getenv(
|
||||
"NILOR_AWS_SECRET_ACCESS_KEY", cfg.worker.aws_secret_access_key
|
||||
)
|
||||
worker_region = os.getenv("NILOR_AWS_DEFAULT_REGION", cfg.worker.aws_region)
|
||||
worker_client_id = os.getenv("NILOR_WORKER_CLIENT_ID", cfg.worker.worker_client_id)
|
||||
worker_workflow_os_norm = os.getenv(
|
||||
"NILOR_WORKFLOW_OS_NORMALIZATION_ENABLED",
|
||||
cfg.worker.workflow_os_normalization_enabled,
|
||||
)
|
||||
cfg.worker = WorkerConfig(
|
||||
sqs_endpoint_url=str(worker_sqs_endpoint_url),
|
||||
jobs_queue=str(worker_jobs_queue),
|
||||
status_queue=str(worker_status_queue),
|
||||
poll_wait_s=worker_poll_wait_s,
|
||||
max_messages=worker_max_messages,
|
||||
aws_access_key_id=str(worker_access_key),
|
||||
aws_secret_access_key=str(worker_secret_key),
|
||||
aws_region=str(worker_region),
|
||||
worker_client_id=str(worker_client_id),
|
||||
workflow_os_normalization_enabled=_coerce_bool(worker_workflow_os_norm),
|
||||
)
|
||||
|
||||
# Re-validate after overrides
|
||||
_validate_comfy_config(cfg.comfy)
|
||||
_validate_worker_config(cfg.worker)
|
||||
# Hygiene (rebuild frozen dataclass)
|
||||
hygiene_enabled = _coerce_bool(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_ENABLED", cfg.hygiene.enabled)
|
||||
)
|
||||
hygiene_idle_poll = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS", cfg.hygiene.idle_poll_seconds
|
||||
)
|
||||
)
|
||||
hygiene_vram_pct = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX", cfg.hygiene.vram_usage_pct_max
|
||||
)
|
||||
)
|
||||
hygiene_ram_pct = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX", cfg.hygiene.ram_usage_pct_max
|
||||
)
|
||||
)
|
||||
hygiene_vram_min = int(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_VRAM_MIN_FREE_MB", cfg.hygiene.vram_min_free_mb)
|
||||
)
|
||||
hygiene_ram_min = int(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_RAM_MIN_FREE_MB", cfg.hygiene.ram_min_free_mb)
|
||||
)
|
||||
hygiene_policy = str(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_ACTION_POLICY", cfg.hygiene.action_policy)
|
||||
).strip()
|
||||
hygiene_max_retries = int(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_MAX_RETRIES", cfg.hygiene.max_retries)
|
||||
)
|
||||
hygiene_cooldown = int(
|
||||
os.getenv("NILOR_MEMORY_HYGIENE_COOLDOWN_SECONDS", cfg.hygiene.cooldown_seconds)
|
||||
)
|
||||
hygiene_sleep_between = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_SLEEP_BETWEEN_ATTEMPTS_SECONDS",
|
||||
cfg.hygiene.sleep_between_attempts_seconds,
|
||||
)
|
||||
)
|
||||
hygiene_max_cycle = int(
|
||||
os.getenv(
|
||||
"NILOR_MEMORY_HYGIENE_MAX_CYCLE_DURATION_SECONDS",
|
||||
cfg.hygiene.max_cycle_duration_seconds,
|
||||
)
|
||||
)
|
||||
cfg.hygiene = MemoryHygieneConfig(
|
||||
enabled=hygiene_enabled,
|
||||
idle_poll_seconds=hygiene_idle_poll,
|
||||
vram_usage_pct_max=hygiene_vram_pct,
|
||||
ram_usage_pct_max=hygiene_ram_pct,
|
||||
vram_min_free_mb=hygiene_vram_min,
|
||||
ram_min_free_mb=hygiene_ram_min,
|
||||
action_policy=hygiene_policy,
|
||||
max_retries=hygiene_max_retries,
|
||||
cooldown_seconds=hygiene_cooldown,
|
||||
sleep_between_attempts_seconds=hygiene_sleep_between,
|
||||
max_cycle_duration_seconds=hygiene_max_cycle,
|
||||
)
|
||||
_validate_hygiene_config(cfg.hygiene)
|
||||
|
||||
|
||||
def _validate_comfy_config(cfg: ComfyApiConfig) -> None:
|
||||
_require_url_scheme(cfg.api_url, {"http", "https"}, "NILOR_COMFYUI_API_URL")
|
||||
_require_url_scheme(cfg.ws_url, {"ws", "wss"}, "NILOR_COMFYUI_WS_URL")
|
||||
if cfg.timeout_s <= 0:
|
||||
raise ValueError(
|
||||
f"NILOR_COMFY_API_TIMEOUT_SECONDS must be a positive integer; got {cfg.timeout_s}"
|
||||
)
|
||||
if (
|
||||
cfg.retry_base_seconds < 0
|
||||
or cfg.retry_multiplier <= 0
|
||||
or cfg.retry_max_sleep_seconds <= 0
|
||||
):
|
||||
raise ValueError("Invalid retry backoff parameters in Comfy client config")
|
||||
if cfg.retry_max_attempts <= 0:
|
||||
raise ValueError("NILOR_COMFY_RETRY_MAX_ATTEMPTS must be a positive integer")
|
||||
if cfg.ws_max_reconnect_attempts < 0 or cfg.ws_max_total_backoff_seconds < 0:
|
||||
raise ValueError(
|
||||
"Invalid websocket reconnect parameters in Comfy client config"
|
||||
)
|
||||
|
||||
|
||||
def _validate_worker_config(cfg: WorkerConfig) -> None:
|
||||
if cfg.poll_wait_s < 0 or cfg.poll_wait_s > 20:
|
||||
raise ValueError(
|
||||
f"NILOR_SQS_POLL_WAIT_TIME must be within [0, 20]; got {cfg.poll_wait_s}"
|
||||
)
|
||||
if cfg.max_messages <= 0:
|
||||
raise ValueError(
|
||||
f"NILOR_SQS_MAX_MESSAGES must be a positive integer; got {cfg.max_messages}"
|
||||
)
|
||||
|
||||
|
||||
def _validate_hygiene_config(cfg: MemoryHygieneConfig) -> None:
|
||||
if cfg.idle_poll_seconds < 0:
|
||||
raise ValueError("NILOR_MEMORY_HYGIENE_IDLE_POLL_SECONDS must be >= 0")
|
||||
if not 0 <= cfg.vram_usage_pct_max <= 100:
|
||||
raise ValueError(
|
||||
"NILOR_MEMORY_HYGIENE_VRAM_USAGE_PCT_MAX must be within [0, 100]"
|
||||
)
|
||||
if not 0 <= cfg.ram_usage_pct_max <= 100:
|
||||
raise ValueError(
|
||||
"NILOR_MEMORY_HYGIENE_RAM_USAGE_PCT_MAX must be within [0, 100]"
|
||||
)
|
||||
if cfg.vram_min_free_mb < 0 or cfg.ram_min_free_mb < 0:
|
||||
raise ValueError("Memory hygiene min free MB must be >= 0")
|
||||
if cfg.max_retries < 0:
|
||||
raise ValueError("NILOR_MEMORY_HYGIENE_MAX_RETRIES must be >= 0")
|
||||
if cfg.cooldown_seconds < 0 or cfg.sleep_between_attempts_seconds < 0:
|
||||
raise ValueError("Memory hygiene cooldown/sleep must be >= 0")
|
||||
if cfg.max_cycle_duration_seconds < 0:
|
||||
raise ValueError("Memory hygiene max cycle duration must be >= 0")
|
||||
allowed = {"free", "unload", "both", "auto"}
|
||||
if cfg.action_policy not in allowed:
|
||||
allowed_str = ", ".join(sorted(allowed))
|
||||
raise ValueError(
|
||||
f"NILOR_MEMORY_HYGIENE_ACTION_POLICY must be one of {{{allowed_str}}}"
|
||||
)
|
||||
|
||||
|
||||
def _coerce_bool(value: object) -> bool:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if value is None:
|
||||
return False
|
||||
text = str(value).strip().lower()
|
||||
if text in {"1", "true", "yes", "on"}:
|
||||
return True
|
||||
if text in {"0", "false", "no", "off"}:
|
||||
return False
|
||||
return bool(text)
|
||||
|
||||
|
||||
def _require_url_scheme(url: str, allowed: set[str], key_name: str) -> None:
|
||||
parsed = urlparse(url)
|
||||
if not parsed.scheme or parsed.scheme.lower() not in allowed:
|
||||
allowed_str = ", ".join(sorted(allowed))
|
||||
raise ValueError(
|
||||
f"{key_name} must start with one of [{allowed_str}]; got: {url!r}"
|
||||
)
|
||||
|
||||
|
||||
def _generate_worker_client_id() -> str:
|
||||
host = socket.gethostname().strip() or "worker"
|
||||
suffix = _random_base36_suffix(5)
|
||||
return f"nilor-worker-{host}-{suffix}"
|
||||
|
||||
|
||||
def _random_base36_suffix(length: int = 5) -> str:
|
||||
import random
|
||||
|
||||
n = random.getrandbits(32)
|
||||
base36 = _to_base36(n)
|
||||
return base36[-length:]
|
||||
|
||||
|
||||
def _to_base36(n: int) -> str:
|
||||
if n == 0:
|
||||
return "0"
|
||||
digits = "0123456789abcdefghijklmnopqrstuvwxyz"
|
||||
sign = "-" if n < 0 else ""
|
||||
n = abs(n)
|
||||
res = []
|
||||
while n:
|
||||
n, r = divmod(n, 36)
|
||||
res.append(digits[r])
|
||||
return sign + "".join(reversed(res))
|
||||
|
||||
|
||||
def load_nilor_nodes_config() -> NilorNodesConfig:
|
||||
"""Convenience loader aligning with brain_rnd config pattern.
|
||||
|
||||
- Loads environment variables via python-dotenv if available
|
||||
- Reads JSON5 defaults and applies env overrides when enabled
|
||||
- Returns a typed `NilorNodesConfig`
|
||||
"""
|
||||
global _CONFIG
|
||||
if _CONFIG is not None:
|
||||
return _CONFIG
|
||||
try:
|
||||
from dotenv import load_dotenv # optional dependency present in sidecar
|
||||
|
||||
try:
|
||||
# Load .env next to nilor-nodes (ComfyUI/custom_nodes/nilor-nodes/.env) first
|
||||
here = Path(__file__).resolve().parent.parent # .../nilor-nodes/
|
||||
dotenv_path = here / ".env"
|
||||
loaded = load_dotenv(dotenv_path=dotenv_path)
|
||||
try:
|
||||
if loaded:
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes: Loaded environment variables from {dotenv_path}"
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"⚠️\u2009 Nilor-Nodes: No .env file found, relying on shell environment variables."
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
# Also allow default search (repo root / current working dir)
|
||||
load_dotenv()
|
||||
except Exception:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
# Use BaseConfig-backed singleton load
|
||||
if json5 is None:
|
||||
raise RuntimeError(
|
||||
"json5 module is required to load configuration from JSON5 file"
|
||||
)
|
||||
cfg = NilorNodesConfig.get_instance() # type: ignore[attr-defined]
|
||||
_apply_env_overrides(cfg)
|
||||
_CONFIG = cfg
|
||||
return _CONFIG
|
||||
|
||||
|
||||
def refresh_nilor_nodes_config() -> NilorNodesConfig:
|
||||
"""Clear the cached config and reload it. Intended for explicit hot-reload."""
|
||||
global _CONFIG
|
||||
_CONFIG = None
|
||||
return load_nilor_nodes_config()
|
||||
@@ -0,0 +1,37 @@
|
||||
import logging
|
||||
import os
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def configure_from_env(
|
||||
primary_env_var: str = "NILOR_LOG_LEVEL", fallback_env_var: str = "LOG_LEVEL"
|
||||
) -> None:
|
||||
# Prefer NILOR_LOG_LEVEL; fall back to LOG_LEVEL for backward compatibility
|
||||
chosen_var = primary_env_var if os.getenv(primary_env_var) else fallback_env_var
|
||||
value = os.getenv(chosen_var)
|
||||
|
||||
default_level = logging.INFO
|
||||
level = getattr(logging, value.upper(), None) if value else default_level
|
||||
if not isinstance(level, int):
|
||||
level = default_level
|
||||
logger.setLevel(level)
|
||||
|
||||
# Announce the effective level to the terminal via the global handlers
|
||||
effective_name = logging.getLevelName(level)
|
||||
if value:
|
||||
if getattr(logging, value.upper(), None) is None:
|
||||
logging.warning(
|
||||
f"⚠️ Nilor-Nodes: {chosen_var}='{value}' is invalid; defaulting to {effective_name}"
|
||||
)
|
||||
else:
|
||||
logging.info(
|
||||
f"ℹ️ Nilor-Nodes: {chosen_var}='{value}' → level set to {effective_name}"
|
||||
)
|
||||
else:
|
||||
logging.info(
|
||||
f"ℹ️ Nilor-Nodes: {primary_env_var} or {fallback_env_var} not set; defaulting to {effective_name}"
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["logger", "configure_from_env"]
|
||||
+121
-86
@@ -7,22 +7,15 @@ import logging
|
||||
import imageio.v2 as imageio
|
||||
import mimetypes
|
||||
import boto3
|
||||
import os
|
||||
import json
|
||||
from dotenv import load_dotenv
|
||||
|
||||
# --- 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)
|
||||
import tempfile
|
||||
import os
|
||||
from .logger import logger
|
||||
from .config.config import load_nilor_nodes_config
|
||||
|
||||
# --- Setup Logging ---
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
# Load shared configuration once
|
||||
_CFG = load_nilor_nodes_config()
|
||||
|
||||
# --- Node Categories ---
|
||||
category = "Nilor Nodes 👺"
|
||||
@@ -63,12 +56,12 @@ class MediaStreamInput:
|
||||
CATEGORY = category + subcategories["streaming"]
|
||||
|
||||
def download(
|
||||
self,
|
||||
presigned_download_url: str,
|
||||
format: str,
|
||||
input_name: str = "default_input",
|
||||
self,
|
||||
presigned_download_url: str,
|
||||
format: str,
|
||||
input_name: str = "default_input",
|
||||
):
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes: MediaStreamInput: Downloading from {presigned_download_url} for input '{input_name}' with format '{format}'"
|
||||
)
|
||||
try:
|
||||
@@ -78,7 +71,7 @@ class MediaStreamInput:
|
||||
manifest_response.raise_for_status()
|
||||
manifest = manifest_response.json()
|
||||
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes: Processing manifest for '{manifest.get('input_name')}' with {len(manifest.get('files', []))} assets."
|
||||
)
|
||||
|
||||
@@ -95,7 +88,7 @@ class MediaStreamInput:
|
||||
resp.raise_for_status()
|
||||
asset_responses.append(resp.content)
|
||||
except requests.RequestException as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes: Failed to download asset {file_info.get('filename')}: {e}"
|
||||
)
|
||||
raise # Re-raise to fail the entire process
|
||||
@@ -103,33 +96,49 @@ 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
|
||||
|
||||
if format == "video":
|
||||
return self._process_video(media_bytes)
|
||||
elif format == "image":
|
||||
return self._process_image(media_bytes)
|
||||
# Stream video to temp file to avoid loading entire video into RAM
|
||||
temp_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
|
||||
try:
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Streaming video to temp file: {temp_file.name}")
|
||||
with requests.get(presigned_download_url, timeout=180, stream=True) as response:
|
||||
response.raise_for_status()
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
temp_file.write(chunk)
|
||||
temp_file.close()
|
||||
return self._process_video(temp_file.name)
|
||||
finally:
|
||||
# Clean up temp file
|
||||
if os.path.exists(temp_file.name):
|
||||
os.unlink(temp_file.name)
|
||||
else:
|
||||
# Should not happen if UI choices are respected
|
||||
raise ValueError(
|
||||
f"[🛑] Nilor-Nodes (MediaStreamInput): Unsupported format '{format}' for single media download."
|
||||
)
|
||||
# For images, load into memory (they're small)
|
||||
response = requests.get(presigned_download_url, timeout=180)
|
||||
response.raise_for_status()
|
||||
media_bytes = response.content
|
||||
|
||||
if format == "image":
|
||||
return self._process_image(media_bytes)
|
||||
else:
|
||||
# Should not happen if UI choices are respected
|
||||
raise ValueError(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Unsupported format '{format}' for single media download."
|
||||
)
|
||||
|
||||
except requests.RequestException as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Failed to download file: {e}"
|
||||
)
|
||||
return (None,)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamInput): Failed to process media: {e}"
|
||||
)
|
||||
return (None,)
|
||||
|
||||
def _process_image_batch(self, image_bytes_list):
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing image batch with {len(image_bytes_list)} images..."
|
||||
)
|
||||
output_images = []
|
||||
@@ -147,13 +156,13 @@ class MediaStreamInput:
|
||||
# Concatenate along the batch dimension (dim=0)
|
||||
images_tensor = torch.cat(output_images, dim=0)
|
||||
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamInput): Image batch processing successful. Batch shape: {images_tensor.shape}"
|
||||
)
|
||||
return (images_tensor,)
|
||||
|
||||
def _process_image(self, image_bytes):
|
||||
logging.info("ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing as image...")
|
||||
logger.info("ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing as image...")
|
||||
image_pil = Image.open(io.BytesIO(image_bytes))
|
||||
|
||||
# Ensure image is in RGB
|
||||
@@ -162,30 +171,48 @@ class MediaStreamInput:
|
||||
np.array(rgb_image_pil).astype(np.float32) / 255.0
|
||||
).unsqueeze(0)
|
||||
|
||||
logging.info("✅ Nilor-Nodes (MediaStreamInput): Image processing successful.")
|
||||
logger.info("✅ Nilor-Nodes (MediaStreamInput): Image processing successful.")
|
||||
return (image_tensor,)
|
||||
|
||||
def _process_video(self, video_bytes):
|
||||
logging.info("ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing as video...")
|
||||
frames = []
|
||||
with imageio.get_reader(io.BytesIO(video_bytes), format="mp4") as reader:
|
||||
for frame in reader:
|
||||
# Convert frame to RGB PIL Image and then to tensor
|
||||
pil_image = Image.fromarray(frame).convert("RGB")
|
||||
numpy_image = np.array(pil_image).astype(np.float32) / 255.0
|
||||
tensor_frame = torch.from_numpy(numpy_image)
|
||||
frames.append(tensor_frame)
|
||||
def _process_video(self, video_path):
|
||||
logger.info(f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Processing video from {video_path}...")
|
||||
|
||||
if not frames:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (MediaStreamInput): No frames could be read from the video."
|
||||
# Open video to get metadata first
|
||||
with imageio.get_reader(video_path, format="mp4") as reader:
|
||||
# Get video metadata
|
||||
metadata = reader.get_meta_data()
|
||||
num_frames = reader.count_frames()
|
||||
|
||||
if num_frames == 0:
|
||||
raise ValueError(
|
||||
"🛑\u2009 Nilor-Nodes (MediaStreamInput): No frames could be read from the video."
|
||||
)
|
||||
|
||||
# Read first frame to get dimensions
|
||||
first_frame = reader.get_data(0)
|
||||
height, width = first_frame.shape[:2]
|
||||
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamInput): Video has {num_frames} frames at {width}x{height}"
|
||||
)
|
||||
|
||||
# Stack frames into a single tensor (batch of images)
|
||||
video_tensor = torch.stack(frames)
|
||||
# Pre-allocate tensor for all frames (N, H, W, 3)
|
||||
video_tensor = torch.empty((num_frames, height, width, 3), dtype=torch.float32)
|
||||
|
||||
logging.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamInput): Video processing successful. Image Shape: {video_tensor.shape}"
|
||||
# Process first frame (already read for dimensions)
|
||||
pil_image = Image.fromarray(first_frame).convert("RGB")
|
||||
numpy_image = np.array(pil_image).astype(np.float32) / 255.0
|
||||
video_tensor[0] = torch.from_numpy(numpy_image)
|
||||
|
||||
# Read remaining frames by explicit index to avoid iterator position ambiguity
|
||||
for i in range(1, num_frames):
|
||||
frame = reader.get_data(i)
|
||||
pil_image = Image.fromarray(frame).convert("RGB")
|
||||
numpy_image = np.array(pil_image).astype(np.float32) / 255.0
|
||||
video_tensor[i] = torch.from_numpy(numpy_image)
|
||||
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamInput): Video processing successful. Tensor shape: {video_tensor.shape}"
|
||||
)
|
||||
return (video_tensor,)
|
||||
|
||||
@@ -231,6 +258,10 @@ class MediaStreamOutput:
|
||||
"STRING",
|
||||
{"multiline": False, "default": "<auto-filled by system>"},
|
||||
),
|
||||
"job_type": (
|
||||
"STRING",
|
||||
{"default": "<auto-filled by system>", "multiline": False},
|
||||
),
|
||||
},
|
||||
"hidden": {
|
||||
"prompt": "PROMPT",
|
||||
@@ -238,30 +269,32 @@ class MediaStreamOutput:
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("uploaded_url",)
|
||||
FUNCTION = "upload_and_notify"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = category + subcategories["streaming"]
|
||||
|
||||
def upload_and_notify(
|
||||
self,
|
||||
images,
|
||||
format,
|
||||
content_id,
|
||||
venue,
|
||||
canvas,
|
||||
scene,
|
||||
presigned_upload_url,
|
||||
job_completions_queue_url,
|
||||
output_object_keys,
|
||||
framerate,
|
||||
output_name: str = "default_output",
|
||||
prompt=None,
|
||||
extra_pnginfo=None,
|
||||
self,
|
||||
images,
|
||||
format,
|
||||
content_id,
|
||||
venue,
|
||||
canvas,
|
||||
scene,
|
||||
presigned_upload_url,
|
||||
job_completions_queue_url,
|
||||
output_object_keys,
|
||||
framerate,
|
||||
output_name: str = "default_output",
|
||||
prompt=None,
|
||||
extra_pnginfo=None,
|
||||
job_type: str | None = None,
|
||||
):
|
||||
if not content_id:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (MediaStreamOutput): content_id is a required input for MediaStreamOutput."
|
||||
"🛑\u2009 Nilor-Nodes (MediaStreamOutput): content_id is a required input for MediaStreamOutput."
|
||||
)
|
||||
|
||||
# The `output_object_keys` is received as a string representation of a dictionary.
|
||||
@@ -271,7 +304,7 @@ class MediaStreamOutput:
|
||||
# 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:
|
||||
logging.error(
|
||||
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.
|
||||
@@ -303,36 +336,38 @@ 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=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"),
|
||||
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,
|
||||
)
|
||||
logging.info(
|
||||
logger.debug(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Sending completion message for content {content_id} to queue: {job_completions_queue_url}"
|
||||
)
|
||||
sqs_client.send_message(
|
||||
QueueUrl=job_completions_queue_url,
|
||||
MessageBody=json.dumps(completion_message),
|
||||
)
|
||||
logging.info(
|
||||
"✅ Nilor-Nodes (MediaStreamOutput): Completion message sent successfully."
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (MediaStreamOutput): Completion message sent successfully for content {content_id} to queue: {job_completions_queue_url}"
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): Failed to send completion message to SQS: {e}"
|
||||
)
|
||||
raise # Re-raise to fail the ComfyUI job
|
||||
|
||||
return {"ui": {"images": []}}
|
||||
return {"ui": {"images": []}, "result": (presigned_upload_url,)}
|
||||
|
||||
def _upload_image(self, image_tensor, url):
|
||||
logging.info(
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Uploading as PNG image..."
|
||||
)
|
||||
i = 255.0 * image_tensor.cpu().numpy()
|
||||
@@ -345,7 +380,7 @@ class MediaStreamOutput:
|
||||
self._perform_upload(buffer, url, "image/png")
|
||||
|
||||
def _upload_video(self, image_batch_tensor, url, framerate):
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Uploading as MP4 video. Frame count: {len(image_batch_tensor)}"
|
||||
)
|
||||
frames = []
|
||||
@@ -362,7 +397,7 @@ class MediaStreamOutput:
|
||||
|
||||
def _perform_upload(self, buffer, url, content_type):
|
||||
try:
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (MediaStreamOutput): Uploading to {url} with Content-Type: {content_type}"
|
||||
)
|
||||
headers = {"Content-Type": content_type}
|
||||
@@ -370,14 +405,14 @@ class MediaStreamOutput:
|
||||
url, data=buffer.read(), headers=headers, timeout=300
|
||||
)
|
||||
response.raise_for_status()
|
||||
logging.info("✅ Nilor-Nodes (MediaStreamOutput): Upload successful.")
|
||||
logger.info("✅ Nilor-Nodes (MediaStreamOutput): Upload successful.")
|
||||
except requests.RequestException as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): Failed to upload media: {e}"
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (MediaStreamOutput): Failed to process and upload media: {e}"
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -0,0 +1,432 @@
|
||||
"""
|
||||
Memory Hygiene (Guardian) — typed skeleton and public API.
|
||||
|
||||
This module provides a lightweight, dependency-injected component that can be
|
||||
invoked between jobs to assess memory pressure and (optionally) remediate via
|
||||
the ComfyUI server's `/free` endpoint. This commit introduces the types and
|
||||
public API only; detailed policy and remediation logic will be implemented in
|
||||
subsequent commits.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import asyncio
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Literal, Any
|
||||
|
||||
from .config.config import MemoryHygieneConfig
|
||||
from .comfyui_client import ComfyUIClientProtocol, SystemStats
|
||||
|
||||
|
||||
RemediationAction = Literal["none", "free", "unload", "both", "auto"]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RemediationResult:
|
||||
"""Outcome of a hygiene check/remediation cycle.
|
||||
|
||||
Attributes:
|
||||
before: Stats captured before any remediation attempt.
|
||||
after: Stats captured after remediation (if attempted), or None.
|
||||
action: Action that was selected/executed for this cycle.
|
||||
attempts: Number of remediation attempts performed.
|
||||
elapsed_seconds: Wall-clock duration of the cycle in seconds.
|
||||
reason: Optional decision rationale (e.g., threshold that triggered).
|
||||
success: True when targets were met or no action was needed.
|
||||
"""
|
||||
|
||||
before: SystemStats
|
||||
after: Optional[SystemStats]
|
||||
action: RemediationAction
|
||||
attempts: int
|
||||
elapsed_seconds: float
|
||||
reason: Optional[str]
|
||||
success: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DerivedStats:
|
||||
"""Computed metrics used by the policy engine.
|
||||
|
||||
Attributes:
|
||||
vram_total: Total VRAM (bytes) if known.
|
||||
vram_free: Free VRAM (bytes) if known.
|
||||
vram_used_pct: Percent VRAM used in [0, 100] when computable.
|
||||
ram_total: Total RAM (bytes) if known.
|
||||
ram_free: Free RAM (bytes) if known.
|
||||
ram_used_pct: Percent RAM used in [0, 100] when computable.
|
||||
"""
|
||||
|
||||
vram_total: Optional[float]
|
||||
vram_free: Optional[float]
|
||||
vram_used_pct: Optional[float]
|
||||
ram_total: Optional[float]
|
||||
ram_free: Optional[float]
|
||||
ram_used_pct: Optional[float]
|
||||
|
||||
|
||||
class MemoryHygiene:
|
||||
"""Memory Guardian component orchestrating checks and remediation between jobs."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
client: ComfyUIClientProtocol,
|
||||
cfg: MemoryHygieneConfig,
|
||||
logger: Optional[Any] = None,
|
||||
) -> None:
|
||||
self._client = client
|
||||
self._cfg = cfg
|
||||
self._logger = logger
|
||||
self._cooldown_until: float = 0.0
|
||||
# Hardened capability state: disable guardian for session after unsupported
|
||||
self._capability_disabled: bool = False
|
||||
self._unsupported_warned: bool = False
|
||||
|
||||
async def check_and_remediate(self) -> RemediationResult:
|
||||
"""Run a single hygiene cycle.
|
||||
|
||||
This skeleton performs capability and enablement checks, then captures a
|
||||
baseline stats snapshot and returns without modification. Subsequent
|
||||
commits implement policy evaluation and remediation.
|
||||
"""
|
||||
start_ts = time.monotonic()
|
||||
|
||||
if not self._cfg.enabled:
|
||||
before, _ = await self._collect_metrics()
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.debug(
|
||||
"⚠️\u2009 Nilor-Nodes (memory_hygiene): disabled; skipping remediation. vram_free=%s ram_free=%s",
|
||||
getattr(before, "vram_free", None),
|
||||
getattr(before, "ram_free", None),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason="disabled",
|
||||
success=True,
|
||||
)
|
||||
|
||||
# Session-wide disable if previously detected unsupported endpoints
|
||||
if self._capability_disabled:
|
||||
return RemediationResult(
|
||||
before=await self._get_stats_safe(),
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason="capability_unsupported",
|
||||
success=True,
|
||||
)
|
||||
|
||||
try:
|
||||
supported = await self._client.supports_hygiene()
|
||||
except Exception:
|
||||
supported = False
|
||||
|
||||
before, derived = await self._collect_metrics()
|
||||
|
||||
if not supported:
|
||||
# Disable for the rest of the session and emit a single warning
|
||||
self._capability_disabled = True
|
||||
if not self._unsupported_warned:
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.warning(
|
||||
"⚠️\u2009 Nilor-Nodes (memory_hygiene): unsupported endpoints; disabling for this session."
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
self._unsupported_warned = True
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason="capability_unsupported",
|
||||
success=True,
|
||||
)
|
||||
|
||||
now = time.monotonic()
|
||||
if now < self._cooldown_until:
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (memory_hygiene): cooldown active for %.2fs; skipping.",
|
||||
max(0.0, self._cooldown_until - now),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, now - start_ts),
|
||||
reason="cooldown_active",
|
||||
success=True,
|
||||
)
|
||||
|
||||
pressure_reason = self._pressure_reason(derived)
|
||||
if pressure_reason is None:
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (memory_hygiene): no pressure; nothing to do."
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=None,
|
||||
action="none",
|
||||
attempts=0,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason="no_pressure",
|
||||
success=True,
|
||||
)
|
||||
|
||||
action = self._choose_initial_action(self._cfg.action_policy)
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.info(
|
||||
"✅ Nilor-Nodes (memory_hygiene): start cycle action=%s reason=%s vram_free=%s ram_free=%s",
|
||||
action,
|
||||
pressure_reason,
|
||||
getattr(before, "vram_free", None),
|
||||
getattr(before, "ram_free", None),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
after, attempts, final_action, outcome_reason, success = (
|
||||
await self._remediate_cycle(
|
||||
initial_action=action,
|
||||
initial_reason=pressure_reason,
|
||||
cycle_start=start_ts,
|
||||
)
|
||||
)
|
||||
# Set cooldown after any attempted cycle (success or not)
|
||||
self._cooldown_until = time.monotonic() + float(self._cfg.cooldown_seconds)
|
||||
try:
|
||||
if self._logger:
|
||||
self._logger.info(
|
||||
"✅ Nilor-Nodes (memory_hygiene): end cycle action=%s attempts=%s success=%s reason=%s vram_free_before=%s vram_free_after=%s ram_free_before=%s ram_free_after=%s elapsed=%.2fs",
|
||||
final_action,
|
||||
attempts,
|
||||
success,
|
||||
outcome_reason,
|
||||
getattr(before, "vram_free", None),
|
||||
getattr(after, "vram_free", None) if after else None,
|
||||
getattr(before, "ram_free", None),
|
||||
getattr(after, "ram_free", None) if after else None,
|
||||
max(0.0, time.monotonic() - start_ts),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return RemediationResult(
|
||||
before=before,
|
||||
after=after,
|
||||
action=final_action,
|
||||
attempts=attempts,
|
||||
elapsed_seconds=max(0.0, time.monotonic() - start_ts),
|
||||
reason=outcome_reason,
|
||||
success=success,
|
||||
)
|
||||
|
||||
async def _get_stats_safe(self) -> SystemStats:
|
||||
try:
|
||||
return await self._client.get_system_stats()
|
||||
except Exception:
|
||||
# Return an empty struct; callers tolerate partial data
|
||||
return SystemStats()
|
||||
|
||||
async def _collect_metrics(self) -> tuple[SystemStats, DerivedStats]:
|
||||
base = await self._get_stats_safe()
|
||||
vram_used_pct = _compute_used_pct(base.vram_total, base.vram_free)
|
||||
ram_used_pct = _compute_used_pct(base.ram_total, base.ram_free)
|
||||
derived = DerivedStats(
|
||||
vram_total=base.vram_total,
|
||||
vram_free=base.vram_free,
|
||||
vram_used_pct=vram_used_pct,
|
||||
ram_total=base.ram_total,
|
||||
ram_free=base.ram_free,
|
||||
ram_used_pct=ram_used_pct,
|
||||
)
|
||||
return base, derived
|
||||
|
||||
def _pressure_reason(self, d: DerivedStats) -> Optional[str]:
|
||||
if _gt_pct(d.vram_used_pct, self._cfg.vram_usage_pct_max):
|
||||
return f"vram_used_pct {d.vram_used_pct}% > {self._cfg.vram_usage_pct_max}%"
|
||||
if _lt_bytes(d.vram_free, self._cfg.vram_min_free_mb):
|
||||
return f"vram_free below {self._cfg.vram_min_free_mb}MB"
|
||||
if _gt_pct(d.ram_used_pct, self._cfg.ram_usage_pct_max):
|
||||
return f"ram_used_pct {d.ram_used_pct}% > {self._cfg.ram_usage_pct_max}%"
|
||||
if _lt_bytes(d.ram_free, self._cfg.ram_min_free_mb):
|
||||
return f"ram_free below {self._cfg.ram_min_free_mb}MB"
|
||||
return None
|
||||
|
||||
def _choose_initial_action(self, policy: str) -> RemediationAction:
|
||||
return _initial_action_for(policy)
|
||||
|
||||
async def _remediate_cycle(
|
||||
self,
|
||||
*,
|
||||
initial_action: RemediationAction,
|
||||
initial_reason: str,
|
||||
cycle_start: float,
|
||||
) -> tuple[Optional[SystemStats], int, RemediationAction, str, bool]:
|
||||
attempts = 0
|
||||
action = initial_action
|
||||
escalated = False
|
||||
last_stats: Optional[SystemStats] = None
|
||||
|
||||
max_retries = max(0, int(self._cfg.max_retries))
|
||||
sleep_between = max(0, int(self._cfg.sleep_between_attempts_seconds))
|
||||
max_cycle_s = max(0, int(self._cfg.max_cycle_duration_seconds))
|
||||
|
||||
def time_budget_exhausted() -> bool:
|
||||
if max_cycle_s <= 0:
|
||||
return False
|
||||
return (time.monotonic() - cycle_start) >= max_cycle_s
|
||||
|
||||
outcome_reason = initial_reason
|
||||
|
||||
while True:
|
||||
# Execute remediation step
|
||||
free_flag, unload_flag = _flags_for_action(action)
|
||||
try:
|
||||
if self._logger:
|
||||
try:
|
||||
self._logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (memory_hygiene): calling /free free_memory=%s unload_models=%s",
|
||||
free_flag,
|
||||
unload_flag,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
await self._client.free(
|
||||
free_memory=free_flag, unload_models=unload_flag
|
||||
)
|
||||
except Exception:
|
||||
# Continue even on errors; treat as unsuccessful attempt
|
||||
pass
|
||||
|
||||
# Wait and re-measure
|
||||
if sleep_between > 0:
|
||||
try:
|
||||
await asyncio.sleep(sleep_between)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
last, derived = await self._collect_metrics()
|
||||
last_stats = last
|
||||
if self._pressure_reason(derived) is None:
|
||||
outcome_reason = "targets_met"
|
||||
return last_stats, attempts + 1, action, outcome_reason, True
|
||||
|
||||
attempts += 1
|
||||
if attempts > max_retries:
|
||||
outcome_reason = "max_retries_exhausted"
|
||||
return last_stats, attempts, action, outcome_reason, False
|
||||
|
||||
if time_budget_exhausted():
|
||||
outcome_reason = "max_duration_reached"
|
||||
return last_stats, attempts, action, outcome_reason, False
|
||||
|
||||
# Escalation logic for auto: free -> unload (once)
|
||||
if initial_action == "auto" and not escalated:
|
||||
action = "unload"
|
||||
escalated = True
|
||||
# For explicit free/unload/both, keep same action for subsequent attempts
|
||||
|
||||
|
||||
def _compute_used_pct(total: Optional[float], free: Optional[float]) -> Optional[float]:
|
||||
try:
|
||||
if total is None or free is None:
|
||||
return None
|
||||
total_f = float(total)
|
||||
free_f = float(free)
|
||||
if total_f <= 0:
|
||||
return None
|
||||
used = max(0.0, min(1.0, (total_f - max(0.0, free_f)) / total_f))
|
||||
return round(used * 100.0, 2)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _mb_to_bytes(mb: Optional[int]) -> Optional[float]:
|
||||
if mb is None:
|
||||
return None
|
||||
try:
|
||||
return float(max(0, int(mb))) * 1024.0 * 1024.0
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _gt_pct(value: Optional[float], threshold_pct: Optional[int]) -> bool:
|
||||
if value is None or threshold_pct is None:
|
||||
return False
|
||||
try:
|
||||
return float(value) > float(threshold_pct)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _lt_bytes(value_bytes: Optional[float], threshold_mb: Optional[int]) -> bool:
|
||||
if value_bytes is None or threshold_mb is None:
|
||||
return False
|
||||
thr = _mb_to_bytes(threshold_mb)
|
||||
if thr is None:
|
||||
return False
|
||||
try:
|
||||
return float(value_bytes) < float(thr)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _normalize_policy(policy: str) -> str:
|
||||
try:
|
||||
return str(policy).strip().lower()
|
||||
except Exception:
|
||||
return "auto"
|
||||
|
||||
|
||||
def _action_literal(policy: str) -> RemediationAction:
|
||||
p = _normalize_policy(policy)
|
||||
if p in ("free", "unload", "both"):
|
||||
return p # type: ignore[return-value]
|
||||
return "auto"
|
||||
|
||||
|
||||
def _initial_action_for(policy: str) -> RemediationAction:
|
||||
lit = _action_literal(policy)
|
||||
if lit == "auto":
|
||||
return "free"
|
||||
return lit
|
||||
|
||||
|
||||
def _flags_for_action(action: RemediationAction) -> tuple[bool, bool]:
|
||||
if action == "free":
|
||||
return True, False
|
||||
if action == "unload":
|
||||
return False, True
|
||||
if action == "both":
|
||||
return True, True
|
||||
# auto is staged; when executing a step, treat like free unless escalated update chooses unload
|
||||
return True, False
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RemediationResult",
|
||||
"RemediationAction",
|
||||
"MemoryHygiene",
|
||||
"DerivedStats",
|
||||
]
|
||||
+219
-52
@@ -18,9 +18,22 @@ from pathlib import Path
|
||||
import cv2
|
||||
import warnings
|
||||
from .utils import pil2tensor, tensor2pil
|
||||
import logging
|
||||
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
|
||||
@@ -171,7 +184,7 @@ class NilorRemapFloatList:
|
||||
# Avoid division by zero
|
||||
if max_input - min_input == 0:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (RemapFloatList): max_input and min_input cannot be the same value."
|
||||
"🛑\u2009 Nilor-Nodes (RemapFloatList): max_input and min_input cannot be the same value."
|
||||
)
|
||||
|
||||
scale = (max_output - min_output) / (max_input - min_input)
|
||||
@@ -228,7 +241,7 @@ class NilorInverseMapFloatList:
|
||||
def inverse_map_float_list(self, list_of_floats):
|
||||
if not list_of_floats:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (InverseMapFloatList): The input list_of_floats cannot be empty."
|
||||
"🛑\u2009 Nilor-Nodes (InverseMapFloatList): The input list_of_floats cannot be empty."
|
||||
)
|
||||
|
||||
min_input = min(list_of_floats)
|
||||
@@ -326,7 +339,7 @@ class NilorCountImagesInDirectory:
|
||||
def count_images_in_directory(self, directory):
|
||||
if not os.path.isdir(directory):
|
||||
raise FileNotFoundError(
|
||||
f"[🛑] Nilor-Nodes (NilorCountImagesInDirectory): Directory '{directory} cannot be found."
|
||||
f"🛑\u2009 Nilor-Nodes (NilorCountImagesInDirectory): Directory '{directory}' cannot be found."
|
||||
)
|
||||
|
||||
list_dir = []
|
||||
@@ -376,7 +389,7 @@ class NilorSelectIndexFromList:
|
||||
# Ensure the index is within bounds
|
||||
if index < 0 or index >= len(actual_list):
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (SelectIndexFromList): Index is outside the bounds of the array."
|
||||
"🛑\u2009 Nilor-Nodes (SelectIndexFromList): Index is outside the bounds of the array."
|
||||
)
|
||||
|
||||
# Returns the value at the given index
|
||||
@@ -415,7 +428,7 @@ class NilorSaveEXRArbitrary:
|
||||
self, channels=None, filename_prefix="output", prompt=None, extra_pnginfo=None
|
||||
):
|
||||
|
||||
logging.info(
|
||||
logger.info(
|
||||
"ℹ️\u2009 Nilor-Nodes (SaveEXRArbitrary): Running save_exr_arbitrary"
|
||||
)
|
||||
# print(f"channels: {channels}")
|
||||
@@ -429,7 +442,7 @@ class NilorSaveEXRArbitrary:
|
||||
try:
|
||||
actual_channels[0]
|
||||
except TypeError:
|
||||
logging.error(
|
||||
logger.error(
|
||||
"🛑\u2009 Nilor-Nodes (SaveEXRArbitrary): actual_channels is not subscriptable"
|
||||
)
|
||||
return
|
||||
@@ -469,7 +482,7 @@ class NilorSaveEXRArbitrary:
|
||||
for tensor in image_channels:
|
||||
if tensor.shape[-2:] != (height, width):
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (SaveEXRArbitrary): All input tensors must have the same dimensions"
|
||||
"🛑\u2009 Nilor-Nodes (SaveEXRArbitrary): All input tensors must have the same dimensions"
|
||||
)
|
||||
|
||||
# Channel naming
|
||||
@@ -520,11 +533,11 @@ class NilorSaveEXRArbitrary:
|
||||
exr_file.writePixels(channel_data)
|
||||
exr_file.close()
|
||||
|
||||
logging.info(
|
||||
f"✅ Nilor-Nodes (SaveEXRArbitrary): EXR file saved successfully to {writepath}"
|
||||
logger.info(
|
||||
f"✅\u2009 Nilor-Nodes (SaveEXRArbitrary): EXR file saved successfully to {writepath}"
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (SaveEXRArbitrary): Failed to write EXR file: {e}"
|
||||
)
|
||||
|
||||
@@ -645,13 +658,13 @@ class NilorShuffleImageBatch:
|
||||
def _check_image_dimensions(self, images):
|
||||
if images.shape[0] == 0:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (ShuffleImageBatch): Input images tensor is empty."
|
||||
"🛑\u2009 Nilor-Nodes (ShuffleImageBatch): Input images tensor is empty."
|
||||
)
|
||||
|
||||
# All images in the batch should have the same dimensions
|
||||
if len(images.shape) != 4:
|
||||
raise ValueError(
|
||||
f"[🛑] Nilor-Nodes (ShuffleImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
f"🛑\u2009 Nilor-Nodes (ShuffleImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
)
|
||||
|
||||
def shuffle_image_batch(self, images: torch.Tensor, seed):
|
||||
@@ -692,13 +705,13 @@ class NilorRepeatTrimImageBatch:
|
||||
def _check_image_dimensions(self, images):
|
||||
if images.shape[0] == 0:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (RepeatTrimImageBatch): Input images tensor is empty."
|
||||
"🛑\u2009 Nilor-Nodes (RepeatTrimImageBatch): Input images tensor is empty."
|
||||
)
|
||||
|
||||
# All images in the batch should have the same dimensions
|
||||
if len(images.shape) != 4:
|
||||
raise ValueError(
|
||||
f"[🛑] Nilor-Nodes (RepeatTrimImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
f"🛑\u2009 Nilor-Nodes (RepeatTrimImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
)
|
||||
|
||||
def repeat_trim_image_batch(self, images: torch.Tensor, count):
|
||||
@@ -737,13 +750,13 @@ class NilorRepeatShuffleTrimImageBatch:
|
||||
def _check_image_dimensions(self, images):
|
||||
if images.shape[0] == 0:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (RepeatShuffleTrimImageBatch): Input images tensor is empty."
|
||||
"🛑\u2009 Nilor-Nodes (RepeatShuffleTrimImageBatch): Input images tensor is empty."
|
||||
)
|
||||
|
||||
# All images in the batch should have the same dimensions
|
||||
if len(images.shape) != 4:
|
||||
raise ValueError(
|
||||
f"[🛑] Nilor-Nodes (RepeatShuffleTrimImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
f"🛑\u2009 Nilor-Nodes (RepeatShuffleTrimImageBatch): Expected 4D tensor (batch, channels, height, width), got shape {images.shape}"
|
||||
)
|
||||
|
||||
def repeat_shuffle_trim_image_batch(self, images: torch.Tensor, seed, count):
|
||||
@@ -809,14 +822,14 @@ class NilorOutputFilenameString:
|
||||
|
||||
if unique_id is not None and extra_pnginfo is not None:
|
||||
if not isinstance(extra_pnginfo, list):
|
||||
logging.error(
|
||||
logger.error(
|
||||
"🛑\u2009 Nilor-Nodes (OutputFilenameString): extra_pnginfo is not a list"
|
||||
)
|
||||
elif (
|
||||
not isinstance(extra_pnginfo[0], dict)
|
||||
or "workflow" not in extra_pnginfo[0]
|
||||
):
|
||||
logging.error(
|
||||
logger.error(
|
||||
"🛑\u2009 Nilor-Nodes (OutputFilenameString): extra_pnginfo[0] is not a dict or missing 'workflow' key"
|
||||
)
|
||||
else:
|
||||
@@ -870,7 +883,7 @@ class NilorNFractionsOfInt:
|
||||
return ([i * numerator // (denominator - 1) for i in range(denominator)],)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"[🛑] Nilor-Nodes (NilorNFractionsOfInt): Unknown type: {type}"
|
||||
f"🛑\u2009 Nilor-Nodes (NilorNFractionsOfInt): Unknown type: {type}"
|
||||
)
|
||||
|
||||
|
||||
@@ -951,17 +964,17 @@ class NilorWanTileResolution:
|
||||
for name, value in dims.items():
|
||||
if value <= 0:
|
||||
raise ValueError(
|
||||
f"[🛑] Nilor-Nodes (NilorWanTileResolution): {name} must be a positive integer."
|
||||
f"🛑\u2009 Nilor-Nodes (NilorWanTileResolution): {name} must be a positive integer."
|
||||
)
|
||||
|
||||
if input_width % 16 != 0 or input_height % 16 != 0:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (NilorWanTileResolution): input_width and input_height must be multiples of 16."
|
||||
"🛑\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(
|
||||
"[🛑] Nilor-Nodes (NilorWanTileResolution): target_width and target_height must be at least the minimum tile size."
|
||||
"🛑\u2009 Nilor-Nodes (NilorWanTileResolution): target_width and target_height must be at least the minimum tile size."
|
||||
)
|
||||
|
||||
min_blocks = self.MIN_TILE_DIM // 16
|
||||
@@ -972,7 +985,7 @@ class NilorWanTileResolution:
|
||||
|
||||
if max_width_blocks < min_blocks or max_height_blocks < min_blocks:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (NilorWanTileResolution): Target dimensions do not allow a tile within the supported range."
|
||||
"🛑\u2009 Nilor-Nodes (NilorWanTileResolution): Target dimensions do not allow a tile within the supported range."
|
||||
)
|
||||
|
||||
aspect_ratio = input_width / input_height
|
||||
@@ -1018,12 +1031,59 @@ class NilorWanTileResolution:
|
||||
if best_dimensions is None:
|
||||
# If no suitable tile resolution was found, raise an error
|
||||
raise RuntimeError(
|
||||
"[🛑] Nilor-Nodes (NilorWanTileResolution): Failed to determine a suitable tile resolution."
|
||||
"🛑\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,)
|
||||
|
||||
|
||||
class NilorCategorizeString:
|
||||
def __init__(self):
|
||||
pass
|
||||
@@ -1144,7 +1204,7 @@ class NilorRandomString:
|
||||
]
|
||||
if not options:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (NilorRandomString): No valid choices provided."
|
||||
"🛑\u2009 Nilor-Nodes (NilorRandomString): No valid choices provided."
|
||||
)
|
||||
|
||||
# Limit to the first 'max_options' entries if there are more options
|
||||
@@ -1188,7 +1248,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"[🛑] Nilor-Nodes (NilorLoadImageByIndex): Image directory {image_directory} does not exist"
|
||||
f"🛑\u2009 Nilor-Nodes (NilorLoadImageByIndex): Image directory {image_directory} does not exist"
|
||||
)
|
||||
|
||||
# Get list of image files
|
||||
@@ -1202,7 +1262,7 @@ class NilorLoadImageByIndex:
|
||||
|
||||
if not files:
|
||||
raise ValueError(
|
||||
f"[🛑] Nilor-Nodes (NilorLoadImageByIndex): No image files found in {image_directory}"
|
||||
f"🛑\u2009 Nilor-Nodes (NilorLoadImageByIndex): No image files found in {image_directory}"
|
||||
)
|
||||
|
||||
# Sort files based on selected mode
|
||||
@@ -1254,7 +1314,7 @@ class NilorExtractFilenameFromPath:
|
||||
# Ensure the input is a valid path
|
||||
if not filepath:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (ExtractFilenameFromPath): Filepath cannot be empty."
|
||||
"🛑\u2009 Nilor-Nodes (ExtractFilenameFromPath): Filepath cannot be empty."
|
||||
)
|
||||
|
||||
path = Path(filepath)
|
||||
@@ -1291,7 +1351,7 @@ class NilorBlurAnalysis:
|
||||
# Ensure images is a 4D tensor.
|
||||
if images.dim() != 4:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (BlurAnalysis): Input images must be a 4D tensor (batch, channels/height, height/width, width/channels)"
|
||||
"🛑\u2009 Nilor-Nodes (BlurAnalysis): Input images must be a 4D tensor (batch, channels/height, height/width, width/channels)"
|
||||
)
|
||||
|
||||
# Detect if using NCHW or NHWC.
|
||||
@@ -1300,7 +1360,7 @@ class NilorBlurAnalysis:
|
||||
images = images.permute(0, 3, 1, 2)
|
||||
else:
|
||||
raise ValueError(
|
||||
"[🛑] Nilor-Nodes (BlurAnalysis): Cannot determine image format (expected channel to be 1 or 3)."
|
||||
"🛑\u2009 Nilor-Nodes (BlurAnalysis): Cannot determine image format (expected channel to be 1 or 3)."
|
||||
)
|
||||
|
||||
output_images = []
|
||||
@@ -1382,6 +1442,7 @@ class NilorToSparseIndexMethod:
|
||||
|
||||
class NilorImageResizeV2:
|
||||
upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
@@ -1390,15 +1451,41 @@ class NilorImageResizeV2:
|
||||
"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}),
|
||||
"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}),
|
||||
"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."}),
|
||||
"per_batch": (
|
||||
"INT",
|
||||
{
|
||||
"default": 16,
|
||||
"min": 0,
|
||||
"max": 4096,
|
||||
"step": 1,
|
||||
"tooltip": "Process images in sub-batches. 0 disables.",
|
||||
},
|
||||
),
|
||||
},
|
||||
"hidden": {"unique_id": "UNIQUE_ID"},
|
||||
}
|
||||
@@ -1411,12 +1498,28 @@ class NilorImageResizeV2:
|
||||
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):
|
||||
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("Lanczos is not supported on the GPU")
|
||||
raise Exception(
|
||||
"🛑\u2009 Nilor-Nodes (NilorImageResizeV2): Lanczos is not supported on the GPU"
|
||||
)
|
||||
device = model_management.get_torch_device()
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
@@ -1427,7 +1530,11 @@ Resizes images with optional aspect preservation, padding/cropping, and sub-batc
|
||||
height = H
|
||||
|
||||
pillarbox_blur = keep_proportion == "pillarbox_blur"
|
||||
if keep_proportion == "resize" or keep_proportion.startswith("pad") or 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)
|
||||
@@ -1484,13 +1591,19 @@ Resizes images with optional aspect preservation, padding/cropping, and sub-batc
|
||||
bytes_per_elem = image.element_size()
|
||||
est_total_bytes = B * height * width * C * bytes_per_elem
|
||||
est_mb = est_total_bytes / (1024 * 1024)
|
||||
print(f"[NilorImageResizeV2] estimated output ~{est_mb:.2f} MB; batching {per_batch}/{B}")
|
||||
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))
|
||||
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]
|
||||
@@ -1522,14 +1635,30 @@ Resizes images with optional aspect preservation, padding/cropping, and sub-batc
|
||||
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)
|
||||
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]
|
||||
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)
|
||||
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):
|
||||
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:
|
||||
@@ -1543,12 +1672,30 @@ Resizes images with optional aspect preservation, padding/cropping, and sub-batc
|
||||
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"
|
||||
"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,
|
||||
)
|
||||
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
|
||||
|
||||
@@ -1567,9 +1714,13 @@ Resizes images with optional aspect preservation, padding/cropping, and sub-batc
|
||||
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)
|
||||
mask_chunks.append(
|
||||
sub_out_mask.cpu() if sub_out_mask is not None else None
|
||||
)
|
||||
try:
|
||||
print(f"[NilorImageResizeV2] batch {current_batch}/{total_batches} · images {end_idx}/{B}")
|
||||
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)
|
||||
@@ -1578,7 +1729,21 @@ Resizes images with optional aspect preservation, padding/cropping, and sub-batc
|
||||
else:
|
||||
out_mask = None
|
||||
|
||||
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))
|
||||
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 = {
|
||||
@@ -1607,6 +1772,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"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
|
||||
@@ -1636,4 +1802,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"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
-1
@@ -4,7 +4,7 @@ aiofiles>=23.2.1
|
||||
aiohttp==3.12.14
|
||||
boto3==1.40.15
|
||||
fastapi==0.110.0
|
||||
huggingface_hub==0.33.4
|
||||
huggingface_hub==0.34.0
|
||||
imageio==2.37.0
|
||||
imageio-ffmpeg==0.6.0
|
||||
numpy>=1.26.4
|
||||
@@ -16,4 +16,5 @@ python-multipart==0.0.9
|
||||
requests==2.31.0
|
||||
uvicorn==0.27.1
|
||||
websockets==11.0.3
|
||||
json5>=0.9.0
|
||||
--prefer-binary
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
"""Shared types and enums for the Nilor-Nodes sidecar.
|
||||
|
||||
This module intentionally contains minimal placeholders to support the
|
||||
configuration loader and future extensions without introducing unnecessary
|
||||
complexity at this stage.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ConfigSource(Enum):
|
||||
"""Represents the origin of a configuration value."""
|
||||
|
||||
ENV = "env"
|
||||
JSON5 = "json5"
|
||||
+1
-1
@@ -59,7 +59,7 @@ class NilorUserInput_Float:
|
||||
"STRING",
|
||||
{"default": "my_float_input", "multiline": False},
|
||||
),
|
||||
"value": ("FLOAT", {"default": 0.0}),
|
||||
"value": ("FLOAT", {"default": 0.0, "step": 0.001}),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+82
-31
@@ -19,45 +19,96 @@ function hideWidgets(node, widgetNames) {
|
||||
});
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "comfy.nilor-nodes.mediaStream",
|
||||
nodeCreated(node) {
|
||||
if (node.comfyClass === "MediaStreamOutput") {
|
||||
// Hide system inputs by default
|
||||
hideWidgets(node, [
|
||||
"content_id",
|
||||
"venue",
|
||||
"canvas",
|
||||
"scene",
|
||||
"presigned_upload_url",
|
||||
"job_completions_queue_url",
|
||||
"output_object_keys"
|
||||
]);
|
||||
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");
|
||||
const formatWidget = node.widgets.find((w) => w.name === "format");
|
||||
if (!formatWidget) return;
|
||||
|
||||
// Initial toggle for framerate based on the default format value
|
||||
toggleFramerateWidget(node, formatWidget.value === "mp4");
|
||||
// Apply current value
|
||||
toggleFramerateWidget(node, formatWidget.value === "mp4");
|
||||
try {
|
||||
const size = node.computeSize();
|
||||
node.onResize?.(size);
|
||||
app.graph?.setDirtyCanvas(true, true);
|
||||
} catch (_) {}
|
||||
|
||||
// Store original callback to chain it
|
||||
const originalCallback = formatWidget.callback;
|
||||
|
||||
formatWidget.callback = function (value) {
|
||||
toggleFramerateWidget(node, value === "mp4");
|
||||
|
||||
// Recalculate node size after toggling widgets
|
||||
// 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);
|
||||
|
||||
if (originalCallback) {
|
||||
return originalCallback.apply(this, arguments);
|
||||
}
|
||||
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 (node.comfyClass === "MediaStreamInput") {
|
||||
// Hide system inputs by default
|
||||
hideWidgets(node, ["presigned_download_url"]);
|
||||
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);
|
||||
},
|
||||
});
|
||||
|
||||
+437
-205
@@ -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` SQS queue for new jobs,
|
||||
Its purpose is to poll the `jobs_to_process-comfyui` SQS queue for new jobs,
|
||||
submit them to the local ComfyUI server, and manage the message lifecycle.
|
||||
It also listens to the ComfyUI websocket to send a "running" status update
|
||||
at the precise moment that job execution begins.
|
||||
@@ -12,50 +12,25 @@ import os
|
||||
import json
|
||||
import logging
|
||||
import asyncio
|
||||
import time
|
||||
import aiohttp
|
||||
import websockets
|
||||
from aiobotocore.session import get_session
|
||||
from dotenv import load_dotenv
|
||||
from botocore.exceptions import EndpointConnectionError
|
||||
|
||||
# --- Load Environment Variables ---
|
||||
# Load from the .env file in the same directory
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
dotenv_path = os.path.join(current_dir, ".env")
|
||||
if os.path.exists(dotenv_path):
|
||||
load_dotenv(dotenv_path=dotenv_path)
|
||||
logging.info(
|
||||
f"✅\u2009 Nilor-Nodes: Loaded environment variables from {dotenv_path}"
|
||||
)
|
||||
else:
|
||||
logging.info(
|
||||
"⚠️\u2009 Nilor-Nodes: No .env file found, relying on shell environment variables."
|
||||
)
|
||||
from aiobotocore.session import get_session
|
||||
from botocore.exceptions import EndpointConnectionError, ClientError
|
||||
from .logger import logger
|
||||
from .comfyui_client import ComfyUILocalClient, ComfyUIClientError
|
||||
from .memory_hygiene import MemoryHygiene
|
||||
from .workflow_normalizer import normalize_comfyui_prompt_for_current_os
|
||||
from .config.config import load_nilor_nodes_config, NilorNodesConfig
|
||||
|
||||
|
||||
# --- Configuration ---
|
||||
SQS_ENDPOINT_URL = os.getenv("SQS_ENDPOINT_URL", "http://localhost:9324")
|
||||
SQS_JOBS_TO_PROCESS_QUEUE_NAME = os.getenv(
|
||||
"SQS_JOBS_TO_PROCESS_QUEUE_NAME", "jobs_to_process"
|
||||
)
|
||||
SQS_JOB_STATUS_UPDATES_QUEUE_NAME = os.getenv(
|
||||
"SQS_JOB_STATUS_UPDATES_QUEUE_NAME", "job_status_updates"
|
||||
)
|
||||
COMFYUI_API_URL = os.getenv("COMFYUI_API_URL", "http://127.0.0.1:8188") + "/prompt"
|
||||
COMFYUI_WS_URL = os.getenv("COMFYUI_WS_URL", "ws://127.0.0.1:8188") + "/ws"
|
||||
AWS_ACCESS_KEY_ID = os.getenv("AWS_ACCESS_KEY_ID", "local")
|
||||
AWS_SECRET_ACCESS_KEY = os.getenv("AWS_SECRET_ACCESS_KEY", "local")
|
||||
AWS_DEFAULT_REGION = os.getenv("AWS_DEFAULT_REGION", "us-east-1")
|
||||
POLL_WAIT_TIME_SECONDS = 20 # SQS Long Polling
|
||||
MAX_MESSAGES = 1
|
||||
|
||||
# --- Setup Logging ---
|
||||
logging.basicConfig(
|
||||
level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s"
|
||||
)
|
||||
# Centralized loader provides precedence env > JSON5 and validation
|
||||
_CFG: NilorNodesConfig = load_nilor_nodes_config()
|
||||
|
||||
|
||||
class WorkerConsumer:
|
||||
def __init__(self):
|
||||
def __init__(self, cfg: NilorNodesConfig):
|
||||
self.session = get_session()
|
||||
self.prompt_id_to_content_id_map = {}
|
||||
self.sent_running_status_prompts = set()
|
||||
@@ -63,32 +38,45 @@ 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=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
region_name=self.cfg.worker.aws_region,
|
||||
endpoint_url=self.cfg.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=self.cfg.worker.aws_access_key_id,
|
||||
aws_secret_access_key=self.cfg.worker.aws_secret_access_key,
|
||||
) as client:
|
||||
try:
|
||||
self.jobs_queue_url = await self._get_queue_url(
|
||||
client, SQS_JOBS_TO_PROCESS_QUEUE_NAME
|
||||
client, self.cfg.worker.jobs_queue
|
||||
)
|
||||
self.status_updates_queue_url = await self._get_queue_url(
|
||||
client, SQS_JOB_STATUS_UPDATES_QUEUE_NAME
|
||||
client, self.cfg.worker.status_queue
|
||||
)
|
||||
return True
|
||||
except EndpointConnectionError as e:
|
||||
# Quiet the noisy traceback by logging a concise warning instead
|
||||
logging.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS endpoint is unreachable at {SQS_ENDPOINT_URL}: {e}. "
|
||||
logger.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS endpoint is unreachable at {self.cfg.worker.sqs_endpoint_url}: {e}. "
|
||||
)
|
||||
return False
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Failed to initialize SQS queues: {e}"
|
||||
)
|
||||
return False
|
||||
@@ -99,7 +87,7 @@ class WorkerConsumer:
|
||||
response = await client.get_queue_url(QueueName=queue_name)
|
||||
return response["QueueUrl"]
|
||||
except client.exceptions.QueueDoesNotExist:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS queue '{queue_name}' does not exist."
|
||||
)
|
||||
raise
|
||||
@@ -107,130 +95,128 @@ class WorkerConsumer:
|
||||
async def listen_for_comfy_events(self):
|
||||
while True:
|
||||
try:
|
||||
async with websockets.connect(COMFYUI_WS_URL) as websocket:
|
||||
logging.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Connected to ComfyUI websocket at {COMFYUI_WS_URL}"
|
||||
# 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."
|
||||
)
|
||||
while True:
|
||||
message = await websocket.recv()
|
||||
if isinstance(message, str):
|
||||
try:
|
||||
event = json.loads(message)
|
||||
event_type = event.get("type")
|
||||
data = event.get("data", {})
|
||||
prompt_id = data.get("prompt_id")
|
||||
await asyncio.sleep(5)
|
||||
continue
|
||||
|
||||
if not prompt_id and "sid" in data:
|
||||
prompt_id = data["sid"]
|
||||
# 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:
|
||||
continue
|
||||
if not prompt_id and "sid" in data:
|
||||
prompt_id = data["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"
|
||||
)
|
||||
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)
|
||||
# Allow 'executing' events through even if missing prompt_id
|
||||
if not prompt_id and event_type != "executing":
|
||||
continue
|
||||
|
||||
# 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)
|
||||
# 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}"
|
||||
)
|
||||
|
||||
# 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)
|
||||
# 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
|
||||
|
||||
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}"
|
||||
)
|
||||
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)
|
||||
|
||||
except json.JSONDecodeError:
|
||||
logging.debug(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Received non-JSON text message from websocket, ignoring."
|
||||
# Handle execution errors
|
||||
elif event_type == "execution_error":
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Received execution error for prompt_id {prompt_id}: {data}"
|
||||
)
|
||||
if prompt_id in self.prompt_id_to_content_id_map:
|
||||
content_id = self.prompt_id_to_content_id_map.pop(prompt_id)
|
||||
ctx = self.content_context_by_content_id.get(content_id, {})
|
||||
policy = ctx.get("status_policy") or {}
|
||||
fail_status = policy.get("fail_status", "failed")
|
||||
try:
|
||||
await self._send_status_update(
|
||||
content_id,
|
||||
fail_status,
|
||||
ctx.get("venue"),
|
||||
ctx.get("canvas"),
|
||||
ctx.get("scene"),
|
||||
ctx.get("job_type"),
|
||||
)
|
||||
else:
|
||||
logging.debug(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Received binary 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",
|
||||
)
|
||||
|
||||
# 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 (
|
||||
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}",
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Websocket listener error: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
await asyncio.sleep(10)
|
||||
await asyncio.sleep(5)
|
||||
|
||||
async def consume_loop(self):
|
||||
"""The main loop to continuously poll for and process messages.
|
||||
@@ -241,6 +227,7 @@ 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:
|
||||
@@ -251,30 +238,48 @@ class WorkerConsumer:
|
||||
)
|
||||
await asyncio.sleep(10)
|
||||
continue
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Starting worker consumer. Polling queue: {self.jobs_queue_url}"
|
||||
)
|
||||
|
||||
logging.debug(
|
||||
# Capacity gate: avoid pulling a new job while local ComfyUI is busy
|
||||
if self.is_busy:
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Skipping poll; worker is busy executing a job."
|
||||
)
|
||||
await asyncio.sleep(5)
|
||||
continue
|
||||
|
||||
# Run memory hygiene between jobs (no-op if disabled/unsupported)
|
||||
try:
|
||||
now = time.monotonic()
|
||||
min_interval = max(0, int(self.cfg.hygiene.idle_poll_seconds))
|
||||
if now - self._last_hygiene_check_ts >= min_interval:
|
||||
await self._run_memory_hygiene(debounce_seconds=0)
|
||||
self._last_hygiene_check_ts = now
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Polling for messages..."
|
||||
)
|
||||
try:
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
region_name=self.cfg.worker.aws_region,
|
||||
endpoint_url=self.cfg.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=self.cfg.worker.aws_access_key_id,
|
||||
aws_secret_access_key=self.cfg.worker.aws_secret_access_key,
|
||||
) as client:
|
||||
response = await client.receive_message(
|
||||
QueueUrl=self.jobs_queue_url,
|
||||
MaxNumberOfMessages=MAX_MESSAGES,
|
||||
WaitTimeSeconds=POLL_WAIT_TIME_SECONDS,
|
||||
MaxNumberOfMessages=self.cfg.worker.max_messages,
|
||||
WaitTimeSeconds=self.cfg.worker.poll_wait_s,
|
||||
)
|
||||
|
||||
messages = response.get("Messages", [])
|
||||
if not messages:
|
||||
logging.debug(
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): No messages received."
|
||||
)
|
||||
continue
|
||||
@@ -286,52 +291,80 @@ class WorkerConsumer:
|
||||
# On successful processing, delete the message
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
region_name=self.cfg.worker.aws_region,
|
||||
endpoint_url=self.cfg.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=self.cfg.worker.aws_access_key_id,
|
||||
aws_secret_access_key=self.cfg.worker.aws_secret_access_key,
|
||||
) as client:
|
||||
await client.delete_message(
|
||||
QueueUrl=self.jobs_queue_url,
|
||||
ReceiptHandle=message["ReceiptHandle"],
|
||||
)
|
||||
logging.info(
|
||||
logger.debug(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Deleted message {message['MessageId']} from queue."
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
# This is a poison pill message, log it but don't retry.
|
||||
# It will be moved to the DLQ after enough failed receives.
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Message {message['MessageId']} is a poison pill (JSON decode failed) and will be ignored."
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Processing failed for message {message['MessageId']}: {e}. It will be returned to the queue for retry."
|
||||
)
|
||||
except ClientError as e:
|
||||
# Some SQS providers/endpoints may sporadically return 503 for ReceiveMessage when the queue is idle
|
||||
error_code = None
|
||||
try:
|
||||
error_code = e.response.get("Error", {}).get("Code")
|
||||
except Exception:
|
||||
pass
|
||||
operation_name = getattr(e, "operation_name", "")
|
||||
if operation_name == "ReceiveMessage" and str(error_code) in (
|
||||
"503",
|
||||
"ServiceUnavailable",
|
||||
):
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Queue is empty or endpoint timed out (ReceiveMessage 503). Polling again shortly..."
|
||||
)
|
||||
# await asyncio.sleep(2)
|
||||
continue
|
||||
|
||||
# Unhandled ClientError; fall back to generic handling
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): SQS client error during ReceiveMessage: {e}"
|
||||
)
|
||||
await asyncio.sleep(10)
|
||||
except EndpointConnectionError as e:
|
||||
# Lost connection to SQS; reset and re-initialize on next loop
|
||||
logging.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Lost connection to SQS at {SQS_ENDPOINT_URL}: {e}. Will retry initialization in 10 seconds."
|
||||
logger.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Lost connection to SQS at {self.cfg.worker.sqs_endpoint_url}: {e}. Will retry initialization in 10 seconds."
|
||||
)
|
||||
self.jobs_queue_url = None
|
||||
self.status_updates_queue_url = None
|
||||
await asyncio.sleep(10)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): An error occurred in the consume loop: {e}"
|
||||
)
|
||||
await asyncio.sleep(10) # Wait before retrying
|
||||
finally:
|
||||
listener_task.cancel()
|
||||
await asyncio.gather(listener_task, return_exceptions=True)
|
||||
logging.info(
|
||||
# Close shared HTTP session if created
|
||||
if self.http_session is not None:
|
||||
try:
|
||||
await self.http_session.close()
|
||||
except Exception:
|
||||
pass
|
||||
logger.info(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Websocket listener stopped."
|
||||
)
|
||||
|
||||
async def process_message(self, message):
|
||||
"""Processes a single SQS message."""
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"ℹ️\u2009 Nilor-Nodes (worker_consumer): Processing message: {message['MessageId']}"
|
||||
)
|
||||
|
||||
@@ -345,9 +378,20 @@ 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:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Invalid message format: missing 'content_id' or 'prompt'. Payload: {job_payload}"
|
||||
)
|
||||
return
|
||||
@@ -361,13 +405,14 @@ 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:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): An unexpected error occurred while processing message: {e}. It will be retried."
|
||||
)
|
||||
# Re-raise to prevent deletion from queue if we want SQS to handle retry
|
||||
@@ -376,35 +421,81 @@ class WorkerConsumer:
|
||||
async def _submit_job_to_comfyui(self, content_id, workflow_data):
|
||||
"""Submits a single job to the ComfyUI API."""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
COMFYUI_API_URL, json=workflow_data, timeout=30
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
response_json = await response.json()
|
||||
prompt_id = response_json.get("prompt_id")
|
||||
logging.info(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Successfully submitted job to ComfyUI. Prompt ID: {prompt_id}"
|
||||
)
|
||||
self.prompt_id_to_content_id_map[prompt_id] = content_id
|
||||
# Attach/override websocket client_id so server targets events to this worker
|
||||
payload = (
|
||||
dict(workflow_data)
|
||||
if isinstance(workflow_data, dict)
|
||||
else workflow_data
|
||||
)
|
||||
|
||||
if isinstance(payload, dict):
|
||||
# Normalize OS-sensitive path formatting inside the ComfyUI prompt graph.
|
||||
# This lets us accept workflows authored on a different OS (e.g. Windows
|
||||
# backslashes) and run them on the current worker OS.
|
||||
enabled = self.cfg.worker.workflow_os_normalization_enabled
|
||||
if enabled and "prompt" in payload:
|
||||
try:
|
||||
normalized_prompt, rewritten = (
|
||||
normalize_comfyui_prompt_for_current_os(payload["prompt"])
|
||||
)
|
||||
if rewritten:
|
||||
payload["prompt"] = normalized_prompt
|
||||
logger.debug(
|
||||
"ℹ️\u2009 Nilor-Nodes (worker_consumer): Normalized %s path-like values in ComfyUI prompt for os=%s.",
|
||||
rewritten,
|
||||
os.name,
|
||||
)
|
||||
except Exception as e:
|
||||
# Don't fail the job if normalization fails; submit as-is.
|
||||
logger.warning(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Workflow normalization failed; submitting original prompt. Error: %s",
|
||||
e,
|
||||
)
|
||||
|
||||
# Force top-level client_id to this worker's stable ID
|
||||
payload["client_id"] = self.worker_client_id
|
||||
# Ensure extra_data exists and force its client_id too
|
||||
extra = payload.get("extra_data") or {}
|
||||
if isinstance(extra, dict):
|
||||
extra["client_id"] = self.worker_client_id
|
||||
payload["extra_data"] = extra
|
||||
|
||||
if self.comfy_client is None:
|
||||
# Construct client using shared session when available; fallback to creating a temp one
|
||||
self.comfy_client = ComfyUILocalClient(
|
||||
base_url=self.cfg.comfy.api_url,
|
||||
ws_url=self.cfg.comfy.ws_url,
|
||||
session=self.http_session,
|
||||
logger=logger,
|
||||
timeout=float(self.cfg.comfy.timeout_s),
|
||||
)
|
||||
|
||||
prompt_id = await self.comfy_client.submit_prompt(payload)
|
||||
logger.debug(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Successfully submitted job to ComfyUI. Prompt ID: {prompt_id}"
|
||||
)
|
||||
self.prompt_id_to_content_id_map[prompt_id] = content_id
|
||||
# Mark worker busy after successful submission to avoid over-queuing on this machine
|
||||
self.is_busy = True
|
||||
self.current_prompt_id = prompt_id
|
||||
|
||||
# No need to delete here, the consume_loop handles message deletion
|
||||
except aiohttp.ClientError as e:
|
||||
logging.error(
|
||||
except ComfyUIClientError as e:
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to submit job to ComfyUI: {e}. Message will be retried."
|
||||
)
|
||||
except (json.JSONDecodeError, KeyError) as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to parse ComfyUI response: {e}. Discarding malformed response."
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): An unexpected error occurred while submitting job to ComfyUI: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
async def _send_status_update(
|
||||
self, content_id, status, venue=None, canvas=None, scene=None
|
||||
self, content_id, status, venue=None, canvas=None, scene=None, job_type=None
|
||||
):
|
||||
try:
|
||||
body = {"content_id": content_id, "status": status}
|
||||
@@ -414,28 +505,169 @@ 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=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
region_name=self.cfg.worker.aws_region,
|
||||
endpoint_url=self.cfg.worker.sqs_endpoint_url,
|
||||
aws_access_key_id=self.cfg.worker.aws_access_key_id,
|
||||
aws_secret_access_key=self.cfg.worker.aws_secret_access_key,
|
||||
) as client:
|
||||
await client.send_message(
|
||||
QueueUrl=self.status_updates_queue_url, MessageBody=message_body
|
||||
)
|
||||
logging.info(
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Sent status update for content {content_id}: {status}"
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(
|
||||
logger.error(
|
||||
f"🛑\u2009 Nilor-Nodes (worker_consumer): Failed to send status update for content {content_id}: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
def _finalize_prompt(self, prompt_id, reason: str | None = None):
|
||||
# No-op if we're not currently busy; avoids redundant work on duplicate signals
|
||||
if not self.is_busy:
|
||||
return
|
||||
try:
|
||||
if reason:
|
||||
logger.info(
|
||||
f"✅ Nilor-Nodes (worker_consumer): Finalizing prompt via '{reason}'. Using prompt_id={prompt_id}"
|
||||
)
|
||||
content_id = self.prompt_id_to_content_id_map.pop(prompt_id, None)
|
||||
if content_id is not None:
|
||||
self.content_context_by_content_id.pop(content_id, None)
|
||||
self.sent_running_status_prompts.discard(prompt_id)
|
||||
finally:
|
||||
# Only clear busy/state if this finalize corresponds to the current in-flight prompt
|
||||
if self.current_prompt_id == prompt_id:
|
||||
self.current_prompt_id = None
|
||||
self.is_busy = False
|
||||
# Debounced hygiene after job completion
|
||||
try:
|
||||
asyncio.create_task(self._run_memory_hygiene(debounce_seconds=1))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def _normalize_node_id(node_id):
|
||||
if node_id is None:
|
||||
return None
|
||||
if isinstance(node_id, str) and node_id.strip() in ("", "None"):
|
||||
return None
|
||||
return node_id
|
||||
|
||||
async def _run_memory_hygiene(self, debounce_seconds: int = 0) -> None:
|
||||
"""Run memory hygiene with optional debounce, guarded by busy state.
|
||||
|
||||
Delegates to the shared hygiene component and ensures we do not
|
||||
pull new work while remediation is running.
|
||||
"""
|
||||
await _run_hygiene_guarded(self, debounce_seconds)
|
||||
|
||||
|
||||
async def consume_jobs():
|
||||
"""Entry point function to be called in a background thread."""
|
||||
consumer = WorkerConsumer()
|
||||
# Create shared HTTP session once and reuse throughout lifecycle
|
||||
consumer = WorkerConsumer(cfg=_CFG)
|
||||
# Initialize shared HTTP session at startup
|
||||
consumer.http_session = aiohttp.ClientSession()
|
||||
|
||||
# Construct ComfyUI client using shared session (WS is always client-driven)
|
||||
consumer.comfy_client = ComfyUILocalClient(
|
||||
base_url=_CFG.comfy.api_url,
|
||||
ws_url=_CFG.comfy.ws_url,
|
||||
session=consumer.http_session,
|
||||
logger=logger,
|
||||
timeout=float(_CFG.comfy.timeout_s),
|
||||
)
|
||||
|
||||
# Initialize Memory Hygiene with the constructed client and loaded config
|
||||
try:
|
||||
consumer.hygiene = MemoryHygiene(
|
||||
client=consumer.comfy_client,
|
||||
cfg=_CFG.hygiene,
|
||||
logger=logger,
|
||||
)
|
||||
except Exception:
|
||||
consumer.hygiene = None
|
||||
|
||||
# Emit concise startup configuration summary (no secrets)
|
||||
try:
|
||||
logger.info(
|
||||
(
|
||||
"ℹ️\u2009 Nilor-Nodes: startup config — Comfy API=%s, WS=%s, timeout_s=%s, "
|
||||
"SQS endpoint=%s, jobs_queue=%s, status_queue=%s"
|
||||
),
|
||||
_CFG.comfy.api_url,
|
||||
_CFG.comfy.ws_url,
|
||||
str(_CFG.comfy.timeout_s),
|
||||
_CFG.worker.sqs_endpoint_url,
|
||||
_CFG.worker.jobs_queue,
|
||||
_CFG.worker.status_queue,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Emit concise hygiene summary (effective values)
|
||||
try:
|
||||
h = _CFG.hygiene
|
||||
logger.info(
|
||||
(
|
||||
"ℹ️\u2009 Nilor-Nodes: startup hygiene — enabled=%s, idle_poll_s=%s, "
|
||||
"vram_used_pct_max=%s, ram_used_pct_max=%s, vram_min_free_mb=%s, ram_min_free_mb=%s, "
|
||||
"policy=%s, max_retries=%s, cooldown_s=%s, sleep_between_s=%s, max_cycle_s=%s"
|
||||
),
|
||||
str(h.enabled),
|
||||
str(h.idle_poll_seconds),
|
||||
str(h.vram_usage_pct_max),
|
||||
str(h.ram_usage_pct_max),
|
||||
str(h.vram_min_free_mb),
|
||||
str(h.ram_min_free_mb),
|
||||
h.action_policy,
|
||||
str(h.max_retries),
|
||||
str(h.cooldown_seconds),
|
||||
str(h.sleep_between_attempts_seconds),
|
||||
str(h.max_cycle_duration_seconds),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
await consumer.consume_loop()
|
||||
|
||||
|
||||
async def _sleep_seconds(seconds: int) -> None:
|
||||
try:
|
||||
await asyncio.sleep(max(0, int(seconds)))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def _maybe_bool(value) -> bool:
|
||||
try:
|
||||
return bool(value)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
async def _run_hygiene_guarded(
|
||||
self_ref: "WorkerConsumer", debounce_seconds: int
|
||||
) -> None:
|
||||
# No-op if not available
|
||||
if not getattr(self_ref, "hygiene", None):
|
||||
return
|
||||
# Optional debounce
|
||||
if debounce_seconds > 0:
|
||||
await _sleep_seconds(debounce_seconds)
|
||||
# Set maintenance busy gate
|
||||
if self_ref.is_busy:
|
||||
return
|
||||
self_ref.is_busy = True
|
||||
try:
|
||||
await self_ref.hygiene.check_and_remediate()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
self_ref.is_busy = False
|
||||
|
||||
@@ -0,0 +1,173 @@
|
||||
"""
|
||||
Workflow normalization helpers.
|
||||
|
||||
Goal: accept ComfyUI "prompt" (API workflow graph) authored on a different OS and
|
||||
rewrite OS-sensitive path formatting (primarily path separators) so that the
|
||||
local ComfyUI instance can resolve model filenames correctly.
|
||||
|
||||
This is intentionally conservative: we only rewrite strings that look like file
|
||||
paths / model names (e.g. end in ".safetensors") and we avoid touching URLs and
|
||||
free-form prompt text.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Iterable, List, Tuple
|
||||
|
||||
__all__ = [
|
||||
"PathRemap",
|
||||
"normalize_comfyui_prompt_for_current_os",
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PathRemap:
|
||||
"""Prefix remap for absolute paths across OSes.
|
||||
|
||||
Example:
|
||||
PathRemap(from_prefix="D:\\ComfyUI\\models", to_prefix="/mnt/models")
|
||||
|
||||
Matching is performed in a canonicalized form (both prefixes and candidate
|
||||
paths have backslashes converted to forward slashes). After remapping, path
|
||||
separators are normalized for the current OS.
|
||||
"""
|
||||
|
||||
from_prefix: str
|
||||
to_prefix: str
|
||||
|
||||
|
||||
_URL_RE = re.compile(r"^[a-zA-Z][a-zA-Z0-9+.-]*://")
|
||||
_WIN_DRIVE_RE = re.compile(r"^[A-Za-z]:[\\/]")
|
||||
|
||||
# Common file extensions encountered in ComfyUI prompts (models + media).
|
||||
_PATH_EXTS = {
|
||||
".safetensors",
|
||||
".pt",
|
||||
".pth",
|
||||
".ckpt",
|
||||
".bin",
|
||||
".onnx",
|
||||
".json",
|
||||
".json5",
|
||||
".yaml",
|
||||
".yml",
|
||||
".txt",
|
||||
".png",
|
||||
".jpg",
|
||||
".jpeg",
|
||||
".webp",
|
||||
".gif",
|
||||
".bmp",
|
||||
".tif",
|
||||
".tiff",
|
||||
".exr",
|
||||
".mp4",
|
||||
".mov",
|
||||
".mkv",
|
||||
".webm",
|
||||
".wav",
|
||||
".mp3",
|
||||
".flac",
|
||||
}
|
||||
|
||||
|
||||
def _canonicalize_for_prefix_match(path: str) -> str:
|
||||
# Use forward slashes for prefix matching regardless of host OS.
|
||||
return path.replace("\\", "/")
|
||||
|
||||
|
||||
def _normalize_separators_for_current_os(path: str) -> str:
|
||||
# ComfyUI often uses OS-native separators in its model name registry.
|
||||
# We normalize to the current OS so lookup keys match.
|
||||
if os.name == "nt":
|
||||
return path.replace("/", "\\")
|
||||
return path.replace("\\", "/")
|
||||
|
||||
|
||||
def _looks_like_path_value(s: str) -> bool:
|
||||
s_stripped = s.strip()
|
||||
if not s_stripped:
|
||||
return False
|
||||
|
||||
# Don't touch URLs (presigned uploads, http inputs, etc.)
|
||||
if _URL_RE.match(s_stripped):
|
||||
return False
|
||||
|
||||
# Don't touch placeholder tokens used by the system.
|
||||
if s_stripped.startswith("<") and s_stripped.endswith(">"):
|
||||
return False
|
||||
|
||||
# Avoid common sentinel.
|
||||
if s_stripped == "None":
|
||||
return False
|
||||
|
||||
lower = s_stripped.lower()
|
||||
if any(lower.endswith(ext) for ext in _PATH_EXTS):
|
||||
return True
|
||||
|
||||
# Absolute Windows paths even without extensions.
|
||||
if _WIN_DRIVE_RE.match(s_stripped):
|
||||
return True
|
||||
|
||||
# Relative paths with explicit prefixes.
|
||||
if s_stripped.startswith(("./", "../", "~/", "~\\")):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def _apply_prefix_remaps(path: str, remaps: Iterable[PathRemap]) -> str:
|
||||
if not remaps:
|
||||
return path
|
||||
|
||||
cand = _canonicalize_for_prefix_match(path)
|
||||
for r in remaps:
|
||||
frm = _canonicalize_for_prefix_match(str(r.from_prefix))
|
||||
if cand.startswith(frm):
|
||||
to = str(r.to_prefix)
|
||||
# Replace using canonical representation then return in that form;
|
||||
# the caller will normalize separators for current OS afterwards.
|
||||
replaced = to + cand[len(frm) :]
|
||||
return replaced
|
||||
return path
|
||||
|
||||
|
||||
def normalize_comfyui_prompt_for_current_os(
|
||||
prompt: Any, *, path_remaps: Iterable[PathRemap] | None = None
|
||||
) -> Tuple[Any, int]:
|
||||
"""Normalize a ComfyUI API `prompt` graph for the current OS.
|
||||
|
||||
Args:
|
||||
prompt: The value of the `/prompt` payload's `prompt` key. Typically a
|
||||
dict mapping node ids to `{class_type, inputs, ...}`.
|
||||
path_remaps: Optional prefix remaps applied before separator
|
||||
normalization (useful for absolute paths).
|
||||
|
||||
Returns:
|
||||
(normalized_prompt, num_rewritten_strings)
|
||||
"""
|
||||
|
||||
remaps: List[PathRemap] = list(path_remaps or [])
|
||||
rewritten = 0
|
||||
|
||||
def walk(x: Any) -> Any:
|
||||
nonlocal rewritten
|
||||
if isinstance(x, dict):
|
||||
return {k: walk(v) for k, v in x.items()}
|
||||
if isinstance(x, list):
|
||||
return [walk(v) for v in x]
|
||||
if isinstance(x, tuple):
|
||||
return tuple(walk(v) for v in x)
|
||||
if isinstance(x, str) and _looks_like_path_value(x):
|
||||
y = _apply_prefix_remaps(x, remaps)
|
||||
y = _normalize_separators_for_current_os(y)
|
||||
if y != x:
|
||||
rewritten += 1
|
||||
return y
|
||||
return x
|
||||
|
||||
return walk(prompt), rewritten
|
||||
|
||||
Reference in New Issue
Block a user