feat(nilor-nodes): integrate typed config loader into worker_consumer
- load Config once with env>JSON5 precedence and pass cfg into WorkerConsumer - replace direct os.getenv reads with cfg.comfy and cfg.worker values - compute Comfy HTTP/WS endpoints from cfg and reuse single aiohttp session - keep behavior identical; no new features introduced
This commit is contained in:
+117
-137
@@ -21,6 +21,16 @@ 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.
|
||||
@@ -63,165 +73,112 @@ class WorkerConfig:
|
||||
worker_client_id: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class NilorNodesConfig:
|
||||
@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
|
||||
|
||||
@classmethod
|
||||
def _get_config_path(cls) -> str:
|
||||
return os.path.join(os.path.dirname(__file__), "config.json5")
|
||||
|
||||
class Config:
|
||||
"""Public configuration loader API for the Nilor-Nodes sidecar.
|
||||
@classmethod
|
||||
def from_dict(cls, config_dict: Dict[str, object]) -> "NilorNodesConfig":
|
||||
allow_env_override = bool(config_dict.get("allow_env_override", True))
|
||||
|
||||
The implementation is intentionally deferred to a subsequent commit. This
|
||||
placeholder establishes the method signature and return type to facilitate
|
||||
incremental refactors in consumers.
|
||||
"""
|
||||
# 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)),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def load(
|
||||
env: os._Environ[str] = os.environ, json5_path: Optional[str] = None
|
||||
) -> NilorNodesConfig:
|
||||
"""Load and return a typed configuration instance.
|
||||
worker_client_id = (
|
||||
str(config_dict.get("NILOR_WORKER_CLIENT_ID", "")).strip()
|
||||
or _generate_worker_client_id()
|
||||
)
|
||||
|
||||
Args:
|
||||
env: Mapping of environment variables used for overrides.
|
||||
json5_path: Optional path to a JSON5 file containing non-secret defaults.
|
||||
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,
|
||||
)
|
||||
|
||||
Returns:
|
||||
NilorNodesConfig: The fully parsed and validated configuration object.
|
||||
"""
|
||||
|
||||
# Read JSON5 file into overrides if provided
|
||||
file_overrides: Dict[str, object] = {}
|
||||
if json5_path:
|
||||
if json5 is None:
|
||||
raise RuntimeError(
|
||||
"json5 module is required to load configuration from JSON5 file"
|
||||
)
|
||||
if os.path.exists(json5_path):
|
||||
try:
|
||||
with open(json5_path, "r", encoding="utf-8") as f:
|
||||
loaded = json5.load(f)
|
||||
if not isinstance(loaded, dict):
|
||||
raise ValueError(
|
||||
"JSON5 configuration must be a top-level object"
|
||||
)
|
||||
file_overrides = {str(k): v for k, v in loaded.items()}
|
||||
except Exception as exc:
|
||||
raise RuntimeError(
|
||||
f"Failed to read JSON5 configuration from '{json5_path}': {type(exc).__name__}: {exc}"
|
||||
)
|
||||
|
||||
# Normalize env into NILOR_* keys with legacy compatibility
|
||||
env_values = _normalize_env_to_nilor(env)
|
||||
|
||||
# Precedence: env > file
|
||||
effective: Dict[str, object] = dict(file_overrides)
|
||||
effective.update(env_values)
|
||||
|
||||
# Parse
|
||||
comfy = _parse_comfy_config(effective)
|
||||
worker = _parse_worker_config(effective)
|
||||
|
||||
# Validate
|
||||
_validate_comfy_config(comfy)
|
||||
_validate_worker_config(worker)
|
||||
|
||||
return NilorNodesConfig(comfy=comfy, worker=worker)
|
||||
cfg = cls(
|
||||
comfy=comfy_cfg, worker=worker_cfg, allow_env_override=allow_env_override
|
||||
)
|
||||
_validate_comfy_config(cfg.comfy)
|
||||
_validate_worker_config(cfg.worker)
|
||||
return cfg
|
||||
|
||||
|
||||
# ---- Internal helpers ----
|
||||
|
||||
|
||||
def _normalize_env_to_nilor(env: os._Environ[str]) -> Dict[str, object]:
|
||||
normalized: Dict[str, object] = {}
|
||||
for key, value in env.items():
|
||||
if key.startswith("NILOR_"):
|
||||
normalized[key] = value
|
||||
return normalized
|
||||
|
||||
|
||||
def _parse_comfy_config(values: Dict[str, object]) -> ComfyApiConfig:
|
||||
api_url = str(values.get("NILOR_COMFYUI_API_URL", "")).strip()
|
||||
ws_url = str(values.get("NILOR_COMFYUI_WS_URL", "")).strip()
|
||||
timeout_raw = values.get("NILOR_COMFY_API_TIMEOUT_SECONDS", 30)
|
||||
try:
|
||||
timeout_s = int(timeout_raw)
|
||||
except Exception:
|
||||
raise ValueError(
|
||||
f"Invalid integer for NILOR_COMFY_API_TIMEOUT_SECONDS: {timeout_raw!r}"
|
||||
)
|
||||
if not api_url:
|
||||
raise ValueError("Missing NILOR_COMFYUI_API_URL (env or JSON5)")
|
||||
if not ws_url:
|
||||
raise ValueError("Missing NILOR_COMFYUI_WS_URL (env or JSON5)")
|
||||
return ComfyApiConfig(api_url=api_url, ws_url=ws_url, timeout_s=timeout_s)
|
||||
|
||||
|
||||
def _parse_worker_config(values: Dict[str, object]) -> WorkerConfig:
|
||||
sqs_endpoint_url = str(values.get("NILOR_SQS_ENDPOINT_URL", "")).strip()
|
||||
jobs_queue = str(values.get("NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME", "")).strip()
|
||||
status_queue = str(
|
||||
values.get("NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME", "")
|
||||
).strip()
|
||||
|
||||
poll_wait_raw = values.get("NILOR_SQS_POLL_WAIT_TIME", 10)
|
||||
max_messages_raw = values.get("NILOR_SQS_MAX_MESSAGES", 1)
|
||||
try:
|
||||
poll_wait_s = int(poll_wait_raw)
|
||||
except Exception:
|
||||
raise ValueError(
|
||||
f"Invalid integer for NILOR_SQS_POLL_WAIT_TIME: {poll_wait_raw!r}"
|
||||
)
|
||||
try:
|
||||
max_messages = int(max_messages_raw)
|
||||
except Exception:
|
||||
raise ValueError(
|
||||
f"Invalid integer for NILOR_SQS_MAX_MESSAGES: {max_messages_raw!r}"
|
||||
)
|
||||
|
||||
aws_access_key_id = str(values.get("NILOR_AWS_ACCESS_KEY_ID", "")).strip()
|
||||
aws_secret_access_key = str(values.get("NILOR_AWS_SECRET_ACCESS_KEY", "")).strip()
|
||||
aws_region = str(values.get("NILOR_AWS_DEFAULT_REGION", "")).strip()
|
||||
|
||||
worker_client_id = str(values.get("NILOR_WORKER_CLIENT_ID", "")).strip()
|
||||
if not worker_client_id:
|
||||
worker_client_id = _generate_worker_client_id()
|
||||
|
||||
if not sqs_endpoint_url:
|
||||
raise ValueError("Missing NILOR_SQS_ENDPOINT_URL (env or JSON5)")
|
||||
if not jobs_queue:
|
||||
raise ValueError("Missing NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME (env or JSON5)")
|
||||
if not status_queue:
|
||||
raise ValueError(
|
||||
"Missing NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME (env or JSON5)"
|
||||
)
|
||||
if not aws_access_key_id:
|
||||
raise ValueError("Missing NILOR_AWS_ACCESS_KEY_ID (set via env)")
|
||||
if not aws_secret_access_key:
|
||||
raise ValueError("Missing NILOR_AWS_SECRET_ACCESS_KEY (set via env)")
|
||||
if not aws_region:
|
||||
raise ValueError("Missing NILOR_AWS_DEFAULT_REGION (env or JSON5)")
|
||||
|
||||
return WorkerConfig(
|
||||
sqs_endpoint_url=sqs_endpoint_url,
|
||||
jobs_queue=jobs_queue,
|
||||
status_queue=status_queue,
|
||||
poll_wait_s=poll_wait_s,
|
||||
max_messages=max_messages,
|
||||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
aws_region=aws_region,
|
||||
worker_client_id=worker_client_id,
|
||||
def _apply_env_overrides(cfg: NilorNodesConfig) -> None:
|
||||
if not cfg.allow_env_override:
|
||||
return
|
||||
# Comfy
|
||||
cfg.comfy.api_url = os.getenv("NILOR_COMFYUI_API_URL", cfg.comfy.api_url)
|
||||
cfg.comfy.ws_url = os.getenv("NILOR_COMFYUI_WS_URL", cfg.comfy.ws_url)
|
||||
cfg.comfy.timeout_s = int(
|
||||
os.getenv("NILOR_COMFY_API_TIMEOUT_SECONDS", cfg.comfy.timeout_s)
|
||||
)
|
||||
|
||||
# Worker
|
||||
cfg.worker.sqs_endpoint_url = os.getenv(
|
||||
"NILOR_SQS_ENDPOINT_URL", cfg.worker.sqs_endpoint_url
|
||||
)
|
||||
cfg.worker.jobs_queue = os.getenv(
|
||||
"NILOR_SQS_JOBS_TO_PROCESS_QUEUE_NAME", cfg.worker.jobs_queue
|
||||
)
|
||||
cfg.worker.status_queue = os.getenv(
|
||||
"NILOR_SQS_JOB_STATUS_UPDATES_QUEUE_NAME", cfg.worker.status_queue
|
||||
)
|
||||
cfg.worker.poll_wait_s = int(
|
||||
os.getenv("NILOR_SQS_POLL_WAIT_TIME", cfg.worker.poll_wait_s)
|
||||
)
|
||||
cfg.worker.max_messages = int(
|
||||
os.getenv("NILOR_SQS_MAX_MESSAGES", cfg.worker.max_messages)
|
||||
)
|
||||
cfg.worker.aws_access_key_id = os.getenv(
|
||||
"NILOR_AWS_ACCESS_KEY_ID", cfg.worker.aws_access_key_id
|
||||
)
|
||||
cfg.worker.aws_secret_access_key = os.getenv(
|
||||
"NILOR_AWS_SECRET_ACCESS_KEY", cfg.worker.aws_secret_access_key
|
||||
)
|
||||
cfg.worker.aws_region = os.getenv("NILOR_AWS_DEFAULT_REGION", cfg.worker.aws_region)
|
||||
cfg.worker.worker_client_id = os.getenv(
|
||||
"NILOR_WORKER_CLIENT_ID", cfg.worker.worker_client_id
|
||||
)
|
||||
|
||||
# Re-validate after overrides
|
||||
_validate_comfy_config(cfg.comfy)
|
||||
_validate_worker_config(cfg.worker)
|
||||
|
||||
|
||||
def _validate_comfy_config(cfg: ComfyApiConfig) -> None:
|
||||
_require_url_scheme(cfg.api_url, {"http", "https"}, "NILOR_COMFYUI_API_URL")
|
||||
@@ -277,3 +234,26 @@ def _to_base36(n: int) -> str:
|
||||
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`
|
||||
"""
|
||||
try:
|
||||
from dotenv import load_dotenv # optional dependency present in sidecar
|
||||
|
||||
load_dotenv()
|
||||
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)
|
||||
return cfg
|
||||
|
||||
+59
-59
@@ -20,6 +20,7 @@ from aiobotocore.session import get_session
|
||||
from dotenv import load_dotenv
|
||||
from botocore.exceptions import EndpointConnectionError, ClientError
|
||||
from .logger import logger
|
||||
from .config.config import load_nilor_nodes_config, NilorNodesConfig
|
||||
|
||||
# --- Load Environment Variables ---
|
||||
# Load from the .env file in the same directory
|
||||
@@ -36,24 +37,12 @@ else:
|
||||
)
|
||||
|
||||
# --- 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-comfyui"
|
||||
)
|
||||
SQS_JOB_STATUS_UPDATES_QUEUE_NAME = os.getenv(
|
||||
"SQS_JOB_STATUS_UPDATES_QUEUE_NAME", "job_status_updates"
|
||||
)
|
||||
SQS_POLL_WAIT_TIME = int(os.getenv("SQS_POLL_WAIT_TIME", "10")) # SQS Long Polling
|
||||
SQS_MAX_MESSAGES = 1
|
||||
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")
|
||||
# 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()
|
||||
@@ -65,35 +54,35 @@ class WorkerConsumer:
|
||||
self.websocket_sid = None
|
||||
|
||||
# Stable client_id for routing events to this worker
|
||||
env_id = os.getenv("NILOR_WORKER_CLIENT_ID")
|
||||
if env_id and env_id.strip():
|
||||
self.worker_client_id = env_id.strip()
|
||||
else:
|
||||
host = socket.gethostname()
|
||||
self.worker_client_id = f"nilor-worker-{host}"
|
||||
self.cfg = cfg
|
||||
self.worker_client_id = cfg.worker.worker_client_id
|
||||
|
||||
# Computed endpoints
|
||||
self.comfy_api_url = cfg.comfy.api_url.rstrip("/") + "/prompt"
|
||||
self.comfy_ws_url = cfg.comfy.ws_url.rstrip("/") + "/ws"
|
||||
self.current_prompt_id = None
|
||||
|
||||
async def _initialize_sqs(self):
|
||||
"""Initializes SQS queue URLs. Returns True on success, False on failure."""
|
||||
async with self.session.create_client(
|
||||
"sqs",
|
||||
region_name=AWS_DEFAULT_REGION,
|
||||
endpoint_url=SQS_ENDPOINT_URL,
|
||||
aws_access_key_id=AWS_ACCESS_KEY_ID,
|
||||
aws_secret_access_key=AWS_SECRET_ACCESS_KEY,
|
||||
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
|
||||
logger.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS endpoint is unreachable at {SQS_ENDPOINT_URL}: {e}. "
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): SQS endpoint is unreachable at {self.cfg.worker.sqs_endpoint_url}: {e}. "
|
||||
)
|
||||
return False
|
||||
except Exception as e:
|
||||
@@ -116,7 +105,7 @@ class WorkerConsumer:
|
||||
async def listen_for_comfy_events(self):
|
||||
while True:
|
||||
try:
|
||||
ws_url = f"{COMFYUI_WS_URL}?clientId={urllib.parse.quote(self.worker_client_id)}"
|
||||
ws_url = f"{self.comfy_ws_url}?clientId={urllib.parse.quote(self.worker_client_id)}"
|
||||
# Allow large preview frames from ComfyUI without dropping the connection (1009: message too big)
|
||||
async with websockets.connect(
|
||||
ws_url,
|
||||
@@ -323,15 +312,15 @@ class WorkerConsumer:
|
||||
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=SQS_MAX_MESSAGES,
|
||||
WaitTimeSeconds=SQS_POLL_WAIT_TIME,
|
||||
MaxNumberOfMessages=self.cfg.worker.max_messages,
|
||||
WaitTimeSeconds=self.cfg.worker.poll_wait_s,
|
||||
)
|
||||
|
||||
messages = response.get("Messages", [])
|
||||
@@ -348,10 +337,10 @@ 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,
|
||||
@@ -396,7 +385,7 @@ class WorkerConsumer:
|
||||
except EndpointConnectionError as e:
|
||||
# Lost connection to SQS; reset and re-initialize on next loop
|
||||
logger.warning(
|
||||
f"⚠️\u2009 Nilor-Nodes (worker_consumer): Lost connection to SQS at {SQS_ENDPOINT_URL}: {e}. Will retry initialization in 10 seconds."
|
||||
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
|
||||
@@ -409,6 +398,12 @@ class WorkerConsumer:
|
||||
finally:
|
||||
listener_task.cancel()
|
||||
await asyncio.gather(listener_task, return_exceptions=True)
|
||||
# Close shared HTTP session if created
|
||||
if self.http_session is not None:
|
||||
try:
|
||||
await self.http_session.close()
|
||||
except Exception:
|
||||
pass
|
||||
logger.info(
|
||||
"⚠️\u2009 Nilor-Nodes (worker_consumer): Websocket listener stopped."
|
||||
)
|
||||
@@ -487,20 +482,22 @@ class WorkerConsumer:
|
||||
extra["client_id"] = self.worker_client_id
|
||||
payload["extra_data"] = extra
|
||||
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.post(
|
||||
COMFYUI_API_URL, json=payload, timeout=30
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
response_json = await response.json()
|
||||
prompt_id = response_json.get("prompt_id")
|
||||
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
|
||||
if self.http_session is None:
|
||||
# Lazily create if not created by the loop yet
|
||||
self.http_session = aiohttp.ClientSession()
|
||||
async with self.http_session.post(
|
||||
self.comfy_api_url, json=payload, timeout=self.cfg.comfy.timeout_s
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
response_json = await response.json()
|
||||
prompt_id = response_json.get("prompt_id")
|
||||
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:
|
||||
@@ -533,10 +530,10 @@ class WorkerConsumer:
|
||||
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
|
||||
@@ -580,5 +577,8 @@ class WorkerConsumer:
|
||||
|
||||
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()
|
||||
await consumer.consume_loop()
|
||||
|
||||
Reference in New Issue
Block a user