diff --git a/invokeai/app/api/routers/remote_workers.py b/invokeai/app/api/routers/remote_workers.py new file mode 100644 index 00000000000..109f51700d3 --- /dev/null +++ b/invokeai/app/api/routers/remote_workers.py @@ -0,0 +1,213 @@ +"""Authenticated per-user remote-worker credential management. + +Only a saved/not-saved indicator and email are returned; no passwords or JWTs. +""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any, Literal + +from fastapi import HTTPException, Query +from fastapi.routing import APIRouter +from pydantic import BaseModel, ConfigDict, Field + +from invokeai.app.api.auth_dependencies import AdminUserOrDefault, CurrentUserOrDefault +from invokeai.app.api.dependencies import ApiDependencies +from invokeai.app.invocations.remote_worker.credential_vault import ( + delete_credentials, + get_saved_credentials, + get_saved_settings, + normalize_url, + save_credentials, + save_settings, +) +from invokeai.app.invocations.remote_worker.diffusers_transfer import ( + cancel_directory_install_job, + get_directory_install_job, + start_directory_install, +) +from invokeai.app.invocations.remote_worker.model_transfer import model_layout_signature +from invokeai.app.invocations.remote_worker.remote_client import RemoteConfig, RemoteInvokeClient, RemoteInvokeError + +remote_workers_router = APIRouter(prefix="/v1/remote_workers", tags=["remote_workers"]) + + +class RemoteWorkerCredentialRequest(BaseModel): + url: str = Field(min_length=1, max_length=2048) + email: str = Field(min_length=1, max_length=320) + password: str = Field(min_length=1, max_length=4096) + remember_me: bool = True + + +class RemoteWorkerCredentialStatus(BaseModel): + saved: bool + email: str | None = None + + +def _status(user_id: str, url: str) -> RemoteWorkerCredentialStatus: + record = get_saved_credentials(user_id, url) + return RemoteWorkerCredentialStatus( + saved=record is not None, + email=str(record["email"]) if record and isinstance(record.get("email"), str) else None, + ) + + +class RemoteWorkerAvailability(BaseModel): + status: Literal["online", "offline", "login_required"] + + +class RemoteWorkersSettings(BaseModel): + model_config = ConfigDict(populate_by_name=True) + + enabled: bool = False + dispatch_mode: Literal["distributed", "remote_only"] = Field(default="distributed", alias="dispatchMode") + worker_urls: str = Field(default="", alias="workerUrls", max_length=32768) + worker_names: dict[str, str] = Field(default_factory=dict, alias="workerNames") + disabled_worker_urls: list[str] = Field(default_factory=list, alias="disabledWorkerUrls") + auto_transfer_missing_models: bool = Field(default=True, alias="autoTransferMissingModels") + keep_remote_copies: bool = Field(default=False, alias="keepRemoteCopies") + model_transfer_host: str = Field(default="", alias="modelTransferHost", max_length=2048) + + +@remote_workers_router.get("/settings", response_model=RemoteWorkersSettings) +def get_remote_worker_settings(current_user: CurrentUserOrDefault) -> RemoteWorkersSettings: + saved = get_saved_settings(current_user.user_id) + return RemoteWorkersSettings.model_validate(saved) if saved is not None else RemoteWorkersSettings() + + +@remote_workers_router.put("/settings", response_model=RemoteWorkersSettings) +def put_remote_worker_settings( + current_user: CurrentUserOrDefault, + body: RemoteWorkersSettings, +) -> RemoteWorkersSettings: + save_settings(current_user.user_id, body.model_dump(by_alias=True)) + return body + + +@remote_workers_router.get("/status", response_model=RemoteWorkerAvailability) +def get_remote_worker_status( + current_user: CurrentUserOrDefault, + url: str = Query(min_length=1, max_length=2048), +) -> RemoteWorkerAvailability: + """Probe this user's ability to reach a configured worker using existing InvokeAI APIs. + + Keep the probe short; credentials stay in the primary's per-user vault. + Do not mistake a reachable worker with rejected credentials for an offline host. + """ + try: + normalized = normalize_url(url) + except ValueError as exc: + raise HTTPException(status_code=422, detail=str(exc)) from exc + client = RemoteInvokeClient( + RemoteConfig.from_environment(base_url=normalized, verify_ssl=False, user_id=current_user.user_id), + request_timeout_seconds=2.5, + ) + try: + client.get_current_item() + return RemoteWorkerAvailability(status="online") + except RemoteInvokeError as exc: + message = str(exc).lower() + if any( + marker in message + for marker in ( + "requires login", + "login failed", + "initial admin setup", + "http 401", + "http 403", + "credentials file", + ) + ): + return RemoteWorkerAvailability(status="login_required") + return RemoteWorkerAvailability(status="offline") + except Exception: + # An unreachable worker should never make the primary's status API fail. + return RemoteWorkerAvailability(status="offline") + + +@remote_workers_router.get("/credentials", response_model=RemoteWorkerCredentialStatus) +def get_remote_worker_credentials_status( + current_user: CurrentUserOrDefault, + url: str = Query(min_length=1, max_length=2048), +) -> RemoteWorkerCredentialStatus: + try: + return _status(current_user.user_id, normalize_url(url)) + except ValueError as exc: + raise HTTPException(status_code=422, detail=str(exc)) from exc + + +@remote_workers_router.put("/credentials", response_model=RemoteWorkerCredentialStatus) +def put_remote_worker_credentials( + current_user: CurrentUserOrDefault, + body: RemoteWorkerCredentialRequest, +) -> RemoteWorkerCredentialStatus: + try: + save_credentials(current_user.user_id, body.url, body.email, body.password, body.remember_me) + return _status(current_user.user_id, body.url) + except ValueError as exc: + raise HTTPException(status_code=422, detail=str(exc)) from exc + + +@remote_workers_router.delete("/credentials", response_model=RemoteWorkerCredentialStatus) +def remove_remote_worker_credentials( + current_user: CurrentUserOrDefault, + url: str = Query(min_length=1, max_length=2048), +) -> RemoteWorkerCredentialStatus: + try: + delete_credentials(current_user.user_id, url) + return RemoteWorkerCredentialStatus(saved=False) + except ValueError as exc: + raise HTTPException(status_code=422, detail=str(exc)) from exc + + +class RemoteModelLayout(BaseModel): + kind: Literal["file", "directory"] + signature: str + + +@remote_workers_router.get("/models/{key}/layout", response_model=RemoteModelLayout) +def get_remote_model_layout(current_user: CurrentUserOrDefault, key: str) -> RemoteModelLayout: + """Return a non-secret signature of one registered model's file layout.""" + services = ApiDependencies.invoker.services + try: + config = services.model_manager.store.get_model(key) + except Exception as exc: + raise HTTPException(status_code=404, detail="Model not found") from exc + + model_path = Path(str(getattr(config, "path", "") or "")) + if not model_path.is_absolute(): + model_path = Path(services.configuration.models_path) / model_path + model_path = model_path.resolve() + if not model_path.exists(): + raise HTTPException(status_code=404, detail="Model files not found") + + kind, signature = model_layout_signature(model_path) + return RemoteModelLayout(kind=kind, signature=signature) + + +@remote_workers_router.post("/diffusers/install") +def install_remote_directory(current_admin: AdminUserOrDefault, body: dict[str, Any]) -> dict[str, Any]: + """Admin-only receiver for short-lived, manifest-verified model transfers.""" + try: + return start_directory_install(body, ApiDependencies.invoker.services) + except ValueError as exc: + raise HTTPException(status_code=422, detail=str(exc)) from exc + + +@remote_workers_router.get("/diffusers/install/{job_id}") +def get_remote_directory_install(current_admin: AdminUserOrDefault, job_id: int) -> dict[str, Any]: + """Poll a directory download and normal model installation.""" + job = get_directory_install_job(job_id) + if job is None: + raise HTTPException(status_code=404, detail="Directory transfer job not found") + return job + + +@remote_workers_router.delete("/diffusers/install/{job_id}") +def cancel_remote_directory_install(current_admin: AdminUserOrDefault, job_id: int) -> dict[str, Any]: + """Cancel only this temporary download, preserving completed models.""" + job = cancel_directory_install_job(job_id) + if job is None: + raise HTTPException(status_code=404, detail="Directory transfer job not found") + return job diff --git a/invokeai/app/api/routers/session_queue.py b/invokeai/app/api/routers/session_queue.py index 6ded39f5ede..e814bb716a8 100644 --- a/invokeai/app/api/routers/session_queue.py +++ b/invokeai/app/api/routers/session_queue.py @@ -10,6 +10,7 @@ from invokeai.app.api.dependencies import ApiDependencies from invokeai.app.api.routers.image_move_maintenance import assert_image_move_maintenance_inactive from invokeai.app.invocations.fields import ImageField, VideoField +from invokeai.app.invocations.remote_worker.early_dispatch import schedule_automatic_remote_dispatches from invokeai.app.services.progress_previews.progress_previews_common import ProgressPreviewDTO from invokeai.app.services.session_processor.session_processor_common import SessionProcessorStatus from invokeai.app.services.session_queue.session_queue_common import ( @@ -256,9 +257,23 @@ async def enqueue_batch( await asyncio.to_thread(assert_image_move_maintenance_inactive) try: - return await ApiDependencies.invoker.services.session_queue.enqueue_batch( + result = await ApiDependencies.invoker.services.session_queue.enqueue_batch( queue_id=queue_id, batch=batch, prepend=prepend, user_id=current_user.user_id ) + # The Remote Worker hook is a no-op for normal batches. It only + # schedules the CPU/network fast lane; remote work still runs off-thread. + try: + schedule_automatic_remote_dispatches( + batch=batch, + item_ids=result.item_ids, + services=ApiDependencies.invoker.services, + ) + except Exception as exc: + # No enqueue failure: the normal queued invocation remains the fallback. + ApiDependencies.invoker.services.logger.warning( + f"IRW early remote dispatch unavailable; using normal queue: {exc}" + ) + return result except EnqueueIdempotencyConflictError as e: raise HTTPException(status_code=409, detail=str(e)) except EnqueueProjectNotFoundError as e: diff --git a/invokeai/app/api_app.py b/invokeai/app/api_app.py index 89513a41040..eb2519c0441 100644 --- a/invokeai/app/api_app.py +++ b/invokeai/app/api_app.py @@ -39,6 +39,7 @@ model_relationships, projects, recall_parameters, + remote_workers, session_queue, style_presets, system_prompts, @@ -753,6 +754,7 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: app.include_router(model_relationships.model_relationships_router, prefix="/api") app.include_router(app_info.app_router, prefix="/api") app.include_router(session_queue.session_queue_router, prefix="/api") +app.include_router(remote_workers.remote_workers_router, prefix="/api") app.include_router(workflows.workflows_router, prefix="/api") app.include_router(style_presets.style_presets_router, prefix="/api") app.include_router(wildcards.wildcards_router, prefix="/api") diff --git a/invokeai/app/invocations/remote_worker/__init__.py b/invokeai/app/invocations/remote_worker/__init__.py new file mode 100644 index 00000000000..49c06b63f6b --- /dev/null +++ b/invokeai/app/invocations/remote_worker/__init__.py @@ -0,0 +1,5 @@ +"""Built-in InvokeAI Remote Worker support. + +The invocation is discovered by InvokeAI core module discovery; installing a +standalone custom-node pack is not required. +""" diff --git a/invokeai/app/invocations/remote_worker/credential_vault.py b/invokeai/app/invocations/remote_worker/credential_vault.py new file mode 100644 index 00000000000..aec2e8a081b --- /dev/null +++ b/invokeai/app/invocations/remote_worker/credential_vault.py @@ -0,0 +1,199 @@ +"""Cross-platform, per-user remote-worker credential vault. + +AES-256-GCM protects credentials at rest. The randomly generated key is stored +separately from the encrypted vault under InvokeAI's configured runtime root. +Protect the entire runtime directory and preserve both files when moving an +installation: anyone with both files can decrypt the credentials. + +Older Windows ``credentials.dpapi`` files are not read or migrated. Re-enter +saved worker logins after upgrading from the DPAPI-only implementation. +""" + +from __future__ import annotations + +import base64 +import json +import os +import threading +import uuid +from pathlib import Path +from urllib.parse import urlsplit + +from cryptography.exceptions import InvalidTag +from cryptography.hazmat.primitives.ciphers.aead import AESGCM + +_LOCK = threading.RLock() +_AAD = b"invokeai-remote-workers-credentials-v1" +_NONCE_BYTES = 12 +_KEY_BYTES = 32 +_SETTINGS_KEY = "__settings__" + + +def normalize_url(value: str) -> str: + url = value.strip().rstrip("/") + parsed = urlsplit(url) + if ( + parsed.scheme not in ("http", "https") + or not parsed.hostname + or parsed.username + or parsed.password + or parsed.query + or parsed.fragment + or not parsed.netloc + ): + raise ValueError( + "Enter a worker URL beginning with http:// or https:// (without login details or query strings)" + ) + try: + _ = parsed.port + except ValueError as exc: + raise ValueError("Invalid port in worker URL") from exc + return url + + +def _vault_path() -> Path: + # Import lazily: remote_client is discovered during core invocation import. + from invokeai.app.services.config.config_default import get_config + + return get_config().root_path / "remote_workers" / "credentials.enc" + + +def _key_path() -> Path: + return _vault_path().with_name("credentials.key") + + +def _ensure_directory(path: Path) -> None: + path.mkdir(mode=0o700, parents=True, exist_ok=True) + if os.name == "posix": + path.chmod(0o700) + + +def _read_key(*, create: bool) -> bytes: + key_path = _key_path() + if create: + _ensure_directory(key_path.parent) + # Exclusive creation: never rotate a key just because a file is damaged. + # In particular, an existing ciphertext must never be paired with a new key. + try: + fd = os.open(key_path, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600) + except FileExistsError: + pass + else: + try: + with os.fdopen(fd, "wb") as key_file: + key_file.write(AESGCM.generate_key(bit_length=256)) + key_file.flush() + os.fsync(key_file.fileno()) + except BaseException: + key_path.unlink(missing_ok=True) + raise + try: + key = key_path.read_bytes() + except FileNotFoundError as exc: + raise RuntimeError( + "Remote Workers encryption key is missing; restore credentials.key with credentials.enc" + ) from exc + if len(key) != _KEY_BYTES: + raise ValueError("Remote Workers encryption key is invalid; restore the original credentials.key") + return key + + +def _read() -> dict[str, dict[str, dict[str, object]]]: + path = _vault_path() + if not path.exists(): + return {} + # A missing/corrupt key is an error, not an empty vault or an excuse to + # replace the vault. This also prevents accidental credential loss. + key = _read_key(create=False) + envelope = json.loads(path.read_text(encoding="utf-8")) + if not isinstance(envelope, dict) or envelope.get("version") != 1: + raise ValueError("Remote Workers credentials file has an unsupported format") + try: + nonce = base64.b64decode(envelope["nonce"], validate=True) + ciphertext = base64.b64decode(envelope["ciphertext"], validate=True) + if len(nonce) != _NONCE_BYTES: + raise ValueError("Invalid AES-GCM nonce") + plaintext = AESGCM(key).decrypt(nonce, ciphertext, _AAD) + except (InvalidTag, KeyError, TypeError, ValueError) as exc: + raise ValueError( + "Remote Workers credentials could not be authenticated; check the key and vault files" + ) from exc + data = json.loads(plaintext) + if not isinstance(data, dict) or data.get("version") != 1 or not isinstance(data.get("users"), dict): + raise ValueError("Remote Workers credentials file has an unsupported format") + return data["users"] + + +def _write(users: dict[str, dict[str, dict[str, object]]]) -> None: + path = _vault_path() + _ensure_directory(path.parent) + key = _read_key(create=True) + plaintext = json.dumps({"version": 1, "users": users}, separators=(",", ":")).encode("utf-8") + nonce = os.urandom(_NONCE_BYTES) + ciphertext = AESGCM(key).encrypt(nonce, plaintext, _AAD) + envelope = json.dumps( + { + "version": 1, + "nonce": base64.b64encode(nonce).decode("ascii"), + "ciphertext": base64.b64encode(ciphertext).decode("ascii"), + }, + separators=(",", ":"), + ).encode("utf-8") + temporary = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp") + fd = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600) + try: + with os.fdopen(fd, "wb") as output: + output.write(envelope) + output.flush() + os.fsync(output.fileno()) + os.replace(temporary, path) + finally: + temporary.unlink(missing_ok=True) + + +def get_saved_credentials(user_id: str, url: str) -> dict[str, object] | None: + with _LOCK: + entry = _read().get(user_id, {}).get(normalize_url(url)) + return dict(entry) if isinstance(entry, dict) else None + + +def get_saved_settings(user_id: str) -> dict[str, object] | None: + """Return this InvokeAI user's encrypted remote-worker settings, if saved.""" + with _LOCK: + entry = _read().get(user_id, {}).get(_SETTINGS_KEY) + return dict(entry) if isinstance(entry, dict) else None + + +def save_settings(user_id: str, settings: dict[str, object]) -> None: + """Persist this InvokeAI user's remote-worker settings in the encrypted vault.""" + if not user_id: + raise ValueError("A user is required") + with _LOCK: + users = _read() + users.setdefault(user_id, {})[_SETTINGS_KEY] = dict(settings) + _write(users) + + +def save_credentials(user_id: str, url: str, email: str, password: str, remember_me: bool = True) -> None: + if not user_id or not email.strip() or not password: + raise ValueError("A user, email, and password are required") + normalized = normalize_url(url) + with _LOCK: + users = _read() + users.setdefault(user_id, {})[normalized] = { + "email": email.strip(), + "password": password, + "remember_me": remember_me, + } + _write(users) + + +def delete_credentials(user_id: str, url: str) -> None: + normalized = normalize_url(url) + with _LOCK: + users = _read() + if normalized in users.get(user_id, {}): + del users[user_id][normalized] + if not users[user_id]: + del users[user_id] + _write(users) diff --git a/invokeai/app/invocations/remote_worker/diffusers_transfer.py b/invokeai/app/invocations/remote_worker/diffusers_transfer.py new file mode 100644 index 00000000000..5f70bfcd5a7 --- /dev/null +++ b/invokeai/app/invocations/remote_worker/diffusers_transfer.py @@ -0,0 +1,463 @@ +"""Authenticated, manifest-based LAN transfer of local model directories. + +The primary serves files from a short-lived, random-token HTTP endpoint. The +remote InvokeAI process stages and verifies every file before asking its normal +model installer to register the directory. No separate worker plugin is needed. +""" + +from __future__ import annotations + +import hashlib +import mimetypes +import secrets +import tempfile +import threading +import time +import urllib.error +import urllib.parse +import urllib.request +from collections.abc import Callable +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path, PurePosixPath +from typing import Any + +from invokeai.app.invocations.remote_worker.model_transfer import ModelTransferError, _route_local_ip + +_MAX_FILES = 8192 +_MAX_BYTES = 2 * 1024**4 +_CHUNK = 1024 * 1024 +_JOB_LOCK = threading.Lock() +_JOBS: dict[int, dict[str, Any]] = {} +_CANCEL_EVENTS: dict[int, threading.Event] = {} +_ACTIVE_HASHES: dict[str, int] = {} +_NEXT_JOB_ID = 0 + + +def _sha256_file(path: Path, should_cancel: Callable[[], bool] | None = None) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + while True: + if should_cancel is not None and should_cancel(): + raise ModelTransferError("Model directory preparation cancelled") + chunk = handle.read(_CHUNK) + if not chunk: + break + digest.update(chunk) + return digest.hexdigest() + + +def _safe_path(text: str) -> PurePosixPath: + if not text or len(text) > 1024 or "\\" in text or ":" in text or "\x00" in text: + raise ValueError("Invalid model file path") + if any(ord(char) < 32 for char in text): + raise ValueError("Control character in model file path") + path = PurePosixPath(text) + if ( + path.is_absolute() + or any(part in {"", ".", ".."} or part.endswith((" ", ".")) for part in text.split("/")) + or path.as_posix() != text + ): + raise ValueError("Unsafe model file path") + return path + + +@dataclass(frozen=True) +class DirectoryFile: + path: str + size: int + sha256: str + + +def inventory_directory( + root: Path, + should_cancel: Callable[[], bool] | None = None, +) -> list[DirectoryFile]: + """Inventory only regular files, refusing symlinks and case-fold collisions.""" + if root.is_symlink() or not root.is_dir(): + raise ModelTransferError("Source must be a real model directory") + children: list[Path] = [] + for child in root.rglob("*"): + if should_cancel is not None and should_cancel(): + raise ModelTransferError("Model directory preparation cancelled") + children.append(child) + children.sort() + + files: list[DirectoryFile] = [] + total = 0 + names: set[str] = set() + for child in children: + if should_cancel is not None and should_cancel(): + raise ModelTransferError("Model directory preparation cancelled") + if child.is_symlink(): + raise ModelTransferError(f"Model contains a symlink: {child}") + if child.is_dir(): + continue + if not child.is_file(): + raise ModelTransferError(f"Model contains a non-file entry: {child}") + relative = child.relative_to(root).as_posix() + try: + _safe_path(relative) + except ValueError as exc: + raise ModelTransferError(f"Unsafe model path: {relative}") from exc + folded = relative.casefold() + if folded in names: + raise ModelTransferError(f"Case-insensitive duplicate model file: {relative}") + names.add(folded) + size = child.stat().st_size + total += size + if len(files) >= _MAX_FILES or total > _MAX_BYTES: + raise ModelTransferError("Model directory exceeds transfer limits") + files.append(DirectoryFile(relative, size, _sha256_file(child, should_cancel))) + if not files: + raise ModelTransferError("Model directory has no files") + return files + + +# Match the filename/configuration signals accepted by InvokeAI's model probe. +# This is an early, deliberately conservative diagnostic, not a replacement for +# the model installer: a file passing this check may still be invalid or incomplete. +_MODEL_WEIGHT_SUFFIXES = frozenset({".bin", ".ckpt", ".gguf", ".onnx", ".pt", ".pth", ".safetensors"}) +_MODEL_CONFIG_NAMES = frozenset({"config.json", "model_index.json", "modular_model_index.json"}) + + +def _validate_transferable_model_files(files: list[DirectoryFile]) -> None: + """Reject tokenizer-only/empty-weight directories before serving them over LAN.""" + if any( + file.size > 0 + and ( + PurePosixPath(file.path).suffix.lower() in _MODEL_WEIGHT_SUFFIXES + or PurePosixPath(file.path).name.lower() in _MODEL_CONFIG_NAMES + ) + for file in files + ): + return + raise ModelTransferError( + "Local model directory appears incomplete: no non-empty model weights or supported model " + "configuration files were found, including in subdirectories. Repair or reinstall the model " + "on the primary InvokeAI instance before retrying remote transfer." + ) + + +class TemporaryDirectoryModelServer: + """Expose only inventoried files at unguessable URLs for a transfer's lifetime.""" + + def __init__( + self, + path: Path, + remote_url: str, + advertise_host: str = "", + should_cancel: Callable[[], bool] | None = None, + ) -> None: + self.path = path + self.remote_url = remote_url + self.advertise_host = advertise_host.strip() + self.should_cancel = should_cancel + self.files: list[DirectoryFile] = [] + self.url = "" + self._token = secrets.token_urlsafe(32) + self._server: ThreadingHTTPServer | None = None + self._thread: threading.Thread | None = None + + def __enter__(self) -> "TemporaryDirectoryModelServer": + self.files = inventory_directory(self.path, self.should_cancel) + _validate_transferable_model_files(self.files) + file_paths = { + f"/{self._token}/{urllib.parse.quote(item.path, safe='/')}": self.path / item.path for item in self.files + } + + class Handler(BaseHTTPRequestHandler): + server_version = "InvokeAIRemoteWorkerDirectoryTransfer/1.0" + + def log_message(self, _format: str, *args: Any) -> None: + return + + def _headers(self) -> tuple[Path, int, int] | None: + file_path = file_paths.get(urllib.parse.urlsplit(self.path).path) + if file_path is None or not file_path.is_file() or file_path.is_symlink(): + self.send_error(404) + return None + size = file_path.stat().st_size + if size == 0: + if self.headers.get("Range"): + self.send_error(416) + return None + start, end, code = 0, -1, 200 + else: + start, end, code = 0, size - 1, 200 + header = self.headers.get("Range", "") + if header: + if not header.startswith("bytes=") or "," in header: + self.send_error(416) + return None + begin, sep, finish = header[6:].partition("-") + if not sep or not begin.isdigit() or (finish and not finish.isdigit()): + self.send_error(416) + return None + start = int(begin) + end = int(finish) if finish else size - 1 + if start >= size or end < start: + self.send_error(416) + return None + end = min(end, size - 1) + code = 206 + self.send_response(code) + self.send_header("Content-Type", mimetypes.guess_type(file_path.name)[0] or "application/octet-stream") + self.send_header("Content-Length", str(end - start + 1)) + self.send_header("Accept-Ranges", "bytes") + self.send_header("Cache-Control", "no-store") + if code == 206: + self.send_header("Content-Range", f"bytes {start}-{end}/{size}") + self.end_headers() + return file_path, start, end + + def do_HEAD(self) -> None: # noqa: N802 + self._headers() + + def do_GET(self) -> None: # noqa: N802 + selected = self._headers() + if selected is None: + return + file_path, start, end = selected + with file_path.open("rb") as source: + source.seek(start) + remaining = end - start + 1 + while remaining: + chunk = source.read(min(_CHUNK, remaining)) + if not chunk: + break + try: + self.wfile.write(chunk) + except (BrokenPipeError, ConnectionResetError): + break + remaining -= len(chunk) + + server = ThreadingHTTPServer(("0.0.0.0", 0), Handler) + server.daemon_threads = True + self._server = server + self._thread = threading.Thread(target=server.serve_forever, daemon=True, name="invokeai-directory-transfer") + self._thread.start() + host = self.advertise_host or _route_local_ip(self.remote_url) + if "://" in host: + host = urllib.parse.urlsplit(host).hostname or host + host = host.strip("[]") + self.url = f"http://{host}:{server.server_port}/{self._token}" + return self + + def manifest(self, *, name: str, model_hash: str) -> dict[str, Any]: + return { + "base_url": self.url, + "name": name, + "model_hash": model_hash, + "files": [{"path": file.path, "size": file.size, "sha256": file.sha256} for file in self.files], + } + + def __exit__(self, _exc_type: Any, _exc: Any, _traceback: Any) -> None: + if self._server is not None: + self._server.shutdown() + self._server.server_close() + if self._thread is not None: + self._thread.join(timeout=2.0) + self._server = None + self._thread = None + + +def _validated_manifest(manifest: dict[str, Any]) -> tuple[str, str, str, list[DirectoryFile], int]: + base_url = manifest.get("base_url") + name = manifest.get("name") + model_hash = manifest.get("model_hash") + raw_files = manifest.get("files") + if not isinstance(base_url, str) or len(base_url) > 2048: + raise ValueError("Invalid model transfer URL") + parsed = urllib.parse.urlsplit(base_url) + if parsed.scheme not in {"http", "https"} or not parsed.hostname or parsed.username or parsed.password: + raise ValueError("Invalid model transfer host") + if parsed.query or parsed.fragment or not parsed.path.strip("/"): + raise ValueError("Invalid model transfer URL path") + if not isinstance(name, str) or not name.strip() or len(name) > 256: + raise ValueError("Invalid model name") + if not isinstance(model_hash, str) or not model_hash or len(model_hash) > 256: + raise ValueError("Invalid model hash") + if not isinstance(raw_files, list) or not 0 < len(raw_files) <= _MAX_FILES: + raise ValueError("Invalid model file count") + files = [] + names: set[str] = set() + total = 0 + for entry in raw_files: + if not isinstance(entry, dict): + raise ValueError("Invalid model file") + path = _safe_path(entry.get("path")) if isinstance(entry.get("path"), str) else None + size, digest = entry.get("size"), entry.get("sha256") + if path is None or type(size) is not int or size < 0 or size > _MAX_BYTES: + raise ValueError("Invalid model file size or path") + if not isinstance(digest, str) or len(digest) != 64 or any(c not in "0123456789abcdef" for c in digest): + raise ValueError("Invalid model file hash") + if path.as_posix().casefold() in names: + raise ValueError("Duplicate model path") + names.add(path.as_posix().casefold()) + total += size + if total > _MAX_BYTES: + raise ValueError("Model exceeds transfer size limit") + files.append(DirectoryFile(path.as_posix(), size, digest)) + return base_url.rstrip("/"), name.strip(), model_hash, files, total + + +def _set_job(job_id: int, **updates: Any) -> None: + with _JOB_LOCK: + _JOBS[job_id].update(updates) + + +def get_directory_install_job(job_id: int) -> dict[str, Any] | None: + with _JOB_LOCK: + current = _JOBS.get(job_id) + return dict(current) if current is not None else None + + +class DirectoryTransferCancelled(Exception): + """The owner no longer needs this temporary model download.""" + + +def _check_cancel(job_id: int) -> None: + if _CANCEL_EVENTS[job_id].is_set(): + raise DirectoryTransferCancelled() + + +def cancel_directory_install_job(job_id: int) -> dict[str, Any] | None: + """Request cooperative cancellation; the worker owns scratch cleanup.""" + with _JOB_LOCK: + job = _JOBS.get(job_id) + if job is None: + return None + if job["status"] in {"completed", "cancelled", "error"}: + return dict(job) + cancel_event = _CANCEL_EVENTS.get(job_id) + if cancel_event is None: + # Cancellation cleanup may have released the event just before the + # terminal status becomes visible. Never turn that race into a 500. + return dict(job) + cancel_event.set() + return dict(job) + + +def _download_file(base_url: str, file: DirectoryFile, dest: Path, job_id: int, completed_bytes: int) -> None: + url = f"{base_url}/{urllib.parse.quote(file.path, safe='/')}" + part = dest.with_name(dest.name + ".part") + dest.parent.mkdir(parents=True, exist_ok=True) + for attempt in range(3): + _check_cancel(job_id) + offset = part.stat().st_size if part.exists() else 0 + if offset > file.size: + part.unlink() + offset = 0 + headers = {"Range": f"bytes={offset}-"} if offset else {} + request = urllib.request.Request(url, headers=headers) + try: + with urllib.request.urlopen(request, timeout=60) as response: + if offset and response.status != 206: + part.unlink(missing_ok=True) + offset = 0 + if offset and response.headers.get("Content-Range", "").split("-", 1)[0] != f"bytes {offset}": + raise ValueError(f"Bad resumed response for {file.path}") + with part.open("ab" if offset else "wb") as target: + while chunk := response.read(_CHUNK): + _check_cancel(job_id) + target.write(chunk) + _set_job(job_id, bytes=completed_bytes + target.tell()) + if target.tell() > file.size: + raise ValueError(f"Oversized model file: {file.path}") + _check_cancel(job_id) + if part.stat().st_size != file.size or _sha256_file(part) != file.sha256: + part.unlink(missing_ok=True) + raise ValueError(f"Hash/size verification failed: {file.path}") + part.replace(dest) + return + except (OSError, ValueError, urllib.error.URLError) as exc: + if attempt == 2: + raise RuntimeError(f"Download failed for {file.path}: {exc}") from exc + time.sleep(0.5 * (attempt + 1)) + raise RuntimeError(f"Download failed for {file.path}") + + +def _run_install(job_id: int, manifest: tuple[str, str, str, list[DirectoryFile], int], services: Any) -> None: + base_url, name, model_hash, files, _total = manifest + was_cancelled = False + try: + _check_cancel(job_id) + _set_job(job_id, status="downloading") + models_path = Path(services.configuration.models_path) + with tempfile.TemporaryDirectory(prefix="tmpinstall_irw_", dir=models_path) as scratch: + model_path = Path(scratch) / "model" + model_path.mkdir() + completed_bytes = 0 + for file in files: + _check_cancel(job_id) + _download_file( + base_url, file, model_path.joinpath(*PurePosixPath(file.path).parts), job_id, completed_bytes + ) + completed_bytes += file.size + _set_job(job_id, bytes=completed_bytes) + _check_cancel(job_id) + _set_job(job_id, status="installing") + from invokeai.app.services.model_records.model_records_base import ModelRecordChanges + + installer = services.model_manager.install + install_job = installer.heuristic_import( + str(model_path), config=ModelRecordChanges(name=name), inplace=False + ) + # Do not interrupt the native installer mid-move: it must finish before + # the temporary directory can be safely removed. + result = installer.wait_for_job(install_job) + if str(getattr(result.status, "value", result.status)) != "completed": + raise RuntimeError(result.error or f"Model installer ended with status {result.status}") + config = result.config_out + installed_hash = str(getattr(config, "hash", "") or "") + if installed_hash != model_hash: + raise RuntimeError( + f"Installed model hash differs from primary: expected {model_hash}, got {installed_hash or 'none'}" + ) + _set_job(job_id, status="completed", model_key=str(getattr(config, "key", ""))) + except DirectoryTransferCancelled: + was_cancelled = True + except Exception as exc: + services.logger.error(f"Remote directory transfer job {job_id}: {exc}") + _set_job(job_id, status="error", error=str(exc)) + finally: + with _JOB_LOCK: + if was_cancelled: + # Publish the terminal state before releasing the cancel event/hash + # lease so another request can never observe a live status with no + # cancellation handle. + _JOBS[job_id].update(status="cancelled", error=None) + if _ACTIVE_HASHES.get(model_hash) == job_id: + _ACTIVE_HASHES.pop(model_hash, None) + _CANCEL_EVENTS.pop(job_id, None) + + +def start_directory_install(manifest: dict[str, Any], services: Any) -> dict[str, Any]: + """Schedule the complete transfer; never block an API request on a large download.""" + global _NEXT_JOB_ID + validated = _validated_manifest(manifest) + model_hash = validated[2] + with _JOB_LOCK: + active_id = _ACTIVE_HASHES.get(model_hash) + if active_id is not None: + if _CANCEL_EVENTS[active_id].is_set(): + raise ValueError("Previous transfer cancellation is still being cleaned up; retry shortly") + return dict(_JOBS[active_id]) + if len(_ACTIVE_HASHES) >= 8: + raise ValueError("Too many simultaneous directory transfers") + _NEXT_JOB_ID += 1 + job_id = _NEXT_JOB_ID + _JOBS[job_id] = {"id": job_id, "status": "waiting", "bytes": 0, "total_bytes": validated[4], "error": None} + _CANCEL_EVENTS[job_id] = threading.Event() + _ACTIVE_HASHES[model_hash] = job_id + if len(_JOBS) > 256: + for old_id in list(_JOBS): + if old_id not in _ACTIVE_HASHES.values() and old_id != job_id: + _JOBS.pop(old_id) + if len(_JOBS) <= 128: + break + threading.Thread( + target=_run_install, args=(job_id, validated, services), daemon=True, name=f"irw-directory-{job_id}" + ).start() + return get_directory_install_job(job_id) or {"id": job_id, "status": "waiting"} diff --git a/invokeai/app/invocations/remote_worker/early_dispatch.py b/invokeai/app/invocations/remote_worker/early_dispatch.py new file mode 100644 index 00000000000..49612b0244f --- /dev/null +++ b/invokeai/app/invocations/remote_worker/early_dispatch.py @@ -0,0 +1,19 @@ +from __future__ import annotations + +from typing import Any + +AUTOMATIC_REMOTE_WORKER_NODE_ID = "__irw_remote_worker_dispatch__" +AUTOMATIC_REMOTE_WORKER_NODE_TYPE = "irw_builtin_remote_worker_dispatch" + + +def schedule_automatic_remote_dispatches(*, batch: Any, item_ids: list[int], services: Any) -> bool: + """Start the backend worker pool for an automatic Remote Workers batch.""" + helper = batch.graph.nodes.get(AUTOMATIC_REMOTE_WORKER_NODE_ID) + if helper is None or helper.get_type() != AUTOMATIC_REMOTE_WORKER_NODE_TYPE: + return False + + from invokeai.app.invocations.remote_worker.worker_pool import schedule_remote_worker_pool + + for item_id in item_ids: + schedule_remote_worker_pool(int(item_id), services) + return True diff --git a/invokeai/app/invocations/remote_worker/model_transfer.py b/invokeai/app/invocations/remote_worker/model_transfer.py new file mode 100644 index 00000000000..eaf6b89e56c --- /dev/null +++ b/invokeai/app/invocations/remote_worker/model_transfer.py @@ -0,0 +1,282 @@ +from __future__ import annotations + +import hashlib +import mimetypes +import secrets +import socket +import threading +import urllib.parse +from dataclasses import dataclass +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from typing import Any + + +class ModelTransferError(RuntimeError): + """Raised when a local model cannot be exposed safely to a remote worker.""" + + +@dataclass(frozen=True) +class LocalModelFile: + path: Path + key: str + name: str + hash: str + base: str + type: str + + +def _enum_value(value: Any) -> str: + raw = getattr(value, "value", value) + return "" if raw is None else str(raw) + + +def model_layout_signature(path: Path) -> tuple[str, str]: + """Return a cheap, cross-platform signature of a model's on-disk layout. + + InvokeAI's model hash intentionally hashes model weight contents but not their + relative paths or tokenizer/config files. For directories, this signature hashes + the sorted relative path of every regular file so two installs with identical + weights but incompatible directory layouts do not compare as equivalent. + """ + if path.is_file(): + return "file", "" + if not path.is_dir(): + raise ModelTransferError(f"Model path is not a file or directory: {path}") + + digest = hashlib.sha256() + for child in sorted((entry for entry in path.rglob("*") if entry.is_file()), key=lambda entry: entry.as_posix()): + relative = child.relative_to(path).as_posix() + digest.update(relative.encode("utf-8")) + digest.update(b"\0") + return "directory", digest.hexdigest() + + +def resolve_local_model_file(services: Any, identifier: dict[str, Any]) -> LocalModelFile: + """Resolve a graph model identifier to a file or directory on this InvokeAI install.""" + key = str(identifier.get("key") or "").strip() + if not key: + raise ModelTransferError("Model identifier has no local model key") + + try: + config = services.model_manager.store.get_model(key) + except Exception as exc: + raise ModelTransferError(f"Could not read local model record {key}: {exc}") from exc + + config_path = getattr(config, "path", None) + if not config_path: + raise ModelTransferError( + f"Local model '{getattr(config, 'name', identifier.get('name', key))}' has no filesystem path" + ) + + model_path = Path(str(config_path)) + if not model_path.is_absolute(): + model_path = Path(services.configuration.models_path) / model_path + model_path = model_path.resolve() + + if not model_path.exists(): + raise ModelTransferError(f"Local model file does not exist: {model_path}") + if not model_path.is_file() and not model_path.is_dir(): + raise ModelTransferError(f"Local model path is not a file or directory: {model_path}") + + model_hash = str(identifier.get("hash") or getattr(config, "hash", "") or "").strip() + if not model_hash: + raise ModelTransferError( + f"Local model '{getattr(config, 'name', model_path.name)}' has no recorded model hash; " + "cannot verify a remote transfer safely" + ) + + return LocalModelFile( + path=model_path, + key=key, + name=str(getattr(config, "name", None) or identifier.get("name") or model_path.stem), + hash=model_hash, + base=_enum_value(getattr(config, "base", None) or identifier.get("base")), + type=_enum_value(getattr(config, "type", None) or identifier.get("type")), + ) + + +def enrich_model_identifier_hashes(graph: dict[str, Any], services: Any) -> int: + """Fill missing graph model hashes from this InvokeAI install's model records.""" + changed = 0 + + def visit(value: Any) -> None: + nonlocal changed + if isinstance(value, list): + for item in value: + visit(item) + return + if not isinstance(value, dict): + return + + if all(field in value for field in ("key", "name", "base", "type")): + key = str(value.get("key") or "") + if key and not value.get("hash"): + try: + config = services.model_manager.store.get_model(key) + model_hash = str(getattr(config, "hash", "") or "").strip() + except Exception: + model_hash = "" + if model_hash: + value["hash"] = model_hash + changed += 1 + + for item in value.values(): + visit(item) + + visit(graph.get("nodes")) + return changed + + +def _route_local_ip(remote_url: str) -> str: + parsed = urllib.parse.urlsplit(remote_url) + host = parsed.hostname + if not host: + raise ModelTransferError(f"Cannot determine remote host from URL: {remote_url}") + port = parsed.port or (443 if parsed.scheme == "https" else 80) + + try: + addresses = socket.getaddrinfo(host, port, socket.AF_INET, socket.SOCK_DGRAM) + except OSError as exc: + raise ModelTransferError(f"Could not resolve remote host '{host}': {exc}") from exc + if not addresses: + raise ModelTransferError(f"Could not resolve an IPv4 route to remote host '{host}'") + + target = addresses[0][4] + sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + try: + sock.connect(target) + local_ip = str(sock.getsockname()[0]) + except OSError as exc: + raise ModelTransferError(f"Could not determine local LAN address for {remote_url}: {exc}") from exc + finally: + sock.close() + + if not local_ip or local_ip.startswith("127.") or local_ip == "0.0.0.0": + raise ModelTransferError( + f"Automatic LAN address detection returned '{local_ip}'. Set Model Transfer Host manually on the node." + ) + return local_ip + + +class TemporaryModelServer: + """Serve exactly one model file on a random-token URL for the life of this context.""" + + def __init__(self, model: LocalModelFile, remote_url: str, advertise_host: str = "") -> None: + self.model = model + self.remote_url = remote_url + self.advertise_host = advertise_host.strip() + self._server: ThreadingHTTPServer | None = None + self._thread: threading.Thread | None = None + token = secrets.token_urlsafe(32) + self._request_path = f"/{token}/{urllib.parse.quote(model.path.name, safe='')}" + self.url = "" + + def __enter__(self) -> "TemporaryModelServer": + file_path = self.model.path + request_path = self._request_path + content_type = mimetypes.guess_type(file_path.name)[0] or "application/octet-stream" + + class Handler(BaseHTTPRequestHandler): + server_version = "InvokeAIRemoteWorkerModelTransfer/0.10" + + def log_message(self, _format: str, *args: Any) -> None: + return + + def _match(self) -> bool: + return urllib.parse.urlsplit(self.path).path == request_path + + def _range(self, size: int) -> tuple[int, int] | None: + header = self.headers.get("Range", "").strip() + if not header: + return None + if not header.startswith("bytes=") or "," in header: + return None + value = header[6:] + start_text, _, end_text = value.partition("-") + try: + if start_text: + start = int(start_text) + end = int(end_text) if end_text else size - 1 + else: + suffix = int(end_text) + start = max(0, size - suffix) + end = size - 1 + except ValueError: + return None + if start < 0 or start >= size or end < start: + return None + return start, min(end, size - 1) + + def _send_headers(self, *, body: bool) -> tuple[int, int] | None: + if not self._match(): + self.send_error(404) + return None + size = file_path.stat().st_size + byte_range = self._range(size) + if self.headers.get("Range") and byte_range is None: + self.send_response(416) + self.send_header("Content-Range", f"bytes */{size}") + self.end_headers() + return None + if byte_range is None: + start, end = 0, size - 1 + self.send_response(200) + else: + start, end = byte_range + self.send_response(206) + self.send_header("Content-Range", f"bytes {start}-{end}/{size}") + self.send_header("Content-Type", content_type) + self.send_header("Content-Length", str(max(0, end - start + 1))) + self.send_header("Accept-Ranges", "bytes") + self.send_header("Cache-Control", "no-store") + self.send_header("Content-Disposition", f'attachment; filename="{file_path.name}"') + self.end_headers() + return start, end + + def do_HEAD(self) -> None: # noqa: N802 - BaseHTTPRequestHandler API + self._send_headers(body=False) + + def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler API + selected = self._send_headers(body=True) + if selected is None: + return + start, end = selected + remaining = end - start + 1 + with file_path.open("rb") as stream: + stream.seek(start) + while remaining > 0: + chunk = stream.read(min(1024 * 1024, remaining)) + if not chunk: + break + try: + self.wfile.write(chunk) + except (BrokenPipeError, ConnectionResetError): + break + remaining -= len(chunk) + + server = ThreadingHTTPServer(("0.0.0.0", 0), Handler) + server.daemon_threads = True + self._server = server + self._thread = threading.Thread( + target=server.serve_forever, + name="invokeai-remote-worker-model-transfer", + daemon=True, + ) + self._thread.start() + + host = self.advertise_host or _route_local_ip(self.remote_url) + if "://" in host: + host = urllib.parse.urlsplit(host).hostname or host + host = host.strip("[]") + self.url = f"http://{host}:{server.server_port}{self._request_path}" + return self + + def __exit__(self, exc_type: Any, exc: Any, traceback: Any) -> None: + if self._server is not None: + self._server.shutdown() + self._server.server_close() + if self._thread is not None: + self._thread.join(timeout=2.0) + self._server = None + self._thread = None diff --git a/invokeai/app/invocations/remote_worker/model_transfer_state.py b/invokeai/app/invocations/remote_worker/model_transfer_state.py new file mode 100644 index 00000000000..22d53b9113f --- /dev/null +++ b/invokeai/app/invocations/remote_worker/model_transfer_state.py @@ -0,0 +1,55 @@ +from __future__ import annotations + +import threading +from dataclasses import dataclass, field + + +@dataclass +class ModelTransferTask: + remote_url: str + model_hash: str + cancel_requested: threading.Event = field(default_factory=threading.Event) + shared_lock: threading.Lock = field(default_factory=threading.Lock) + + +_LOCK = threading.Lock() +_TRANSFER_TASKS: dict[int, ModelTransferTask] = {} +_TRANSFER_SEQUENCE = 0 + + +def register_model_transfer(remote_url: str, model_hash: str) -> tuple[int, ModelTransferTask]: + """Register one generation participating in a possibly shared worker-side model install.""" + global _TRANSFER_SEQUENCE + task = ModelTransferTask(remote_url=remote_url, model_hash=model_hash) + + with _LOCK: + task.shared_lock = next( + ( + active.shared_lock + for active in _TRANSFER_TASKS.values() + if active.remote_url == remote_url and active.model_hash == model_hash + ), + task.shared_lock, + ) + _TRANSFER_SEQUENCE += 1 + transfer_id = _TRANSFER_SEQUENCE + _TRANSFER_TASKS[transfer_id] = task + + return transfer_id, task + + +def unregister_model_transfer(transfer_id: int) -> None: + with _LOCK: + _TRANSFER_TASKS.pop(transfer_id, None) + + +def another_generation_needs_model(task: ModelTransferTask) -> bool: + """Preserve a shared install while another live generation still requires it.""" + with _LOCK: + return any( + other is not task + and not other.cancel_requested.is_set() + and other.remote_url == task.remote_url + and other.model_hash == task.model_hash + for other in _TRANSFER_TASKS.values() + ) diff --git a/invokeai/app/invocations/remote_worker/remote_client.py b/invokeai/app/invocations/remote_worker/remote_client.py new file mode 100644 index 00000000000..8ebd0b6db73 --- /dev/null +++ b/invokeai/app/invocations/remote_worker/remote_client.py @@ -0,0 +1,893 @@ +from __future__ import annotations + +import json +import os +import secrets +import ssl +import time +import urllib.error +import urllib.parse +import urllib.request +from copy import deepcopy +from dataclasses import dataclass +from io import BytesIO +from pathlib import Path +from typing import Any, Callable + +from PIL import Image + + +class RemoteInvokeError(RuntimeError): + """Raised when communication with the remote InvokeAI worker fails.""" + + +@dataclass(frozen=True) +class RemoteCredentials: + email: str + password: str + remember_me: bool = True + + +@dataclass(frozen=True) +class RemoteConfig: + base_url: str + api_key: str + auth_header: str = "Authorization" + auth_prefix: str = "Bearer " + verify_ssl: bool = True + credentials_file: str = "" + user_id: str = "" + + @classmethod + def from_environment(cls, base_url: str = "", verify_ssl: bool = True, user_id: str = "") -> "RemoteConfig": + url = (base_url or os.getenv("INVOKE_REMOTE_URL", "")).strip().rstrip("/") + if not url: + raise RemoteInvokeError( + "Remote InvokeAI URL is not configured. Set INVOKE_REMOTE_URL or fill in Remote URL on the node." + ) + + default_credentials_file = Path(__file__).with_name("remote_auth.json") + credentials_file = os.getenv("INVOKE_REMOTE_AUTH_FILE", str(default_credentials_file)).strip() + + return cls( + base_url=url, + api_key=os.getenv("INVOKE_REMOTE_API_KEY", "").strip(), + auth_header=os.getenv("INVOKE_REMOTE_AUTH_HEADER", "Authorization").strip() or "Authorization", + auth_prefix=os.getenv("INVOKE_REMOTE_AUTH_PREFIX", "Bearer "), + verify_ssl=verify_ssl, + credentials_file=credentials_file, + user_id=user_id, + ) + + +class RemoteInvokeClient: + def __init__(self, config: RemoteConfig, request_timeout_seconds: float = 30.0): + self.config = config + self.request_timeout_seconds = request_timeout_seconds + self._auth_checked = False + self._multiuser = False + self._token = config.api_key or "" + self._credentials: RemoteCredentials | None = None + + def _ssl_context(self): + if self.config.verify_ssl: + return None + return ssl._create_unverified_context() # noqa: SLF001 - explicitly user-controlled for LAN/self-signed use + + @staticmethod + def _normalise_url(url: str) -> str: + return str(url).strip().rstrip("/") + + def _load_credentials(self) -> RemoteCredentials: + if self._credentials is not None: + return self._credentials + + if self.config.user_id: + # Per-local-user saved credentials; import lazily during graph bootstrap. + from invokeai.app.invocations.remote_worker.credential_vault import get_saved_credentials + + entry = get_saved_credentials(self.config.user_id, self.config.base_url) + if entry is None: + raise RemoteInvokeError( + f"Remote InvokeAI at {self.config.base_url} requires login. " + "Open Remote Workers, enter this worker's email and password, and save them." + ) + self._credentials = RemoteCredentials( + email=str(entry["email"]), + password=str(entry["password"]), + remember_me=bool(entry.get("remember_me", True)), + ) + return self._credentials + + # Legacy file path for non-panel callers without a local user identity. + path = Path(self.config.credentials_file) + if not path.is_file(): + raise RemoteInvokeError( + "Remote InvokeAI has multi-user mode enabled, but no credentials file was found. " + f"Create '{path}' from remote_auth.example.json and add credentials for {self.config.base_url}, " + "or set INVOKE_REMOTE_AUTH_FILE to another JSON file." + ) + + try: + data = json.loads(path.read_text(encoding="utf-8")) + except Exception as exc: + raise RemoteInvokeError(f"Could not read remote auth file '{path}': {exc}") from exc + + if not isinstance(data, dict): + raise RemoteInvokeError(f"Remote auth file '{path}' must contain a JSON object") + + servers = data.get("servers", data) + if not isinstance(servers, dict): + raise RemoteInvokeError(f"Remote auth file '{path}' must contain a 'servers' object") + + wanted = self._normalise_url(self.config.base_url) + entry = None + for key, value in servers.items(): + if self._normalise_url(str(key)) == wanted: + entry = value + break + + if not isinstance(entry, dict): + raise RemoteInvokeError( + f"Remote InvokeAI has multi-user mode enabled, but '{path}' has no credentials for {wanted}" + ) + + email = str(entry.get("email", "")).strip() + password = str(entry.get("password", "")) + remember_me = bool(entry.get("remember_me", True)) + if not email or not password: + raise RemoteInvokeError( + f"Credentials for {wanted} in '{path}' must include non-empty 'email' and 'password' values" + ) + + self._credentials = RemoteCredentials(email=email, password=password, remember_me=remember_me) + return self._credentials + + def _headers(self, json_body: bool = False, include_auth: bool = True) -> dict[str, str]: + headers = {"Accept": "application/json"} + if json_body: + headers["Content-Type"] = "application/json" + if include_auth and self._token: + headers[self.config.auth_header] = f"{self.config.auth_prefix}{self._token}" + return headers + + def _request_raw( + self, + method: str, + path: str, + body: Any = None, + json_body: bool = False, + include_auth: bool = True, + content_type: str | None = None, + content_length: int | None = None, + ) -> bytes: + url = f"{self.config.base_url}{path}" + headers = self._headers(json_body=json_body, include_auth=include_auth) + if content_type is not None: + headers["Content-Type"] = content_type + if content_length is not None: + headers["Content-Length"] = str(content_length) + request = urllib.request.Request( + url=url, + data=body, + headers=headers, + method=method, + ) + try: + with urllib.request.urlopen( + request, + timeout=self.request_timeout_seconds, + context=self._ssl_context(), + ) as response: + return response.read() + except urllib.error.URLError as exc: + if isinstance(exc, urllib.error.HTTPError): + raise + raise RemoteInvokeError(f"Could not reach remote InvokeAI at {url}: {exc}") from exc + + def _login(self) -> None: + credentials = self._load_credentials() + payload = json.dumps( + { + "email": credentials.email, + "password": credentials.password, + "remember_me": credentials.remember_me, + }, + separators=(",", ":"), + ).encode("utf-8") + try: + raw = self._request_raw( + "POST", + "/api/v1/auth/login", + body=payload, + json_body=True, + include_auth=False, + ) + except urllib.error.HTTPError as exc: + try: + detail = exc.read().decode("utf-8", errors="replace") + except Exception: + detail = str(exc) + raise RemoteInvokeError( + f"Remote InvokeAI login failed with HTTP {exc.code} for {self.config.base_url}: {detail[:2000]}" + ) from exc + + try: + data = json.loads(raw.decode("utf-8")) + token = str(data.get("token", "")).strip() if isinstance(data, dict) else "" + except Exception as exc: + raise RemoteInvokeError("Remote InvokeAI login returned invalid JSON") from exc + if not token: + raise RemoteInvokeError("Remote InvokeAI login succeeded but did not return a token") + self._token = token + + def _ensure_auth_mode(self) -> None: + if self._auth_checked: + if self._multiuser and not self._token: + self._login() + return + + # InvokeAI v7 always exposes /api/v1/auth/status. The response explicitly + # tells us whether multi-user mode is enabled; endpoint existence alone is + # not a valid signal. + try: + raw = self._request_raw("GET", "/api/v1/auth/status", include_auth=False) + except urllib.error.HTTPError as exc: + # Keep a compatibility fallback for unusual v7 builds/proxies that do not + # expose the status endpoint: treat a missing endpoint as single-user. + if exc.code in (404, 405): + self._multiuser = False + self._auth_checked = True + self._token = self.config.api_key or "" + return + try: + detail = exc.read().decode("utf-8", errors="replace") + except Exception: + detail = str(exc) + raise RemoteInvokeError( + f"Could not determine remote InvokeAI authentication mode: HTTP {exc.code} from " + f"/api/v1/auth/status: {detail[:2000]}" + ) from exc + + try: + status = json.loads(raw.decode("utf-8")) + except Exception as exc: + raise RemoteInvokeError("Remote InvokeAI auth status returned invalid JSON") from exc + if not isinstance(status, dict) or "multiuser_enabled" not in status: + raise RemoteInvokeError("Remote InvokeAI auth status did not include the v7 'multiuser_enabled' field") + + self._multiuser = bool(status.get("multiuser_enabled")) + self._auth_checked = True + if self._multiuser and bool(status.get("setup_required")): + raise RemoteInvokeError( + f"Remote InvokeAI at {self.config.base_url} has multi-user mode enabled but initial admin setup is still required" + ) + if self._multiuser and not self._token: + self._login() + + def _request( + self, + method: str, + path: str, + body: Any = None, + json_body: bool = False, + content_type: str | None = None, + content_length: int | None = None, + ) -> bytes: + self._ensure_auth_mode() + try: + return self._request_raw( + method, + path, + body=body, + json_body=json_body, + include_auth=True, + content_type=content_type, + content_length=content_length, + ) + except urllib.error.HTTPError as exc: + # In multi-user mode, a 401 usually means the cached JWT expired. Re-login + # once and replay the exact request. Do not loop indefinitely on bad credentials. + if exc.code == 401 and self._multiuser: + self._token = "" + self._login() + try: + return self._request_raw( + method, + path, + body=body, + json_body=json_body, + include_auth=True, + content_type=content_type, + content_length=content_length, + ) + except urllib.error.HTTPError as retry_exc: + exc = retry_exc + try: + detail = exc.read().decode("utf-8", errors="replace") + except Exception: + detail = str(exc) + raise RemoteInvokeError(f"Remote InvokeAI returned HTTP {exc.code} for {path}: {detail[:2000]}") from exc + + def _request_json_value(self, method: str, path: str, payload: dict[str, Any] | None = None) -> Any: + body = None + json_body = payload is not None + if payload is not None: + body = json.dumps(payload, separators=(",", ":")).encode("utf-8") + raw = self._request(method, path, body=body, json_body=json_body) + try: + return json.loads(raw.decode("utf-8")) + except Exception as exc: + preview = raw[:500].decode("utf-8", errors="replace") + raise RemoteInvokeError(f"Expected JSON from remote InvokeAI for {path}, got: {preview}") from exc + + def _request_json(self, method: str, path: str, payload: dict[str, Any] | None = None) -> dict[str, Any]: + data = self._request_json_value(method, path, payload) + if not isinstance(data, dict): + raise RemoteInvokeError(f"Expected a JSON object from remote InvokeAI for {path}") + return data + + def list_models(self) -> list[dict[str, Any]]: + data = self._request_json_value("GET", "/api/v2/models/") + if isinstance(data, list): + return [x for x in data if isinstance(x, dict)] + if isinstance(data, dict): + for key in ("models", "items", "data"): + value = data.get(key) + if isinstance(value, list): + return [x for x in value if isinstance(x, dict)] + raise RemoteInvokeError("Remote /api/v2/models/ response did not contain a model list") + + def get_model(self, key: str) -> dict[str, Any]: + encoded = urllib.parse.quote(str(key), safe="") + return self._request_json("GET", f"/api/v2/models/i/{encoded}") + + def get_model_layout(self, key: str) -> dict[str, Any]: + """Return the Remote Workers layout signature for one registered model.""" + encoded = urllib.parse.quote(str(key), safe="") + return self._request_json("GET", f"/api/v1/remote_workers/models/{encoded}/layout") + + def get_model_by_hash(self, model_hash: str) -> dict[str, Any] | None: + """Return the remote model record with this content hash, or None when it is absent.""" + encoded = urllib.parse.quote(str(model_hash), safe="") + try: + return self._request_json("GET", f"/api/v2/models/get_by_hash?hash={encoded}") + except RemoteInvokeError as exc: + if "HTTP 404" in str(exc): + return None + raise + + def install_model_from_url(self, source_url: str, *, name: str = "") -> dict[str, Any]: + """Ask the remote InvokeAI model manager to download/probe/register one model URL.""" + encoded = urllib.parse.quote(str(source_url), safe="") + config: dict[str, Any] = {} + if name.strip(): + config["name"] = name.strip() + return self._request_json( + "POST", + f"/api/v2/models/install?source={encoded}&inplace=false", + payload=config, + ) + + def install_directory_from_manifest(self, manifest: dict[str, Any]) -> dict[str, Any]: + """Wait briefly if the previous cancelled transfer is still cleaning up.""" + deadline = time.monotonic() + 20 + while True: + try: + return self._request_json("POST", "/api/v1/remote_workers/diffusers/install", payload=manifest) + except RemoteInvokeError as exc: + if ( + "HTTP 422" not in str(exc) + or "Previous transfer cancellation is still being cleaned up" not in str(exc) + or time.monotonic() >= deadline + ): + raise + time.sleep(0.25) + + def get_directory_install_job(self, job_id: int) -> dict[str, Any]: + return self._request_json("GET", f"/api/v1/remote_workers/diffusers/install/{int(job_id)}") + + def cancel_model_install(self, job_id: int, *, directory: bool = False) -> None: + """Stop one install job; never delete a completed installed model.""" + path = ( + f"/api/v1/remote_workers/diffusers/install/{int(job_id)}" + if directory + else f"/api/v2/models/install/{int(job_id)}" + ) + self._request("DELETE", path) + + def get_model_install_job(self, job_id: int) -> dict[str, Any]: + return self._request_json("GET", f"/api/v2/models/install/{int(job_id)}") + + @staticmethod + def _model_payload(detail: dict[str, Any]) -> dict[str, Any]: + nested = detail.get("model") + return nested if isinstance(nested, dict) else detail + + def remap_model_identifiers( + self, + graph: dict[str, Any], + missing_model_handler: Callable[[dict[str, Any]], None] | None = None, + model_match_validator: Callable[[dict[str, Any], dict[str, Any]], bool] | None = None, + ) -> list[str]: + """Replace local model keys with the matching remote installation keys. + + Resolve by model hash first. The hash is portable across InvokeAI installs, + while the database key/UUID is installation-local. Name+base+type remains a + compatibility fallback for model identifiers without a usable hash. + + If ``missing_model_handler`` is supplied, it is called once for a required model + that cannot be found remotely. The handler may transfer/install the model; resolution + is then retried before the graph is rejected. + + ``model_match_validator`` can reject a same-hash candidate when installation-local + details (such as directory layout) are incompatible with the primary model. + """ + remote_models = self.list_models() + details_cache: dict[str, dict[str, Any]] = {} + messages: list[str] = [] + handled_missing: set[str] = set() + + def detail_for(model: dict[str, Any]) -> dict[str, Any]: + key = str(model.get("key", "")) + if not key: + return model + if key not in details_cache: + try: + details_cache[key] = self.get_model(key) + except RemoteInvokeError: + details_cache[key] = model + return details_cache[key] + + def refresh_models() -> None: + nonlocal remote_models + remote_models = self.list_models() + details_cache.clear() + + def identifier_from(detail: dict[str, Any], fallback: dict[str, Any]) -> dict[str, Any]: + detail = self._model_payload(detail) + result = deepcopy(fallback) + for k in ("key", "hash", "name", "base", "type", "submodel_type"): + if k in detail and detail[k] is not None: + result[k] = detail[k] + return result + + def candidate_is_compatible(value: dict[str, Any], detail: dict[str, Any]) -> bool: + local_hash = str(value.get("hash") or "").strip() + if model_match_validator is None or not local_hash: + return True + return bool(model_match_validator(deepcopy(value), self._model_payload(detail))) + + def resolve(value: dict[str, Any]) -> dict[str, Any] | None: + local_hash = str(value.get("hash") or "").strip() + if local_hash: + # Do not use /get_by_hash here: that endpoint returns only the first + # record, but multiple installs can legitimately share a weight hash + # while having different directory layouts. + for model in remote_models: + payload = self._model_payload(model) + if str(payload.get("hash") or "").strip() != local_hash: + continue + detail = detail_for(model) + if candidate_is_compatible(value, detail): + return detail + + local_name = value.get("name") + local_base = value.get("base") + local_type = value.get("type") + candidates = [ + model + for model in remote_models + if self._model_payload(model).get("name") == local_name + and self._model_payload(model).get("base") == local_base + and self._model_payload(model).get("type") == local_type + ] + verified: list[dict[str, Any]] = [] + for candidate in candidates: + detail = detail_for(candidate) + payload = self._model_payload(detail) + remote_hash = str(payload.get("hash") or "").strip() + if local_hash and remote_hash and local_hash != remote_hash: + continue + if not candidate_is_compatible(value, detail): + continue + verified.append(detail) + + if len(verified) == 1: + return verified[0] + if len(verified) > 1: + raise RemoteInvokeError( + f"Remote worker has multiple matching models named '{local_name}'. " + "Hash-based resolution could not disambiguate them." + ) + return None + + def remap(value: Any) -> Any: + if isinstance(value, list): + return [remap(x) for x in value] + if not isinstance(value, dict): + return value + + if all(k in value for k in ("key", "name", "base", "type")): + local_name = str(value.get("name") or value.get("key") or "model") + local_hash = str(value.get("hash") or "").strip() + resolved = resolve(value) + + missing_identity = local_hash or str(value.get("key") or local_name) + if resolved is None and missing_model_handler is not None and missing_identity not in handled_missing: + handled_missing.add(missing_identity) + missing_model_handler(deepcopy(value)) + refresh_models() + resolved = resolve(value) + + if resolved is None: + if local_hash: + raise RemoteInvokeError( + f"Remote worker does not have required model '{local_name}' with hash {local_hash}" + ) + raise RemoteInvokeError( + f"Remote worker does not have required model '{local_name}' " + f"(base={value.get('base')}, type={value.get('type')})" + ) + + old_key = str(value.get("key")) + mapped = identifier_from(resolved, value) + new_key = str(mapped.get("key")) + remote_hash = str(mapped.get("hash") or "").strip() + if local_hash and remote_hash and local_hash != remote_hash: + raise RemoteInvokeError( + f"Remote model '{local_name}' resolved to hash {remote_hash}, expected {local_hash}" + ) + if new_key != old_key: + method = "hash" if local_hash else "name/base/type" + messages.append(f"{local_name}: {old_key} -> {new_key} ({method})") + return mapped + + return {k: remap(v) for k, v in value.items()} + + nodes = graph.get("nodes") + if not isinstance(nodes, dict): + raise RemoteInvokeError("Executable graph has no nodes object") + for node_id, node in list(nodes.items()): + nodes[node_id] = remap(node) + return messages + + def cancel_queue_item(self, item_id: int | str, queue_id: str = "default") -> dict[str, Any]: + """Cancel only this remote queue item, with this client's saved user credentials.""" + queue_path = urllib.parse.quote(queue_id, safe="") + return self._request_json("PUT", f"/api/v1/queue/{queue_path}/i/{int(item_id)}/cancel") + + def delete_queue_item(self, item_id: int | str, queue_id: str = "default") -> None: + """Remove one finished worker queue record; do not prune anyone else's history.""" + queue_path = urllib.parse.quote(queue_id, safe="") + self._request("DELETE", f"/api/v1/queue/{queue_path}/i/{int(item_id)}") + + def get_item(self, item_id: int | str, queue_id: str = "default") -> dict[str, Any]: + return self._request_json("GET", f"/api/v1/queue/{urllib.parse.quote(queue_id)}/i/{int(item_id)}") + + def get_current_item(self, queue_id: str = "default") -> dict[str, Any] | None: + data = self._request_json_value("GET", f"/api/v1/queue/{urllib.parse.quote(queue_id)}/current") + if data is None: + return None + if not isinstance(data, dict): + raise RemoteInvokeError("Expected current queue item to be a JSON object or null") + return data + + def list_progress_previews(self, queue_id: str = "default") -> list[dict[str, Any]]: + """Return InvokeAI v7's retained live progress previews for this authenticated user.""" + data = self._request_json_value("GET", f"/api/v1/queue/{urllib.parse.quote(queue_id)}/previews") + if not isinstance(data, list): + raise RemoteInvokeError("Remote progress previews endpoint did not return a JSON list") + return [entry for entry in data if isinstance(entry, dict)] + + def get_progress_preview(self, item_id: int, queue_id: str = "default") -> dict[str, Any] | None: + """Return the latest retained progress preview for one remote queue item, if available.""" + wanted = int(item_id) + for entry in self.list_progress_previews(queue_id=queue_id): + try: + if int(entry.get("item_id")) == wanted: + return entry + except (TypeError, ValueError): + continue + return None + + def enqueue_graph( + self, + graph: dict[str, Any], + queue_id: str = "default", + origin: str = "invokeai-remote-worker-node", + destination: str | None = None, + workflow: dict[str, Any] | None = None, + ) -> int: + batch: dict[str, Any] = { + "graph": graph, + "runs": 1, + "origin": origin, + } + if destination is not None: + batch["destination"] = destination + if workflow is not None: + batch["workflow"] = workflow + payload = {"batch": batch} + result = self._request_json( + "POST", + f"/api/v1/queue/{urllib.parse.quote(queue_id)}/enqueue_batch", + payload=payload, + ) + item_ids = result.get("item_ids") + if not isinstance(item_ids, list) or not item_ids: + raise RemoteInvokeError(f"Remote enqueue succeeded but no item_ids were returned: {result}") + return int(item_ids[0]) + + def wait_for_item( + self, + item_id: int, + queue_id: str = "default", + poll_interval_seconds: float = 0.75, + timeout_seconds: float = 1800.0, + ) -> dict[str, Any]: + started = time.monotonic() + while True: + item = self.get_item(item_id=item_id, queue_id=queue_id) + status = str(item.get("status", "")).lower() + if status == "completed": + return item + if status in {"failed", "canceled", "cancelled"}: + errors = item.get("session", {}).get("errors", {}) + raise RemoteInvokeError( + f"Remote queue item {item_id} ended with status '{status}'. Errors: {json.dumps(errors)[:3000]}" + ) + if time.monotonic() - started > timeout_seconds: + raise RemoteInvokeError(f"Remote queue item {item_id} timed out after {timeout_seconds:g} seconds") + time.sleep(poll_interval_seconds) + + @staticmethod + def extract_image_names( + item: dict[str, Any], + non_intermediate_only: bool = True, + allowed_result_node_ids: set[str] | list[str] | tuple[str, ...] | None = None, + allow_empty: bool = False, + ) -> list[str]: + """Return unique image names from completed session results. + + If ``allowed_result_node_ids`` is provided, only results from those exact source + graph nodes are considered. Remote result filtering uses this to respect InvokeAI's + Save to Gallery setting (``is_intermediate == False``) instead of importing every + temporary image produced by a multi-stage workflow. + + Without an explicit allow-list, intermediate image-producing nodes are skipped by + default. The legacy fallback to all image outputs is retained only in that mode. + """ + session = item.get("session", {}) + results = session.get("results", {}) if isinstance(session, dict) else {} + if not isinstance(results, dict): + raise RemoteInvokeError("Remote queue item has no session.results object") + + graph_nodes = {} + graph = session.get("graph") if isinstance(session, dict) else None + if isinstance(graph, dict) and isinstance(graph.get("nodes"), dict): + graph_nodes = graph["nodes"] + + allowed_ids = None if allowed_result_node_ids is None else {str(x) for x in allowed_result_node_ids} + + def collect(filter_intermediate: bool) -> list[str]: + image_names: list[str] = [] + seen: set[str] = set() + for result_id, result in results.items(): + if not isinstance(result, dict): + continue + if allowed_ids is not None and str(result_id) not in allowed_ids: + continue + if filter_intermediate and graph_nodes: + graph_node = graph_nodes.get(str(result_id)) + if isinstance(graph_node, dict) and bool(graph_node.get("is_intermediate", False)): + continue + + candidates: list[str] = [] + image = result.get("image") + if isinstance(image, dict) and image.get("image_name"): + candidates.append(str(image["image_name"])) + images = result.get("images") + if isinstance(images, list): + for entry in images: + if isinstance(entry, dict) and entry.get("image_name"): + candidates.append(str(entry["image_name"])) + + for name in candidates: + if name not in seen: + seen.add(name) + image_names.append(name) + return image_names + + image_names = collect(non_intermediate_only) + # Do not fall back to temporary images when the caller supplied an explicit + # Save-to-Gallery allow-list. Importing extra intermediate images would violate + # the workflow's gallery settings. + if not image_names and non_intermediate_only and allowed_ids is None: + image_names = collect(False) + if not image_names and not allow_empty: + if allowed_ids is not None: + raise RemoteInvokeError( + "Remote render completed but none of the Save-to-Gallery nodes produced an image output" + ) + raise RemoteInvokeError("Remote render completed but no image output image_name was found") + return image_names + + @staticmethod + def extract_video_names(item: dict[str, Any]) -> list[str]: + """Unique video outputs from completed InvokeAI v7 sessions (including video lists).""" + session = item.get("session", {}) + results = session.get("results", {}) if isinstance(session, dict) else {} + if not isinstance(results, dict): + raise RemoteInvokeError("Remote queue item has no session.results object") + names: list[str] = [] + seen: set[str] = set() + for result in results.values(): + if not isinstance(result, dict): + continue + candidates: list[Any] = [result.get("video")] + videos = result.get("videos") + if isinstance(videos, list): + candidates.extend(videos) + # Some video-producing integrations return the field at the top level. + candidates.append(result) + for value in candidates: + if not isinstance(value, dict): + continue + name = value.get("video_name") + if isinstance(name, str) and name and name not in seen: + seen.add(name) + names.append(name) + return names + + def get_video_metadata(self, video_name: str) -> str | None: + """Fetch the remote video record metadata before the native local import.""" + encoded_name = urllib.parse.quote(video_name, safe="") + value = self._request_json_value("GET", f"/api/v1/videos/i/{encoded_name}/metadata") + if value is None: + return None + if not isinstance(value, dict): + raise RemoteInvokeError(f"Remote video '{video_name}' returned invalid metadata") + return json.dumps(value, separators=(",", ":"), ensure_ascii=False) + + def get_video_dto(self, video_name: str) -> dict[str, Any]: + encoded_name = urllib.parse.quote(video_name, safe="") + return self._request_json("GET", f"/api/v1/videos/i/{encoded_name}") + + def filter_gallery_video_names(self, video_names: list[str]) -> list[str]: + """Use the remote's actual persisted Save-to-Gallery flag, as for images.""" + gallery_names: list[str] = [] + for video_name in video_names: + if self.get_video_dto(video_name).get("is_intermediate") is False: + gallery_names.append(video_name) + return gallery_names + + def download_video(self, video_name: str) -> bytes: + encoded_name = urllib.parse.quote(video_name, safe="") + return self._request("GET", f"/api/v1/videos/i/{encoded_name}/full") + + def delete_video(self, video_name: str) -> dict[str, Any]: + encoded_name = urllib.parse.quote(video_name, safe="") + return self._request_json("DELETE", f"/api/v1/videos/i/{encoded_name}") + + def upload_input_image(self, image: Image.Image) -> str: + """Upload a local input as an InvokeAI-owned remote image and return its NEW name. + + Intermediate+user category prevents source inputs appearing as finished Gallery + results. This uses the same authenticated client (including the 401 retry). + """ + png = BytesIO() + try: + image.save(png, format="PNG") + except Exception as exc: + raise RemoteInvokeError(f"Could not encode source image as PNG: {exc}") from exc + boundary = f"irw-{secrets.token_hex(16)}" + prefix = ( + f"--{boundary}\r\n" + 'Content-Disposition: form-data; name="file"; filename="remote-input.png"\r\n' + "Content-Type: image/png\r\n\r\n" + ).encode("utf-8") + body = prefix + png.getvalue() + f"\r\n--{boundary}--\r\n".encode("utf-8") + query = urllib.parse.urlencode({"image_category": "user", "is_intermediate": "true"}) + path = f"/api/v1/images/upload?{query}" + raw = self._request("POST", path, body=body, content_type=f"multipart/form-data; boundary={boundary}") + try: + record = json.loads(raw.decode("utf-8")) + except (UnicodeError, ValueError) as exc: + raise RemoteInvokeError("Remote image upload did not return valid JSON") from exc + remote_name = record.get("image_name") if isinstance(record, dict) else None + if not isinstance(remote_name, str) or not remote_name: + raise RemoteInvokeError("Remote image upload did not return image_name") + return remote_name + + def upload_input_video(self, video_path: Path) -> str: + """Stream a local source video to the remote as an intermediate input.""" + path = Path(video_path) + try: + file_size = path.stat().st_size + except OSError as exc: + raise RemoteInvokeError(f"Could not read source video '{path}': {exc}") from exc + + boundary = f"irw-{secrets.token_hex(16)}" + prefix = ( + f"--{boundary}\r\n" + 'Content-Disposition: form-data; name="file"; filename="remote-input.mp4"\r\n' + "Content-Type: video/mp4\r\n\r\n" + ).encode("utf-8") + suffix = f"\r\n--{boundary}--\r\n".encode("utf-8") + + class MultipartVideoBody: + def __iter__(self): + yield prefix + try: + with path.open("rb") as source: + while chunk := source.read(1024 * 1024): + yield chunk + except OSError as exc: + raise RemoteInvokeError(f"Could not read source video '{path}': {exc}") from exc + yield suffix + + query = urllib.parse.urlencode({"video_category": "user", "is_intermediate": "true"}) + raw = self._request( + "POST", + f"/api/v1/videos/upload?{query}", + body=MultipartVideoBody(), + content_type=f"multipart/form-data; boundary={boundary}", + content_length=len(prefix) + file_size + len(suffix), + ) + try: + record = json.loads(raw.decode("utf-8")) + except (UnicodeError, ValueError) as exc: + raise RemoteInvokeError("Remote video upload did not return valid JSON") from exc + remote_name = record.get("video_name") if isinstance(record, dict) else None + if not isinstance(remote_name, str) or not remote_name: + raise RemoteInvokeError("Remote video upload did not return video_name") + return remote_name + + def get_image_metadata(self, image_name: str) -> str | None: + """Fetch the exact image record metadata, not PNG/PIL metadata.""" + encoded_name = urllib.parse.quote(image_name, safe="") + value = self._request_json_value("GET", f"/api/v1/images/i/{encoded_name}/metadata") + if value is None: + return None + if not isinstance(value, dict): + raise RemoteInvokeError(f"Remote image '{image_name}' returned invalid metadata") + return json.dumps(value, separators=(",", ":"), ensure_ascii=False) + + def get_image_dto(self, image_name: str) -> dict[str, Any]: + encoded_name = urllib.parse.quote(image_name, safe="") + return self._request_json("GET", f"/api/v1/images/i/{encoded_name}") + + def filter_gallery_image_names(self, image_names: list[str]) -> list[str]: + """Return only images that InvokeAI itself marks as non-intermediate. + + InvokeAI's Save to Gallery state is persisted on the generated image record as + ``is_intermediate``. Filtering the completed session's image names through the + remote ImageDTO avoids relying on graph/result node-id correspondence, which can + differ after graph expansion and execution. + """ + gallery_names: list[str] = [] + for image_name in image_names: + dto = self.get_image_dto(image_name) + if dto.get("is_intermediate") is False: + gallery_names.append(image_name) + return gallery_names + + def delete_image(self, image_name: str) -> dict[str, Any]: + """Delete an image through InvokeAI's normal image service API. + + This removes the image record and lets InvokeAI clean up its associated + stored/thumbnail/cache files through the same path the UI uses. + """ + encoded_name = urllib.parse.quote(image_name, safe="") + return self._request_json("DELETE", f"/api/v1/images/i/{encoded_name}") + + def download_image(self, image_name: str) -> Image.Image: + encoded_name = urllib.parse.quote(image_name, safe="") + raw = self._request("GET", f"/api/v1/images/i/{encoded_name}/full") + try: + with Image.open(BytesIO(raw)) as image: + image.load() + return image.copy() + except Exception as exc: + raise RemoteInvokeError(f"Downloaded remote image '{image_name}' could not be decoded") from exc diff --git a/invokeai/app/invocations/remote_worker/remote_media.py b/invokeai/app/invocations/remote_worker/remote_media.py new file mode 100644 index 00000000000..90e2980afc3 --- /dev/null +++ b/invokeai/app/invocations/remote_worker/remote_media.py @@ -0,0 +1,174 @@ +from __future__ import annotations + +import json +import tempfile +from pathlib import Path +from typing import Any + +from invokeai.app.invocations.remote_worker.remote_client import RemoteInvokeClient, RemoteInvokeError +from invokeai.app.services.board_records.board_records_common import BoardVisibility +from invokeai.app.services.image_records.image_records_common import ImageCategory, ResourceOrigin +from invokeai.app.services.session_processor.session_processor_common import ProgressImage +from invokeai.app.util.video_thumbnails import probe_video_with_codec + + +def _model_json(value: Any) -> str | None: + if value is None: + return None + try: + return value.model_dump_json() + except Exception: + try: + return json.dumps(value, separators=(",", ":")) + except Exception: + return None + + +def progress_image_from_remote(preview: dict[str, Any]) -> ProgressImage | None: + raw = preview.get("image") + if not isinstance(raw, dict): + return None + data_url = raw.get("dataURL") + if not isinstance(data_url, str) or not data_url: + return None + try: + return ProgressImage( + width=int(raw.get("width")), + height=int(raw.get("height")), + dataURL=data_url, + ) + except Exception: + return None + + +def _assert_background_save_access(services: Any, user_id: str | None, board_id: str | None) -> None: + if not getattr(services.configuration, "multiuser", False): + return + user = services.users.get(user_id) + if user is None or not user.is_active: + raise PermissionError("Queue user is not authorized to save returned remote images") + if board_id: + board = services.boards.get_dto(board_id) + if not user.is_admin and board.user_id != user_id and board.board_visibility != BoardVisibility.Public: + raise PermissionError("Queue user is not authorized to save returned remote images to this board") + + +def save_local_image( + *, + services: Any, + queue_item: Any, + invocation: Any, + image: Any, + metadata: str | None, + board_id: str, + result_destination: str, + source_node_id: str | None = None, +) -> Any: + user_id = getattr(queue_item, "user_id", None) + target_board_id = (board_id.strip() or None) if result_destination == "gallery" else None + _assert_background_save_access(services, user_id, target_board_id) + + workflow_json = _model_json(getattr(queue_item, "workflow", None)) + session = getattr(queue_item, "session", None) + graph_json = _model_json(getattr(session, "graph", None)) + + return services.images.create( + image=image, + is_intermediate=False, + image_category=ImageCategory.OTHER if result_destination == "canvas" else ImageCategory.GENERAL, + board_id=target_board_id, + metadata=metadata, + image_origin=ResourceOrigin.INTERNAL, + workflow=workflow_json, + graph=graph_json, + session_id=getattr(queue_item, "session_id", None), + node_id=source_node_id or getattr(invocation, "id", None), + user_id=user_id, + ) + + +def save_local_video( + *, + services: Any, + queue_item: Any, + invocation: Any, + video_bytes: bytes, + metadata: str | None, + board_id: str, + result_destination: str, + source_node_id: str | None = None, +) -> Any: + """Stage one MP4 alongside outputs/videos; the native service moves it into storage.""" + user_id = getattr(queue_item, "user_id", None) + target_board_id = (board_id.strip() or None) if result_destination == "gallery" else None + _assert_background_save_access(services, user_id, target_board_id) + + outputs_path = services.configuration.outputs_path + if outputs_path is None: + raise RemoteInvokeError("Primary InvokeAI has no configured outputs path for video import") + + stage_dir = Path(outputs_path) / "videos" + stage_dir.mkdir(parents=True, exist_ok=True) + stage_path: Path | None = None + try: + with tempfile.NamedTemporaryFile(prefix=".irw_remote_", suffix=".mp4", dir=stage_dir, delete=False) as file: + stage_path = Path(file.name) + file.write(video_bytes) + + width, height, duration, fps, _codec = probe_video_with_codec(stage_path) + workflow_json = _model_json(getattr(queue_item, "workflow", None)) + session = getattr(queue_item, "session", None) + graph_json = _model_json(getattr(session, "graph", None)) + + return services.videos.create( + source_path=stage_path, + width=width, + height=height, + duration=duration, + fps=fps, + video_origin=ResourceOrigin.INTERNAL, + video_category=ImageCategory.OTHER if result_destination == "canvas" else ImageCategory.GENERAL, + board_id=target_board_id, + is_intermediate=False, + metadata=metadata, + workflow=workflow_json, + graph=graph_json, + session_id=getattr(queue_item, "session_id", None), + node_id=source_node_id or getattr(invocation, "id", None), + user_id=user_id, + ) + finally: + if stage_path is not None: + stage_path.unlink(missing_ok=True) + + +def cleanup_remote_media( + *, + client: RemoteInvokeClient, + image_names: list[str], + video_names: list[str], + services: Any, + reason: str, +) -> None: + """Best-effort cleanup after the primary no longer needs worker media.""" + deleted_images = 0 + deleted_videos = 0 + + for remote_name in image_names: + try: + client.delete_image(remote_name) + deleted_images += 1 + except Exception as exc: + services.logger.warning(f"Remote Workers: {reason}; remote image cleanup failed for {remote_name}: {exc}") + + for remote_name in video_names: + try: + client.delete_video(remote_name) + deleted_videos += 1 + except Exception as exc: + services.logger.warning(f"Remote Workers: {reason}; remote video cleanup failed for {remote_name}: {exc}") + + services.logger.info( + f"Remote Workers: remote cleanup deleted {deleted_images}/{len(image_names)} image(s), " + f"{deleted_videos}/{len(video_names)} video(s)" + ) diff --git a/invokeai/app/invocations/remote_worker/remote_nodes.py b/invokeai/app/invocations/remote_worker/remote_nodes.py new file mode 100644 index 00000000000..9c46d6ae9ab --- /dev/null +++ b/invokeai/app/invocations/remote_worker/remote_nodes.py @@ -0,0 +1,718 @@ +import json +import time +from typing import Any, Literal + +from invokeai.app.invocations.remote_worker.diffusers_transfer import TemporaryDirectoryModelServer +from invokeai.app.invocations.remote_worker.early_dispatch import AUTOMATIC_REMOTE_WORKER_NODE_TYPE +from invokeai.app.invocations.remote_worker.model_transfer import ( + ModelTransferError, + TemporaryModelServer, + model_layout_signature, + resolve_local_model_file, +) +from invokeai.app.invocations.remote_worker.model_transfer_state import ( + another_generation_needs_model, + register_model_transfer, + unregister_model_transfer, +) +from invokeai.app.invocations.remote_worker.remote_client import RemoteConfig, RemoteInvokeClient, RemoteInvokeError +from invokeai.app.services.session_processor.session_processor_common import CanceledException +from invokeai.invocation_api import ( + BaseInvocation, + BaseInvocationOutput, + InputField, + InvocationContext, + OutputField, + invocation, + invocation_output, +) + +# Strip the internal dispatch helper before the graph is sent to a worker. +_HELPER_NODE_TYPES = {AUTOMATIC_REMOTE_WORKER_NODE_TYPE} + + +def _remote_client(remote_url: str, user_id: str = "") -> RemoteInvokeClient: + return RemoteInvokeClient(RemoteConfig.from_environment(base_url=remote_url, verify_ssl=False, user_id=user_id)) + + +def _emit_model_transfer_progress( + *, + context: InvocationContext, + invocation: Any, + remote_index: int, + model: Any, + directory: bool, + phase: str, + job: dict[str, Any] | None = None, + error: str = "", +) -> None: + """Send transfer-only progress to the queue item's owner; never expose LAN URLs or credentials.""" + queue_item = context._data.queue_item + backend_item_id = int(getattr(queue_item, "item_id", 0) or 0) + if backend_item_id < 1: + return + job = job or {} + payload = { + "backend_item_id": backend_item_id, + "remote_index": remote_index, + "model_hash": model.hash, + "name": model.name, + "directory": directory, + "phase": phase, + "bytes": max(0, int(job.get("bytes") or 0)), + "total_bytes": max(0, int(job.get("total_bytes") or 0)), + } + if error: + payload["error"] = str(error)[:2048] + + try: + from invokeai.app.services.events.events_common import InvocationProgressEvent + + source_id = queue_item.session.prepared_source_mapping.get(invocation.id, invocation.id) + context._services.events.dispatch( + InvocationProgressEvent( + queue_id=queue_item.queue_id, + item_id=queue_item.item_id, + batch_id=queue_item.batch_id, + origin=queue_item.origin, + destination=queue_item.destination, + user_id=queue_item.user_id, + session_id=queue_item.session_id, + invocation=invocation.get_event_invocation(), + invocation_source_id=source_id, + message="[[IRW_MODEL_TRANSFER]]" + json.dumps(payload, separators=(",", ":")), + percentage=None, + ) + ) + except Exception as exc: + # A failed UI notification must never fail a model transfer. + context.logger.debug(f"Could not emit remote model transfer progress: {exc}") + + +class RemoteModelTransferCancelled(CanceledException): + """The local generation was canceled while preparing a worker model.""" + + +def _remote_model_matches_local( + remote_client: RemoteInvokeClient, + model: Any, + candidate: dict[str, Any], + expected_layout: tuple[str, str] | None = None, +) -> bool: + """True only when a same-hash remote record also has a compatible on-disk layout.""" + payload = candidate.get("model") if isinstance(candidate.get("model"), dict) else candidate + if str(payload.get("hash") or "").strip() != model.hash: + return False + + # A single model file has no subdirectory layout to validate; the content hash + # already identifies its bytes. + if model.path.is_file(): + return True + + key = str(payload.get("key") or "").strip() + if not key: + return False + local_kind, local_signature = expected_layout or model_layout_signature(model.path) + try: + remote_layout = remote_client.get_model_layout(key) + except RemoteInvokeError as exc: + if "HTTP 404" in str(exc): + return False + raise + return ( + str(remote_layout.get("kind") or "") == local_kind + and str(remote_layout.get("signature") or "") == local_signature + ) + + +def _find_compatible_remote_model(remote_client: RemoteInvokeClient, model: Any) -> dict[str, Any] | None: + """Find a remote copy matching the model hash and, for directories, its on-disk layout.""" + if model.path.is_file(): + return remote_client.get_model_by_hash(model.hash) + + expected_layout = model_layout_signature(model.path) + for candidate in remote_client.list_models(): + payload = candidate.get("model") if isinstance(candidate.get("model"), dict) else candidate + if str(payload.get("hash") or "").strip() != model.hash: + continue + if _remote_model_matches_local(remote_client, model, candidate, expected_layout): + return candidate + return None + + +def _transfer_missing_model_to_remote( + *, + context: InvocationContext, + invocation: Any, + remote_client: RemoteInvokeClient, + remote_index: int, + identifier: dict[str, Any], + transfer_host: str, + timeout_seconds: int, +) -> None: + """Transfer one missing model, observing native queue cancellation throughout preparation and installation.""" + try: + model = resolve_local_model_file(context._services, identifier) + except ModelTransferError as exc: + raise RemoteInvokeError(str(exc)) from exc + + existing = _find_compatible_remote_model(remote_client, model) + if existing is not None: + return + + transfer_id, transfer = register_model_transfer(remote_client.config.base_url, model.hash) + is_directory = model.path.is_dir() + job_id: int | None = None + status = "" + cancel_sent = False + cancel_notified = False + cancel_started: float | None = None + cancel_error_logged = False + lock_acquired = False + + def cancelled() -> bool: + if transfer.cancel_requested.is_set(): + return True + try: + item = context._services.session_queue.get_queue_item(int(context._data.queue_item.item_id)) + except Exception: + return False + is_cancelled = str(getattr(item.status, "value", item.status)).lower() in {"canceled", "cancelled"} + if is_cancelled: + transfer.cancel_requested.set() + return is_cancelled + + def notify_cancelled() -> None: + nonlocal cancel_notified + if cancel_notified: + return + cancel_notified = True + _emit_model_transfer_progress( + context=context, + invocation=invocation, + remote_index=remote_index, + model=model, + directory=is_directory, + phase="cancelled", + ) + + def preparation_should_cancel() -> bool: + if not cancelled(): + return False + notify_cancelled() + return not another_generation_needs_model(transfer) + + def check_cancellation() -> None: + nonlocal cancel_started, cancel_sent, cancel_error_logged + if not cancelled(): + return + notify_cancelled() + if job_id is None: + # Before a worker-side install exists, this caller still owns the + # singleflight preparation. Keep preparing for another live waiter. + if another_generation_needs_model(transfer): + return + raise RemoteModelTransferCancelled("Remote model transfer cancelled") + if status in {"completed", "error", "cancelled", "canceled", "failed"}: + raise RemoteModelTransferCancelled("Remote model transfer cancelled") + # A second live generation may be using the SAME worker-side install. + # Keep this primary HTTP server alive until that job reaches a terminal state. + if another_generation_needs_model(transfer): + return + if cancel_started is None: + cancel_started = time.monotonic() + # Do not interrupt InvokeAI while it moves/registers an already-downloaded model. + if status in {"running", "installing"}: + if time.monotonic() - cancel_started > 90: + raise RemoteModelTransferCancelled("Remote model transfer cancelled during installation") + return + if not cancel_sent: + path = ( + f"/api/v1/remote_workers/diffusers/install/{job_id}" + if is_directory + else f"/api/v2/models/install/{job_id}" + ) + try: + cancel_install = getattr(remote_client, "cancel_model_install", None) + if callable(cancel_install): + cancel_install(job_id, directory=is_directory) + else: + remote_client._request("DELETE", path) + except Exception as exc: + if not cancel_error_logged: + context.logger.warning( + f"Remote #{remote_index}: model install job {job_id} cancellation request failed: {exc}" + ) + cancel_error_logged = True + else: + cancel_sent = True + # A missing/unresponsive worker must not keep the primary server alive indefinitely. + if time.monotonic() - cancel_started > 30: + context.logger.warning( + f"Remote #{remote_index}: model install job {job_id} cancellation could not be confirmed" + ) + raise RemoteModelTransferCancelled("Remote model transfer cancellation not confirmed") + + def report_install_progress(job: dict[str, Any]) -> None: + if cancelled(): + return + phase = { + "waiting": "waiting", + "downloading": "downloading", + "downloads_done": "downloading", + "running": "installing", + "installing": "installing", + "completed": "verifying", + "error": "failed", + "cancelled": "cancelled", + "canceled": "cancelled", + }.get(status, "waiting") + _emit_model_transfer_progress( + context=context, + invocation=invocation, + remote_index=remote_index, + model=model, + directory=is_directory, + phase=phase, + job=job, + error=str(job.get("error") or "") if phase == "failed" else "", + ) + + try: + while not lock_acquired: + if cancelled(): + notify_cancelled() + raise RemoteModelTransferCancelled("Remote model transfer cancelled") + lock_acquired = transfer.shared_lock.acquire(timeout=0.25) + check_cancellation() + if _find_compatible_remote_model(remote_client, model) is not None: + return + check_cancellation() + _emit_model_transfer_progress( + context=context, + invocation=invocation, + remote_index=remote_index, + model=model, + directory=is_directory, + phase="preparing", + ) + server = ( + TemporaryDirectoryModelServer( + path=model.path, + remote_url=remote_client.config.base_url, + advertise_host=transfer_host, + should_cancel=preparation_should_cancel, + ) + if is_directory + else TemporaryModelServer( + model=model, remote_url=remote_client.config.base_url, advertise_host=transfer_host + ) + ) + with server: + check_cancellation() + size = ( + sum(file.size for file in server.files) + if isinstance(server, TemporaryDirectoryModelServer) + else model.path.stat().st_size + ) + context.logger.warning( + f"Remote #{remote_index}: required model '{model.name}' is missing; " + f"serving {model.path.name} ({size / (1024**3):.2f} GiB) directly from this InvokeAI host over the LAN" + ) + context.logger.info(f"Remote #{remote_index}: temporary model transfer endpoint ready at {server.url}") + if is_directory: + assert isinstance(server, TemporaryDirectoryModelServer) + job = remote_client.install_directory_from_manifest( + server.manifest(name=model.name, model_hash=model.hash) + ) + else: + job = remote_client.install_model_from_url(server.url, name=model.name) + try: + job_id = int(job.get("id")) + except (TypeError, ValueError) as exc: + raise RemoteInvokeError(f"Remote model installer returned no usable job id: {job}") from exc + status = str(job.get("status") or "").lower() + context.logger.info( + f"Remote #{remote_index}: InvokeAI model install job {job_id} started for '{model.name}'" + ) + + started = time.monotonic() + while True: + check_cancellation() + try: + job = ( + remote_client.get_directory_install_job(job_id) + if is_directory + else remote_client.get_model_install_job(job_id) + ) + except RemoteInvokeError as exc: + if cancelled() and cancel_sent and "HTTP 404" in str(exc): + raise RemoteModelTransferCancelled("Remote model transfer cancelled") from exc + raise + status = str(job.get("status") or "").lower() + report_install_progress(job) + check_cancellation() + if status == "completed": + break + if status in {"error", "canceled", "cancelled"}: + detail = job.get("error") or job.get("error_type") or "unknown installation error" + raise RemoteInvokeError(f"Remote model install job {job_id} ended with status '{status}': {detail}") + if time.monotonic() - started > float(timeout_seconds): + raise RemoteInvokeError( + f"Remote model install job {job_id} timed out after {timeout_seconds:g} seconds" + ) + time.sleep(1.0) + context.logger.info(f"Remote #{remote_index}: model install job {job_id} completed with status {status}") + + check_cancellation() + installed = _find_compatible_remote_model(remote_client, model) + check_cancellation() + if installed is None: + raise RemoteInvokeError( + f"Remote #{remote_index}: '{model.name}' finished installing but hash {model.hash} " + "was not found in the remote model manager" + ) + _emit_model_transfer_progress( + context=context, + invocation=invocation, + remote_index=remote_index, + model=model, + directory=is_directory, + phase="completed", + ) + remote_payload = installed.get("model") if isinstance(installed.get("model"), dict) else installed + context.logger.info( + f"Remote #{remote_index}: verified transferred model '{model.name}' by hash/layout; " + f"remote key={remote_payload.get('key', 'unknown')}" + ) + except RemoteModelTransferCancelled: + notify_cancelled() + raise + except ModelTransferError as exc: + if cancelled(): + notify_cancelled() + raise RemoteModelTransferCancelled("Remote model transfer cancelled") from exc + _emit_model_transfer_progress( + context=context, + invocation=invocation, + remote_index=remote_index, + model=model, + directory=is_directory, + phase="failed", + error=str(exc), + ) + raise RemoteInvokeError(str(exc)) from exc + except Exception as exc: + if cancelled(): + notify_cancelled() + raise RemoteModelTransferCancelled("Remote model transfer cancelled") from exc + _emit_model_transfer_progress( + context=context, + invocation=invocation, + remote_index=remote_index, + model=model, + directory=is_directory, + phase="failed", + error=str(exc), + ) + raise + finally: + if lock_acquired: + transfer.shared_lock.release() + unregister_model_transfer(transfer_id) + + +def _strip_helper_nodes(graph: dict[str, Any]) -> list[str]: + nodes = graph.get("nodes") + if not isinstance(nodes, dict): + raise RemoteInvokeError("Current workflow graph does not contain a nodes object") + + removed = { + str(node_id) + for node_id, node in nodes.items() + if isinstance(node, dict) and str(node.get("type", "")) in _HELPER_NODE_TYPES + } + for node_id in removed: + nodes.pop(node_id, None) + + edges = graph.get("edges") + if isinstance(edges, list) and removed: + kept_edges = [] + for edge in edges: + if not isinstance(edge, dict): + kept_edges.append(edge) + continue + source = edge.get("source") if isinstance(edge.get("source"), dict) else {} + destination = edge.get("destination") if isinstance(edge.get("destination"), dict) else {} + if str(source.get("node_id")) in removed or str(destination.get("node_id")) in removed: + continue + kept_edges.append(edge) + graph["edges"] = kept_edges + + if not nodes: + raise RemoteInvokeError("Nothing remains after removing the internal Remote Worker dispatch helper.") + return sorted(removed) + + +def _disable_graph_cache(graph: dict[str, Any]) -> int: + nodes = graph.get("nodes") + if not isinstance(nodes, dict): + return 0 + count = 0 + for node in nodes.values(): + if isinstance(node, dict): + node["use_cache"] = False + count += 1 + return count + + +def _find_media_references(value: Any, found: set[str]) -> None: + if isinstance(value, list): + for item in value: + _find_media_references(item, found) + return + if not isinstance(value, dict): + return + image_name = value.get("image_name") + video_name = value.get("video_name") + if isinstance(image_name, str) and image_name: + found.add(f"image:{image_name}") + if isinstance(video_name, str) and video_name: + found.add(f"video:{video_name}") + for child in value.values(): + _find_media_references(child, found) + + +def _graph_media_references(graph: dict[str, Any]) -> list[str]: + found: set[str] = set() + nodes = graph.get("nodes") + if isinstance(nodes, dict): + for node in nodes.values(): + _find_media_references(node, found) + return sorted(found) + + +def _remap_graph_image_names(value: Any, mapped: dict[str, str]) -> int: + """Rewrite image references (including nested ImageField/list inputs), not other strings.""" + changed = 0 + if isinstance(value, list): + for entry in value: + changed += _remap_graph_image_names(entry, mapped) + elif isinstance(value, dict): + original = value.get("image_name") + if isinstance(original, str) and original in mapped: + value["image_name"] = mapped[original] + changed += 1 + for entry in value.values(): + changed += _remap_graph_image_names(entry, mapped) + return changed + + +def _remap_graph_video_names(value: Any, mapped: dict[str, str]) -> int: + """Rewrite video references (including nested VideoField/list inputs), not other strings.""" + changed = 0 + if isinstance(value, list): + for entry in value: + changed += _remap_graph_video_names(entry, mapped) + elif isinstance(value, dict): + original = value.get("video_name") + if isinstance(original, str) and original in mapped: + value["video_name"] = mapped[original] + changed += 1 + for entry in value.values(): + changed += _remap_graph_video_names(entry, mapped) + return changed + + +def _transfer_source_videos_to_remote( + *, + context: InvocationContext, + remote_client: RemoteInvokeClient, + graph: dict[str, Any], + video_names: list[str], + remote_index: int, + uploaded_names: list[str], +) -> None: + """Copy every distinct local video once per worker before enqueueing the graph.""" + mapped: dict[str, str] = {} + for local_name in video_names: + try: + # InvocationContext performs the authenticated queue owner's read-access check. + local_path = context.videos.get_path(local_name) + except Exception as exc: + raise RemoteInvokeError( + f"Remote #{remote_index}: cannot read primary source video '{local_name}': {exc}" + ) from exc + try: + mapped[local_name] = remote_client.upload_input_video(local_path) + uploaded_names.append(mapped[local_name]) + except Exception as exc: + raise RemoteInvokeError( + f"Remote #{remote_index}: could not transfer source video '{local_name}': {exc}" + ) from exc + context.logger.debug( + f"Remote #{remote_index}: transferred input video '{local_name}' -> '{mapped[local_name]}'" + ) + + changed = _remap_graph_video_names(graph.get("nodes", {}), mapped) + context.logger.info( + f"Remote #{remote_index}: remapped {changed} video field(s) from {len(mapped)} transferred source video(s)" + ) + + +def _transfer_source_images_to_remote( + *, + context: InvocationContext, + remote_client: RemoteInvokeClient, + graph: dict[str, Any], + image_names: list[str], + remote_index: int, + uploaded_names: list[str], +) -> None: + """Copy every distinct local image once per worker before enqueueing the graph.""" + mapped: dict[str, str] = {} + for local_name in image_names: + try: + # InvocationContext checks access for the authenticated queue owner. + local_image = context.images.get_pil(local_name) + except Exception as exc: + raise RemoteInvokeError( + f"Remote #{remote_index}: cannot read primary source image '{local_name}': {exc}" + ) from exc + try: + mapped[local_name] = remote_client.upload_input_image(local_image) + uploaded_names.append(mapped[local_name]) + except Exception as exc: + raise RemoteInvokeError( + f"Remote #{remote_index}: could not transfer source image '{local_name}': {exc}" + ) from exc + context.logger.debug( + f"Remote #{remote_index}: transferred input image '{local_name}' -> '{mapped[local_name]}'" + ) + # Never change the local source graph; `graph` is a per-worker deepcopy. + changed = _remap_graph_image_names(graph.get("nodes", {}), mapped) + context.logger.info( + f"Remote #{remote_index}: remapped {changed} image field(s) from {len(mapped)} transferred source image(s)" + ) + + +def _strip_remote_board_assignments(graph: dict[str, Any]) -> list[str]: + nodes = graph.get("nodes") + if not isinstance(nodes, dict): + return [] + removed_from: list[str] = [] + for node_id, node in nodes.items(): + if not isinstance(node, dict): + continue + changed = False + if isinstance(node.get("board"), dict): + node["board"] = None + changed = True + if "board_id" in node and node.get("board_id") is not None: + node["board_id"] = None + changed = True + if changed: + removed_from.append(str(node_id)) + return removed_from + + +@invocation_output("irw_builtin_remote_worker_dispatch_output") +class RemoteWorkerDispatchOutput(BaseInvocationOutput): + started: bool = OutputField(description="Whether the backend Remote Worker pool was started for this queue item.") + + +# IMPORTANT: the Python class name intentionally begins with AAA. InvokeAI-7 +# groups ready nodes by Python class name and, absent ready_order, selects classes +# alphabetically. This lets the internal Remote Worker helper run before normal +# workflow classes when Local wins the dequeue race. +@invocation( + AUTOMATIC_REMOTE_WORKER_NODE_TYPE, + title="Remote Workers - Built-in Worker Pool", + tags=["remote", "invokeai", "worker", "parallel", "workflow"], + category="Remote Invoke", + version="1.0.0", + use_cache=False, +) +class AAARemoteWorkerDispatchInvocation(BaseInvocation): + """Internal helper that connects a normal InvokeAI queue item to the backend Remote Worker pool.""" + + result_destination: Literal["gallery", "canvas"] = InputField( + default="gallery", + ui_hidden=True, + description="Internal: captured InvokeAI result destination (Gallery or Canvas).", + ) + local_gallery_board_id: str = InputField( + default="", + ui_hidden=True, + description="Internal: Gallery board captured when this workflow was queued.", + ) + remote_url: str = InputField( + default="", + ui_hidden=True, + description="Internal: primary remote InvokeAI URL.", + ) + additional_remote_urls: str = InputField( + default="", + ui_hidden=True, + description="Internal: additional remote InvokeAI URLs.", + ) + dispatch_mode: Literal["Distributed", "Remote Only"] = InputField( + default="Distributed", + ui_hidden=True, + description="Internal automatic worker-pool mode.", + ) + remote_worker_names: str = InputField( + default="[]", + ui_hidden=True, + description="Internal JSON list of user-defined worker names aligned with the configured URLs.", + ) + keep_remote_copies: bool = InputField( + default=False, + ui_hidden=True, + description="Keep generated media on remote workers after successful import.", + ) + auto_transfer_missing_models: bool = InputField( + default=True, + ui_hidden=True, + description="Transfer supported missing models to a worker before rendering.", + ) + model_transfer_host: str = InputField( + default="", + ui_hidden=True, + description="Optional primary LAN host used by workers during model transfer.", + ) + model_transfer_timeout_seconds: int = InputField( + default=7200, + ge=60, + le=86400, + ui_hidden=True, + description="Maximum time to wait for a remote model transfer/install job.", + ) + collector_poll_interval_seconds: float = InputField( + default=0.75, + ge=0.25, + le=30.0, + ui_hidden=True, + description="How often worker-pool lanes poll remote progress and status.", + ) + collector_timeout_seconds: int = InputField( + default=14400, + ge=10, + le=86400, + ui_hidden=True, + description="Maximum wait per remote queued/rendering phase.", + ) + + def invoke(self, context: InvocationContext) -> RemoteWorkerDispatchOutput: + from invokeai.app.invocations.remote_worker.worker_pool import ( + ensure_remote_worker_pool, + run_current_remote_only, + ) + + item_id = int(context._data.queue_item.item_id) + if self.dispatch_mode == "Remote Only": + run_current_remote_only(context, self) + else: + ensure_remote_worker_pool(context._services, item_id) + + return RemoteWorkerDispatchOutput(started=True) diff --git a/invokeai/app/invocations/remote_worker/worker_pool.py b/invokeai/app/invocations/remote_worker/worker_pool.py new file mode 100644 index 00000000000..a4c22959b7a --- /dev/null +++ b/invokeai/app/invocations/remote_worker/worker_pool.py @@ -0,0 +1,1306 @@ +from __future__ import annotations + +import json +import re +import threading +import time +import uuid +from copy import deepcopy +from dataclasses import dataclass +from typing import Any, Literal + +from invokeai.app.invocations.primitives import ImageOutput, VideoOutput +from invokeai.app.invocations.remote_worker.remote_client import RemoteInvokeError +from invokeai.app.services.shared.invocation_context import InvocationContextData, build_invocation_context + +_AUTOMATIC_NODE_ID = "__irw_remote_worker_dispatch__" +_AUTOMATIC_NODE_TYPE = "irw_builtin_remote_worker_dispatch" + +_POOL_LOCK = threading.Lock() +_POOLS: dict[tuple[str, str], "_RemoteWorkerPool"] = {} +_SLOT_LOCKS_GUARD = threading.Lock() +_SLOT_LOCKS: dict[str, threading.Lock] = {} +_LOCAL_HANDOFFS_GUARD = threading.Lock() +_LOCAL_HANDOFF_ITEMS: set[int] = set() + + +@dataclass(frozen=True) +class WorkerSpec: + url: str + name: str + slot: int + + +@dataclass(frozen=True) +class PoolSettings: + mode: Literal["Distributed", "Remote Only"] + workers: tuple[WorkerSpec, ...] + result_destination: Literal["gallery", "canvas"] + local_gallery_board_id: str + keep_remote_copies: bool + auto_transfer_missing_models: bool + model_transfer_host: str + model_transfer_timeout_seconds: int + poll_interval_seconds: float + timeout_seconds: int + + +def _node_value(node: Any, name: str, default: Any = None) -> Any: + if isinstance(node, dict): + return node.get(name, default) + return getattr(node, name, default) + + +def _node_type(node: Any) -> str: + if isinstance(node, dict): + return str(node.get("type") or "") + get_type = getattr(node, "get_type", None) + if callable(get_type): + try: + return str(get_type()) + except Exception: + pass + return str(getattr(node, "type", "") or "") + + +def _helper_for_queue_item(queue_item: Any) -> Any | None: + try: + nodes = queue_item.session.graph.nodes + except Exception: + return None + if not isinstance(nodes, dict): + return None + direct = nodes.get(_AUTOMATIC_NODE_ID) + if direct is not None and _node_type(direct) == _AUTOMATIC_NODE_TYPE: + return direct + for node in nodes.values(): + if _node_type(node) == _AUTOMATIC_NODE_TYPE: + return node + return None + + +def _split_urls(primary: str, additional: str) -> list[str]: + raw: list[str] = [] + if primary.strip(): + raw.append(primary.strip()) + raw.extend(part.strip() for part in re.split(r"[,;\n\r]+", additional or "") if part.strip()) + + seen: set[str] = set() + result: list[str] = [] + for value in raw: + normalized = value.rstrip("/") + key = normalized.casefold() + if normalized and key not in seen: + seen.add(key) + result.append(normalized) + return result + + +def _settings_from_helper(helper: Any) -> PoolSettings | None: + mode = str(_node_value(helper, "dispatch_mode", "Distributed")) + if mode not in {"Distributed", "Remote Only"}: + return None + + urls = _split_urls( + str(_node_value(helper, "remote_url", "") or ""), + str(_node_value(helper, "additional_remote_urls", "") or ""), + ) + if not urls: + return None + + try: + decoded = json.loads(str(_node_value(helper, "remote_worker_names", "[]") or "[]")) + except Exception: + decoded = [] + names = decoded if isinstance(decoded, list) else [] + + workers = tuple( + WorkerSpec( + url=url, + name=( + str(names[index]).strip() if index < len(names) and str(names[index]).strip() else f"Remote {index + 1}" + ), + slot=index + 1, + ) + for index, url in enumerate(urls) + ) + + destination = str(_node_value(helper, "result_destination", "gallery")) + if destination not in {"gallery", "canvas"}: + destination = "gallery" + + return PoolSettings( + mode=mode, # type: ignore[arg-type] + workers=workers, + result_destination=destination, # type: ignore[arg-type] + local_gallery_board_id=str(_node_value(helper, "local_gallery_board_id", "") or ""), + keep_remote_copies=bool(_node_value(helper, "keep_remote_copies", False)), + auto_transfer_missing_models=bool(_node_value(helper, "auto_transfer_missing_models", True)), + model_transfer_host=str(_node_value(helper, "model_transfer_host", "") or ""), + model_transfer_timeout_seconds=max( + 60, int(_node_value(helper, "model_transfer_timeout_seconds", 7200) or 7200) + ), + poll_interval_seconds=max(0.25, float(_node_value(helper, "collector_poll_interval_seconds", 0.75) or 0.75)), + timeout_seconds=max(10, int(_node_value(helper, "collector_timeout_seconds", 14400) or 14400)), + ) + + +def _settings_for_queue_item(queue_item: Any) -> PoolSettings | None: + helper = _helper_for_queue_item(queue_item) + return _settings_from_helper(helper) if helper is not None else None + + +def _worker_allowed(settings: PoolSettings, worker: WorkerSpec) -> bool: + return any(candidate.url.casefold() == worker.url.casefold() for candidate in settings.workers) + + +def _slot_lock(worker: WorkerSpec) -> threading.Lock: + key = worker.url.rstrip("/").casefold() + with _SLOT_LOCKS_GUARD: + lock = _SLOT_LOCKS.get(key) + if lock is None: + lock = threading.Lock() + _SLOT_LOCKS[key] = lock + return lock + + +def _register_local_handoff(item_id: int) -> None: + with _LOCAL_HANDOFFS_GUARD: + _LOCAL_HANDOFF_ITEMS.add(int(item_id)) + + +def _unregister_local_handoff(item_id: int) -> None: + with _LOCAL_HANDOFFS_GUARD: + _LOCAL_HANDOFF_ITEMS.discard(int(item_id)) + + +def _is_local_handoff(item_id: int) -> bool: + with _LOCAL_HANDOFFS_GUARD: + return int(item_id) in _LOCAL_HANDOFF_ITEMS + + +def _queue_item_ids( + services: Any, + *, + queue_id: str, + user_id: str, + statuses: tuple[str, ...], +) -> list[int]: + if not statuses: + return [] + + db = getattr(services.session_queue, "_db", None) + if db is None: + raise RuntimeError("Remote worker pool requires InvokeAI's SQLite session queue") + + placeholders = ", ".join("?" for _ in statuses) + with db.transaction() as cursor: + cursor.execute( + f"""--sql + SELECT item_id + FROM session_queue + WHERE queue_id = ? + AND user_id = ? + AND status IN ({placeholders}) + AND parent_item_id IS NULL + ORDER BY priority DESC, item_id ASC + """, + (queue_id, user_id, *statuses), + ) + return [int(row[0]) for row in cursor.fetchall()] + + +def _status(services: Any, item_id: int) -> str: + try: + return str(services.session_queue.get_queue_item(item_id).status or "").lower() + except Exception: + return "" + + +def _park_remote_only_items(services: Any, queue_id: str, user_id: str) -> None: + queue = services.session_queue + dequeue_lock = getattr(queue, "_dequeue_lock", None) + set_status = getattr(queue, "_set_queue_item_status", None) + if dequeue_lock is None or set_status is None: + raise RuntimeError("Remote worker pool requires queue claim primitives") + + for item_id in _queue_item_ids( + services, + queue_id=queue_id, + user_id=user_id, + statuses=("pending",), + ): + try: + item = queue.get_queue_item(item_id) + except Exception: + continue + + settings = _settings_for_queue_item(item) + if settings is None or settings.mode != "Remote Only": + continue + + with dequeue_lock: + try: + fresh = queue.get_queue_item(item_id) + except Exception: + continue + fresh_settings = _settings_for_queue_item(fresh) + if fresh.status != "pending" or fresh_settings is None or fresh_settings.mode != "Remote Only": + continue + set_status(item_id=item_id, status="waiting", queue_item=fresh) + + +def _claim_for_worker( + services: Any, + *, + queue_id: str, + user_id: str, + worker: WorkerSpec, + excluded: set[int], +) -> tuple[Any, PoolSettings] | None: + queue = services.session_queue + dequeue_lock = getattr(queue, "_dequeue_lock", None) + set_status = getattr(queue, "_set_queue_item_status", None) + if dequeue_lock is None or set_status is None: + raise RuntimeError("Remote worker pool requires queue claim primitives") + + for item_id in _queue_item_ids( + services, + queue_id=queue_id, + user_id=user_id, + statuses=("pending", "waiting"), + ): + if _is_local_handoff(item_id): + continue + if item_id in excluded: + continue + + try: + candidate = queue.get_queue_item(item_id) + except Exception: + continue + settings = _settings_for_queue_item(candidate) + if settings is None or not _worker_allowed(settings, worker): + continue + if settings.mode == "Distributed" and candidate.status != "pending": + continue + if settings.mode == "Remote Only" and candidate.status not in {"pending", "waiting"}: + continue + + with dequeue_lock: + try: + fresh = queue.get_queue_item(item_id) + except Exception: + continue + fresh_settings = _settings_for_queue_item(fresh) + if fresh_settings is None or not _worker_allowed(fresh_settings, worker): + continue + if fresh_settings.mode == "Distributed" and fresh.status != "pending": + continue + if fresh_settings.mode == "Remote Only" and fresh.status not in {"pending", "waiting"}: + continue + + claimed = set_status( + item_id=item_id, + status="in_progress", + device=f"remote:{worker.name}", + queue_item=fresh, + ) + if claimed.status == "in_progress": + return claimed, fresh_settings + + return None + + +def _has_eligible_item( + services: Any, + *, + queue_id: str, + user_id: str, + worker: WorkerSpec, +) -> bool: + for item_id in _queue_item_ids( + services, + queue_id=queue_id, + user_id=user_id, + statuses=("pending", "waiting"), + ): + try: + item = services.session_queue.get_queue_item(item_id) + except Exception: + continue + settings = _settings_for_queue_item(item) + if settings is not None and _worker_allowed(settings, worker): + return True + return False + + +def _requeue(services: Any, item_id: int, settings: PoolSettings, reason: str) -> None: + if _status(services, item_id) != "in_progress": + return + + target = "waiting" if settings.mode == "Remote Only" else "pending" + set_status = getattr(services.session_queue, "_set_queue_item_status", None) + if set_status is None: + raise RuntimeError("Remote worker pool requires queue status transitions") + set_status(item_id=item_id, status=target) + services.logger.warning(f"Remote Workers: returned local item {item_id} to {target}: {reason}") + + +def _is_remote_oom_error(errors: Any) -> bool: + """Return True only for clear remote device-memory exhaustion failures.""" + try: + detail = json.dumps(errors, ensure_ascii=False).casefold() + except Exception: + detail = str(errors).casefold() + return "outofmemoryerror" in detail or "cuda out of memory" in detail or "would exceed allowed memory" in detail + + +def _fail_remote_oom(services: Any, queue_item: Any, worker: WorkerSpec, errors: Any) -> None: + item_id = int(queue_item.item_id) + if _status(services, item_id) not in {"in_progress", "waiting", "pending"}: + return + + detail = json.dumps(errors, ensure_ascii=False)[:4000] + services.session_queue.fail_queue_item( + item_id=item_id, + error_type="OutOfMemoryError", + error_message=f"{worker.name} ran out of memory: {detail}", + error_traceback="", + ) + services.logger.warning( + f"Remote Workers [{worker.name}]: remote item failed with an out-of-memory error; " + f"marked local item {item_id} failed instead of requeueing it" + ) + + +def _event_item(queue_item: Any, invocation: Any) -> Any: + item = queue_item.model_copy(deep=True) + invocation_id = str(invocation.id) + item.session.prepared_source_mapping.setdefault(invocation_id, invocation_id) + return item + + +def _emit_started(services: Any, queue_item: Any, invocation: Any, worker: WorkerSpec) -> Any: + item = _event_item(queue_item, invocation) + services.events.emit_invocation_started(queue_item=item, invocation=invocation) + services.events.emit_invocation_progress( + queue_item=item, + invocation=invocation, + message=f"{worker.name} · Preparing", + percentage=None, + image=None, + ) + return item + + +def _emit_progress( + services: Any, + queue_item: Any, + invocation: Any, + worker: WorkerSpec, + *, + message: str, + percentage: float | None = None, + image: Any = None, + revision: int | None = None, +) -> None: + safe = str(message or "Rendering").replace("\n", " ").strip() + services.events.emit_invocation_progress( + queue_item=queue_item, + invocation=invocation, + message=f"{worker.name} · {safe}", + percentage=percentage, + image=image, + revision=revision, + ) + + +def _emit_result( + services: Any, + queue_item: Any, + event_item: Any, + invocation: Any, + source_id: str, + output: Any, +) -> None: + synthetic_id = str(uuid.uuid4()) + synthetic = invocation.model_copy(update={"id": synthetic_id}) + source_id = str(source_id or invocation.id) + + # Persist remote results under the REAL source node id. Do not add the + # synthetic event id to prepared_source_mapping: GraphExecutionState requires + # every key in that mapping to exist in execution_graph. + # + # If the same source produces more than one imported output, keep the prior + # value under an unmapped synthetic result key and leave the newest/final + # output at source_id. Unfiltered history still retains both, while consumers + # filtering for canvas_output/video_output naturally get the final result. + previous = queue_item.session.results.get(source_id) + if previous is not None: + queue_item.session.results[synthetic_id] = previous + queue_item.session.results[source_id] = output + + # event_item is an ephemeral deep copy used only for live events. Mapping the + # synthetic event invocation here preserves live source routing without ever + # persisting a fake execution node. + event_item.session.prepared_source_mapping[synthetic_id] = source_id + event_item.session.results[synthetic_id] = output + + services.events.emit_invocation_started(queue_item=event_item, invocation=synthetic) + services.events.emit_invocation_complete(queue_item=event_item, invocation=synthetic, output=output) + + +def _persist_remote_results( + services: Any, + queue_item: Any, + worker: WorkerSpec, + *, + phase: str, +) -> None: + try: + services.session_queue.save_queue_item_session(int(queue_item.item_id), queue_item.session) + except Exception as exc: + services.logger.warning( + f"Remote Workers [{worker.name}]: completed item {queue_item.item_id}, " + f"but could not persist imported result history {phase}: {exc}" + ) + + +def _capture_local_output_board(graph: dict[str, Any]) -> str | None: + """Capture only board assignments from nodes that save visible output media.""" + nodes = graph.get("nodes") + if not isinstance(nodes, dict): + return None + + for node in nodes.values(): + if not isinstance(node, dict) or node.get("is_intermediate") is not False: + continue + + board = node.get("board") + if isinstance(board, dict): + value = board.get("board_id") or board.get("id") + if isinstance(value, str) and value.strip(): + return value.strip() + + value = node.get("board_id") + if isinstance(value, str) and value.strip(): + return value.strip() + + return None + + +def _build_remote_graph( + queue_item: Any, + settings: PoolSettings, + services: Any, +) -> tuple[dict[str, Any], str]: + from invokeai.app.invocations.remote_worker.model_transfer import enrich_model_identifier_hashes + from invokeai.app.invocations.remote_worker.remote_nodes import ( + _disable_graph_cache, + _strip_helper_nodes, + _strip_remote_board_assignments, + ) + + try: + dumped = queue_item.model_dump(mode="json") + source_graph = dumped["session"]["graph"] + except Exception as exc: + raise RemoteInvokeError("Could not read executable graph from queue item") from exc + + if not isinstance(source_graph, dict) or not isinstance(source_graph.get("nodes"), dict): + raise RemoteInvokeError("Queue item did not contain a usable executable graph") + + remote_graph = deepcopy(source_graph) + remote_graph["id"] = str(uuid.uuid4()) + + explicit_board = _capture_local_output_board(source_graph) + _strip_helper_nodes(remote_graph) + _strip_remote_board_assignments(remote_graph) + _disable_graph_cache(remote_graph) + enrich_model_identifier_hashes(remote_graph, services) + + board_id = explicit_board or settings.local_gallery_board_id.strip() + if board_id.lower() == "none" or settings.result_destination == "canvas": + board_id = "" + + return remote_graph, board_id + + +def _build_context(services: Any, queue_item: Any, invocation: Any): + item_id = int(queue_item.item_id) + + def canceled() -> bool: + return _status(services, item_id) in {"canceled", "cancelled", "failed", "completed"} + + return build_invocation_context( + services=services, + data=InvocationContextData( + queue_item=queue_item, + invocation=invocation, + source_invocation_id=invocation.id, + ), + is_canceled=canceled, + ) + + +def _cleanup_remote_inputs( + client: Any, + services: Any, + settings: PoolSettings, + image_names: list[str], + video_names: list[str], + *, + reason: str, +) -> None: + if settings.keep_remote_copies or (not image_names and not video_names): + return + + from invokeai.app.invocations.remote_worker.remote_media import cleanup_remote_media + + cleanup_remote_media( + client=client, + image_names=list(dict.fromkeys(image_names)), + video_names=list(dict.fromkeys(video_names)), + services=services, + reason=reason, + ) + + +def _dispatch_remote( + services: Any, + queue_item: Any, + invocation: Any, + settings: PoolSettings, + worker: WorkerSpec, +) -> tuple[Any, int, str, list[str], list[str]]: + from invokeai.app.invocations.remote_worker.model_transfer import resolve_local_model_file + from invokeai.app.invocations.remote_worker.remote_nodes import ( + RemoteModelTransferCancelled, + _graph_media_references, + _remote_client, + _remote_model_matches_local, + _transfer_missing_model_to_remote, + _transfer_source_images_to_remote, + _transfer_source_videos_to_remote, + ) + + context = _build_context(services, queue_item, invocation) + graph, board_id = _build_remote_graph(queue_item, settings, services) + client = _remote_client(worker.url, str(queue_item.user_id)) + + # Probe reachability/authentication before expensive preparation. + client.get_current_item() + + missing_handler = None + if settings.auto_transfer_missing_models: + + def missing_handler(identifier: dict[str, Any]) -> None: + _transfer_missing_model_to_remote( + context=context, + invocation=invocation, + remote_client=client, + remote_index=worker.slot, + identifier=identifier, + transfer_host=settings.model_transfer_host, + timeout_seconds=settings.model_transfer_timeout_seconds, + ) + + local_models: dict[str, Any] = {} + + def model_match_validator(identifier: dict[str, Any], candidate: dict[str, Any]) -> bool: + local_key = str(identifier.get("key") or "") + model = local_models.get(local_key) + if model is None: + model = resolve_local_model_file(context._services, identifier) + local_models[local_key] = model + return _remote_model_matches_local(client, model, candidate) + + try: + client.remap_model_identifiers( + graph, + missing_model_handler=missing_handler, + model_match_validator=model_match_validator, + ) + except RemoteModelTransferCancelled: + raise + + media_refs = _graph_media_references(graph) + image_names = [ref[6:] for ref in media_refs if ref.startswith("image:")] + video_names = [ref[6:] for ref in media_refs if ref.startswith("video:")] + uploaded_image_names: list[str] = [] + uploaded_video_names: list[str] = [] + + try: + if image_names: + _transfer_source_images_to_remote( + context=context, + remote_client=client, + graph=graph, + image_names=image_names, + remote_index=worker.slot, + uploaded_names=uploaded_image_names, + ) + + if video_names: + _transfer_source_videos_to_remote( + context=context, + remote_client=client, + graph=graph, + video_names=video_names, + remote_index=worker.slot, + uploaded_names=uploaded_video_names, + ) + + origin = str(getattr(queue_item, "origin", "") or "invokeai") + remote_item_id = client.enqueue_graph( + graph=graph, + queue_id="default", + origin=f"{origin}:remote-worker:{worker.slot}", + ) + except Exception: + _cleanup_remote_inputs( + client, + services, + settings, + uploaded_image_names, + uploaded_video_names, + reason="remote dispatch failed before queueing", + ) + raise + + return client, int(remote_item_id), board_id, uploaded_image_names, uploaded_video_names + + +def _remote_media_source_ids(completed_item: dict[str, Any]) -> tuple[dict[str, str], dict[str, str]]: + """Map worker media names back to their original graph source node ids.""" + session = completed_item.get("session") + if not isinstance(session, dict): + return {}, {} + + results = session.get("results") + if not isinstance(results, dict): + return {}, {} + + raw_mapping = session.get("prepared_source_mapping") + prepared_source_mapping = raw_mapping if isinstance(raw_mapping, dict) else {} + image_sources: dict[str, str] = {} + video_sources: dict[str, str] = {} + + for result_id, result in results.items(): + if not isinstance(result, dict): + continue + + result_key = str(result_id) + mapped = prepared_source_mapping.get(result_key) + source_id = str(mapped) if isinstance(mapped, str) and mapped else result_key + + image = result.get("image") + if isinstance(image, dict): + image_name = image.get("image_name") + if isinstance(image_name, str) and image_name: + image_sources[image_name] = source_id + + images = result.get("images") + if isinstance(images, list): + for entry in images: + if not isinstance(entry, dict): + continue + image_name = entry.get("image_name") + if isinstance(image_name, str) and image_name: + image_sources[image_name] = source_id + + video_candidates: list[Any] = [result.get("video")] + videos = result.get("videos") + if isinstance(videos, list): + video_candidates.extend(videos) + video_candidates.append(result) + + for entry in video_candidates: + if not isinstance(entry, dict): + continue + video_name = entry.get("video_name") + if isinstance(video_name, str) and video_name: + video_sources[video_name] = source_id + + return image_sources, video_sources + + +def _import_completed( + services: Any, + queue_item: Any, + invocation: Any, + settings: PoolSettings, + client: Any, + completed_item: dict[str, Any], + board_id: str, +) -> tuple[list[Any], list[Any]]: + from invokeai.app.invocations.remote_worker.remote_media import ( + cleanup_remote_media, + save_local_image, + save_local_video, + ) + + image_source_ids, video_source_ids = _remote_media_source_ids(completed_item) + all_images = client.extract_image_names( + completed_item, + non_intermediate_only=False, + allow_empty=True, + ) + all_videos = client.extract_video_names(completed_item) + images = client.filter_gallery_image_names(all_images) + videos = client.filter_gallery_video_names(all_videos) + + if settings.result_destination == "canvas": + if not images and all_images: + images = [all_images[-1]] + if not videos and all_videos: + videos = [all_videos[-1]] + + if not images and not videos: + raise RemoteInvokeError("Remote render completed without an importable image or video output") + + image_dtos: list[Any] = [] + video_dtos: list[Any] = [] + + for remote_name in images: + image_metadata = client.get_image_metadata(remote_name) + image = client.download_image(remote_name) + image_dtos.append( + save_local_image( + services=services, + queue_item=queue_item, + invocation=invocation, + image=image, + metadata=image_metadata, + board_id=board_id, + result_destination=settings.result_destination, + source_node_id=image_source_ids.get(remote_name), + ) + ) + + for remote_name in videos: + video_metadata = client.get_video_metadata(remote_name) + video_bytes = client.download_video(remote_name) + video_dtos.append( + save_local_video( + services=services, + queue_item=queue_item, + invocation=invocation, + video_bytes=video_bytes, + metadata=video_metadata, + board_id=board_id, + result_destination=settings.result_destination, + source_node_id=video_source_ids.get(remote_name), + ) + ) + + if not settings.keep_remote_copies: + cleanup_remote_media( + client=client, + image_names=all_images, + video_names=all_videos, + services=services, + reason="local import succeeded", + ) + + return image_dtos, video_dtos + + +def _cancel_remote( + client: Any, + remote_item_id: int, + services: Any, + worker: WorkerSpec, +) -> bool: + try: + client.cancel_queue_item(remote_item_id, "default") + except Exception as exc: + services.logger.warning(f"Remote Workers [{worker.name}]: cancellation request failed: {exc}") + + for _ in range(61): + try: + item = client.get_item(remote_item_id, "default") + except Exception: + return False + + status = str(item.get("status") or "").lower() + if status in {"completed", "canceled", "cancelled", "failed"}: + if status != "failed": + try: + client.delete_queue_item(remote_item_id, "default") + except Exception as exc: + services.logger.warning( + f"Remote Workers [{worker.name}]: could not delete canceled remote queue item " + f"{remote_item_id}: {exc}" + ) + return True + + time.sleep(0.5) + + return False + + +def _run_remote_job( + services: Any, + queue_item: Any, + settings: PoolSettings, + worker: WorkerSpec, + *, + complete_local: bool, +) -> str: + invocation = _helper_for_queue_item(queue_item) + if invocation is None: + raise RemoteInvokeError("Remote worker helper node was not found") + + client, remote_item_id, board_id, input_image_names, input_video_names = _dispatch_remote( + services, + queue_item, + invocation, + settings, + worker, + ) + event_item = _emit_started(services, queue_item, invocation, worker) + + services.logger.info( + f"Remote Workers [{worker.name}]: local item {queue_item.item_id} -> remote item {remote_item_id}" + ) + + started = time.monotonic() + last_signature: tuple[Any, Any, Any] | None = None + last_network_log = 0.0 + + while True: + local_status = _status(services, int(queue_item.item_id)) + if local_status in {"canceled", "cancelled"}: + remote_stopped = _cancel_remote(client, remote_item_id, services, worker) + if remote_stopped: + _cleanup_remote_inputs( + client, + services, + settings, + input_image_names, + input_video_names, + reason="remote job canceled", + ) + return "canceled" + + if time.monotonic() - started > settings.timeout_seconds: + remote_stopped = _cancel_remote(client, remote_item_id, services, worker) + if remote_stopped: + _cleanup_remote_inputs( + client, + services, + settings, + input_image_names, + input_video_names, + reason="remote job timed out", + ) + raise RemoteInvokeError( + f"{worker.name} remote item {remote_item_id} timed out after {settings.timeout_seconds:g}s" + ) + + try: + item = client.get_item(remote_item_id, "default") + except Exception as exc: + now = time.monotonic() + if now - last_network_log >= 10: + services.logger.warning( + f"Remote Workers [{worker.name}]: temporarily unavailable while monitoring " + f"item {remote_item_id}; will retry: {exc}" + ) + last_network_log = now + time.sleep(settings.poll_interval_seconds) + continue + + remote_status = str(item.get("status") or "").lower() + + if remote_status == "completed": + try: + image_dtos, video_dtos = _import_completed( + services, + queue_item, + invocation, + settings, + client, + item, + board_id, + ) + finally: + _cleanup_remote_inputs( + client, + services, + settings, + input_image_names, + input_video_names, + reason="remote job completed", + ) + + for dto in image_dtos: + _emit_result( + services, + queue_item, + event_item, + invocation, + source_id=str(getattr(dto, "node_id", None) or invocation.id), + output=ImageOutput.build(dto), + ) + for dto in video_dtos: + _emit_result( + services, + queue_item, + event_item, + invocation, + source_id=str(getattr(dto, "node_id", None) or invocation.id), + output=VideoOutput.build(dto), + ) + + _emit_progress( + services, + event_item, + invocation, + worker, + message="Complete", + percentage=1.0, + ) + + try: + client.delete_queue_item(remote_item_id, "default") + except Exception as exc: + services.logger.warning( + f"Remote Workers [{worker.name}]: imported result but could not delete " + f"remote queue item {remote_item_id}: {exc}" + ) + + if complete_local and _status(services, int(queue_item.item_id)) == "in_progress": + # Persist before the terminal event: the frontend immediately fetches this + # queue row on completion and filters results by prepared_source_mapping. + _persist_remote_results(services, queue_item, worker, phase="before completion") + services.session_queue.complete_queue_item(int(queue_item.item_id)) + # Keep the previous post-completion write as a best-effort safeguard. + _persist_remote_results(services, queue_item, worker, phase="after completion") + + return "completed" + + if remote_status in {"failed", "canceled", "cancelled"}: + _cleanup_remote_inputs( + client, + services, + settings, + input_image_names, + input_video_names, + reason=f"remote job ended with status {remote_status}", + ) + errors = item.get("session", {}).get("errors", {}) if isinstance(item.get("session"), dict) else {} + if remote_status == "failed" and _is_remote_oom_error(errors): + _fail_remote_oom(services, queue_item, worker, errors) + raise RemoteInvokeError( + f"{worker.name} remote item {remote_item_id} ended with status " + f"'{remote_status}': {json.dumps(errors)[:2000]}" + ) + + if remote_status in {"in_progress", "running"}: + try: + preview = client.get_progress_preview(remote_item_id, "default") + except Exception: + preview = None + + if preview: + signature = ( + preview.get("revision"), + preview.get("percentage"), + preview.get("message"), + ) + if signature != last_signature: + from invokeai.app.invocations.remote_worker.remote_media import progress_image_from_remote + + try: + percentage = float(preview["percentage"]) if preview.get("percentage") is not None else None + except (TypeError, ValueError): + percentage = None + + try: + revision = int(preview["revision"]) if preview.get("revision") is not None else None + except (TypeError, ValueError): + revision = None + + _emit_progress( + services, + event_item, + invocation, + worker, + message=str(preview.get("message") or "Rendering"), + percentage=percentage, + image=progress_image_from_remote(preview), + revision=revision, + ) + last_signature = signature + elif last_signature is None: + _emit_progress( + services, + event_item, + invocation, + worker, + message="Rendering", + ) + + time.sleep(settings.poll_interval_seconds) + + +def _mark_remaining_local_nodes_skipped(context: Any) -> None: + session = context._data.queue_item.session + current_exec_id = str(context._data.invocation.id) + current_source_id = str(session.prepared_source_mapping.get(current_exec_id, current_exec_id)) + + for exec_id in list(session.execution_graph.nodes.keys()): + exec_id = str(exec_id) + if exec_id == current_exec_id: + continue + session.executed.add(exec_id) + try: + session._set_prepared_exec_state(exec_id, "skipped") + except Exception: + pass + try: + session._remove_from_ready_queues(exec_id) + except Exception: + pass + + for source_id in list(session.graph.nodes.keys()): + source_id = str(source_id) + if source_id == current_source_id: + continue + session.executed.add(source_id) + if source_id not in session.executed_history: + session.executed_history.append(source_id) + + +class _RemoteWorkerPool: + def __init__(self, services: Any, queue_id: str, user_id: str) -> None: + self.services = services + self.queue_id = queue_id + self.user_id = user_id + self._guard = threading.Lock() + self._lanes: dict[str, threading.Thread] = {} + + def ensure_workers(self, workers: tuple[WorkerSpec, ...]) -> None: + with self._guard: + for worker in workers: + key = worker.url.casefold() + thread = self._lanes.get(key) + if thread is not None and thread.is_alive(): + continue + + thread = threading.Thread( + target=self._run_lane, + args=(worker,), + daemon=True, + name=f"invokeai-remote-pool-{worker.name}", + ) + self._lanes[key] = thread + thread.start() + + def _run_lane(self, worker: WorkerSpec) -> None: + excluded: set[int] = set() + idle_checks = 0 + last_unavailable_log = 0.0 + + while True: + _park_remote_only_items(self.services, self.queue_id, self.user_id) + + if not _has_eligible_item( + self.services, + queue_id=self.queue_id, + user_id=self.user_id, + worker=worker, + ): + idle_checks += 1 + if idle_checks >= 4: + return + time.sleep(0.25) + continue + + idle_checks = 0 + lock = _slot_lock(worker) + if not lock.acquire(blocking=False): + time.sleep(0.05) + continue + + try: + # Do not claim the real local row until this remote is reachable. + try: + from invokeai.app.invocations.remote_worker.remote_nodes import _remote_client + + _remote_client(worker.url, self.user_id).get_current_item() + except Exception as exc: + now = time.monotonic() + if last_unavailable_log == 0.0 or now - last_unavailable_log >= 30.0: + self.services.logger.warning( + f"Remote Workers [{worker.name}]: unavailable; " + f"will retry while eligible work remains: {exc}" + ) + last_unavailable_log = now + time.sleep(2.0) + continue + + claimed = _claim_for_worker( + self.services, + queue_id=self.queue_id, + user_id=self.user_id, + worker=worker, + excluded=excluded, + ) + if claimed is None: + time.sleep(0.05) + continue + + queue_item, settings = claimed + try: + outcome = _run_remote_job( + self.services, + queue_item, + settings, + worker, + complete_local=True, + ) + except Exception as exc: + excluded.add(int(queue_item.item_id)) + _requeue( + self.services, + int(queue_item.item_id), + settings, + f"{worker.name} could not complete it ({exc})", + ) + return + + if outcome == "canceled": + continue + finally: + lock.release() + + +def ensure_remote_worker_pool(services: Any, item_id: int) -> bool: + try: + item = services.session_queue.get_queue_item(int(item_id)) + except Exception: + return False + + settings = _settings_for_queue_item(item) + if settings is None: + return False + + key = (str(item.queue_id), str(item.user_id)) + with _POOL_LOCK: + pool = _POOLS.get(key) + if pool is None: + pool = _RemoteWorkerPool(services, key[0], key[1]) + _POOLS[key] = pool + + pool.ensure_workers(settings.workers) + return True + + +def schedule_remote_worker_pool(item_id: int, services: Any) -> None: + # Starting lanes is CPU/network-only and must not block the enqueue API. + threading.Thread( + target=ensure_remote_worker_pool, + args=(services, int(item_id)), + daemon=True, + name=f"invokeai-remote-pool-start-{item_id}", + ).start() + + +def run_current_remote_only(context: Any, invocation: Any) -> None: + queue_item = context._data.queue_item + item_id = int(queue_item.item_id) + settings = _settings_from_helper(invocation) + if settings is None or settings.mode != "Remote Only": + raise RemoteInvokeError("Remote Only worker settings are invalid") + + set_status = getattr(context._services.session_queue, "_set_queue_item_status", None) + if set_status is None: + raise RemoteInvokeError("Remote Only requires queue status transitions") + + _register_local_handoff(item_id) + try: + if _status(context._services, item_id) == "in_progress": + set_status( + item_id=item_id, + status="waiting", + queue_item=queue_item, + ) + context.logger.info(f"Remote Workers: Remote Only local item {item_id} is waiting for a free remote worker") + + ensure_remote_worker_pool(context._services, item_id) + + last_error: Exception | None = None + attempted: set[str] = set() + + while not context.util.is_canceled(): + made_attempt = False + + for worker in settings.workers: + key = worker.url.casefold() + if key in attempted: + continue + + lock = _slot_lock(worker) + if not lock.acquire(blocking=False): + continue + + made_attempt = True + attempted.add(key) + try: + if _status(context._services, item_id) == "waiting": + set_status( + item_id=item_id, + status="in_progress", + device=f"remote:{worker.name}", + queue_item=queue_item, + ) + + try: + outcome = _run_remote_job( + context._services, + queue_item, + settings, + worker, + complete_local=False, + ) + except Exception as exc: + last_error = exc + if not context.util.is_canceled() and _status(context._services, item_id) == "in_progress": + set_status( + item_id=item_id, + status="waiting", + queue_item=queue_item, + ) + context.logger.warning( + f"Remote Workers [{worker.name}]: could not accept Local-held Remote Only item: {exc}" + ) + continue + + if outcome == "completed": + _mark_remaining_local_nodes_skipped(context) + context.util.signal_progress("Remote Only completed", 1.0) + return + if outcome == "canceled": + return + finally: + lock.release() + + if len(attempted) >= len(settings.workers): + break + + if not made_attempt: + time.sleep(0.1) + + if context.util.is_canceled(): + return + if last_error is not None: + raise RemoteInvokeError( + f"Remote Only completed with no successful remote worker: {last_error}" + ) from last_error + raise RemoteInvokeError("Remote Only completed with no available remote worker") + finally: + _unregister_local_handoff(item_id) diff --git a/invokeai/frontend/webv2/performance/architecture-baseline.json b/invokeai/frontend/webv2/performance/architecture-baseline.json index 5dc3c67c054..3055ee961b6 100644 --- a/invokeai/frontend/webv2/performance/architecture-baseline.json +++ b/invokeai/frontend/webv2/performance/architecture-baseline.json @@ -2,15 +2,15 @@ "build": { "launchpad": { "baseline": { - "brotliBytes": 893586, + "brotliBytes": 893888, "cssRawBytes": 2159, "fontRawBytes": 219480, - "gzipBytes": 892942, + "gzipBytes": 893144, "imageRawBytes": 156376, - "initialRawBytes": 2118609, + "initialRawBytes": 2119771, "largestAssetRawBytes": 749035, "otherAssetRawBytes": 0, - "ownedRawBytes": 93082, + "ownedRawBytes": 93037, "requestCount": 36, "scriptRequestCount": 20, "sourceOwners": [ @@ -334,17 +334,17 @@ }, "editor": { "baseline": { - "brotliBytes": 1243382, + "brotliBytes": 1246113, "cssRawBytes": 2159, "fontRawBytes": 219480, - "gzipBytes": 1237668, + "gzipBytes": 1240880, "imageRawBytes": 156376, - "initialRawBytes": 3203516, + "initialRawBytes": 3215988, "largestAssetRawBytes": 749035, "otherAssetRawBytes": 0, - "ownedRawBytes": 144128, - "requestCount": 74, - "scriptRequestCount": 58, + "ownedRawBytes": 144054, + "requestCount": 73, + "scriptRequestCount": 57, "sourceOwners": [ "package:@ark-ui/react", "package:@chakra-ui/react", @@ -642,6 +642,10 @@ "source:src/features/queue/data/progressStore.ts", "source:src/features/queue/data/queries.ts", "source:src/features/queue/data/realtimeRuntime.ts", + "source:src/features/queue/data/remoteWorkersGraph.ts", + "source:src/features/queue/data/remoteWorkersGraphContract.ts", + "source:src/features/queue/data/remoteWorkersHealth.ts", + "source:src/features/queue/data/remoteWorkersStore.ts", "source:src/features/queue/data/serverApi.ts", "source:src/features/queue/data/submissionApi.ts", "source:src/features/queue/devices.ts", @@ -653,6 +657,7 @@ "source:src/features/queue/runtime.ts", "source:src/features/queue/runtime/coordinator.ts", "source:src/features/queue/runtime/receiptAcknowledgements.ts", + "source:src/features/queue/runtime/remoteModelTransferToasts.ts", "source:src/features/queue/ui/QueueUiContext.tsx", "source:src/features/queue/ui/currentBatchItems.ts", "source:src/features/queue/ui/queueConfirmationStore.ts", @@ -1053,6 +1058,7 @@ "source:src/workbench/widgets/project/manifest.ts", "source:src/workbench/widgets/queue-status/manifest.ts", "source:src/workbench/widgets/queue/manifest.ts", + "source:src/workbench/widgets/remote-workers/manifest.ts", "source:src/workbench/widgets/server-status/manifest.ts", "source:src/workbench/widgets/upscale/manifest.ts", "source:src/workbench/widgets/video/manifest.ts", diff --git a/invokeai/frontend/webv2/public/locales/en.json b/invokeai/frontend/webv2/public/locales/en.json index 025608e4354..c793e8e759d 100644 --- a/invokeai/frontend/webv2/public/locales/en.json +++ b/invokeai/frontend/webv2/public/locales/en.json @@ -2179,7 +2179,75 @@ "color": "Color", "swatches": "Swatches", "history": "History", - "overview": "Overview" + "overview": "Overview", + "remoteWorkers": "Remote Workers" + }, + "remoteWorkers": { + "title": "Distributed rendering", + "enabled": "Enabled", + "disabled": "Disabled", + "enable": "Enable distributed rendering", + "description": "Send new generations to Local, remote workers, or both. Worker logins are private to each InvokeAI user.", + "dispatch": { + "title": "Dispatch", + "distributed": "Distributed", + "remoteOnly": "Remote Only", + "distributedDescription": "Local and enabled remotes share the real InvokeAI queue. Whichever worker is free takes the next job.", + "remoteOnlyDescription": "Only enabled remotes consume these queue items. Local does not render them." + }, + "status": { + "disabled": "Disabled", + "paused": "Paused", + "online": "Online", + "offline": "Offline", + "loginRequired": "Login required", + "checking": "Checking" + }, + "worker": { + "settings": "Remote worker {{slot}} settings", + "disableForNewJobs": "Disable remote worker {{slot}} for new jobs", + "enableForNewJobs": "Enable remote worker {{slot}} for new jobs", + "collapseSettings": "Collapse remote worker {{slot}} settings", + "expandSettings": "Expand remote worker {{slot}} settings", + "nameLabel": "Remote worker {{slot}} name", + "defaultName": "Remote {{slot}}" + }, + "login": { + "checking": "Checking login", + "saved": "Login saved", + "none": "No saved login", + "description": "Optional. For this worker's multi-user login, enter your own account credentials.", + "emailPlaceholder": "Remote InvokeAI email", + "passwordPlaceholder": "Remote InvokeAI password", + "replacePasswordPlaceholder": "New password (to replace saved login)", + "save": "Save login", + "remove": "Remove login", + "savedMessage": "Login saved on this InvokeAI server.", + "removedMessage": "Saved login removed.", + "loadError": "Could not load saved login status", + "saveError": "Could not save worker login", + "removeError": "Could not remove worker login" + }, + "workers": { + "title": "Workers", + "configured_one": "{{count}} configured", + "configured_other": "{{count}} configured", + "none": "Add a valid http(s) worker URL to use distributed rendering.", + "editAddresses": "Edit worker addresses", + "doneEditingAddresses": "Done editing addresses", + "editorDescription": "One URL per line. Worker names can be changed in each worker's expanded settings.", + "urlsLabel": "Remote worker URLs", + "statusHelp": "Paused means no availability checks while distributed rendering is off. When enabled: green = online, red = offline, orange = login required. These badges are display-only; the backend verifies reachability before a worker claims a job. The power icon excludes a worker from new jobs without stopping running ones. Offline workers remain configured and can rejoin pending work when they recover. Passwords are encrypted on the primary instance; protect its runtime directory and use HTTPS across untrusted networks." + }, + "modelTransfer": { + "title": "Model transfer", + "transferMissing": "Transfer missing models", + "keepCopies": "Keep copies on remote workers", + "advanced": "Advanced model transfer settings", + "hostLabel": "Primary host address for model transfers (optional)", + "hostPlaceholder": "Auto-detect primary host LAN IP" + }, + "destinationHelp": "Results follow your selected destination: Gallery uses the board selected when you invoke; Canvas receives staging candidates you can accept." }, "generate": { "activeCount_one": "{{count}} active", diff --git a/invokeai/frontend/webv2/scripts/widget-sources.mjs b/invokeai/frontend/webv2/scripts/widget-sources.mjs index 25ffd628197..8e85c6724f7 100644 --- a/invokeai/frontend/webv2/scripts/widget-sources.mjs +++ b/invokeai/frontend/webv2/scripts/widget-sources.mjs @@ -13,6 +13,7 @@ export const WIDGET_SOURCES = new Map([ ['src/workbench/widgets/notifications/implementation.ts', 'notifications'], ['src/workbench/widgets/preview/implementation.ts', 'preview'], ['src/workbench/widgets/project/implementation.ts', 'project'], + ['src/workbench/widgets/remote-workers/implementation.ts', 'remote-workers'], ['src/workbench/widgets/queue-status/implementation.ts', 'queue-status'], ['src/workbench/widgets/server-status/implementation.ts', 'server-status'], ['src/features/gallery/widget.ts', 'gallery'], diff --git a/invokeai/frontend/webv2/src/app/GalleryUiAdapter.test.tsx b/invokeai/frontend/webv2/src/app/GalleryUiAdapter.test.tsx index 588bf60a1a6..c59e8421a31 100644 --- a/invokeai/frontend/webv2/src/app/GalleryUiAdapter.test.tsx +++ b/invokeai/frontend/webv2/src/app/GalleryUiAdapter.test.tsx @@ -14,6 +14,7 @@ let store: ReturnType; let adapter: GalleryUiAdapter; const noop = () => undefined; const livePreviewFollow = vi.fn(); +const livePreviewShowSaved = vi.fn(); const openWorkbenchWidget = vi.fn(); vi.mock('@features/gallery/react', () => ({ @@ -31,6 +32,8 @@ vi.mock('@workbench/widgets/preview/livePreviewFollow', () => ({ follow: livePreviewFollow, pin: vi.fn(), showAll: vi.fn(), + showSaved: livePreviewShowSaved, + viewingSaved: false, }), })); vi.mock('@workbench/projects/useProjectFileActions', () => ({ useExportLibraryProject: () => noop })); @@ -77,6 +80,36 @@ describe('Gallery live-follow adapter', () => { }); }); +describe('Gallery saved-media navigation', () => { + const image = { + kind: 'image', + name: 'existing.png', + boardId: 'none', + category: 'general', + createdAt: '2026-09-21T00:00:00.000Z', + fullUrl: '/existing.png', + height: 64, + isIntermediate: false, + starred: false, + thumbnailUrl: '/existing-thumb.png', + width: 64, + } as Parameters[0]; + + it('hands even a reselected Gallery image to Preview while live progress exists', () => { + const owner = renderAdapter(); + owner.gallery.selectItem(image); + owner.gallery.selectItem(image); + expect(livePreviewShowSaved).toHaveBeenCalledTimes(2); + }); + + it('hands range and toggle selections to Preview too', () => { + const owner = renderAdapter(); + owner.gallery.setItemMultiSelection(['image:existing.png'], image); + owner.gallery.toggleItemSelection(image, null); + expect(livePreviewShowSaved).toHaveBeenCalledTimes(2); + }); +}); + describe('Gallery settings adapter ownership', () => { it('updates the captured active project through the Gallery command', () => { const owner = renderAdapter(); diff --git a/invokeai/frontend/webv2/src/app/GalleryUiAdapter.tsx b/invokeai/frontend/webv2/src/app/GalleryUiAdapter.tsx index 2f2d7107c98..cfd3cff7983 100644 --- a/invokeai/frontend/webv2/src/app/GalleryUiAdapter.tsx +++ b/invokeai/frontend/webv2/src/app/GalleryUiAdapter.tsx @@ -55,6 +55,22 @@ export const GalleryUiAdapterProvider = ({ children }: { children: ReactNode }) exportProject, gallery: { ...gallery, + selectItem: (item) => { + livePreview.showSaved(); + gallery.selectItem(item); + }, + selectImage: (image) => { + livePreview.showSaved(); + gallery.selectImage(image); + }, + setItemMultiSelection: (itemKeys, primaryItem) => { + livePreview.showSaved(); + gallery.setItemMultiSelection(itemKeys, primaryItem); + }, + toggleItemSelection: (item, nextPrimaryItem) => { + livePreview.showSaved(); + gallery.toggleItemSelection(item, nextPrimaryItem); + }, updateSettings: (settings) => { if (isAccountScopeCurrent(accountScope) && queries.isActiveProject(projectId)) { gallery.updateSettings(settings, projectId); diff --git a/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraph.test.ts b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraph.test.ts new file mode 100644 index 00000000000..3eb2e9e188e --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraph.test.ts @@ -0,0 +1,180 @@ +import type { QueueBackendGraph } from '@features/queue/core/types'; + +import { beforeEach, describe, expect, it, vi } from 'vitest'; + +const mocks = vi.hoisted(() => ({ + getRemoteWorkerName: vi.fn(), + getRemoteWorkerUrls: vi.fn(), + getRemoteWorkersSettings: vi.fn(), + isRemoteWorkerEnabled: vi.fn(), +})); + +vi.mock('./remoteWorkersStore', () => ({ + getRemoteWorkerName: mocks.getRemoteWorkerName, + getRemoteWorkerUrls: mocks.getRemoteWorkerUrls, + getRemoteWorkersSettings: mocks.getRemoteWorkersSettings, + isRemoteWorkerEnabled: mocks.isRemoteWorkerEnabled, +})); + +import { applyRemoteWorkersToGraph } from './remoteWorkersGraph'; +import { AUTOMATIC_REMOTE_WORKER_NODE_ID, AUTOMATIC_REMOTE_WORKER_NODE_TYPE } from './remoteWorkersGraphContract'; + +const worker1 = 'http://worker-1:9090'; +const worker2 = 'http://worker-2:9090'; + +const makeGraph = (): QueueBackendGraph => ({ + id: 'graph-1', + edges: [ + { + source: { node_id: 'prompt', field: 'text' }, + destination: { node_id: 'denoise', field: 'positive_conditioning' }, + }, + ], + nodes: { + prompt: { id: 'prompt', type: 'string', text: 'hello' }, + denoise: { id: 'denoise', type: 'denoise' }, + }, +}); + +describe('applyRemoteWorkersToGraph', () => { + beforeEach(() => { + mocks.getRemoteWorkersSettings.mockReturnValue({ + enabled: true, + workerUrls: `${worker1}\n${worker2}`, + workerNames: {}, + disabledWorkerUrls: [], + dispatchMode: 'distributed', + autoTransferMissingModels: true, + modelTransferHost: ' primary-host ', + keepRemoteCopies: false, + }); + mocks.getRemoteWorkerUrls.mockReturnValue([worker1, worker2]); + mocks.isRemoteWorkerEnabled.mockReturnValue(true); + mocks.getRemoteWorkerName.mockImplementation((_url: string, index: number) => + index === 0 ? 'RTX5080' : 'Server GPU' + ); + }); + + it('leaves the graph untouched when Remote Workers are disabled', () => { + const graph = makeGraph(); + mocks.getRemoteWorkersSettings.mockReturnValue({ enabled: false }); + + expect(applyRemoteWorkersToGraph(graph, null, 'gallery')).toBe(graph); + }); + + it('does not overwrite a pre-existing node that uses the automatic helper id', () => { + const graph = makeGraph(); + graph.nodes[AUTOMATIC_REMOTE_WORKER_NODE_ID] = { + id: AUTOMATIC_REMOTE_WORKER_NODE_ID, + type: 'some_other_node', + }; + + expect(applyRemoteWorkersToGraph(graph, null, 'gallery')).toBe(graph); + }); + + it('injects one Distributed helper while preserving the executable graph', () => { + const graph = makeGraph(); + const result = applyRemoteWorkersToGraph(graph, 'board-1', 'gallery'); + const helper = result.nodes[AUTOMATIC_REMOTE_WORKER_NODE_ID]; + + expect(result).not.toBe(graph); + expect(result.edges).toEqual(graph.edges); + expect(result.nodes.prompt).toEqual(graph.nodes.prompt); + expect(helper).toMatchObject({ + id: AUTOMATIC_REMOTE_WORKER_NODE_ID, + type: AUTOMATIC_REMOTE_WORKER_NODE_TYPE, + remote_url: worker1, + additional_remote_urls: worker2, + remote_worker_names: JSON.stringify(['RTX5080', 'Server GPU']), + dispatch_mode: 'Distributed', + auto_transfer_missing_models: true, + model_transfer_host: 'primary-host', + keep_remote_copies: false, + result_destination: 'gallery', + local_gallery_board_id: 'board-1', + }); + expect(helper).not.toHaveProperty('source_graph_json'); + }); + + it('keeps the executable graph intact for Remote Only', () => { + const graph = makeGraph(); + mocks.getRemoteWorkersSettings.mockReturnValue({ + enabled: true, + workerUrls: `${worker1}\n${worker2}`, + workerNames: {}, + disabledWorkerUrls: [], + dispatchMode: 'remote_only', + autoTransferMissingModels: false, + modelTransferHost: '', + keepRemoteCopies: true, + }); + + const result = applyRemoteWorkersToGraph(graph, 'board-1', 'gallery'); + const helper = result.nodes[AUTOMATIC_REMOTE_WORKER_NODE_ID]; + + expect(result.edges).toEqual(graph.edges); + expect(result.nodes.prompt).toEqual(graph.nodes.prompt); + expect(helper).toMatchObject({ + dispatch_mode: 'Remote Only', + result_destination: 'gallery', + local_gallery_board_id: 'board-1', + keep_remote_copies: true, + }); + expect(helper).not.toHaveProperty('source_graph_json'); + }); + + it('never carries a Gallery board into a Canvas dispatch', () => { + const graph = makeGraph(); + const result = applyRemoteWorkersToGraph(graph, 'board-1', 'canvas'); + + expect(result.nodes[AUTOMATIC_REMOTE_WORKER_NODE_ID]?.local_gallery_board_id).toBe(''); + expect(result.nodes[AUTOMATIC_REMOTE_WORKER_NODE_ID]?.result_destination).toBe('canvas'); + }); + + it('excludes workers disabled with the power toggle without consulting browser health', () => { + const graph = makeGraph(); + mocks.isRemoteWorkerEnabled.mockImplementation((url: string) => url === worker2); + + const result = applyRemoteWorkersToGraph(graph, null, 'gallery'); + const helper = result.nodes[AUTOMATIC_REMOTE_WORKER_NODE_ID]; + + expect(helper).toMatchObject({ + remote_url: worker2, + additional_remote_urls: '', + remote_worker_names: JSON.stringify(['Server GPU']), + }); + }); + + it('leaves Distributed local-only when every configured worker is disabled', () => { + const graph = makeGraph(); + mocks.isRemoteWorkerEnabled.mockReturnValue(false); + + expect(applyRemoteWorkersToGraph(graph, null, 'gallery')).toBe(graph); + }); + + it('keeps a Remote Only helper when every configured worker is disabled', () => { + const graph = makeGraph(); + mocks.getRemoteWorkersSettings.mockReturnValue({ + enabled: true, + workerUrls: `${worker1}\n${worker2}`, + workerNames: {}, + disabledWorkerUrls: [worker1, worker2], + dispatchMode: 'remote_only', + autoTransferMissingModels: false, + modelTransferHost: '', + keepRemoteCopies: false, + }); + mocks.isRemoteWorkerEnabled.mockReturnValue(false); + + const result = applyRemoteWorkersToGraph(graph, null, 'gallery'); + const helper = result.nodes[AUTOMATIC_REMOTE_WORKER_NODE_ID]; + + expect(result).not.toBe(graph); + expect(helper).toMatchObject({ + dispatch_mode: 'Remote Only', + additional_remote_urls: '', + remote_worker_names: '[]', + }); + expect(helper?.remote_url).toBeUndefined(); + }); +}); diff --git a/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraph.ts b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraph.ts new file mode 100644 index 00000000000..d7d2f50902d --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraph.ts @@ -0,0 +1,64 @@ +import type { QueueBackendGraph, QueueResultDestination } from '@features/queue/core/types'; + +import { + AUTOMATIC_REMOTE_WORKER_NODE_ID, + AUTOMATIC_REMOTE_WORKER_NODE_TYPE, + hasRemoteWorkerDispatchNode, +} from './remoteWorkersGraphContract'; +import { + getRemoteWorkerName, + getRemoteWorkerUrls, + getRemoteWorkersSettings, + isRemoteWorkerEnabled, +} from './remoteWorkersStore'; + +/** Only the submitted graph is modified. The saved Canvas/workflow stays untouched. */ +export const applyRemoteWorkersToGraph = ( + graph: QueueBackendGraph, + galleryBoardId: string | null, + destination: QueueResultDestination +): QueueBackendGraph => { + const settings = getRemoteWorkersSettings(); + if (!settings.enabled || (destination !== 'gallery' && destination !== 'canvas')) { + return graph; + } + + const configuredUrls = getRemoteWorkerUrls(settings.workerUrls); + // Availability is backend-owned; browser health is display-only. + const urls = configuredUrls.filter(isRemoteWorkerEnabled); + if ((urls.length === 0 && settings.dispatchMode !== 'remote_only') || hasRemoteWorkerDispatchNode(graph)) { + return graph; + } + if (Object.hasOwn(graph.nodes, AUTOMATIC_REMOTE_WORKER_NODE_ID)) { + return graph; + } + + const names = urls.map((url) => getRemoteWorkerName(url, configuredUrls.indexOf(url))); + const [remoteUrl, ...additionalUrls] = urls; + + const kickoff = { + id: AUTOMATIC_REMOTE_WORKER_NODE_ID, + type: AUTOMATIC_REMOTE_WORKER_NODE_TYPE, + use_cache: false, + is_intermediate: true, + remote_url: remoteUrl, + additional_remote_urls: additionalUrls.join('\n'), + remote_worker_names: JSON.stringify(names), + dispatch_mode: settings.dispatchMode === 'remote_only' ? 'Remote Only' : 'Distributed', + auto_transfer_missing_models: settings.autoTransferMissingModels, + model_transfer_host: settings.modelTransferHost.trim(), + keep_remote_copies: settings.keepRemoteCopies, + result_destination: destination, + local_gallery_board_id: + destination === 'gallery' && galleryBoardId && galleryBoardId !== 'none' ? galleryBoardId : '', + }; + + // Both modes keep the real executable graph in the local queue row. Distributed + // lets Local race the backend worker pool for it. Remote Only is parked by the + // backend; if Local wins the enqueue race, the AAA helper hands it to a remote + // and skips the remaining local nodes after the remote completes. + return { + ...graph, + nodes: { ...graph.nodes, [AUTOMATIC_REMOTE_WORKER_NODE_ID]: kickoff }, + }; +}; diff --git a/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraphContract.ts b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraphContract.ts new file mode 100644 index 00000000000..69d71a7fafe --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersGraphContract.ts @@ -0,0 +1,7 @@ +import type { QueueBackendGraph } from '@features/queue/core/types'; + +export const AUTOMATIC_REMOTE_WORKER_NODE_ID = '__irw_remote_worker_dispatch__'; +export const AUTOMATIC_REMOTE_WORKER_NODE_TYPE = 'irw_builtin_remote_worker_dispatch'; + +export const hasRemoteWorkerDispatchNode = (graph: QueueBackendGraph): boolean => + Object.values(graph.nodes).some((node) => node.type === AUTOMATIC_REMOTE_WORKER_NODE_TYPE); diff --git a/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersHealth.test.ts b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersHealth.test.ts new file mode 100644 index 00000000000..0f6f8d7d5d5 --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersHealth.test.ts @@ -0,0 +1,147 @@ +import { accountLifecycle } from '@platform/state/accountLifecycle'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; + +const transport = vi.hoisted(() => ({ apiFetchJson: vi.fn() })); +vi.mock('@platform/transport/http', () => transport); + +import { + invalidateRemoteWorkerHealth, + refreshRemoteWorkerHealth, + remoteWorkersHealthStore, +} from './remoteWorkersHealth'; +import { + DEFAULT_REMOTE_WORKERS_SETTINGS, + getRemoteWorkerUrls, + getRemoteWorkersSettings, + isRemoteWorkerEnabled, + setRemoteWorkerEnabled, + setRemoteWorkersSettings, +} from './remoteWorkersStore'; + +const worker1 = 'http://192.168.1.101:9090'; +const worker2 = 'http://192.168.1.102:9090'; + +describe('Remote worker health', () => { + beforeEach(() => { + accountLifecycle.activate('worker-health-test-user'); + setRemoteWorkersSettings({ + disabledWorkerUrls: [], + enabled: true, + workerUrls: `${worker1}\n${worker2}`, + }); + remoteWorkersHealthStore.setSnapshot({ byUrl: {} }); + transport.apiFetchJson.mockReset(); + }); + + it('does not ping any worker while distributed rendering is disabled', async () => { + setRemoteWorkersSettings({ enabled: false }); + + await refreshRemoteWorkerHealth([worker1, worker2]); + + expect(transport.apiFetchJson).not.toHaveBeenCalled(); + expect(remoteWorkersHealthStore.getSnapshot().byUrl).toEqual({}); + }); + + it('does not ping a worker disabled with its power toggle', async () => { + setRemoteWorkerEnabled(worker1, false); + transport.apiFetchJson.mockResolvedValue({ status: 'online' }); + + await refreshRemoteWorkerHealth([worker1, worker2]); + + expect(transport.apiFetchJson).toHaveBeenCalledTimes(1); + expect(String(transport.apiFetchJson.mock.calls[0]?.[0])).toContain(encodeURIComponent(worker2)); + expect(remoteWorkersHealthStore.getSnapshot().byUrl[worker1]).toBeUndefined(); + expect(remoteWorkersHealthStore.getSnapshot().byUrl[worker2]?.status).toBe('online'); + }); + + it('ignores an in-flight result if the worker is disabled before it completes', async () => { + let finish!: (value: { status: string }) => void; + transport.apiFetchJson.mockImplementation( + () => + new Promise((resolve) => { + finish = resolve; + }) + ); + + const pending = refreshRemoteWorkerHealth([worker1]); + await vi.waitFor(() => expect(transport.apiFetchJson).toHaveBeenCalledTimes(1)); + setRemoteWorkerEnabled(worker1, false); + finish({ status: 'online' }); + await pending; + + expect(remoteWorkersHealthStore.getSnapshot().byUrl[worker1]?.status).not.toBe('online'); + }); + + it('reports an offline worker online after recovery', async () => { + transport.apiFetchJson.mockResolvedValueOnce({ status: 'offline' }).mockResolvedValueOnce({ status: 'online' }); + + await refreshRemoteWorkerHealth([worker1]); + expect(remoteWorkersHealthStore.getSnapshot().byUrl[worker1]?.status).toBe('offline'); + + remoteWorkersHealthStore.setSnapshot({ byUrl: { [worker1]: { status: 'offline', checkedAt: 0 } } }); + await refreshRemoteWorkerHealth([worker1]); + + expect(remoteWorkersHealthStore.getSnapshot().byUrl[worker1]?.status).toBe('online'); + }); + + it('stores the disabled choice per URL and retains it across the global toggle', () => { + setRemoteWorkerEnabled(worker1, false); + + expect(getRemoteWorkersSettings().disabledWorkerUrls).toEqual([worker1]); + expect(getRemoteWorkerUrls(`${worker2}\n${worker1}`).filter(isRemoteWorkerEnabled)).toEqual([worker2]); + + setRemoteWorkersSettings({ enabled: false }); + setRemoteWorkersSettings({ enabled: true }); + + expect(isRemoteWorkerEnabled(worker1)).toBe(false); + setRemoteWorkerEnabled(worker1, true); + expect(getRemoteWorkersSettings().disabledWorkerUrls).toEqual([]); + }); + + it('forces a fresh worker probe after credential health is invalidated', async () => { + remoteWorkersHealthStore.setSnapshot({ + byUrl: { [worker1]: { status: 'login_required', checkedAt: Date.now() } }, + }); + transport.apiFetchJson.mockResolvedValue({ status: 'online' }); + + await refreshRemoteWorkerHealth([worker1]); + expect(transport.apiFetchJson).not.toHaveBeenCalled(); + + invalidateRemoteWorkerHealth(worker1); + await refreshRemoteWorkerHealth([worker1]); + + expect(transport.apiFetchJson).toHaveBeenCalledTimes(1); + expect(remoteWorkersHealthStore.getSnapshot().byUrl[worker1]?.status).toBe('online'); + }); + + it('does not mark a worker online when its saved login is rejected', async () => { + transport.apiFetchJson.mockResolvedValue({ status: 'login_required' }); + + await refreshRemoteWorkerHealth([worker1]); + + expect(remoteWorkersHealthStore.getSnapshot().byUrl[worker1]?.status).toBe('login_required'); + }); + + it('discards results from a previous account', async () => { + let finish!: (value: { status: string }) => void; + transport.apiFetchJson.mockImplementation((path: string) => { + if (path === '/api/v1/remote_workers/settings') { + return Promise.resolve({ + ...DEFAULT_REMOTE_WORKERS_SETTINGS, + disabledWorkerUrls: [], + workerNames: {}, + }); + } + return new Promise((resolve) => { + finish = resolve; + }); + }); + + const pending = refreshRemoteWorkerHealth([worker1]); + accountLifecycle.activate('another-user'); + finish({ status: 'online' }); + await pending; + + expect(remoteWorkersHealthStore.getSnapshot().byUrl[worker1]).toBeUndefined(); + }); +}); diff --git a/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersHealth.ts b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersHealth.ts new file mode 100644 index 00000000000..bb31197b599 --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersHealth.ts @@ -0,0 +1,112 @@ +import { + captureAccountScope, + isAccountScopeCurrent, + registerAccountOwnedResource, +} from '@platform/state/accountLifecycle'; +import { createExternalStore } from '@platform/state/externalStore'; +import { apiFetchJson } from '@platform/transport/http'; + +import { getRemoteWorkersSettings, isRemoteWorkerEnabled } from './remoteWorkersStore'; + +export type RemoteWorkerStatus = 'checking' | 'online' | 'offline' | 'login_required'; +type RemoteWorkerHealth = { status: RemoteWorkerStatus; checkedAt: number }; + +const HEALTH_STALE_MS = 15_000; +export const remoteWorkersHealthStore = createExternalStore<{ byUrl: Record }>({ + byUrl: {}, +}); +const inFlight = new Map>(); +const healthRevisionByUrl = new Map(); + +registerAccountOwnedResource({ + name: 'remote-workers-health', + clear: () => { + inFlight.clear(); + healthRevisionByUrl.clear(); + remoteWorkersHealthStore.setSnapshot({ byUrl: {} }); + }, +}); + +/** + * Discards cached/in-flight health for one worker after its authentication changes. + * Any older probe may still finish at the transport layer, but its result cannot + * overwrite the next probe because the per-URL revision has changed. + */ +export const invalidateRemoteWorkerHealth = (url: string): void => { + healthRevisionByUrl.set(url, (healthRevisionByUrl.get(url) ?? 0) + 1); + inFlight.delete(url); + + const { byUrl } = remoteWorkersHealthStore.getSnapshot(); + if (!(url in byUrl)) { + return; + } + const next = { ...byUrl }; + delete next[url]; + remoteWorkersHealthStore.setSnapshot({ byUrl: next }); +}; + +/** Probes the primary's authenticated status endpoint; the browser never fetches a worker URL directly. */ +export const refreshRemoteWorkerHealth = async (urls: readonly string[]): Promise => { + const owner = captureAccountScope(); + if (!owner.accountId || !getRemoteWorkersSettings().enabled) { + return; + } + await Promise.all( + [...new Set(urls)].map(async (url) => { + // Settings can change between starting the batch and reaching this URL. + if (!getRemoteWorkersSettings().enabled || !isRemoteWorkerEnabled(url)) { + return; + } + const running = inFlight.get(url); + if (running) { + await running; + return; + } + const prior = remoteWorkersHealthStore.getSnapshot().byUrl[url]; + if (prior && prior.status !== 'checking' && Date.now() - prior.checkedAt < HEALTH_STALE_MS) { + return; + } + if (!prior) { + remoteWorkersHealthStore.setSnapshot({ + byUrl: { ...remoteWorkersHealthStore.getSnapshot().byUrl, [url]: { status: 'checking', checkedAt: 0 } }, + }); + } + const revision = healthRevisionByUrl.get(url) ?? 0; + const task = (async () => { + let status: RemoteWorkerStatus; + try { + if (!getRemoteWorkersSettings().enabled || !isRemoteWorkerEnabled(url)) { + return; + } + const result = await apiFetchJson<{ status: RemoteWorkerStatus }>( + `/api/v1/remote_workers/status?url=${encodeURIComponent(url)}`, + { signal: owner.signal } + ); + status = result.status === 'online' || result.status === 'login_required' ? result.status : 'offline'; + } catch { + status = 'offline'; + } + // An already-running request may finish after the user switches the feature off. + // Never mark such a worker online from a result received while disabled. + if ( + isAccountScopeCurrent(owner) && + getRemoteWorkersSettings().enabled && + isRemoteWorkerEnabled(url) && + (healthRevisionByUrl.get(url) ?? 0) === revision + ) { + remoteWorkersHealthStore.setSnapshot({ + byUrl: { ...remoteWorkersHealthStore.getSnapshot().byUrl, [url]: { status, checkedAt: Date.now() } }, + }); + } + })(); + inFlight.set(url, task); + try { + await task; + } finally { + if (inFlight.get(url) === task) { + inFlight.delete(url); + } + } + }) + ); +}; diff --git a/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersStore.test.ts b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersStore.test.ts new file mode 100644 index 00000000000..b1fda181268 --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersStore.test.ts @@ -0,0 +1,97 @@ +import { accountLifecycle } from '@platform/state/accountLifecycle'; +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; + +const mocks = vi.hoisted(() => ({ + apiFetchJson: vi.fn(), +})); + +vi.mock('@platform/transport/http', () => ({ + apiFetchJson: mocks.apiFetchJson, +})); + +import { + DEFAULT_REMOTE_WORKERS_SETTINGS, + getRemoteWorkerName, + getRemoteWorkerUrls, + getRemoteWorkersSettings, + loadRemoteWorkersSettings, + setRemoteWorkerName, + setRemoteWorkersSettings, +} from './remoteWorkersStore'; + +let resetCounter = 0; + +const activateSingleUser = (): void => { + accountLifecycle.activate(`remote-worker-store-reset-${++resetCounter}`); + accountLifecycle.activate('single-user'); +}; + +describe('remote worker server settings', () => { + beforeEach(async () => { + mocks.apiFetchJson.mockReset(); + mocks.apiFetchJson.mockImplementation((_path: string, init?: RequestInit) => { + if (init?.method === 'PUT') { + return Promise.resolve(JSON.parse(String(init.body))); + } + return Promise.resolve({ ...DEFAULT_REMOTE_WORKERS_SETTINGS, disabledWorkerUrls: [], workerNames: {} }); + }); + activateSingleUser(); + await loadRemoteWorkersSettings(); + }); + + afterEach(() => { + accountLifecycle.invalidate(); + }); + + it('loads settings from the authenticated primary server', async () => { + mocks.apiFetchJson.mockImplementation((_path: string, init?: RequestInit) => { + if (init?.method === 'PUT') { + return Promise.resolve(JSON.parse(String(init.body))); + } + return Promise.resolve({ + ...DEFAULT_REMOTE_WORKERS_SETTINGS, + enabled: true, + dispatchMode: 'remote_only', + workerUrls: 'https://example.test/invoke', + }); + }); + + await loadRemoteWorkersSettings(); + + expect(getRemoteWorkersSettings().dispatchMode).toBe('remote_only'); + expect(getRemoteWorkersSettings().workerUrls).toBe('https://example.test/invoke'); + }); + + it('does not depend on browser localStorage and strips embedded URL credentials', () => { + Object.defineProperty(globalThis, 'localStorage', { + configurable: true, + get: () => { + throw new Error('remote worker settings must not read browser storage'); + }, + }); + + setRemoteWorkersSettings({ + workerUrls: 'http://alice:secret@192.168.1.101:9090\nhttp://192.168.1.102:9090', + }); + + expect(getRemoteWorkersSettings().workerUrls).toBe('http://192.168.1.101:9090\nhttp://192.168.1.102:9090'); + expect(getRemoteWorkerUrls(getRemoteWorkersSettings().workerUrls)).toEqual([ + 'http://192.168.1.101:9090', + 'http://192.168.1.102:9090', + ]); + + Reflect.deleteProperty(globalThis, 'localStorage'); + }); + + it('stores a user-defined worker name by normalized URL', () => { + const url = 'http://192.168.1.101:9090'; + + setRemoteWorkerName(url, 'RTX5080'); + + expect(getRemoteWorkerName(url, 0)).toBe('RTX5080'); + expect(getRemoteWorkersSettings().workerNames[url]).toBe('RTX5080'); + + setRemoteWorkerName(url, ' '); + expect(getRemoteWorkerName(url, 0)).toBe('Remote 1'); + }); +}); diff --git a/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersStore.ts b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersStore.ts new file mode 100644 index 00000000000..163ac83b855 --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/data/remoteWorkersStore.ts @@ -0,0 +1,271 @@ +import { + captureAccountScope, + isAccountScopeCurrent, + registerAccountOwnedResource, + type AccountScope, +} from '@platform/state/accountLifecycle'; +import { createExternalStore } from '@platform/state/externalStore'; +import { apiFetchJson } from '@platform/transport/http'; + +export type RemoteDispatchMode = 'distributed' | 'remote_only'; + +export interface RemoteWorkersSettings { + enabled: boolean; + dispatchMode: RemoteDispatchMode; + workerUrls: string; + /** Optional user-defined display names keyed by normalized worker URL. */ + workerNames: Record; + /** Disabled by normalized URL so reordering workers never changes which worker is paused. */ + disabledWorkerUrls: string[]; + autoTransferMissingModels: boolean; + keepRemoteCopies: boolean; + modelTransferHost: string; +} + +export const DEFAULT_REMOTE_WORKERS_SETTINGS: RemoteWorkersSettings = { + enabled: false, + dispatchMode: 'distributed', + workerUrls: '', + workerNames: {}, + disabledWorkerUrls: [], + autoTransferMissingModels: true, + keepRemoteCopies: false, + modelTransferHost: '', +}; + +const makeDefaultSettings = (): RemoteWorkersSettings => ({ + ...DEFAULT_REMOTE_WORKERS_SETTINGS, + workerNames: {}, + disabledWorkerUrls: [], +}); + +const stripEmbeddedUrlCredentials = (raw: string): string => + raw.replace(/[^;,\r\n\s]+/g, (candidate) => { + let url: URL; + try { + url = new URL(candidate); + } catch { + return candidate; + } + if (!['http:', 'https:'].includes(url.protocol) || (!url.username && !url.password)) { + return candidate; + } + const hadTrailingSlash = candidate.endsWith('/'); + url.username = ''; + url.password = ''; + const sanitized = url.toString(); + return !hadTrailingSlash && url.pathname === '/' && !url.search && !url.hash + ? sanitized.replace(/\/$/, '') + : sanitized; + }); + +const normalizeSettings = (value: unknown): RemoteWorkersSettings => { + if (typeof value !== 'object' || value === null) { + return makeDefaultSettings(); + } + const entry = value as Record; + return { + enabled: entry.enabled === true, + dispatchMode: entry.dispatchMode === 'remote_only' ? 'remote_only' : 'distributed', + workerUrls: stripEmbeddedUrlCredentials(typeof entry.workerUrls === 'string' ? entry.workerUrls : ''), + workerNames: + typeof entry.workerNames === 'object' && entry.workerNames !== null + ? Object.fromEntries( + Object.entries(entry.workerNames as Record) + .filter((item): item is [string, string] => typeof item[1] === 'string') + .map(([url, name]) => [url.toLowerCase(), name]) + ) + : {}, + disabledWorkerUrls: Array.isArray(entry.disabledWorkerUrls) + ? entry.disabledWorkerUrls.filter((url): url is string => typeof url === 'string').map((url) => url.toLowerCase()) + : [], + autoTransferMissingModels: entry.autoTransferMissingModels !== false, + keepRemoteCopies: entry.keepRemoteCopies === true, + modelTransferHost: typeof entry.modelTransferHost === 'string' ? entry.modelTransferHost : '', + }; +}; + +export const remoteWorkersStore = createExternalStore(makeDefaultSettings()); + +let loadedEpoch: number | null = null; +let pendingBeforeLoad: Partial = {}; +let pendingSave: { owner: AccountScope; revision: number; settings: RemoteWorkersSettings } | undefined; +let saveTimer: ReturnType | undefined; +let settingsRevision = 0; + +const cancelPendingSave = (): void => { + if (saveTimer !== undefined) { + globalThis.clearTimeout(saveTimer); + } + saveTimer = undefined; + pendingSave = undefined; +}; + +const scheduleSave = (settings: RemoteWorkersSettings): void => { + const owner = captureAccountScope(); + if (!owner.accountId) { + return; + } + + settingsRevision += 1; + const revision = settingsRevision; + pendingSave = { owner, revision, settings }; + if (saveTimer !== undefined) { + globalThis.clearTimeout(saveTimer); + } + saveTimer = globalThis.setTimeout(() => { + saveTimer = undefined; + const pending = pendingSave; + pendingSave = undefined; + if (!pending || !isAccountScopeCurrent(pending.owner)) { + return; + } + void apiFetchJson('/api/v1/remote_workers/settings', { + method: 'PUT', + body: JSON.stringify(pending.settings), + signal: pending.owner.signal, + }).catch(() => { + if (isAccountScopeCurrent(pending.owner) && settingsRevision === pending.revision) { + void loadRemoteWorkersSettings(); + } + }); + }, 250); +}; + +export const loadRemoteWorkersSettings = async (): Promise => { + const owner = captureAccountScope(); + if (!owner.accountId) { + remoteWorkersStore.setSnapshot(makeDefaultSettings()); + loadedEpoch = null; + return; + } + + const loadRevision = settingsRevision; + try { + const saved = await apiFetchJson('/api/v1/remote_workers/settings', { + signal: owner.signal, + }); + if (!isAccountScopeCurrent(owner) || loadRevision !== settingsRevision) { + return; + } + const pending = pendingBeforeLoad; + pendingBeforeLoad = {}; + loadedEpoch = owner.epoch; + const next = { ...normalizeSettings(saved), ...pending }; + remoteWorkersStore.setSnapshot(next); + if (Object.keys(pending).length > 0) { + scheduleSave(next); + } + } catch { + if (!isAccountScopeCurrent(owner) || loadRevision !== settingsRevision) { + return; + } + const pending = pendingBeforeLoad; + pendingBeforeLoad = {}; + loadedEpoch = owner.epoch; + const next = { ...makeDefaultSettings(), ...pending }; + remoteWorkersStore.setSnapshot(next); + if (Object.keys(pending).length > 0) { + scheduleSave(next); + } + } +}; + +const resetForAccountChange = (): void => { + settingsRevision += 1; + loadedEpoch = null; + pendingBeforeLoad = {}; + cancelPendingSave(); + remoteWorkersStore.setSnapshot(makeDefaultSettings()); + if (captureAccountScope().accountId) { + void loadRemoteWorkersSettings(); + } +}; + +registerAccountOwnedResource({ + name: 'remote-workers-server-settings', + clear: resetForAccountChange, +}); + +if (captureAccountScope().accountId) { + void loadRemoteWorkersSettings(); +} + +export const getRemoteWorkersSettings = (): RemoteWorkersSettings => + captureAccountScope().accountId ? remoteWorkersStore.getSnapshot() : makeDefaultSettings(); + +export const setRemoteWorkersSettings = (patch: Partial): void => { + const owner = captureAccountScope(); + if (!owner.accountId) { + return; + } + const safePatch = + patch.workerUrls === undefined ? patch : { ...patch, workerUrls: stripEmbeddedUrlCredentials(patch.workerUrls) }; + const next = { ...remoteWorkersStore.getSnapshot(), ...safePatch }; + remoteWorkersStore.setSnapshot(next); + + if (loadedEpoch !== owner.epoch) { + pendingBeforeLoad = { ...pendingBeforeLoad, ...safePatch }; + return; + } + scheduleSave(next); +}; + +/** Worker selection is per account; changing a slot or worker address never changes which worker is paused. */ +export const isRemoteWorkerEnabled = (url: string): boolean => + !getRemoteWorkersSettings().disabledWorkerUrls.includes(url.toLowerCase()); + +export const setRemoteWorkerEnabled = (url: string, enabled: boolean): void => { + const key = url.toLowerCase(); + const disabled = getRemoteWorkersSettings().disabledWorkerUrls; + if (disabled.includes(key) === !enabled) { + return; + } + setRemoteWorkersSettings({ + disabledWorkerUrls: enabled ? disabled.filter((value) => value !== key) : [...disabled, key], + }); +}; + +export const getRemoteWorkerName = (url: string, index: number): string => { + const saved = getRemoteWorkersSettings().workerNames[url.toLowerCase()]?.trim(); + return saved || `Remote ${index + 1}`; +}; + +export const setRemoteWorkerName = (url: string, name: string): void => { + const key = url.toLowerCase(); + const names = { ...getRemoteWorkersSettings().workerNames }; + const trimmed = name.trim(); + if (trimmed) { + names[key] = name; + } else { + delete names[key]; + } + setRemoteWorkersSettings({ workerNames: names }); +}; + +/** Whitespace/newlines/commas/semicolons separate remotes; order determines slot. */ +export const getRemoteWorkerUrls = (raw: string): string[] => { + const seen = new Set(); + const urls: string[] = []; + for (const part of raw.split(/[;,\r\n\s]+/)) { + const candidate = part.trim().replace(/\/+$/, ''); + if (!candidate) { + continue; + } + let url: URL; + try { + url = new URL(candidate); + } catch { + continue; + } + if (!['http:', 'https:'].includes(url.protocol) || !url.hostname || url.username || url.password) { + continue; + } + const key = candidate.toLowerCase(); + if (!seen.has(key)) { + seen.add(key); + urls.push(candidate); + } + } + return urls; +}; diff --git a/invokeai/frontend/webv2/src/features/queue/index.ts b/invokeai/frontend/webv2/src/features/queue/index.ts index 77dd39b0d64..0e0fa8ee6b3 100644 --- a/invokeai/frontend/webv2/src/features/queue/index.ts +++ b/invokeai/frontend/webv2/src/features/queue/index.ts @@ -67,3 +67,19 @@ export { } from './publicApi'; export type { QueueRunLockPort } from './runtime'; export { hasPendingWorkflowQueueItem } from './ui/queueViewModel'; + +export { + getRemoteWorkerName, + getRemoteWorkerUrls, + isRemoteWorkerEnabled, + remoteWorkersStore, + setRemoteWorkerEnabled, + setRemoteWorkerName, + setRemoteWorkersSettings, +} from './data/remoteWorkersStore'; +export { + invalidateRemoteWorkerHealth, + refreshRemoteWorkerHealth, + remoteWorkersHealthStore, +} from './data/remoteWorkersHealth'; +export type { RemoteDispatchMode } from './data/remoteWorkersStore'; diff --git a/invokeai/frontend/webv2/src/features/queue/runtime.ts b/invokeai/frontend/webv2/src/features/queue/runtime.ts index 42925226cda..4ea0aadb616 100644 --- a/invokeai/frontend/webv2/src/features/queue/runtime.ts +++ b/invokeai/frontend/webv2/src/features/queue/runtime.ts @@ -35,6 +35,8 @@ import { mapWithConcurrency } from '@platform/core/concurrency'; import { captureAccountScope, isAccountScopeCurrent } from '@platform/state/accountLifecycle'; import { ApiError, getApiErrorMessage } from '@platform/transport/http'; +import { applyRemoteWorkersToGraph } from './data/remoteWorkersGraph'; + export interface QueueResultDestinationPort { addImagesToGalleryBoard(boardId: string, imageNames: string[]): Promise; addVideosToGalleryBoard(boardId: string, videoNames: string[]): Promise; @@ -261,6 +263,14 @@ export const createQueueItemBackendSubmission = ( return { error: 'Queue item backend submission has an invalid batch count.', kind: 'invalid' }; } + // Only the submitted graph is modified: never write automatic nodes into + // the user's saved workflow document or change a recovered queue snapshot. + const graph = applyRemoteWorkersToGraph( + submission.graph, + queueItem.snapshot.galleryBoardId, + queueItem.snapshot.destination + ); + if (submission.kind === 'generate') { const seedStep = readSubmissionSeedStep(submission); @@ -289,6 +299,7 @@ export const createQueueItemBackendSubmission = ( kind: 'generate', request: { ...compiled, + graph, destination: queueItem.snapshot.destination, ...(isQueueSeedStep(submission.seedStep) ? {} : { legacySeedPlan: true as const }), projectId: project.id, @@ -323,6 +334,7 @@ export const createQueueItemBackendSubmission = ( kind: 'workflow', request: { ...compiled, + graph, destination: queueItem.snapshot.destination, projectId: project.id, sourceQueueItemId: queueItem.id, diff --git a/invokeai/frontend/webv2/src/features/queue/runtime/coordinator.test.ts b/invokeai/frontend/webv2/src/features/queue/runtime/coordinator.test.ts index d75124b95c4..7fb6c212aa5 100644 --- a/invokeai/frontend/webv2/src/features/queue/runtime/coordinator.test.ts +++ b/invokeai/frontend/webv2/src/features/queue/runtime/coordinator.test.ts @@ -18,8 +18,13 @@ import { import { accountLifecycle } from '@platform/state/accountLifecycle'; import { ApiError } from '@platform/transport/http'; import { createSocketHub, type BackendSocket } from '@platform/transport/socketHub'; +import { toaster } from '@platform/ui'; import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +vi.mock('@platform/ui', () => ({ + toaster: { create: vi.fn(), dismiss: vi.fn(), update: vi.fn() }, +})); + import { createQueueCoordinator, QueueEnqueueNotAcceptedError, @@ -568,6 +573,38 @@ describe('queueCoordinator', () => { expect(harness.callbacks.onGalleryRefresh).not.toHaveBeenCalled(); }); + it('ignores model-transfer progress from another account', () => { + harness.coordinator.dispose(); + harness.hub.disconnect(); + + accountLifecycle.invalidate(); + accountLifecycle.activate('user-1'); + harness = createHarness(); + harness.coordinator.connect(); + + vi.mocked(toaster.create).mockClear(); + + const message = `[[IRW_MODEL_TRANSFER]]${JSON.stringify({ + backend_item_id: 10, + remote_index: 2, + model_hash: 'hash', + name: 'Example', + directory: true, + phase: 'downloading', + bytes: 40, + total_bytes: 100, + })}`; + + harness.socket.fire('invocation_progress', { message, user_id: 'user-2' }); + expect(toaster.create).not.toHaveBeenCalled(); + + harness.socket.fire('invocation_progress', { message, user_id: 'user-1' }); + expect(toaster.create).toHaveBeenCalledTimes(1); + + accountLifecycle.invalidate(); + accountLifecycle.activate('single-user'); + }); + it('routes progress events to the tracked item and clears them on completion', async () => { harness.coordinator.connect(); diff --git a/invokeai/frontend/webv2/src/features/queue/runtime/coordinator.ts b/invokeai/frontend/webv2/src/features/queue/runtime/coordinator.ts index ce0ca2e6a47..8ac99e21f11 100644 --- a/invokeai/frontend/webv2/src/features/queue/runtime/coordinator.ts +++ b/invokeai/frontend/webv2/src/features/queue/runtime/coordinator.ts @@ -37,6 +37,8 @@ import { createLogger } from '@platform/logging/logger'; import { captureAccountScope, isAccountScopeCurrent } from '@platform/state/accountLifecycle'; import { ApiError } from '@platform/transport/http'; +import { createRemoteModelTransferToasts, parseRemoteModelTransferProgress } from './remoteModelTransferToasts'; + const GALLERY_REFRESH_COALESCE_MS = 400; const SAFETY_SWEEP_INTERVAL_MS = 30_000; /** Node-level detail supports the terminal queue-item failure the history owner records. */ @@ -242,6 +244,7 @@ export const createQueueCoordinator = ( const progressImage = options.progressImage ?? progressImageStore; const galleryRefreshCoalesceMs = options.galleryRefreshCoalesceMs ?? GALLERY_REFRESH_COALESCE_MS; const sweepIntervalMs = options.sweepIntervalMs ?? SAFETY_SWEEP_INTERVAL_MS; + const remoteModelTransferToasts = createRemoteModelTransferToasts(); const runs = new Map(); const runProgress = new Map(); @@ -869,6 +872,15 @@ export const createQueueCoordinator = ( return; } + const transfer = parseRemoteModelTransferProgress(event.message); + if (transfer) { + if (owner.accountId !== 'single-user' && event.user_id !== owner.accountId) { + return; + } + remoteModelTransferToasts.receive(transfer); + return; + } + const backendItemId = getTrackedBackendItemId(event); const wait = waits.get(backendItemId); @@ -1000,6 +1012,7 @@ export const createQueueCoordinator = ( /** Detach generation listeners; the hub keeps the socket alive. */ const dispose = (): void => { + remoteModelTransferToasts.dispose(); isDisposed = true; activeProgressTarget.clear(); progressImage.clear(); diff --git a/invokeai/frontend/webv2/src/features/queue/runtime/remoteModelTransferToasts.test.ts b/invokeai/frontend/webv2/src/features/queue/runtime/remoteModelTransferToasts.test.ts new file mode 100644 index 00000000000..18fa496df84 --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/runtime/remoteModelTransferToasts.test.ts @@ -0,0 +1,98 @@ +import { describe, expect, it, vi } from 'vitest'; + +vi.mock('@platform/ui', () => ({ + toaster: { create: vi.fn(), update: vi.fn(), dismiss: vi.fn() }, +})); + +import { toaster } from '@platform/ui'; + +import { createRemoteModelTransferToasts, parseRemoteModelTransferProgress } from './remoteModelTransferToasts'; + +const envelope = (overrides: Record = {}) => + `[[IRW_MODEL_TRANSFER]]${JSON.stringify({ + backend_item_id: 10, + remote_index: 2, + model_hash: 'hash', + name: 'Example', + directory: true, + phase: 'downloading', + bytes: 40, + total_bytes: 100, + ...overrides, + })}`; + +describe('remote model transfer notifications', () => { + it('validates the dedicated progress envelope', () => { + expect(parseRemoteModelTransferProgress('regular progress')).toBeNull(); + expect(parseRemoteModelTransferProgress('[[IRW_MODEL_TRANSFER]]{')).toBeNull(); + expect(parseRemoteModelTransferProgress(envelope({ remote_index: -1 }))).toBeNull(); + expect(parseRemoteModelTransferProgress(envelope({ bytes: -1 }))).toBeNull(); + expect(parseRemoteModelTransferProgress(envelope({ phase: 'unexpected' }))).toBeNull(); + expect(parseRemoteModelTransferProgress(envelope())?.phase).toBe('downloading'); + }); + + it('bounds remembered terminal transfer ids while keeping active state separate', () => { + vi.clearAllMocks(); + const sink = createRemoteModelTransferToasts(); + + for (let index = 0; index < 257; index += 1) { + sink.receive({ + backendItemId: 1000 + index, + remoteIndex: 1, + modelHash: `hash-${index}`, + name: `Model ${index}`, + directory: true, + phase: 'completed', + bytes: 1, + totalBytes: 1, + }); + } + + expect(toaster.create).toHaveBeenCalledTimes(257); + + sink.receive({ + backendItemId: 1000, + remoteIndex: 1, + modelHash: 'hash-0', + name: 'Model 0', + directory: true, + phase: 'downloading', + bytes: 0, + totalBytes: 1, + }); + expect(toaster.create).toHaveBeenCalledTimes(258); + + sink.receive({ + backendItemId: 1256, + remoteIndex: 1, + modelHash: 'hash-256', + name: 'Model 256', + directory: true, + phase: 'downloading', + bytes: 0, + totalBytes: 1, + }); + expect(toaster.create).toHaveBeenCalledTimes(258); + }); + + it('updates one loading toast, settles it and does not resurrect it', () => { + vi.clearAllMocks(); + const sink = createRemoteModelTransferToasts(); + const progress = parseRemoteModelTransferProgress(envelope()); + const completed = parseRemoteModelTransferProgress(envelope({ phase: 'completed', bytes: 100 })); + expect(progress).not.toBeNull(); + expect(completed).not.toBeNull(); + if (!progress || !completed) { + return; + } + sink.receive(progress); + sink.receive({ ...progress, bytes: 60 }); + sink.receive(completed); + sink.receive(progress); + expect(toaster.create).toHaveBeenCalledTimes(1); + expect(toaster.update).toHaveBeenCalledTimes(2); + expect(toaster.update).toHaveBeenLastCalledWith(expect.any(String), expect.objectContaining({ type: 'success' })); + sink.dispose(); + expect(toaster.dismiss).toHaveBeenCalledTimes(1); + }); +}); diff --git a/invokeai/frontend/webv2/src/features/queue/runtime/remoteModelTransferToasts.ts b/invokeai/frontend/webv2/src/features/queue/runtime/remoteModelTransferToasts.ts new file mode 100644 index 00000000000..60c6ece1c22 --- /dev/null +++ b/invokeai/frontend/webv2/src/features/queue/runtime/remoteModelTransferToasts.ts @@ -0,0 +1,185 @@ +import { toaster } from '@platform/ui'; + +const PREFIX = '[[IRW_MODEL_TRANSFER]]'; + +type Phase = + | 'preparing' + | 'waiting' + | 'downloading' + | 'installing' + | 'verifying' + | 'completed' + | 'failed' + | 'cancelled'; + +export interface RemoteModelTransferProgress { + backendItemId: number; + remoteIndex: number; + modelHash: string; + name: string; + directory: boolean; + phase: Phase; + bytes: number; + totalBytes: number; + error?: string; +} + +const MAX_TERMINAL_IDS = 256; + +const PHASES = new Set([ + 'preparing', + 'waiting', + 'downloading', + 'installing', + 'verifying', + 'completed', + 'failed', + 'cancelled', +]); + +/** Model-transfer events are emitted by the primary for the queue item's owner. */ +export const parseRemoteModelTransferProgress = (message: string): RemoteModelTransferProgress | null => { + if (!message.startsWith(PREFIX)) { + return null; + } + let data: unknown; + try { + data = JSON.parse(message.slice(PREFIX.length)); + } catch { + return null; + } + if (data === null || typeof data !== 'object' || Array.isArray(data)) { + return null; + } + const value = data as Record; + if ( + !Number.isSafeInteger(value.backend_item_id) || + (value.backend_item_id as number) < 1 || + !Number.isSafeInteger(value.remote_index) || + (value.remote_index as number) < 1 || + (value.remote_index as number) > 64 || + typeof value.model_hash !== 'string' || + value.model_hash.length === 0 || + value.model_hash.length > 256 || + typeof value.name !== 'string' || + value.name.length === 0 || + value.name.length > 256 || + typeof value.directory !== 'boolean' || + typeof value.phase !== 'string' || + !PHASES.has(value.phase as Phase) || + !Number.isSafeInteger(value.bytes) || + (value.bytes as number) < 0 || + !Number.isSafeInteger(value.total_bytes) || + (value.total_bytes as number) < 0 || + (value.error !== undefined && (typeof value.error !== 'string' || value.error.length > 2048)) + ) { + return null; + } + return { + backendItemId: value.backend_item_id as number, + remoteIndex: value.remote_index as number, + modelHash: value.model_hash as string, + name: value.name as string, + directory: value.directory as boolean, + phase: value.phase as Phase, + bytes: value.bytes as number, + totalBytes: value.total_bytes as number, + ...(typeof value.error === 'string' ? { error: value.error } : {}), + }; +}; + +const readableBytes = (bytes: number): string => { + if (bytes >= 1024 ** 3) { + return `${(bytes / 1024 ** 3).toFixed(2)} GiB`; + } + return `${(bytes / 1024 ** 2).toFixed(1)} MiB`; +}; + +const descriptionFor = (value: RemoteModelTransferProgress): string => { + switch (value.phase) { + case 'preparing': + return value.directory + ? 'Preparing model directory and file checksums…' + : 'Preparing single-file model transfer…'; + case 'waiting': + return 'Waiting for the remote model installer…'; + case 'downloading': { + const progress = + value.totalBytes > 0 ? `${Math.min(100, Math.floor((value.bytes / value.totalBytes) * 100))}% · ` : ''; + const size = + value.totalBytes > 0 + ? `${readableBytes(value.bytes)} / ${readableBytes(value.totalBytes)}` + : `${readableBytes(value.bytes)} transferred`; + return `Transferring ${progress}${size}`; + } + case 'installing': + return 'Installing and registering model on remote worker…'; + case 'verifying': + return 'Verifying installed model hash…'; + case 'completed': + return 'Model transferred, installed and verified.'; + case 'failed': + return value.error ? `Model transfer failed: ${value.error.slice(0, 240)}` : 'Model transfer failed.'; + case 'cancelled': + return 'Model transfer cancelled.'; + } +}; + +/** One updating toast per primary queue item, remote worker and model. */ +export const createRemoteModelTransferToasts = () => { + const active = new Set(); + const terminal = new Set(); + + const receive = (value: RemoteModelTransferProgress): void => { + const id = `irw-model-transfer:${value.backendItemId}:${value.remoteIndex}:${value.modelHash}`; + if (terminal.has(id)) { + return; + } + + const isTerminal = value.phase === 'completed' || value.phase === 'failed' || value.phase === 'cancelled'; + const title = `Remote ${value.remoteIndex} · ${value.name}`; + const description = descriptionFor(value); + const type = + value.phase === 'completed' + ? 'success' + : value.phase === 'failed' + ? 'error' + : value.phase === 'cancelled' + ? 'info' + : 'loading'; + const duration = isTerminal ? 6000 : Infinity; + + if (active.has(id)) { + toaster.update(id, { title, description, type, duration }); + } else { + toaster.create({ id, title, description, type, duration }); + } + + if (isTerminal) { + active.delete(id); + terminal.add(id); + while (terminal.size > MAX_TERMINAL_IDS) { + const oldest = terminal.values().next().value; + if (oldest === undefined) { + break; + } + terminal.delete(oldest); + } + } else { + active.add(id); + } + }; + + const dispose = (): void => { + for (const id of active) { + toaster.dismiss(id); + } + for (const id of terminal) { + toaster.dismiss(id); + } + active.clear(); + terminal.clear(); + }; + + return { receive, dispose }; +}; diff --git a/invokeai/frontend/webv2/src/features/queue/ui/QueueItemDetails.tsx b/invokeai/frontend/webv2/src/features/queue/ui/QueueItemDetails.tsx index 7a406758839..70a9037e2d5 100644 --- a/invokeai/frontend/webv2/src/features/queue/ui/QueueItemDetails.tsx +++ b/invokeai/frontend/webv2/src/features/queue/ui/QueueItemDetails.tsx @@ -37,7 +37,8 @@ export const QueueItemDetails = ({ item }: { item: QueueItemReadModel }) => { const { ItemActions } = useQueueUi(); const meta = extractGenerationMeta(item); const duration = formatDuration(item.startedAt, item.completedAt); - const deviceLabel = useDeviceLabel(item.device); + const remoteWorkerName = item.device?.startsWith('remote:') ? item.device.slice('remote:'.length) : null; + const deviceLabel = useDeviceLabel(remoteWorkerName ? null : item.device); return ( @@ -61,7 +62,9 @@ export const QueueItemDetails = ({ item }: { item: QueueItemReadModel }) => { {item.userDisplayName || item.userEmail ? ( {item.userDisplayName ?? item.userEmail} ) : null} - {deviceLabel ? {deviceLabel.name} : null} + {remoteWorkerName || deviceLabel ? ( + {remoteWorkerName ?? deviceLabel?.name} + ) : null} {item.errorMessage ? ( diff --git a/invokeai/frontend/webv2/src/features/queue/ui/QueueItemRow.tsx b/invokeai/frontend/webv2/src/features/queue/ui/QueueItemRow.tsx index abb677a1a87..9cd8bccc4e7 100644 --- a/invokeai/frontend/webv2/src/features/queue/ui/QueueItemRow.tsx +++ b/invokeai/frontend/webv2/src/features/queue/ui/QueueItemRow.tsx @@ -61,7 +61,9 @@ export const QueueItemRow = memo( const statusLabel = t(getStatusMeta(item.status).labelKey); // A running item's device arrives on the progress event before the row's DTO is // refetched, so prefer the live value and fall back to the persisted one. - const deviceLabel = useDeviceLabel(progress?.device ?? item.device); + const rawDevice = progress?.device ?? item.device; + const remoteWorkerName = rawDevice?.startsWith('remote:') ? rawDevice.slice('remote:'.length) : null; + const deviceLabel = useDeviceLabel(remoteWorkerName ? null : rawDevice); const showBorder = expanded || isFailed; const borderColor = showBorder ? (isFailed ? 'fg.error' : 'border') : 'transparent'; @@ -76,13 +78,17 @@ export const QueueItemRow = memo( {[ statusLabel, ageLabel, - deviceLabel ? t('widgets.queue.device.shortLabel', { index: deviceLabel.index }) : null, + remoteWorkerName ?? + (deviceLabel ? t('widgets.queue.device.shortLabel', { index: deviceLabel.index }) : null), ] .filter(Boolean) .join(' · ')} diff --git a/invokeai/frontend/webv2/src/features/queue/ui/currentBatchItems.test.ts b/invokeai/frontend/webv2/src/features/queue/ui/currentBatchItems.test.ts index 86edf7d801a..d7821bd78f7 100644 --- a/invokeai/frontend/webv2/src/features/queue/ui/currentBatchItems.test.ts +++ b/invokeai/frontend/webv2/src/features/queue/ui/currentBatchItems.test.ts @@ -35,4 +35,24 @@ describe('getCurrentBatchItems', () => { expect(getCurrentBatchItems({ current: null, items, next }).map((item) => item.id)).toEqual([5, 6]); }); + + it('keeps concurrently running items from another batch out of recent history', () => { + const current = createItem(10, 'batch-local', 'in_progress'); + const remoteRunning = createItem(11, 'batch-remote', 'in_progress'); + const unrelatedPending = createItem(12, 'batch-remote', 'pending'); + + expect( + getCurrentBatchItems({ + current, + items: [remoteRunning, unrelatedPending], + next: null, + }).map((item) => item.id) + ).toEqual([10, 11]); + }); + + it('shows running items even when the backend has no single current or next item', () => { + const running = createItem(20, 'batch-remote', 'in_progress'); + + expect(getCurrentBatchItems({ current: null, items: [running], next: null }).map((item) => item.id)).toEqual([20]); + }); }); diff --git a/invokeai/frontend/webv2/src/features/queue/ui/currentBatchItems.ts b/invokeai/frontend/webv2/src/features/queue/ui/currentBatchItems.ts index f951cc03ca3..6ec59134571 100644 --- a/invokeai/frontend/webv2/src/features/queue/ui/currentBatchItems.ts +++ b/invokeai/frontend/webv2/src/features/queue/ui/currentBatchItems.ts @@ -1,8 +1,5 @@ import type { QueueItemReadModel } from '@features/queue/core/types'; -const isCurrentBatchStatus = (status: QueueItemReadModel['status']): boolean => - status === 'pending' || status === 'in_progress'; - export const getCurrentBatchItems = ({ current, items, @@ -13,15 +10,17 @@ export const getCurrentBatchItems = ({ next: QueueItemReadModel | null; }): QueueItemReadModel[] => { const batchId = current?.batchId ?? next?.batchId ?? null; - - if (!batchId) { - return []; - } - const itemsById = new Map(); for (const item of [current, next, ...items]) { - if (item && item.batchId === batchId && isCurrentBatchStatus(item.status)) { + if (!item) { + continue; + } + + const isRunning = item.status === 'in_progress'; + const isPendingCurrentBatch = batchId !== null && item.status === 'pending' && item.batchId === batchId; + + if (isRunning || isPendingCurrentBatch) { itemsById.set(item.id, item); } } diff --git a/invokeai/frontend/webv2/src/workbench/widgetRegistry.test.ts b/invokeai/frontend/webv2/src/workbench/widgetRegistry.test.ts index a3a425da4d5..5b48467a32c 100644 --- a/invokeai/frontend/webv2/src/workbench/widgetRegistry.test.ts +++ b/invokeai/frontend/webv2/src/workbench/widgetRegistry.test.ts @@ -32,8 +32,10 @@ describe('widget registry', () => { it('registers first-party widget manifests without icon validation failures', () => { const widgets = registerFirstPartyWidgets(); - expect(widgets).toHaveLength(16); - expect(widgets.map((widget) => widget.manifest.id)).toEqual(expect.arrayContaining(['layers', 'image-map'])); + expect(widgets).toHaveLength(17); + expect(widgets.map((widget) => widget.manifest.id)).toEqual( + expect.arrayContaining(['layers', 'image-map', 'remote-workers']) + ); expect(widgets.flatMap((widget) => widget.failure ?? [])).toEqual([]); expect(widgets.every((widget) => widget.status === 'enabled')).toBe(true); }); diff --git a/invokeai/frontend/webv2/src/workbench/widgets/manifests.ts b/invokeai/frontend/webv2/src/workbench/widgets/manifests.ts index ab7ceab3230..d9f762af456 100644 --- a/invokeai/frontend/webv2/src/workbench/widgets/manifests.ts +++ b/invokeai/frontend/webv2/src/workbench/widgets/manifests.ts @@ -12,6 +12,7 @@ import { previewWidgetManifest } from './preview/manifest'; import { projectWidgetManifest } from './project/manifest'; import { queueStatusWidgetManifest } from './queue-status/manifest'; import { queueWidgetManifest } from './queue/manifest'; +import { remoteWorkersWidgetManifest } from './remote-workers/manifest'; import { serverStatusWidgetManifest } from './server-status/manifest'; import { upscaleWidgetManifest } from './upscale/manifest'; import { videoWidgetManifest } from './video/manifest'; @@ -30,6 +31,7 @@ export const firstPartyWidgetManifests: WidgetManifest[] = [ projectWidgetManifest, layersWidgetManifest, queueWidgetManifest, + remoteWorkersWidgetManifest, notificationsWidgetManifest, serverStatusWidgetManifest, queueStatusWidgetManifest, diff --git a/invokeai/frontend/webv2/src/workbench/widgets/preview/PreviewWidgetView.tsx b/invokeai/frontend/webv2/src/workbench/widgets/preview/PreviewWidgetView.tsx index b0e2ba19f6b..438af0bf238 100644 --- a/invokeai/frontend/webv2/src/workbench/widgets/preview/PreviewWidgetView.tsx +++ b/invokeai/frontend/webv2/src/workbench/widgets/preview/PreviewWidgetView.tsx @@ -181,7 +181,7 @@ export const PreviewWidgetView = ({ region, runtime }: WidgetViewProps) => { const selectedItemKey = selectedItem ? toGalleryItemKey(selectedItem) : null; const activeGalleryPlaceholder = livePreview.sessions.find((session) => session.id === livePreview.followedSessionId) ?? null; - const shouldFollowLive = activeGalleryPlaceholder !== null; + const shouldFollowLive = activeGalleryPlaceholder !== null && !livePreview.viewingSaved; const isComparing = !shouldFollowLive && selectedItem?.kind === 'image' && @@ -223,11 +223,12 @@ export const PreviewWidgetView = ({ region, runtime }: WidgetViewProps) => { const selectGalleryItemAtPage = useCallback( (item: GalleryItem, selectionPage: number) => { + livePreview.showSaved(); gallery.selectItem(item, undefined, selectionPage, true); // Deliberate navigation: the grid follows it, unlike auto-selection. requestGalleryItemReveal(toGalleryItemKey(item)); }, - [gallery] + [gallery, livePreview] ); const { boardItems, diff --git a/invokeai/frontend/webv2/src/workbench/widgets/preview/livePreviewFollow.tsx b/invokeai/frontend/webv2/src/workbench/widgets/preview/livePreviewFollow.tsx index 8d4ad57fa2c..8afa0e0e490 100644 --- a/invokeai/frontend/webv2/src/workbench/widgets/preview/livePreviewFollow.tsx +++ b/invokeai/frontend/webv2/src/workbench/widgets/preview/livePreviewFollow.tsx @@ -15,15 +15,19 @@ interface LivePreviewFollow { sessions: QueueActiveSession[]; gallerySessions: QueueProgressSession[]; pinnedSessionId: string | null; + /** A deliberate saved Gallery selection temporarily takes priority over live rendering. */ + viewingSaved: boolean; /** * Shared live target: pinned, then newest-started running session, then first settling session in gallery order, * else null. */ followedSessionId: string | null; - /** Turns live-follow on and pins `sessionId`: a tile click, or an arrow step onto a tile. */ + /** Turns live-follow on and pins `sessionId` when a live thumbnail or navigation step is selected. */ follow(sessionId: string): void; pin(sessionId: string): void; showAll(): void; + /** A Gallery click opens saved media without stopping the running generations. */ + showSaved(): void; } const LivePreviewFollowContext = createContext(null); @@ -45,25 +49,31 @@ export const LivePreviewFollowProvider = ({ children }: { children: ReactNode }) () => getQueueProgressSessions(items.filter(isGalleryProgressItem), sessions), [items, sessions] ); - const [selection, setSelection] = useState<{ projectId: string; sessionId: string | null }>({ + const [selection, setSelection] = useState<{ + projectId: string; + sessionId: string | null; + viewingSaved: boolean; + }>({ projectId, sessionId: null, + viewingSaved: false, }); const isStale = selection.projectId !== projectId || - !enabled || + (!selection.viewingSaved && !enabled) || + (selection.viewingSaved && sessions.length === 0) || (selection.sessionId !== null && !sessions.some((session) => session.id === selection.sessionId && session.state === 'running')); - if (isStale && (selection.sessionId !== null || selection.projectId !== projectId)) { - setSelection({ projectId, sessionId: null }); + if (isStale && (selection.sessionId !== null || selection.viewingSaved || selection.projectId !== projectId)) { + setSelection({ projectId, sessionId: null, viewingSaved: false }); } - const pinnedSessionId = isStale ? null : selection.sessionId; + const viewingSaved = !isStale && selection.viewingSaved; + const pinnedSessionId = isStale || viewingSaved ? null : selection.sessionId; // Follow the highest-id running session (FIFO start order), not the latest progress frame; pins take precedence. const newestRunningSessionId = sessions.filter((session) => session.state === 'running').at(-1)?.id ?? null; const preferredSessionId = pinnedSessionId ?? newestRunningSessionId; - const followedSessionId = enabled - ? (getFollowedProgressSession(gallerySessions, preferredSessionId)?.id ?? null) - : null; + const followedSessionId = + enabled && !viewingSaved ? (getFollowedProgressSession(gallerySessions, preferredSessionId)?.id ?? null) : null; const { account } = useWorkbenchCommands(); const value = useMemo( () => ({ @@ -71,14 +81,16 @@ export const LivePreviewFollowProvider = ({ children }: { children: ReactNode }) gallerySessions, pinnedSessionId, followedSessionId, + viewingSaved, follow: (sessionId) => { account.updateProjectPreferences({ showProgressImagesInViewer: true }); - setSelection({ projectId, sessionId }); + setSelection({ projectId, sessionId, viewingSaved: false }); }, - pin: (sessionId) => setSelection({ projectId, sessionId }), - showAll: () => setSelection({ projectId, sessionId: null }), + pin: (sessionId) => setSelection({ projectId, sessionId, viewingSaved: false }), + showAll: () => setSelection({ projectId, sessionId: null, viewingSaved: false }), + showSaved: () => setSelection({ projectId, sessionId: null, viewingSaved: true }), }), - [account, sessions, gallerySessions, followedSessionId, pinnedSessionId, projectId] + [account, sessions, gallerySessions, followedSessionId, pinnedSessionId, projectId, viewingSaved] ); return {children}; }; diff --git a/invokeai/frontend/webv2/src/workbench/widgets/remote-workers/RemoteWorkersWidgetView.tsx b/invokeai/frontend/webv2/src/workbench/widgets/remote-workers/RemoteWorkersWidgetView.tsx new file mode 100644 index 00000000000..dff46d66810 --- /dev/null +++ b/invokeai/frontend/webv2/src/workbench/widgets/remote-workers/RemoteWorkersWidgetView.tsx @@ -0,0 +1,488 @@ +import type { WidgetViewProps } from '@workbench/widgetContracts'; +import type { ChangeEvent } from 'react'; + +import { Badge, Box, Button, HStack, Input, NativeSelect, Stack, Switch, Text, Textarea } from '@chakra-ui/react'; +import { + getRemoteWorkerUrls, + invalidateRemoteWorkerHealth, + refreshRemoteWorkerHealth, + remoteWorkersHealthStore, + remoteWorkersStore, + setRemoteWorkerEnabled, + setRemoteWorkerName, + setRemoteWorkersSettings, + type RemoteDispatchMode, +} from '@features/queue'; +import { captureAccountScope } from '@platform/state/accountLifecycle'; +import { apiFetchJson, getApiErrorMessage } from '@platform/transport/http'; +import { ChevronDownIcon, ChevronUpIcon, PencilIcon, PowerIcon } from 'lucide-react'; +import { useCallback, useEffect, useMemo, useState } from 'react'; +import { useTranslation } from 'react-i18next'; + +const handleEnabledChange = (details: { checked: boolean }): void => { + setRemoteWorkersSettings({ enabled: details.checked }); +}; + +const handleDispatchModeChange = (event: ChangeEvent): void => { + setRemoteWorkersSettings({ dispatchMode: event.target.value as RemoteDispatchMode }); +}; + +const handleWorkerUrlsChange = (event: ChangeEvent): void => { + setRemoteWorkersSettings({ workerUrls: event.target.value }); +}; + +const handleAutoTransferChange = (details: { checked: boolean }): void => { + setRemoteWorkersSettings({ autoTransferMissingModels: details.checked }); +}; + +const handleKeepCopiesChange = (details: { checked: boolean }): void => { + setRemoteWorkersSettings({ keepRemoteCopies: details.checked }); +}; + +const handleTransferHostChange = (event: ChangeEvent): void => { + setRemoteWorkersSettings({ modelTransferHost: event.target.value }); +}; + +interface CredentialStatus { + saved: boolean; + email: string | null; +} + +/** The password never enters queue settings, localStorage, or the workflow graph. */ +const WorkerAuthRow = ({ enabled, slot, url }: { enabled: boolean; slot: number; url: string }) => { + const { t } = useTranslation(); + const [expanded, setExpanded] = useState(false); + const workerEnabled = remoteWorkersStore.useSelector( + (settings) => !settings.disabledWorkerUrls.includes(url.toLowerCase()) + ); + const savedName = remoteWorkersStore.useSelector((settings) => settings.workerNames[url.toLowerCase()]?.trim() ?? ''); + const defaultName = t('widgets.remoteWorkers.worker.defaultName', { slot }); + const name = savedName || defaultName; + const availability = remoteWorkersHealthStore.useSnapshot().byUrl[url]?.status ?? 'checking'; + const availabilityLabel = !workerEnabled + ? t('widgets.remoteWorkers.status.disabled') + : !enabled + ? t('widgets.remoteWorkers.status.paused') + : availability === 'online' + ? t('widgets.remoteWorkers.status.online') + : availability === 'offline' + ? t('widgets.remoteWorkers.status.offline') + : availability === 'login_required' + ? t('widgets.remoteWorkers.status.loginRequired') + : t('widgets.remoteWorkers.status.checking'); + const availabilityColor = + !workerEnabled || !enabled + ? 'gray' + : availability === 'online' + ? 'green' + : availability === 'offline' + ? 'red' + : availability === 'login_required' + ? 'orange' + : 'gray'; + const [email, setEmail] = useState(''); + const [password, setPassword] = useState(''); + const [saved, setSaved] = useState(false); + const [checkingLogin, setCheckingLogin] = useState(true); + const [busy, setBusy] = useState(false); + const [message, setMessage] = useState(''); + + useEffect(() => { + let active = true; + void apiFetchJson(`/api/v1/remote_workers/credentials?url=${encodeURIComponent(url)}`) + .then((status) => { + if (!active) { + return; + } + setSaved(status.saved); + setEmail(status.email ?? ''); + setCheckingLogin(false); + setMessage(''); + }) + .catch((error: unknown) => { + if (active) { + setCheckingLogin(false); + setMessage(getApiErrorMessage(error, t('widgets.remoteWorkers.login.loadError'))); + } + }); + return () => { + active = false; + }; + }, [t, url]); + + const handleEmailChange = useCallback((event: ChangeEvent) => { + setEmail(event.target.value); + }, []); + const handlePasswordChange = useCallback((event: ChangeEvent) => { + setPassword(event.target.value); + }, []); + const handleNameBlur = useCallback( + (event: ChangeEvent) => { + const value = event.currentTarget.value; + setRemoteWorkerName(url, value); + if (!value.trim()) { + event.currentTarget.value = defaultName; + } + }, + [defaultName, url] + ); + const refreshHealthAfterCredentialChange = useCallback(() => { + invalidateRemoteWorkerHealth(url); + if (enabled) { + void refreshRemoteWorkerHealth([url]); + } + }, [enabled, url]); + const handleSave = useCallback(async () => { + setBusy(true); + setMessage(''); + try { + const status = await apiFetchJson('/api/v1/remote_workers/credentials', { + method: 'PUT', + body: JSON.stringify({ url, email, password, remember_me: true }), + }); + setSaved(status.saved); + setEmail(status.email ?? ''); + setPassword(''); + refreshHealthAfterCredentialChange(); + setMessage(t('widgets.remoteWorkers.login.savedMessage')); + } catch (error) { + setMessage(getApiErrorMessage(error, t('widgets.remoteWorkers.login.saveError'))); + } finally { + setBusy(false); + } + }, [url, email, password, refreshHealthAfterCredentialChange, t]); + const handleRemove = useCallback(async () => { + setBusy(true); + setMessage(''); + try { + await apiFetchJson(`/api/v1/remote_workers/credentials?url=${encodeURIComponent(url)}`, { + method: 'DELETE', + }); + setSaved(false); + setEmail(''); + setPassword(''); + refreshHealthAfterCredentialChange(); + setMessage(t('widgets.remoteWorkers.login.removedMessage')); + } catch (error) { + setMessage(getApiErrorMessage(error, t('widgets.remoteWorkers.login.removeError'))); + } finally { + setBusy(false); + } + }, [url, refreshHealthAfterCredentialChange, t]); + + const toggleExpanded = useCallback(() => setExpanded((open) => !open), []); + const toggleWorkerEnabled = useCallback(() => setRemoteWorkerEnabled(url, !workerEnabled), [url, workerEnabled]); + + return ( + + + + + + + {expanded ? ( + + + + {checkingLogin + ? t('widgets.remoteWorkers.login.checking') + : saved + ? t('widgets.remoteWorkers.login.saved') + : t('widgets.remoteWorkers.login.none')} + + + {t('widgets.remoteWorkers.login.description')} + + + + + + + + + ) : null} + {message ? ( + + {message} + + ) : null} + + ); +}; + +/** Availability polling is display-only; backend workers decide job eligibility. */ +export const RemoteWorkersWidgetView = (_props: WidgetViewProps) => { + const { t } = useTranslation(); + const settings = remoteWorkersStore.useSnapshot(); + const urls = useMemo(() => getRemoteWorkerUrls(settings.workerUrls), [settings.workerUrls]); + const accountId = captureAccountScope().accountId; + const [showAdvanced, setShowAdvanced] = useState(false); + const [showWorkerEditor, setShowWorkerEditor] = useState(false); + + useEffect(() => { + if (!settings.enabled) { + if (Object.keys(remoteWorkersHealthStore.getSnapshot().byUrl).length > 0) { + remoteWorkersHealthStore.setSnapshot({ byUrl: {} }); + } + return; + } + + const refresh = (): void => { + void refreshRemoteWorkerHealth(urls); + }; + + refresh(); + const timer = globalThis.setInterval(refresh, 15_000); + return () => globalThis.clearInterval(timer); + }, [settings.enabled, urls]); + const handleAdvancedChange = useCallback((details: { checked: boolean }) => { + setShowAdvanced(details.checked); + }, []); + const toggleWorkerEditor = useCallback(() => setShowWorkerEditor((open) => !open), []); + + return ( + + + + {t('widgets.remoteWorkers.title')} + + {settings.enabled ? t('widgets.remoteWorkers.enabled') : t('widgets.remoteWorkers.disabled')} + + + + + + + + {t('widgets.remoteWorkers.enable')} + + + {t('widgets.remoteWorkers.description')} + + + + + + + {t('widgets.remoteWorkers.dispatch.title')} + + + + + + + + + + {settings.dispatchMode === 'remote_only' + ? t('widgets.remoteWorkers.dispatch.remoteOnlyDescription') + : t('widgets.remoteWorkers.dispatch.distributedDescription')} + + + + + + + + + {t('widgets.remoteWorkers.workers.title')} + + + {t('widgets.remoteWorkers.workers.configured', { count: urls.length })} + + + {urls.length === 0 ? ( + + {t('widgets.remoteWorkers.workers.none')} + + ) : ( + + {urls.map((url, index) => ( + + ))} + + )} + + {showWorkerEditor ? ( + + + {t('widgets.remoteWorkers.workers.editorDescription')} + +