From 1038f12baad9857cfe7adbc307d283e278014bf5 Mon Sep 17 00:00:00 2001 From: Souptik Chakraborty <62941615+Souptik96@users.noreply.github.com> Date: Fri, 14 Aug 2026 11:26:30 +0530 Subject: [PATCH 01/57] fix(extensions): recover worker after failed texture setup A generation with enable_texture builds the texture pipeline lazily, and extensions free the shape pipeline first to make room for it. When that setup failed (missing xatlas being the common case) the generator was left with _model = None while the worker process stayed alive, and nothing reset the loaded state: - runner.py called gen.generate() unconditionally, with no loaded check. - ExtensionProcess._loaded stayed True, because it is only cleared by unload(), stop() and the cancel hard-kill path, not by a failed run. - GeneratorRegistry.get_active() therefore skipped load(). Every later generation then raised "TypeError: 'NoneType' object is not callable" until the worker was killed by hand. The runner now ensures the model is loaded before inference, mirroring what get_active() already does host-side, and reports its post-failure loaded state so ExtensionProcess can drop its cached flag and reload on the next run. The original failure is still surfaced unchanged. Fixes #239 --- api/runner.py | 50 ++++++- api/services/extension_process.py | 7 + api/tests/test_extension_process.py | 45 ++++++ api/tests/test_runner.py | 204 ++++++++++++++++++++++++++++ 4 files changed, 305 insertions(+), 1 deletion(-) diff --git a/api/runner.py b/api/runner.py index 12f2bf23..9afccdd9 100644 --- a/api/runner.py +++ b/api/runner.py @@ -98,6 +98,46 @@ def _resolve_ready_schema(GenClass, node: dict, manifest: dict) -> list: return node.get("params_schema") or manifest.get("params_schema", []) +def _generator_is_loaded(gen) -> bool: + """ + Best-effort read of the generator's loaded state. + + Never raises: this is used on the error path, where a generator left in a + half-initialised state can make is_loaded() itself blow up. An unreadable + state is reported as "not loaded" so the caller reloads rather than reusing + a worker that may be broken. + """ + try: + return bool(gen.is_loaded()) + except Exception: + return False + + +def _ensure_model_loaded(gen) -> bool: + """ + Guarantees the model is in memory before inference. Returns True if a + reload was needed. + + Texture pipelines are built lazily on first use, and extensions typically + free the shape pipeline first to make room for them. When that setup fails + (a missing `xatlas` being the common case) the generator is left with + _model = None, yet the worker process stays alive and both ExtensionProcess + and GeneratorRegistry still consider it loaded — so load() is never called + again and every later generate() dies with + + TypeError: 'NoneType' object is not callable + + until the process is killed by hand. Re-loading here mirrors what + GeneratorRegistry.get_active() already does host-side, which makes a failed + setup recoverable on the next attempt instead of permanently poisoning the + worker. + """ + if _generator_is_loaded(gen): + return False + gen.load() + return True + + def _apply_manifest_metadata(gen, manifest: dict, node: dict) -> None: gen.hf_repo = node.get("hf_repo") or manifest.get("hf_repo", "") gen.hf_skip_prefixes = node.get("hf_skip_prefixes") or manifest.get("hf_skip_prefixes", []) @@ -167,6 +207,10 @@ def progress_cb(pct: int, step: str = "") -> None: send({"type": "progress", "id": rid, "pct": pct, "step": step}) try: + if _ensure_model_loaded(gen): + send({"type": "log", "level": "warning", + "message": ("Model was not loaded (earlier setup failure?); " + "reloaded before generating.")}) output_path = gen.generate(image_bytes, params, progress_cb, cancel_evt) send({"type": "done", "id": rid, "output_path": str(output_path)}) except Exception as exc: @@ -174,9 +218,13 @@ def progress_cb(pct: int, step: str = "") -> None: if type(exc).__name__ == "GenerationCancelled": send({"type": "cancelled", "id": rid}) else: + # Report whether the model survived the failure so the + # host can drop its cached "loaded" flag and reload + # instead of reusing a worker whose _model is None. send({"type": "error", "id": rid, "message": str(exc), - "traceback": traceback.format_exc()}) + "traceback": traceback.format_exc(), + "loaded": _generator_is_loaded(gen)}) finally: _cancel.pop(rid, None) diff --git a/api/services/extension_process.py b/api/services/extension_process.py index 341df6d6..dff48886 100644 --- a/api/services/extension_process.py +++ b/api/services/extension_process.py @@ -351,6 +351,13 @@ def generate( return Path(msg["output_path"]) elif t == "error": + # A failed generation can leave the worker without a model: a + # lazy texture-setup failure frees the shape pipeline and then + # raises. The worker reports its post-failure state, so drop + # our cached flag and let GeneratorRegistry.get_active() reload + # before the next run rather than reusing a broken worker. + if msg.get("loaded") is False: + self._loaded = False raise RuntimeError(msg.get("traceback") or msg.get("message", "Unknown error")) elif t == "cancelled": diff --git a/api/tests/test_extension_process.py b/api/tests/test_extension_process.py index 9cd5c5d6..e82e170f 100644 --- a/api/tests/test_extension_process.py +++ b/api/tests/test_extension_process.py @@ -142,5 +142,50 @@ def test_empty_queue_raises_timeout_error(self) -> None: proc._recv(timeout=0.05) +class GenerateErrorLoadedFlagTests(unittest.TestCase): + """ + Issue #239: a generation that fails during lazy texture setup leaves the + worker without a model. If _loaded stays True, GeneratorRegistry.get_active() + skips load() forever and every later run reuses the broken worker. + """ + + def _failing_generate(self, error_msg: dict) -> ExtensionProcess: + proc = _make_proc() + proc._loaded = True + proc._send = lambda msg: None # type: ignore[assignment] + proc._queue.put(error_msg) + with self.assertRaises(RuntimeError): + proc.generate(b"", {}) + return proc + + def test_clears_loaded_when_worker_reports_model_lost(self) -> None: + proc = self._failing_generate( + {"type": "error", "message": "No module named 'xatlas'", "loaded": False} + ) + self.assertFalse(proc._loaded) + + def test_keeps_loaded_when_worker_still_has_its_model(self) -> None: + proc = self._failing_generate( + {"type": "error", "message": "bad input image", "loaded": True} + ) + self.assertTrue(proc._loaded) + + def test_keeps_loaded_when_worker_reports_no_state(self) -> None: + proc = self._failing_generate({"type": "error", "message": "boom"}) + self.assertTrue(proc._loaded) + + def test_error_still_propagates_the_original_cause(self) -> None: + proc = _make_proc() + proc._loaded = True + proc._send = lambda msg: None # type: ignore[assignment] + proc._queue.put( + {"type": "error", "message": "short", "traceback": "full traceback here", + "loaded": False} + ) + with self.assertRaises(RuntimeError) as ctx: + proc.generate(b"", {}) + self.assertIn("full traceback here", str(ctx.exception)) + + if __name__ == "__main__": unittest.main() diff --git a/api/tests/test_runner.py b/api/tests/test_runner.py index f11a9c22..8fce3d31 100644 --- a/api/tests/test_runner.py +++ b/api/tests/test_runner.py @@ -154,5 +154,209 @@ def test_send_writes_single_json_line(self) -> None: self.assertEqual(json.loads(written), {"type": "ready", "params_schema": []}) +_FAKE_TEXGEN_GENERATOR = ''' +from pathlib import Path + +INSTANCES = [] + + +class FakeTexGen: + """ + Mimics an extension whose texture pipeline is built lazily on first use and + frees the shape pipeline first to make room for it. + + The first texture setup fails (missing `xatlas`), which is what issue #239 + reports; a later attempt succeeds once the dependency is available. + """ + + def __init__(self, model_dir, outputs_dir): + self.model_dir = model_dir + self.outputs_dir = outputs_dir + self._model = None + self.load_calls = 0 + self.texgen_attempts = 0 + INSTANCES.append(self) + + def is_loaded(self): + return self._model is not None + + def load(self): + self.load_calls += 1 + self._model = lambda image: "mesh" + + def unload(self): + self._model = None + + def _setup_texgen(self): + # Free the shape pipeline before building the texture pipeline. + self._model = None + self.texgen_attempts += 1 + if self.texgen_attempts == 1: + raise RuntimeError("No module named 'xatlas'") + # Texture setup succeeded: shape pipeline comes back. + self.load() + + def generate(self, image_bytes, params, progress_cb=None, cancel_event=None): + # Stands in for `self._model(image)` on a generator whose model is gone. + if self._model is None: + raise TypeError("'NoneType' object is not callable") + self._model(image_bytes) + if params.get("enable_texture"): + self._setup_texgen() + return Path("out.glb") +''' + + +class _RunnerDriver: + """Runs runner.main() against a throwaway extension dir.""" + + def __init__(self, generator_src: str, generator_class: str) -> None: + self.ext_dir = Path(tempfile.mkdtemp(prefix="modly-texgen-test-")) + (self.ext_dir / "generator.py").write_text(generator_src, encoding="utf-8") + (self.ext_dir / "manifest.json").write_text( + json.dumps({"id": "demo-ext", "generator_class": generator_class}), + encoding="utf-8", + ) + + def run(self, actions: list) -> list: + """Feeds actions on stdin, returns the parsed messages runner emitted.""" + original_ext_dir = runner.EXT_DIR + original_stdin = sys.stdin + original_module = sys.modules.pop("generator", None) + runner.EXT_DIR = self.ext_dir + sys.stdin = io.StringIO("".join(json.dumps(a) + "\n" for a in actions)) + out = io.StringIO() + try: + with redirect_stdout(out): + runner.main() + self.generator_module = sys.modules["generator"] + finally: + runner.EXT_DIR = original_ext_dir + sys.stdin = original_stdin + sys.modules.pop("generator", None) + if original_module is not None: + sys.modules["generator"] = original_module + return [json.loads(line) for line in out.getvalue().splitlines() if line.strip()] + + +class GeneratorLoadedStateTests(unittest.TestCase): + def test_reports_loaded_state(self) -> None: + gen = type("Gen", (), {"is_loaded": lambda self: True})() + self.assertTrue(runner._generator_is_loaded(gen)) + + def test_treats_raising_is_loaded_as_not_loaded(self) -> None: + class Gen: + def is_loaded(self): + raise RuntimeError("half-initialised") + + self.assertFalse(runner._generator_is_loaded(Gen())) + + def test_ensure_model_loaded_is_a_noop_when_already_loaded(self) -> None: + class Gen: + load_calls = 0 + + def is_loaded(self): + return True + + def load(self): + self.load_calls += 1 + + gen = Gen() + self.assertFalse(runner._ensure_model_loaded(gen)) + self.assertEqual(gen.load_calls, 0) + + def test_ensure_model_loaded_reloads_when_model_is_gone(self) -> None: + class Gen: + def __init__(self): + self._model = None + self.load_calls = 0 + + def is_loaded(self): + return self._model is not None + + def load(self): + self.load_calls += 1 + self._model = object() + + gen = Gen() + self.assertTrue(runner._ensure_model_loaded(gen)) + self.assertEqual(gen.load_calls, 1) + self.assertIsNotNone(gen._model) + + +class TextureSetupRecoveryTests(unittest.TestCase): + """ + Regression tests for issue #239: a failed lazy texture setup left the worker + with _model = None and no recovery path, so every later generation raised + "TypeError: 'NoneType' object is not callable" until the process was killed. + """ + + def setUp(self) -> None: + self.driver = _RunnerDriver(_FAKE_TEXGEN_GENERATOR, "FakeTexGen") + self.texture_run = { + "action": "generate", + "image_b64": "", + "params": {"enable_texture": True}, + } + + def test_worker_recovers_after_failed_texture_setup(self) -> None: + messages = self.driver.run([ + {"action": "load"}, + dict(self.texture_run, id="run-1"), + dict(self.texture_run, id="run-2"), + ]) + + by_id = {m.get("id"): m for m in messages if m.get("type") in ("done", "error")} + + # First run fails during texture setup and surfaces the real cause. + self.assertEqual(by_id["run-1"]["type"], "error") + self.assertIn("xatlas", by_id["run-1"]["message"]) + + # The retry must succeed instead of dying on a None model. + self.assertEqual( + by_id["run-2"]["type"], "done", + msg=f"retry did not recover: {by_id['run-2']}", + ) + self.assertNotIn("NoneType", json.dumps(by_id["run-2"])) + + # …and the worker's model is genuinely back. + gen = self.driver.generator_module.INSTANCES[0] + self.assertIsNotNone(gen._model) + self.assertTrue(gen.is_loaded()) + + def test_failed_run_reports_that_the_model_was_lost(self) -> None: + messages = self.driver.run([ + {"action": "load"}, + dict(self.texture_run, id="run-1"), + ]) + + error = next(m for m in messages if m.get("type") == "error") + self.assertIs(error["loaded"], False) + + def test_reload_before_generate_is_logged(self) -> None: + messages = self.driver.run([ + {"action": "load"}, + dict(self.texture_run, id="run-1"), + dict(self.texture_run, id="run-2"), + ]) + + logs = [m for m in messages if m.get("type") == "log"] + self.assertTrue( + any("reloaded before generating" in m.get("message", "") for m in logs), + msg=f"expected a reload log, got {logs}", + ) + + def test_successful_run_does_not_reload_the_model(self) -> None: + messages = self.driver.run([ + {"action": "load"}, + {"action": "generate", "id": "run-1", "image_b64": "", "params": {}}, + ]) + + self.assertEqual( + [m["type"] for m in messages if m.get("id") == "run-1"], ["done"] + ) + self.assertEqual(self.driver.generator_module.INSTANCES[0].load_calls, 1) + + if __name__ == "__main__": unittest.main() From f9a828fbc7745cb3caa663bde481615aa3e818ae Mon Sep 17 00:00:00 2001 From: Lion Rayonnant <106342136+lionrayonnant@users.noreply.github.com> Date: Mon, 17 Aug 2026 20:25:00 +0200 Subject: [PATCH 02/57] feat: add AMD ROCm support (Linux + Windows) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds AMD GPU detection and a PyTorch/ROCm redirect for extension setup, on top of the existing NVIDIA/MPS/CPU paths. - electron/main/gpu-detect.ts: detects AMD GPUs (KFD topology on Linux, Win32_VideoController on Windows), resolves the ROCm pip index and requirements. NVIDIA keeps detection priority; explicit overrides (MODLY_TORCH_FLAVOR, MODLY_ROCM_GFX, MODLY_ROCM_INDEX, MODLY_ROCM_TORCH_SPEC) are available for machines the auto-detection gets wrong. - electron/main/setup-launcher.ts: extracted from ipc-handlers.ts, adds a compatibility shim that redirects an extension's pip torch install to ROCm wheels — needed because most third-party extension setup.py scripts predate AMD support and hardcode a CUDA index. - api/routers/extensions.py: the FastAPI-side GPU detection no longer mistakes a ROCm build's device capability for CUDA compute capability (both answer torch.cuda.get_device_capability the same way), and reads the AMD compute target from the same KFD topology as the Electron side. - electron/main/copy-runtime.ts: fixes an unrelated but blocking AppImage bug found while verifying this end-to-end — fs.cp rewrote the bundled Python runtime's relative symlinks into absolute paths pointing at the ephemeral AppImage mount, so every extension venv died on the next launch. verbatimSymlinks keeps them relative. - docs/running-on-amd-rocm.md, arch/decisions/AMD-ROCM-SUPPORT.md: usage, verified configuration, and known limitations. Verified end-to-end on a Radeon RX 9060 XT (gfx1200): detection, ROCm wheel install (torch 2.13.0+rocm7.2), and a full image-to-3D generation through hunyuan3d-mini all complete successfully on the GPU. --- README.md | 3 + api/routers/extensions.py | 81 +++++- arch/decisions/AMD-ROCM-SUPPORT.md | 99 +++++++ arch/decisions/README.md | 1 + docs/running-on-amd-rocm.md | 193 +++++++++++++ electron/main/copy-runtime.test.mjs | 93 ++++++ electron/main/copy-runtime.ts | 35 +++ electron/main/gpu-detect.test.mjs | 230 +++++++++++++++ electron/main/gpu-detect.ts | 390 ++++++++++++++++++++++++++ electron/main/ipc-handlers.ts | 191 +++---------- electron/main/python-setup.ts | 5 +- electron/main/setup-launcher.test.mjs | 224 +++++++++++++++ electron/main/setup-launcher.ts | 231 +++++++++++++++ 13 files changed, 1611 insertions(+), 165 deletions(-) create mode 100644 arch/decisions/AMD-ROCM-SUPPORT.md create mode 100644 docs/running-on-amd-rocm.md create mode 100644 electron/main/copy-runtime.test.mjs create mode 100644 electron/main/copy-runtime.ts create mode 100644 electron/main/gpu-detect.test.mjs create mode 100644 electron/main/gpu-detect.ts create mode 100644 electron/main/setup-launcher.test.mjs create mode 100644 electron/main/setup-launcher.ts diff --git a/README.md b/README.md index 11184ddc..293368e3 100644 --- a/README.md +++ b/README.md @@ -69,6 +69,9 @@ npm run build ## Platform notes +- AMD GPUs are supported through ROCm on Linux and Windows: a Radeon card is detected + automatically and extensions are steered to ROCm PyTorch wheels, with no ROCm install + required. See [docs/running-on-amd-rocm.md](docs/running-on-amd-rocm.md). - macOS support targets Apple Silicon only. - macOS uses native window controls. Windows and Linux keep the existing custom controls. - The top bar includes a live RAM indicator sourced from the main process. diff --git a/api/routers/extensions.py b/api/routers/extensions.py index 2313d3b3..ce549035 100644 --- a/api/routers/extensions.py +++ b/api/routers/extensions.py @@ -1,6 +1,9 @@ import asyncio +import json +import re import subprocess import sys +from pathlib import Path from fastapi import APIRouter, HTTPException router = APIRouter(tags=["extensions"]) @@ -43,14 +46,29 @@ async def setup_extension(ext_id: str): return {"status": "skipped", "reason": "no setup.py"} # Detect GPU compute capability - gpu_sm = _detect_gpu_sm() + gpu_sm = _detect_gpu_sm() + gfx_target = _detect_gfx_target() + + # Pass arguments as JSON so setup.py sees torch_flavor. Note this endpoint is + # a fallback: Electron normally runs setup.py itself, and only that path gets + # the ROCm index rewriting for extensions that ignore torch_flavor. + args = json.dumps({ + "python_exe": sys.executable, + "ext_dir": str(ext_dir), + "gpu_sm": gpu_sm, + "cuda_version": 0, + "accelerator": "rocm" if gfx_target else ("cuda" if gpu_sm else "cpu"), + "torch_flavor": "rocm" if gfx_target else ("cuda" if gpu_sm else "cpu"), + "gfx_target": gfx_target, + "platform": sys.platform, + }) # Run setup.py using Modly's embedded Python (sys.executable) loop = asyncio.get_running_loop() result = await loop.run_in_executor( None, lambda: subprocess.run( - [sys.executable, str(setup_py), sys.executable, str(ext_dir), str(gpu_sm)], + [sys.executable, str(setup_py), args], capture_output=True, text=True, ) @@ -60,9 +78,10 @@ async def setup_extension(ext_id: str): raise HTTPException(500, f"setup.py failed:\n{result.stderr}") return { - "status": "ok", - "gpu_sm": gpu_sm, - "output": result.stdout, + "status": "ok", + "gpu_sm": gpu_sm, + "gfx_target": gfx_target, + "output": result.stdout, } @@ -74,12 +93,62 @@ async def extension_errors(): def _detect_gpu_sm() -> int: - """Returns GPU compute capability as integer (e.g. 86 for SM 8.6), or 0 if no GPU.""" + """ + Returns GPU compute capability as integer (e.g. 86 for SM 8.6), or 0 if no GPU. + + Returns 0 on ROCm too. PyTorch's HIP build answers the whole torch.cuda API, + so get_device_capability() happily reports (12, 0) for a gfx1200 Radeon — + indistinguishable from an sm_120 Blackwell, which would send setup.py to the + CUDA 12.8 index. AMD cards are identified by _detect_gfx_target() instead. + """ try: import torch + if torch.version.hip: + return 0 if torch.cuda.is_available(): major, minor = torch.cuda.get_device_capability(0) return major * 10 + minor except Exception: pass return 0 + + +def _detect_gfx_target() -> str: + """ + Returns the ROCm compute target (e.g. "gfx1200"), or "" when there is no AMD GPU. + + Reads the kernel's KFD topology rather than asking torch: this process runs + in Modly's main venv, which has no torch at all (see api/requirements.txt). + The amdgpu driver publishes the target on its own, so no ROCm install is + needed either. Mirrors electron/main/gpu-detect.ts. + """ + kfd_nodes = Path("/sys/class/kfd/kfd/topology/nodes") + if not Path("/dev/kfd").exists() or not kfd_nodes.is_dir(): + return "" + + def _prop(text: str, key: str) -> int: + match = re.search(rf"^{key}\s+(\d+)\s*$", text, re.M) + return int(match.group(1)) if match else 0 + + try: + nodes = sorted(kfd_nodes.iterdir(), key=lambda p: int(p.name) if p.name.isdigit() else 0) + except OSError: + return "" + + for node in nodes: + try: + text = (node / "properties").read_text() + except OSError: + continue + # Node 0 is the CPU node (simd_count 0) and carries no compute target. + if _prop(text, "simd_count") <= 0: + continue + # major*10000 + minor*100 + step, minor and step read as hex digits. + version = _prop(text, "gfx_target_version") + if version <= 0: + continue + major, minor, step = version // 10000, (version % 10000) // 100, version % 100 + if major <= 0 or minor > 15 or step > 15: + continue + return f"gfx{major}{minor:x}{step:x}" + return "" diff --git a/arch/decisions/AMD-ROCM-SUPPORT.md b/arch/decisions/AMD-ROCM-SUPPORT.md new file mode 100644 index 00000000..7cd4ee13 --- /dev/null +++ b/arch/decisions/AMD-ROCM-SUPPORT.md @@ -0,0 +1,99 @@ +# AMD-ROCM-SUPPORT + +- Status: proposed +- Date: 2026-08-17 + +## Decision + +Modly supports AMD Radeon GPUs through ROCm on Linux and Windows. Detection is +automatic and requires no ROCm installation on the user's machine. + +Scope and operating rules: + +- Detection is centralised in `electron/main/gpu-detect.ts` and produces, for AMD + machines, a compute target plus the pip index and requirements an extension's + torch install must end up using. +- NVIDIA keeps detection priority. On a machine with both vendors the existing + CUDA behaviour is unchanged. +- AMD machines report `gpu_sm = 0` and `cuda_version = 0`, never a synthesised + compute capability. +- Extensions are told the flavour via a `torch_flavor` setup argument. Extensions + that ignore it are corrected by a rewrite shim in `electron/main/setup-launcher.ts`. +- The ROCm wheel source differs by platform: `download.pytorch.org/whl/rocm7.2` + on Linux, `repo.amd.com/rocm/whl-multi-arch/` on Windows. +- Every automatic choice has an environment-variable override. See + `docs/running-on-amd-rocm.md`. + +## Context + +Modly never installs PyTorch itself. Each extension ships a `setup.py` that +creates its own venv and installs torch from an index it hardcodes — and those +scripts are third-party code in separate GitHub repositories that Modly cannot +edit. Before this work `detectGpuInfo()` only probed `nvidia-smi`, so an AMD +machine was reported as `accelerator: 'cpu'` with `gpu_sm: 0`, which sent every +extension down its legacy CUDA 11.8 branch and installed a torch that cannot see +the GPU at all. + +That leaves two distinct problems, and both have to be solved: + +- Extensions that *do* understand AMD were never told. The official + `modly-hunyuan3d-mini-extension` has accepted a `torch_flavor: "rocm"` argument + for some time; Modly simply never sent it. +- Extensions that don't understand AMD — `triposg`, `trellis2`, and the rest — + hardcode `--index-url .../whl/cu124` and have no branch to select. Passing an + argument achieves nothing for them. + +The wheel sources are also not symmetric across platforms. `download.pytorch.org` +publishes no ROCm wheels for Windows at all; AMD's own multi-arch index does, but +there the compute target is selected by a pip extra (`torch[device-gfx1200]`) +rather than by the index URL, which means Windows needs the compute target +*before* the install, not after. + +## Consequences + +- **A rewrite shim is unavoidable.** Correcting extensions we cannot edit means + intercepting their pip calls. The launcher already patched `subprocess` for two + other compatibility fixes, so ROCm redirection joins those rather than + introducing a new mechanism. +- **The decision is made in TypeScript, applied in Python.** The launcher is an + inline Python string that cannot be unit-tested in isolation, so index and + requirement resolution lives in `gpu-detect.ts` and reaches the launcher as + environment variables. The launcher is separately exercised end-to-end by + `setup-launcher.test.mjs`, which runs it against the command shapes the + official extensions actually use. +- **`gpu_sm = 0` is load-bearing, not a placeholder.** Extensions written before + `torch_flavor` branch on that number, and 0 selects their most conservative + path. It also keeps them off `rembg[gpu]`, whose `onnxruntime-gpu` is + CUDA-only. Reporting a synthesised capability instead would break both. + PyTorch's HIP build answers the whole `torch.cuda` API, so + `get_device_capability()` reports `(12, 0)` for a gfx1200 Radeon — + indistinguishable from an sm_120 Blackwell. `api/routers/extensions.py` guards + on `torch.version.hip` for this reason. +- **Compute-target discovery is platform-specific.** Linux reads + `gfx_target_version` from the kernel's KFD topology, which needs no ROCm + install and no external binary. Windows has no equivalent, so it maps PCI + device ids from `Win32_VideoController` through a table keyed by silicon. That + table is a maintenance surface: new AMD silicon needs an entry, and an unmapped + AMD card falls back to CPU with an actionable message rather than guessing a + wheel. +- **The Linux and Windows torch versions diverge.** Linux gets unpinned wheels + from the pytorch.org ROCm index (currently torch 2.11+); Windows gets a pinned + pair from AMD's index. Extension code written against torch 2.6/2.7 may not + survive that jump, which is why `MODLY_ROCM_INDEX` and `MODLY_ROCM_TORCH_SPEC` + exist as first-class escape hatches rather than debug affordances. +- **Linux is verified, Windows is not.** On a Radeon RX 9060 XT (gfx1200), + `torch 2.13.0+rocm7.2` loads, rocBLAS and MIOpen kernels execute, 14 GB of the + card's 16 GB allocates and reads back cleanly ([ROCm #6295](https://github.com/ROCm/ROCm/issues/6295), + which reports this card capped near 8 GB, did not reproduce), and a full + image-to-3D generation completes through the normal `ExtensionProcess` path in + 221 s. For Windows the wheel URLs, `cp311` availability and index layout were + checked, but no end-to-end run has been performed. +- **This work also required fixing an unrelated AppImage bug** to be verifiable + at all. `ensureStableEmbeddedPython()` copied the bundled runtime with + `fs.cp`, which rewrites relative symlinks into absolute paths pointing back at + the ephemeral `/tmp/.mount_Modly-XXXXXX/` mount — so the "stable" copy was not + stable, and every extension venv built from it died on the next launch with a + misleading `No module named 'PIL'`. See `electron/main/copy-runtime.ts`. +- **Texture generation is out of scope.** `api/texture_baker` already carries a + HIP build path, but those native extensions are not built as part of standard + extension setup, so texture generation is not covered by this ADR. diff --git a/arch/decisions/README.md b/arch/decisions/README.md index 063df577..925441ea 100644 --- a/arch/decisions/README.md +++ b/arch/decisions/README.md @@ -8,3 +8,4 @@ single reviewable document. Current ADRs: - [APPLE-SILICON-SUPPORT](./APPLE-SILICON-SUPPORT.md) +- [AMD-ROCM-SUPPORT](./AMD-ROCM-SUPPORT.md) diff --git a/docs/running-on-amd-rocm.md b/docs/running-on-amd-rocm.md new file mode 100644 index 00000000..5fe60f8a --- /dev/null +++ b/docs/running-on-amd-rocm.md @@ -0,0 +1,193 @@ +# Running Modly on an AMD GPU (ROCm) + +Modly's default GPU path is NVIDIA/CUDA (plus Metal/MPS on Apple Silicon). This page +covers the AMD path: how Modly detects a Radeon card, which PyTorch wheels it steers +extensions to, and what to do when the automatic choice is wrong. + +Nothing here needs a manual setup: install the app, install an extension, and the AMD +path is taken automatically when an AMD GPU is present. + +--- + +## 1. Requirements + +| | Linux | Windows | +|---|---|---| +| GPU | RDNA 2 or newer discrete Radeon (see the table below) | same | +| Driver | in-tree `amdgpu` kernel driver — any recent distro kernel | Adrenalin with the ROCm runtime (26.2.2 or newer) | +| ROCm install | **not required** — the PyTorch ROCm wheels bundle their own runtime | not required | +| Device access | `/dev/kfd` and `/dev/dri/renderD*` must be readable by your user | n/a | + +On most distributions `/dev/kfd` is world-accessible. If it is not, add yourself to the +`render` group and log back in: + +```bash +ls -l /dev/kfd # crw-rw-rw- → nothing to do +sudo usermod -aG render "$USER" +``` + +Integrated (APU) graphics are not mapped by the Windows detection table. They work on +Linux whenever the kernel publishes a compute target, but they are not a target Modly +tests against. + +--- + +## 2. How detection works + +Detection lives in `electron/main/gpu-detect.ts` and runs before an extension's +`setup.py`, in this order: + +1. `MODLY_TORCH_FLAVOR` (`cuda` / `rocm` / `cpu`) — an explicit override wins over everything. +2. Apple Silicon → MPS. +3. `nvidia-smi` → CUDA. **NVIDIA keeps priority**: on a machine with both vendors, + nothing about the existing CUDA behaviour changes. +4. AMD: + - **Linux** — reads `gfx_target_version` from the kernel's KFD topology + (`/sys/class/kfd/kfd/topology/nodes/*/properties`). No ROCm install and no external + binary involved; the `amdgpu` driver publishes this on its own. + - **Windows** — reads the PCI device id from `Win32_VideoController` via PowerShell + and maps it to a compute target. +5. Otherwise → CPU. + +You can check what was detected in the logs (Settings → Logs), on the line beginning +`[ext-setup] accelerator=`: + +``` +[ext-setup] accelerator=rocm gfx=gfx1200 torch_index=https://download.pytorch.org/whl/rocm7.2 +``` + +### Supported compute targets + +| Silicon | Compute target | Cards | +|---|---|---| +| Navi 44 | `gfx1200` | RX 9060 XT | +| Navi 48 | `gfx1201` | RX 9070, RX 9070 XT, RX 9070 GRE, AI PRO R9700 | +| Navi 31 | `gfx1100` | RX 7900 XT/XTX/GRE, PRO W7800/W7900 | +| Navi 32 | `gfx1101` | RX 7700 XT, RX 7800 XT, PRO W7700 | +| Navi 33 | `gfx1102` | RX 7600 series, PRO W7500/W7600 | +| Navi 21 | `gfx1030` | RX 6800/6800 XT/6900 XT/6950 XT, PRO W6800 | +| Navi 22/23/24 | `gfx1031` / `gfx1032` / `gfx1034` | RX 6700 / 6600 / 6400 series | + +On Linux the target is read from the kernel, so any AMD GPU the kernel knows about is +picked up — the table above only bounds the *Windows* mapping. AMD officially supports +`gfx1030`, `gfx110x` and `gfx120x`; the RDNA 2 mid-range entries work in practice but +are not part of AMD's supported matrix. + +--- + +## 3. Which wheels get installed + +Modly does not install PyTorch itself — each extension's `setup.py` does, from an index +it hardcodes. On an AMD machine Modly corrects that choice two ways: + +- It passes `torch_flavor: "rocm"` in the setup arguments. Extensions that know about + AMD (the official `hunyuan3d-mini` one does) branch on it themselves. +- For every extension that doesn't, a compatibility shim in the setup launcher rewrites + the pip command: the CUDA (or CPU) index is swapped for the ROCm one, and the pinned + `torch==…` requirements are replaced. Non-torch installs are left untouched. + +| Platform | Index | Requirements | +|---|---|---| +| Linux | `https://download.pytorch.org/whl/rocm7.2` | `torch`, `torchvision`, unpinned | +| Windows | `https://repo.amd.com/rocm/whl-multi-arch/` | `torch[device-gfxNNNN]==2.11.0+rocm7.14.0` and matching `torchvision` | + +The two differ because `download.pytorch.org` publishes no ROCm wheels for Windows at +all. AMD's multi-arch index does, and there the compute target is selected by a pip +extra rather than by the index URL. + +`sm` and `cuda_version` are deliberately reported as `0` for AMD. Extensions written +before `torch_flavor` existed branch on those numbers, and `0` sends them down their +most conservative path — which also keeps them off `rembg[gpu]`, whose `onnxruntime-gpu` +is CUDA-only. + +--- + +## 4. Verifying an install + +After installing an extension (or running **Repair** on it from the Models page), check +what actually landed in its venv: + +```bash +# Linux; adjust to your extensions directory +EXT=~/Documents/Modly/extensions/hunyuan3d-mini +"$EXT/venv/bin/python" -c "import torch; print(torch.__version__, '| hip', torch.version.hip, '| avail', torch.cuda.is_available(), '|', torch.cuda.get_device_name(0))" +``` + +Expected: a `+rocm…` version, a non-null `hip`, `True`, and your card's name. A version +ending in `+cu118` or `+cu124` means the AMD path was not taken — check the detection +line in the logs. + +Then confirm real VRAM is usable, not just that the device opens: + +```bash +"$EXT/venv/bin/python" -c "import torch; x = torch.empty(5_000_000_000, dtype=torch.float16, device='cuda'); torch.cuda.synchronize(); print('10 GB allocated OK')" +``` + +--- + +## 5. Escape hatches + +All of these are environment variables read at detection time — set them before +launching Modly. + +| Variable | Effect | +|---|---| +| `MODLY_TORCH_FLAVOR` | `cuda` / `rocm` / `cpu`. Forces the path, skipping detection entirely. | +| `MODLY_ROCM_GFX` | Forces the compute target (e.g. `gfx1201`). Needed for an AMD card the Windows table doesn't map. | +| `MODLY_ROCM_INDEX` | Overrides the pip index, e.g. `https://download.pytorch.org/whl/rocm6.4` to fall back to torch 2.8/2.9. | +| `MODLY_ROCM_TORCH_SPEC` | Overrides the requirements entirely, space-separated: `"torch==2.9.1 torchvision==0.24.1"`. | +| `HSA_OVERRIDE_GFX_VERSION` | ROCm's own override, for cards without native kernels (e.g. `10.3.0` on an unsupported RDNA 2 part). Not needed on RDNA 3/4. | + +After changing any of these, run **Repair** on the extension so its venv is rebuilt. + +--- + +## 6. Troubleshooting + +| Symptom | Cause | Fix | +|---|---|---| +| Logs say `accelerator=cpu` on an AMD machine | `/dev/kfd` missing or unreadable (Linux), or an unmapped PCI id (Windows) | Check `ls -l /dev/kfd`; on Windows set `MODLY_ROCM_GFX` | +| `torch.__version__` ends in `+cu118` | Extension venv predates AMD support | Run **Repair** on the extension | +| `torch.cuda.is_available()` is `False` with a ROCm build | Device nodes not accessible from the process | Add your user to the `render` group, log back in | +| `HIP error: invalid device function` | Wheel has no kernels for your card | Set `HSA_OVERRIDE_GFX_VERSION` to a supported nearby target | +| Extension imports fail after install (`diffusers`/`transformers`) | The ROCm index only carries recent torch (2.11+), newer than some extensions expect | `MODLY_ROCM_INDEX=https://download.pytorch.org/whl/rocm6.4`, then **Repair** | +| Allocations above ~8 GB segfault on a 16 GB RX 9060 XT | [ROCm issue #6295](https://github.com/ROCm/ROCm/issues/6295) — did not reproduce on torch 2.13.0+rocm7.2 (see below) | If you hit it, pin an older stack via `MODLY_ROCM_INDEX` | + +--- + +## 7. Verified configuration and limitations + +The Linux path was measured end to end on a Radeon RX 9060 XT (Navi 44, gfx1200, +16 GB), CachyOS kernel 7.1.8, with the stack this page installs by default: + +- `torch 2.13.0+rocm7.2` / `torchvision 0.28.0+rocm7.2`, HIP runtime 7.2.53211 +- Detection reported `gfx1200`; `torch.cuda.is_available()` is `True` and + `get_device_properties(0).gcnArchName` is `gfx1200` +- rocBLAS (fp16 matmul) and MIOpen (conv2d) kernels both execute — no + `no kernel image is available` failures +- Allocations of 4/6/8/10/12/14 GB all succeeded, filled and read back. + [ROCm #6295](https://github.com/ROCm/ROCm/issues/6295), which reports this exact + card capped near 8 GB with a segfault, **did not reproduce** on this stack. + It remains open upstream, so it may still affect older ROCm builds. +- A full image-to-3D generation with `hunyuan3d-mini` completed through the normal + `ExtensionProcess` subprocess path: model load 18 s, generation 221 s, 15.8 MB GLB. + +### Attention kernels + +During generation PyTorch emits: + +> Mem Efficient attention on Current AMD GPU is still experimental. Enable it with +> `TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL=1` + +Generation works without it — `scaled_dot_product_attention` falls back to a +slower path. Setting `TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL=1` before launching +Modly enables the memory-efficient kernels, at the cost of running code AMD still +labels experimental on RDNA 3/4. Modly does not set it for you. + +Limitations: + +- **Windows is untested** by the Modly maintainers. The wheel URLs and the `cp311` + availability were verified, but no end-to-end run has been done on that path. +- **Texture generation** relies on optional native extensions (`texture_baker`, + `uv_unwrapper`) whose CUDA kernels have a HIP path but are not built as part of the + standard extension setup. diff --git a/electron/main/copy-runtime.test.mjs b/electron/main/copy-runtime.test.mjs new file mode 100644 index 00000000..fd8b9ba4 --- /dev/null +++ b/electron/main/copy-runtime.test.mjs @@ -0,0 +1,93 @@ +/** + * Guards the symlink property that makes the AppImage's "stable" Python copy + * actually stable. Getting this wrong is silent at copy time and only breaks on + * the *next* launch, once the source mount is gone — so it needs a test that + * checks the links rather than the copy succeeding. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { + mkdtempSync, mkdirSync, writeFileSync, symlinkSync, readlinkSync, rmSync, existsSync, +} from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve, isAbsolute } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-copy-test-')), 'copy-runtime.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/copy-runtime.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { copyRuntimeTree } = loadModule() + +/** Builds a miniature of the bundled runtime: a real binary plus relative links. */ +function makeRuntime(root) { + mkdirSync(join(root, 'bin'), { recursive: true }) + mkdirSync(join(root, 'lib'), { recursive: true }) + writeFileSync(join(root, 'bin', 'python3.11'), '#!/bin/sh\n', 'utf8') + symlinkSync('python3.11', join(root, 'bin', 'python3')) + symlinkSync('python3.11', join(root, 'bin', 'python')) + writeFileSync(join(root, 'lib', 'libpython3.11.so.1.0'), '', 'utf8') + symlinkSync('libpython3.11.so.1.0', join(root, 'lib', 'libpython3.11.so')) +} + +test('copyRuntimeTree keeps relative symlinks relative', async () => { + const dir = mkdtempSync(join(tmpdir(), 'modly-runtime-')) + const source = join(dir, 'mount', 'python-embed') + const dest = join(dir, 'stable', 'python-embed') + makeRuntime(source) + + await copyRuntimeTree(source, dest) + + for (const link of ['bin/python3', 'bin/python', 'lib/libpython3.11.so']) { + const target = readlinkSync(join(dest, link)) + assert.ok( + !isAbsolute(target), + `${link} was rewritten to an absolute path (${target}); it would point back at the source mount`, + ) + assert.ok(!target.includes(source), `${link} still references the source tree`) + } + + rmSync(dir, { recursive: true, force: true }) +}) + +test('the copy survives the source being deleted', async () => { + // This is the actual failure mode: the AppImage mount disappears between + // launches, and every venv built from the copy dies with it. + const dir = mkdtempSync(join(tmpdir(), 'modly-runtime-')) + const source = join(dir, 'mount', 'python-embed') + const dest = join(dir, 'stable', 'python-embed') + makeRuntime(source) + + await copyRuntimeTree(source, dest) + rmSync(join(dir, 'mount'), { recursive: true, force: true }) + + // existsSync follows symlinks, so this is false for a dangling link — exactly + // the check generator_registry.py's _venv_python(...).exists() performs. + assert.ok(existsSync(join(dest, 'bin', 'python3')), 'bin/python3 is dangling after the source went away') + assert.ok(existsSync(join(dest, 'lib', 'libpython3.11.so')), 'libpython3.11.so is dangling') +}) + +test('copyRuntimeTree copies regular files and directory structure', async () => { + const dir = mkdtempSync(join(tmpdir(), 'modly-runtime-')) + const source = join(dir, 'mount', 'python-embed') + const dest = join(dir, 'stable', 'python-embed') + makeRuntime(source) + + await copyRuntimeTree(source, dest) + + assert.ok(existsSync(join(dest, 'bin', 'python3.11'))) + assert.ok(existsSync(join(dest, 'lib', 'libpython3.11.so.1.0'))) + + rmSync(dir, { recursive: true, force: true }) +}) diff --git a/electron/main/copy-runtime.ts b/electron/main/copy-runtime.ts new file mode 100644 index 00000000..eb062f5c --- /dev/null +++ b/electron/main/copy-runtime.ts @@ -0,0 +1,35 @@ +/** + * Copying the bundled Python runtime out of an ephemeral AppImage mount. + * + * Kept apart from python-setup.ts (which imports electron) so the symlink + * behaviour this depends on can be tested — it is subtle, silent when wrong, + * and only shows up on the *next* launch. + */ + +import { cp } from 'fs/promises' + +/** + * Copies a self-contained runtime tree, keeping relative symlinks relative. + * + * `verbatimSymlinks` is the whole point. Without it `fs.cp` resolves a relative + * link (`bin/python3 -> python3.11`) into an absolute path pointing back at the + * *source* tree. When the source is an AppImage mount at + * /tmp/.mount_Modly-XXXXXX/ — which is a different path on every launch — the + * copy silently keeps a hard dependency on a directory that is about to vanish: + * + * - `bin/python3` in the "stable" copy points into the old mount + * - `sys._base_executable` of any venv made from it inherits that path + * - every extension venv records it in pyvenv.cfg and as a bin/python symlink + * - on the next launch the mount is gone, so every extension venv is dead + * + * The user-visible symptom is remote from the cause: the extension registry + * finds no usable venv, falls back to importing generator.py in the main API + * process, and reports a missing third-party module such as `No module named 'PIL'`. + */ +export async function copyRuntimeTree(source: string, destination: string): Promise { + await cp(source, destination, { + recursive: true, + preserveTimestamps: true, + verbatimSymlinks: true, + }) +} diff --git a/electron/main/gpu-detect.test.mjs b/electron/main/gpu-detect.test.mjs new file mode 100644 index 00000000..c4a2edcb --- /dev/null +++ b/electron/main/gpu-detect.test.mjs @@ -0,0 +1,230 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-gpu-test-')), 'gpu-detect.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/gpu-detect.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const mod = loadModule() + +// ─── KFD topology parsing ───────────────────────────────────────────────────── + +test('formatGfxTarget decodes gfx_target_version across GPU generations', () => { + // Encoding is major*10000 + minor*100 + step, minor and step read as hex + // digits. 120000 is the value this machine's RX 9060 XT actually reports. + assert.equal(mod.formatGfxTarget(120000), 'gfx1200') // Navi 44 / RX 9060 XT + assert.equal(mod.formatGfxTarget(120001), 'gfx1201') // Navi 48 + assert.equal(mod.formatGfxTarget(110000), 'gfx1100') // Navi 31 + assert.equal(mod.formatGfxTarget(110002), 'gfx1102') // Navi 33 + assert.equal(mod.formatGfxTarget(100300), 'gfx1030') // Navi 21 + assert.equal(mod.formatGfxTarget(90402), 'gfx942') // MI300, hex step + assert.equal(mod.formatGfxTarget(90010), 'gfx90a') // MI200, hex step +}) + +test('formatGfxTarget rejects the "no GPU" and malformed encodings', () => { + assert.equal(mod.formatGfxTarget(0), null) + assert.equal(mod.formatGfxTarget(-1), null) + assert.equal(mod.formatGfxTarget(1.5), null) +}) + +test('parseKfdGfxTarget skips the CPU node and reads the first GPU', () => { + // Verbatim shape of /sys/class/kfd/kfd/topology/nodes/*/properties + const cpuNode = 'cpu_cores_count 32\nsimd_count 0\ngfx_target_version 0\n' + const gpuNode = 'cpu_cores_count 0\nsimd_count 64\ngfx_target_version 120000\n' + + assert.equal(mod.parseKfdGfxTarget([cpuNode, gpuNode]), 'gfx1200') +}) + +test('parseKfdGfxTarget returns null without a usable GPU node', () => { + assert.equal(mod.parseKfdGfxTarget([]), null) + assert.equal(mod.parseKfdGfxTarget(['simd_count 0\ngfx_target_version 0\n']), null) + // A node advertising SIMDs but no compute target is not something we can target + assert.equal(mod.parseKfdGfxTarget(['simd_count 64\ngfx_target_version 0\n']), null) +}) + +test('parseKfdGfxTarget does not confuse a key with its prefix', () => { + // simd_count must not be satisfied by e.g. "max_simd_count" + const node = 'max_simd_count 999\nsimd_count 64\ngfx_target_version 110002\n' + assert.equal(mod.parseKfdGfxTarget([node]), 'gfx1102') +}) + +// ─── nvidia-smi parsing (non-regression) ────────────────────────────────────── + +test('parseNvidiaSmi maps compute cap and driver version to CUDA version', () => { + assert.deepEqual(mod.parseNvidiaSmi('8.6, 551.61\n'), { sm: 86, cudaVersion: 124 }) + assert.deepEqual(mod.parseNvidiaSmi('12.0, 572.16\n'), { sm: 120, cudaVersion: 128 }) + assert.deepEqual(mod.parseNvidiaSmi('6.1, 470.82\n'), { sm: 61, cudaVersion: 118 }) +}) + +test('parseNvidiaSmi falls back to sm 86 on an unparseable compute cap', () => { + assert.deepEqual(mod.parseNvidiaSmi('N/A, 551.61\n'), { sm: 86, cudaVersion: 124 }) +}) + +test('parseNvidiaSmi returns null on empty output', () => { + assert.equal(mod.parseNvidiaSmi(''), null) + assert.equal(mod.parseNvidiaSmi(' \n'), null) +}) + +// ─── Windows adapter mapping ────────────────────────────────────────────────── + +test('parseWindowsVideoControllers accepts both the object and array JSON shapes', () => { + const single = mod.parseWindowsVideoControllers( + '{"Name":"AMD Radeon RX 9060 XT","PNPDeviceID":"PCI\\\\VEN_1002&DEV_7590&SUBSYS_06391043&REV_C0\\\\4&1"}', + ) + assert.deepEqual(single, [ + { name: 'AMD Radeon RX 9060 XT', pnpDeviceId: 'PCI\\VEN_1002&DEV_7590&SUBSYS_06391043&REV_C0\\4&1' }, + ]) + + const many = mod.parseWindowsVideoControllers( + '[{"Name":"A","PNPDeviceID":"PCI\\\\VEN_1002&DEV_7550"},{"Name":"B","PNPDeviceID":"PCI\\\\VEN_8086&DEV_1234"}]', + ) + assert.equal(many.length, 2) +}) + +test('parseWindowsVideoControllers survives malformed PowerShell output', () => { + assert.deepEqual(mod.parseWindowsVideoControllers('not json'), []) + assert.deepEqual(mod.parseWindowsVideoControllers('{"Name":"No id"}'), []) +}) + +test('parseAmdPciDeviceId only matches AMD vendor ids', () => { + assert.equal(mod.parseAmdPciDeviceId('PCI\\VEN_1002&DEV_7590&SUBSYS_0'), '7590') + assert.equal(mod.parseAmdPciDeviceId('PCI\\VEN_10DE&DEV_2684'), null) +}) + +test('resolveWindowsGfxTarget maps device ids by silicon, not by marketing range', () => { + const target = (deviceId) => + mod.resolveWindowsGfxTarget([{ name: 'x', pnpDeviceId: `PCI\\VEN_1002&DEV_${deviceId}` }]).gfxTarget + + assert.equal(target('7590'), 'gfx1200') // Navi 44 + assert.equal(target('7550'), 'gfx1201') // Navi 48 + assert.equal(target('744C'), 'gfx1100') // Navi 31, uppercase from WMI + assert.equal(target('747e'), 'gfx1101') // Navi 32 + // 0x73f0 sells as "RX 7600M XT" but is Navi 33 — it must not land with its + // 0x73xx RDNA2 neighbours. + assert.equal(target('73f0'), 'gfx1102') + assert.equal(target('73bf'), 'gfx1030') // Navi 21 +}) + +test('resolveWindowsGfxTarget reports unmapped AMD cards instead of guessing', () => { + const result = mod.resolveWindowsGfxTarget([ + { name: 'AMD Radeon RX 9999', pnpDeviceId: 'PCI\\VEN_1002&DEV_FFFF' }, + ]) + assert.equal(result.gfxTarget, null) + assert.deepEqual(result.amdAdapters, ['AMD Radeon RX 9999']) +}) + +test('resolveWindowsGfxTarget ignores non-AMD adapters', () => { + const result = mod.resolveWindowsGfxTarget([ + { name: 'NVIDIA RTX 4090', pnpDeviceId: 'PCI\\VEN_10DE&DEV_2684' }, + ]) + assert.equal(result.gfxTarget, null) + assert.deepEqual(result.amdAdapters, []) +}) + +// ─── ROCm wheel source resolution ───────────────────────────────────────────── + +test('resolveRocmTorchSpec uses the unpinned pytorch.org index on Linux', () => { + const { indexUrl, specs } = mod.resolveRocmTorchSpec('linux', 'gfx1200', {}) + assert.equal(indexUrl, 'https://download.pytorch.org/whl/rocm7.2') + assert.deepEqual(specs, ['torch', 'torchvision']) +}) + +test('resolveRocmTorchSpec uses AMD\'s index with a device extra on Windows', () => { + // download.pytorch.org publishes no ROCm wheels for Windows at all. + const { indexUrl, specs } = mod.resolveRocmTorchSpec('win32', 'gfx1200', {}) + assert.equal(indexUrl, 'https://repo.amd.com/rocm/whl-multi-arch/') + assert.deepEqual(specs, [ + 'torch[device-gfx1200]==2.11.0+rocm7.14.0', + 'torchvision[device-gfx1200]==0.26.0+rocm7.14.0', + ]) +}) + +test('resolveRocmTorchSpec honours MODLY_ROCM_INDEX and MODLY_ROCM_TORCH_SPEC', () => { + const rolledBack = mod.resolveRocmTorchSpec('linux', 'gfx1200', { + MODLY_ROCM_INDEX: 'https://download.pytorch.org/whl/rocm6.4', + MODLY_ROCM_TORCH_SPEC: 'torch==2.8.0 torchvision==0.23.0', + }) + assert.equal(rolledBack.indexUrl, 'https://download.pytorch.org/whl/rocm6.4') + assert.deepEqual(rolledBack.specs, ['torch==2.8.0', 'torchvision==0.23.0']) +}) + +test('resolveRocmTorchSpec stays unpinned on Windows without a compute target', () => { + // No target means no device extra to ask for; better an install that fails + // loudly than one silently pinned to the wrong architecture. + const { specs } = mod.resolveRocmTorchSpec('win32', null, {}) + assert.deepEqual(specs, ['torch', 'torchvision']) +}) + +test('torchFlavorFor maps accelerators to the setup.py argument', () => { + assert.equal(mod.torchFlavorFor('rocm'), 'rocm') + assert.equal(mod.torchFlavorFor('cuda'), 'cuda') + assert.equal(mod.torchFlavorFor('cpu'), 'cpu') + assert.equal(mod.torchFlavorFor('mps'), 'cpu') +}) + +// ─── Detection precedence ───────────────────────────────────────────────────── + +test('detectGpuInfo keeps Apple Silicon on MPS', async () => { + const info = await mod.detectGpuInfo({ env: {}, platform: 'darwin', arch: 'arm64' }) + assert.equal(info.accelerator, 'mps') +}) + +test('detectGpuInfo forced to rocm resolves wheels without probing hardware', async () => { + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'rocm', MODLY_ROCM_GFX: 'gfx1201' }, + platform: 'linux', + arch: 'x64', + }) + assert.equal(info.accelerator, 'rocm') + assert.equal(info.gfxTarget, 'gfx1201') + assert.equal(info.torchIndexUrl, 'https://download.pytorch.org/whl/rocm7.2') + // sm/cudaVersion stay at 0 so CUDA-era extensions take their most + // conservative branch (and keep off rembg[gpu], which is CUDA-only). + assert.equal(info.sm, 0) + assert.equal(info.cudaVersion, 0) +}) + +test('detectGpuInfo forced to rocm overrides MPS on Apple Silicon', async () => { + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'rocm', MODLY_ROCM_GFX: 'gfx1200' }, + platform: 'darwin', + arch: 'arm64', + }) + assert.equal(info.accelerator, 'rocm') +}) + +test('detectGpuInfo forced to cpu short-circuits everything', async () => { + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'cpu', MODLY_ROCM_GFX: 'gfx1200' }, + platform: 'linux', + arch: 'x64', + }) + assert.deepEqual(info, { sm: 0, cudaVersion: 0, accelerator: 'cpu' }) +}) + +test('detectGpuInfo falls back to CPU when forced to rocm on Windows with no target', async () => { + const logs = [] + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'rocm' }, + platform: 'win32', + arch: 'x64', + onLog: (line) => logs.push(line), + }) + assert.equal(info.accelerator, 'cpu') + assert.ok(logs.some((l) => l.includes('MODLY_ROCM_GFX'))) +}) diff --git a/electron/main/gpu-detect.ts b/electron/main/gpu-detect.ts new file mode 100644 index 00000000..62201151 --- /dev/null +++ b/electron/main/gpu-detect.ts @@ -0,0 +1,390 @@ +/** + * GPU detection and PyTorch flavour resolution. + * + * Deliberately free of electron imports: the parsing and resolution helpers are + * pure and get unit-tested by bundling this file directly (gpu-detect.test.mjs). + * + * Modly never installs PyTorch itself — each extension's setup.py does, from an + * index it picks on its own. What we produce here is the information that lets + * that choice land on the right wheels: the accelerator, and for AMD the ROCm + * compute target plus the pip index/requirements the setup must end up using. + */ + +import { spawn } from 'child_process' +import { existsSync, readdirSync, readFileSync } from 'fs' +import { join } from 'path' + +export type Accelerator = 'cuda' | 'rocm' | 'mps' | 'cpu' +export type TorchFlavor = 'cuda' | 'rocm' | 'cpu' + +export interface GpuInfo { + sm: number + cudaVersion: number + accelerator: Accelerator + /** ROCm compute target ("gfx1200"). Only set when accelerator is 'rocm'. */ + gfxTarget?: string + /** pip --index-url the torch install has to come from (ROCm only). */ + torchIndexUrl?: string + /** pip requirements replacing whatever torch pins an extension hardcodes. */ + torchSpecs?: string[] +} + +// ─── ROCm wheel sources ─────────────────────────────────────────────────────── +// +// Linux and Windows need different indexes. download.pytorch.org publishes no +// ROCm wheels for Windows at all, so Windows has to go through AMD's multi-arch +// index, where the compute target is selected by a pip extra +// (torch[device-gfx1200]) instead of by the index URL. + +const ROCM_LINUX_INDEX = 'https://download.pytorch.org/whl/rocm7.2' +const ROCM_WINDOWS_INDEX = 'https://repo.amd.com/rocm/whl-multi-arch/' + +// Pinned because AMD's index carries several ROCm builds side by side; this is +// the newest pair published for cp311 (Modly's embedded Python) on Windows. +const ROCM_WINDOWS_TORCH = '2.11.0+rocm7.14.0' +const ROCM_WINDOWS_TORCHVISION = '0.26.0+rocm7.14.0' + +const KFD_TOPOLOGY_DIR = '/sys/class/kfd/kfd/topology/nodes' + +/** + * AMD PCI device id → ROCm compute target, for Windows where there is no KFD + * topology to read. Keyed by silicon rather than by marketing name: 0x73f0 is + * sold as an "RX 7600M XT" but is Navi 33, so it belongs with gfx1102, not with + * its 0x73xx neighbours. Device ids come from the pci.ids database. + */ +const WINDOWS_PCI_GFX_TARGETS: Record = { + // Navi 21 (RDNA 2) + '73a1': 'gfx1030', '73a2': 'gfx1030', '73a3': 'gfx1030', '73a5': 'gfx1030', + '73ab': 'gfx1030', '73ae': 'gfx1030', '73af': 'gfx1030', '73bf': 'gfx1030', + // Navi 22 (RDNA 2) + '73c3': 'gfx1031', '73ce': 'gfx1031', '73df': 'gfx1031', + // Navi 23 (RDNA 2) + '73e0': 'gfx1032', '73e1': 'gfx1032', '73e3': 'gfx1032', '73ef': 'gfx1032', + '73ff': 'gfx1032', + // Navi 24 (RDNA 2) + '7421': 'gfx1034', '7422': 'gfx1034', '7423': 'gfx1034', '7424': 'gfx1034', + '743f': 'gfx1034', + // Navi 31 (RDNA 3) + '7448': 'gfx1100', '7449': 'gfx1100', '744a': 'gfx1100', '744b': 'gfx1100', + '744c': 'gfx1100', '745e': 'gfx1100', + // Navi 32 (RDNA 3) + '7460': 'gfx1101', '7461': 'gfx1101', '7470': 'gfx1101', '747e': 'gfx1101', + // Navi 33 (RDNA 3) + '73f0': 'gfx1102', '7480': 'gfx1102', '7481': 'gfx1102', '7483': 'gfx1102', + '7487': 'gfx1102', '7489': 'gfx1102', '748b': 'gfx1102', '7499': 'gfx1102', + '749f': 'gfx1102', + // Navi 44 / Navi 48 (RDNA 4) + '7590': 'gfx1200', + '7550': 'gfx1201', '7551': 'gfx1201', +} + +// ─── Pure parsing helpers ───────────────────────────────────────────────────── + +/** + * Parses `nvidia-smi --query-gpu=compute_cap,driver_version --format=csv,noheader`. + * Returns null when the output carries no usable line. + */ +export function parseNvidiaSmi(stdout: string): { sm: number; cudaVersion: number } | null { + const line = stdout.trim().split('\n')[0]?.trim() // e.g. "8.6, 551.61" + if (!line) return null + + const parts = line.split(',').map((s) => s.trim()) + const sm = Math.round(parseFloat(parts[0] ?? '') * 10) // → 86 + + // Derive max supported CUDA version from driver version + // Driver ≥ 520 → CUDA 11.8, ≥ 525 → 12.0, ≥ 530 → 12.1, ≥ 535 → 12.2, + // ≥ 545 → 12.3, ≥ 550 → 12.4, ≥ 555 → 12.5, ≥ 560 → 12.6 + const driverMajor = parseInt((parts[1] ?? '').split('.')[0] ?? '0', 10) + let cudaVersion = 118 // safe minimum + if (driverMajor >= 570) cudaVersion = 128 // Blackwell (RTX 50xx, sm_120) + else if (driverMajor >= 560) cudaVersion = 126 + else if (driverMajor >= 555) cudaVersion = 125 + else if (driverMajor >= 550) cudaVersion = 124 + else if (driverMajor >= 545) cudaVersion = 123 + else if (driverMajor >= 535) cudaVersion = 122 + else if (driverMajor >= 530) cudaVersion = 121 + else if (driverMajor >= 525) cudaVersion = 120 + else if (driverMajor >= 520) cudaVersion = 118 + + return { sm: isNaN(sm) ? 86 : sm, cudaVersion } +} + +/** + * Decodes the KFD `gfx_target_version` integer (major*10000 + minor*100 + step, + * with minor and step read as hex digits) into a compute target name: + * 120000 → gfx1200, 90402 → gfx942, 90010 → gfx90a. + */ +export function formatGfxTarget(version: number): string | null { + if (!Number.isInteger(version) || version <= 0) return null + const major = Math.floor(version / 10000) + const minor = Math.floor((version % 10000) / 100) + const step = version % 100 + if (major <= 0 || minor > 15 || step > 15) return null + return `gfx${major}${minor.toString(16)}${step.toString(16)}` +} + +/** + * Picks the compute target out of the KFD topology node `properties` files. + * Node 0 is the CPU node (simd_count 0) and is skipped; the first real GPU wins. + */ +export function parseKfdGfxTarget(nodeProperties: string[]): string | null { + for (const text of nodeProperties) { + if (readKfdProperty(text, 'simd_count') <= 0) continue + const target = formatGfxTarget(readKfdProperty(text, 'gfx_target_version')) + if (target) return target + } + return null +} + +function readKfdProperty(text: string, key: string): number { + const match = new RegExp(`^${key}\\s+(\\d+)\\s*$`, 'm').exec(text) + return match ? parseInt(match[1], 10) : 0 +} + +export interface VideoController { + name: string + pnpDeviceId: string +} + +/** + * Parses the JSON emitted by `Get-CimInstance Win32_VideoController | ConvertTo-Json`. + * PowerShell emits a bare object rather than an array when there is one adapter. + */ +export function parseWindowsVideoControllers(stdout: string): VideoController[] { + let parsed: unknown + try { + parsed = JSON.parse(stdout) + } catch { + return [] + } + const list = Array.isArray(parsed) ? parsed : [parsed] + return list + .filter((entry): entry is Record => !!entry && typeof entry === 'object') + .map((entry) => ({ + name: typeof entry['Name'] === 'string' ? entry['Name'] : '', + pnpDeviceId: typeof entry['PNPDeviceID'] === 'string' ? entry['PNPDeviceID'] : '', + })) + .filter((c) => c.pnpDeviceId !== '') +} + +/** Extracts the PCI device id from a PNPDeviceID, e.g. `PCI\VEN_1002&DEV_7590&…` → "7590". */ +export function parseAmdPciDeviceId(pnpDeviceId: string): string | null { + const match = /VEN_1002&DEV_([0-9A-F]{4})/i.exec(pnpDeviceId) + return match ? match[1].toLowerCase() : null +} + +/** + * Resolves a compute target from the installed display adapters. Returns the + * AMD adapters it saw as well, so an unmapped card can be reported by name + * instead of silently falling back to CPU. + */ +export function resolveWindowsGfxTarget( + controllers: VideoController[], +): { gfxTarget: string | null; amdAdapters: string[] } { + const amdAdapters: string[] = [] + let gfxTarget: string | null = null + + for (const controller of controllers) { + const deviceId = parseAmdPciDeviceId(controller.pnpDeviceId) + if (!deviceId) continue + amdAdapters.push(controller.name || `PCI 1002:${deviceId}`) + gfxTarget ??= WINDOWS_PCI_GFX_TARGETS[deviceId] ?? null + } + + return { gfxTarget, amdAdapters } +} + +/** + * The pip index and requirements an extension's torch install has to end up + * using on this platform. `MODLY_ROCM_INDEX` and `MODLY_ROCM_TORCH_SPEC` + * (space-separated requirements) override either half — the escape hatch when a + * newer torch breaks an extension and you need to drop back to, say, rocm6.4. + */ +export function resolveRocmTorchSpec( + platform: string, + gfxTarget: string | null, + env: NodeJS.ProcessEnv = process.env, +): { indexUrl: string; specs: string[] } { + const indexOverride = env['MODLY_ROCM_INDEX']?.trim() + const specOverride = env['MODLY_ROCM_TORCH_SPEC']?.trim() + + const isWindows = platform === 'win32' + const indexUrl = indexOverride || (isWindows ? ROCM_WINDOWS_INDEX : ROCM_LINUX_INDEX) + + if (specOverride) return { indexUrl, specs: specOverride.split(/\s+/).filter(Boolean) } + + // AMD's multi-arch index ships one torch per compute target, selected by a + // pip extra. The pytorch.org ROCm index bakes the targets into a single + // wheel, so there is nothing to select and nothing to pin. + if (isWindows && gfxTarget) { + return { + indexUrl, + specs: [ + `torch[device-${gfxTarget}]==${ROCM_WINDOWS_TORCH}`, + `torchvision[device-${gfxTarget}]==${ROCM_WINDOWS_TORCHVISION}`, + ], + } + } + return { indexUrl, specs: ['torch', 'torchvision'] } +} + +/** The `torch_flavor` value extension setup.py scripts branch on. */ +export function torchFlavorFor(accelerator: Accelerator): TorchFlavor { + if (accelerator === 'rocm') return 'rocm' + if (accelerator === 'cuda') return 'cuda' + return 'cpu' +} + +export function describeGpuInfo(info: GpuInfo): string { + const bits = [`accelerator=${info.accelerator}`] + if (info.accelerator === 'cuda') bits.push(`sm=${info.sm}`, `cuda=${info.cudaVersion}`) + if (info.gfxTarget) bits.push(`gfx=${info.gfxTarget}`) + if (info.torchIndexUrl) bits.push(`torch_index=${info.torchIndexUrl}`) + return bits.join(' ') +} + +// ─── Detection ──────────────────────────────────────────────────────────────── + +function cpuInfo(): GpuInfo { + return { sm: 0, cudaVersion: 0, accelerator: 'cpu' } +} + +/** + * AMD keeps sm/cudaVersion at 0 on purpose. Extensions that predate `torch_flavor` + * branch on those two numbers, and 0 sends them down their most conservative + * path — which also keeps them off `rembg[gpu]`, whose onnxruntime-gpu is + * CUDA-only. Their torch install is then corrected by the ROCm setup shim. + */ +function rocmInfo( + gfxTarget: string | null, + platform: string, + env: NodeJS.ProcessEnv, +): GpuInfo { + const { indexUrl, specs } = resolveRocmTorchSpec(platform, gfxTarget, env) + return { + sm: 0, + cudaVersion: 0, + accelerator: 'rocm', + ...(gfxTarget ? { gfxTarget } : {}), + torchIndexUrl: indexUrl, + torchSpecs: specs, + } +} + +function readFlavorOverride(env: NodeJS.ProcessEnv): TorchFlavor | null { + const raw = env['MODLY_TORCH_FLAVOR']?.trim().toLowerCase() + return raw === 'cuda' || raw === 'rocm' || raw === 'cpu' ? raw : null +} + +function queryNvidiaSmi(): Promise<{ sm: number; cudaVersion: number } | null> { + return new Promise((resolve) => { + // Query compute cap + driver version in one call + const proc = spawn('nvidia-smi', ['--query-gpu=compute_cap,driver_version', '--format=csv,noheader'], { + stdio: ['ignore', 'pipe', 'ignore'], + }) + let out = '' + proc.stdout?.on('data', (d: Buffer) => { out += d.toString() }) + proc.on('close', (code) => resolve(code === 0 ? parseNvidiaSmi(out) : null)) + proc.on('error', () => resolve(null)) + }) +} + +/** + * Reads the compute target straight out of the kernel's KFD topology. This + * needs no ROCm installation and no external binary — the amdgpu driver alone + * publishes it, which is exactly the state of a machine that has only ever run + * PyTorch ROCm wheels (they bundle their own runtime). + */ +function readKfdGfxTarget(): string | null { + if (!existsSync('/dev/kfd')) return null + try { + const nodes = readdirSync(KFD_TOPOLOGY_DIR).sort((a, b) => Number(a) - Number(b)) + const properties = nodes.map((node) => { + try { + return readFileSync(join(KFD_TOPOLOGY_DIR, node, 'properties'), 'utf-8') + } catch { + return '' + } + }) + return parseKfdGfxTarget(properties) + } catch { + return null + } +} + +function queryWindowsVideoControllers(): Promise { + return new Promise((resolve) => { + const proc = spawn('powershell', [ + '-NoProfile', '-NonInteractive', '-Command', + 'Get-CimInstance Win32_VideoController | Select-Object Name,PNPDeviceID | ConvertTo-Json -Compress', + ], { stdio: ['ignore', 'pipe', 'ignore'] }) + let out = '' + proc.stdout?.on('data', (d: Buffer) => { out += d.toString() }) + proc.on('close', (code) => resolve(code === 0 ? parseWindowsVideoControllers(out) : [])) + proc.on('error', () => resolve([])) + }) +} + +export interface DetectOptions { + env?: NodeJS.ProcessEnv + platform?: string + arch?: string + onLog?: (line: string) => void +} + +export async function detectGpuInfo(options: DetectOptions = {}): Promise { + const env = options.env ?? process.env + const platform = options.platform ?? process.platform + const arch = options.arch ?? process.arch + const log = options.onLog ?? (() => {}) + + const forced = readFlavorOverride(env) + if (forced) log(`[gpu-detect] MODLY_TORCH_FLAVOR=${forced} — skipping auto-detection`) + + if (forced === 'cpu') return cpuInfo() + + if (forced === 'rocm') { + const gfxTarget = await resolveGfxTarget(env, platform) + if (!gfxTarget && platform === 'win32') { + log('[gpu-detect] ROCm forced on Windows but no compute target found — set MODLY_ROCM_GFX (e.g. gfx1200)') + return cpuInfo() + } + return rocmInfo(gfxTarget, platform, env) + } + + if (platform === 'darwin' && arch === 'arm64') { + return { sm: 0, cudaVersion: 0, accelerator: 'mps' } + } + + // NVIDIA keeps priority: on a machine with both, nothing about the existing + // CUDA behaviour changes. + const nvidia = await queryNvidiaSmi() + if (nvidia) return { ...nvidia, accelerator: 'cuda' } + if (forced === 'cuda') return { sm: 0, cudaVersion: 0, accelerator: 'cuda' } + + const gfxTarget = await resolveGfxTarget(env, platform) + if (gfxTarget) { + log(`[gpu-detect] AMD GPU detected — compute target ${gfxTarget}`) + return rocmInfo(gfxTarget, platform, env) + } + + if (platform === 'win32') { + const { amdAdapters } = resolveWindowsGfxTarget(await queryWindowsVideoControllers()) + if (amdAdapters.length > 0) { + log( + `[gpu-detect] AMD GPU found (${amdAdapters.join(', ')}) but its ROCm compute target is unknown. ` + + 'Falling back to CPU — set MODLY_ROCM_GFX (e.g. gfx1201) to force one.', + ) + } + } + + return cpuInfo() +} + +async function resolveGfxTarget(env: NodeJS.ProcessEnv, platform: string): Promise { + const override = env['MODLY_ROCM_GFX']?.trim() + if (override) return override + if (platform === 'linux') return readKfdGfxTarget() + if (platform === 'win32') return resolveWindowsGfxTarget(await queryWindowsVideoControllers()).gfxTarget + return null +} diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 14fc984a..f4ad43e3 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -33,67 +33,20 @@ import { resolvePathWithinRoot, } from './extension-path-guard' import { validateInstallManifest } from './extension-install-utils' +import { detectGpuInfo, describeGpuInfo, torchFlavorFor, type GpuInfo } from './gpu-detect' +import { SETUP_LAUNCHER_SOURCE } from './setup-launcher' import { registerWorkspaceAssetLibraryIpcHandlers } from './artifact-registry-service' import { updatesSupported } from './updater' type WindowGetter = () => BrowserWindow | null const pExecFile = promisify(execFile) -// ─── GPU detect (best-effort, no Python required) ───────────────────────────── - -interface GpuInfo { - sm: number - cudaVersion: number - accelerator: 'cuda' | 'mps' | 'cpu' -} - -function detectGpuInfo(): Promise { - if (process.platform === 'darwin' && process.arch === 'arm64') { - return Promise.resolve({ sm: 0, cudaVersion: 0, accelerator: 'mps' }) - } - - return new Promise((resolve) => { - // Query compute cap + driver version in one call - const proc = spawn('nvidia-smi', ['--query-gpu=compute_cap,driver_version', '--format=csv,noheader'], { - stdio: ['ignore', 'pipe', 'ignore'], - }) - let out = '' - proc.stdout?.on('data', (d: Buffer) => { out += d.toString() }) - proc.on('close', (code) => { - if (code === 0) { - const line = out.trim().split('\n')[0].trim() // e.g. "8.6, 551.61" - const parts = line.split(',').map(s => s.trim()) - const sm = Math.round(parseFloat(parts[0] ?? '') * 10) // → 86 - // Derive max supported CUDA version from driver version - // Driver ≥ 520 → CUDA 11.8, ≥ 525 → 12.0, ≥ 530 → 12.1, ≥ 535 → 12.2, - // ≥ 545 → 12.3, ≥ 550 → 12.4, ≥ 555 → 12.5, ≥ 560 → 12.6 - const driverMajor = parseInt((parts[1] ?? '').split('.')[0] ?? '0', 10) - let cudaVersion = 118 // safe minimum - if (driverMajor >= 570) cudaVersion = 128 // Blackwell (RTX 50xx, sm_120) - else if (driverMajor >= 560) cudaVersion = 126 - else if (driverMajor >= 555) cudaVersion = 125 - else if (driverMajor >= 550) cudaVersion = 124 - else if (driverMajor >= 545) cudaVersion = 123 - else if (driverMajor >= 535) cudaVersion = 122 - else if (driverMajor >= 530) cudaVersion = 121 - else if (driverMajor >= 525) cudaVersion = 120 - else if (driverMajor >= 520) cudaVersion = 118 - resolve({ sm: isNaN(sm) ? 86 : sm, cudaVersion, accelerator: 'cuda' }) - } else { - resolve({ sm: 0, cudaVersion: 0, accelerator: 'cpu' }) - } - }) - proc.on('error', () => resolve({ sm: 0, cudaVersion: 0, accelerator: 'cpu' })) - }) -} - // ─── Run an extension's setup.py directly (no FastAPI needed) ───────────────── function runExtensionSetup( - extDir: string, - gpuSm: number, - cudaVersion: number, - onLog?: (line: string) => void, + extDir: string, + gpu: GpuInfo, + onLog?: (line: string) => void, ): Promise { return new Promise((resolve, reject) => { const userData = app.getPath('userData') @@ -107,113 +60,34 @@ function runExtensionSetup( const pipCacheDir = join(getSettings(userData).dependenciesDir, 'pip-cache') try { mkdirSync(pipCacheDir, { recursive: true }) } catch { /* pip creates it too */ } - const accelerator = process.platform === 'darwin' && process.arch === 'arm64' ? 'mps' : gpuSm > 0 ? 'cuda' : 'cpu' + const torchFlavor = torchFlavorFor(gpu.accelerator) const args = JSON.stringify({ python_exe: pythonExe, ext_dir: extDir, - gpu_sm: gpuSm, - cuda_version: cudaVersion, - accelerator, + gpu_sm: gpu.sm, + cuda_version: gpu.cudaVersion, + accelerator: gpu.accelerator, + // Extensions that know about AMD branch on torch_flavor (the official + // hunyuan3d-mini one does). Those that don't get corrected by the ROCm + // shim in setup-launcher.ts instead. + torch_flavor: torchFlavor, + gfx_target: gpu.gfxTarget ?? '', + torch_index_url: gpu.torchIndexUrl ?? '', platform: process.platform, arch: process.arch, }) - const launcher = ` -import runpy -import subprocess -import sys - -setup_py = sys.argv[1] -setup_args = sys.argv[2:] - -_original_run = subprocess.run -_original_check_call = subprocess.check_call -_original_check_output = subprocess.check_output - -def _is_cuda_torch_index(value): - return isinstance(value, str) and value.startswith("https://download.pytorch.org/whl/cu") - -def _mentions_torch(command): - if not isinstance(command, (list, tuple)): - return False - return any(str(part).startswith(("torch==", "torchvision==", "torchaudio==")) for part in command) - -def _rewrite_command(command): - if sys.platform != "darwin" or not _mentions_torch(command): - return command - if not isinstance(command, (list, tuple)): - return command - - rewritten = [] - changed = False - i = 0 - while i < len(command): - part = command[i] - text = str(part) - if text in ("--index-url", "-i", "--extra-index-url") and i + 1 < len(command) and _is_cuda_torch_index(str(command[i + 1])): - changed = True - i += 2 - continue - if text.startswith("--index-url=") or text.startswith("--extra-index-url="): - value = text.split("=", 1)[1] - if _is_cuda_torch_index(value): - changed = True - i += 1 - continue - rewritten.append(part) - i += 1 - - if changed: - print("[Modly setup compat] Removed CUDA-only PyTorch index on macOS; pip will use macOS wheels.", file=sys.stderr) - return rewritten - return command - -def _is_pip_command(command): - if not isinstance(command, (list, tuple)): - return False - return any("pip" in str(part).lower() for part in command[:3]) - -def _strip_no_cache(command): - # Extension setup scripts often hardcode --no-cache-dir, which forces pip to - # re-download multi-GB wheels on every retry. Modly provides a shared cache - # via PIP_CACHE_DIR, so drop the flag and let pip use it. - if not _is_pip_command(command): - return command - if not any(str(part) == "--no-cache-dir" for part in command): - return command - print("[Modly setup compat] Removed --no-cache-dir so pip reuses the shared wheel cache.", file=sys.stderr) - return [part for part in command if str(part) != "--no-cache-dir"] - -def _transform_command(command): - return _strip_no_cache(_rewrite_command(command)) - -def _patched_run(*args, **kwargs): - args = list(args) - if args: - args[0] = _transform_command(args[0]) - return _original_run(*args, **kwargs) - -def _patched_check_call(*args, **kwargs): - args = list(args) - if args: - args[0] = _transform_command(args[0]) - return _original_check_call(*args, **kwargs) - -def _patched_check_output(*args, **kwargs): - args = list(args) - if args: - args[0] = _transform_command(args[0]) - return _original_check_output(*args, **kwargs) - -subprocess.run = _patched_run -subprocess.check_call = _patched_check_call -subprocess.check_output = _patched_check_output - -sys.argv = [setup_py] + setup_args -runpy.run_path(setup_py, run_name="__main__") -` + const launcher = SETUP_LAUNCHER_SOURCE + // The rewrite decision itself is made in gpu-detect.ts (and unit-tested + // there); the launcher above only applies what these carry. const proc = spawn(pythonExe, ['-c', launcher, setupPy, args], { stdio: ['ignore', 'pipe', 'pipe'], - env: { ...process.env, PIP_CACHE_DIR: pipCacheDir }, + env: { + ...process.env, + PIP_CACHE_DIR: pipCacheDir, + MODLY_TORCH_FLAVOR: torchFlavor, + MODLY_TORCH_INDEX_URL: gpu.torchIndexUrl ?? '', + MODLY_TORCH_SPECS: JSON.stringify(gpu.torchSpecs ?? []), + }, }) const handleLine = (line: string) => { if (line) onLog?.(line) } @@ -1149,8 +1023,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // 7a. Python process extension: run setup.py if present (same as model extensions) if (existsSync(join(destDir, 'setup.py'))) { emit({ step: 'setting_up', message: 'Setting up Python environment…' }) - const { sm: gpuSm, cudaVersion } = await detectGpuInfo() - await runExtensionSetup(destDir, gpuSm, cudaVersion, (line) => { + const gpu = await detectGpuInfo({ onLog: (line) => logger.info(line) }) + logger.info(`[ext-setup] ${describeGpuInfo(gpu)}`) + await runExtensionSetup(destDir, gpu, (line) => { logger.info(`[ext-setup] ${line}`) emit({ step: 'setting_up', message: line }) }) @@ -1185,8 +1060,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // 7c. Model extension: run setup.py directly (no FastAPI required) if (existsSync(join(destDir, 'setup.py'))) { emit({ step: 'setting_up', message: 'Setting up Python environment…' }) - const { sm: gpuSm, cudaVersion } = await detectGpuInfo() - await runExtensionSetup(destDir, gpuSm, cudaVersion, (line) => { + const gpu = await detectGpuInfo({ onLog: (line) => logger.info(line) }) + logger.info(`[ext-setup] ${describeGpuInfo(gpu)}`) + await runExtensionSetup(destDir, gpu, (line) => { logger.info(`[ext-setup] ${line}`) emit({ step: 'setting_up', message: line }) }) @@ -1279,8 +1155,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe if (!existsSync(join(extDir, 'setup.py'))) { return { success: false, error: 'setup.py is missing from the extension folder — the install looks incomplete. Uninstall the extension and install it again.' } } - const { sm: gpuSm, cudaVersion } = await detectGpuInfo() - await runExtensionSetup(extDir, gpuSm, cudaVersion, (line) => logger.info(`[ext-repair] ${line}`)) + const gpu = await detectGpuInfo({ onLog: (line) => logger.info(line) }) + logger.info(`[ext-repair] ${describeGpuInfo(gpu)}`) + await runExtensionSetup(extDir, gpu, (line) => logger.info(`[ext-repair] ${line}`)) try { await axios.post(`${API_BASE_URL}/extensions/reload`, {}, { timeout: 10_000 }) } catch { /* ignore if Python is not running yet */ } diff --git a/electron/main/python-setup.ts b/electron/main/python-setup.ts index 3b85db1b..13e86a25 100644 --- a/electron/main/python-setup.ts +++ b/electron/main/python-setup.ts @@ -1,10 +1,11 @@ import { BrowserWindow, app } from 'electron' import { existsSync, readFileSync, writeFileSync } from 'fs' -import { cp, rm, mkdir } from 'fs/promises' +import { rm, mkdir } from 'fs/promises' import { join } from 'path' import { spawn, execSync } from 'child_process' import { createHash } from 'crypto' import { getSettings } from './settings-store' +import { copyRuntimeTree } from './copy-runtime' const SETUP_VERSION = 3 @@ -176,7 +177,7 @@ async function ensureStableEmbeddedPython(userData: string, win: BrowserWindow): await rm(stableDir, { recursive: true, force: true }) } await mkdir(stableDir, { recursive: true }) - await cp(getEmbeddedPythonDir(), stableDir, { recursive: true, preserveTimestamps: true }) + await copyRuntimeTree(getEmbeddedPythonDir(), stableDir) writeFileSync(versionFile, currentVersion, 'utf-8') console.log('[PythonSetup] Python runtime ready at:', stableDir) } diff --git a/electron/main/setup-launcher.test.mjs b/electron/main/setup-launcher.test.mjs new file mode 100644 index 00000000..c4a40524 --- /dev/null +++ b/electron/main/setup-launcher.test.mjs @@ -0,0 +1,224 @@ +/** + * Runs the setup launcher for real, with a stub standing in for pip, and checks + * what the extension's pip invocation was rewritten into. + * + * The commands exercised below are copied from the official extensions' + * setup.py, so a change that breaks AMD installs fails here rather than after a + * multi-gigabyte download on a user's machine. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync, readFileSync, existsSync } from 'node:fs' +import { spawnSync } from 'node:child_process' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-launcher-test-')), 'setup-launcher.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/setup-launcher.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { SETUP_LAUNCHER_SOURCE } = loadModule() + +function findPython() { + for (const candidate of ['python3', 'python']) { + const probe = spawnSync(candidate, ['--version'], { stdio: 'ignore' }) + if (probe.status === 0) return candidate + } + return null +} + +const PYTHON = findPython() + +// The launcher's _is_pip_command looks for "pip" in the first three tokens, so +// the stub is named pip.py — matching how extensions invoke "/bin/pip". +const FAKE_PIP = ` +import json, sys +with open(sys.argv[1], "w") as handle: + json.dump(sys.argv[2:], handle) +` + +/** + * Runs a pip command through the launcher and returns the argv the stub + * actually received, i.e. the command after every rewrite. + */ +function runThroughLauncher(pipArgs, env = {}) { + const dir = mkdtempSync(join(tmpdir(), 'modly-launcher-run-')) + const fakePip = join(dir, 'pip.py') + const capture = join(dir, 'captured.json') + writeFileSync(fakePip, FAKE_PIP, 'utf8') + + // Stands in for the extension's setup.py: issues one pip call, exactly as the + // real ones do, and lets the launcher's patched subprocess handle it. + const setupPy = join(dir, 'setup.py') + writeFileSync(setupPy, [ + 'import subprocess, sys', + `PIP = ${JSON.stringify([fakePip, capture])}`, + `subprocess.run([sys.executable] + PIP + ${JSON.stringify(pipArgs)}, check=True)`, + ].join('\n'), 'utf8') + + const result = spawnSync(PYTHON, ['-c', SETUP_LAUNCHER_SOURCE, setupPy, '{}'], { + encoding: 'utf8', + env: { ...process.env, ...env }, + }) + assert.equal(result.status, 0, `launcher failed:\n${result.stderr}`) + assert.ok(existsSync(capture), `pip stub was never invoked:\n${result.stderr}`) + + // Drop the stub's own two arguments; keep the pip command itself. + return { argv: JSON.parse(readFileSync(capture, 'utf8')), stderr: result.stderr } +} + +const ROCM_ENV = { + MODLY_TORCH_FLAVOR: 'rocm', + MODLY_TORCH_INDEX_URL: 'https://download.pytorch.org/whl/rocm7.2', + MODLY_TORCH_SPECS: JSON.stringify(['torch', 'torchvision']), +} + +test('ROCm shim redirects a CUDA-pinned install (triposg / trellis2 shape)', { skip: !PYTHON }, () => { + const { argv, stderr } = runThroughLauncher( + ['install', 'torch==2.6.0', 'torchvision==0.21.0', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + ]) + assert.match(stderr, /Redirected PyTorch to ROCm wheels/) +}) + +test('ROCm shim rescues the CPU fallback hunyuan3d-mini forces on Windows', { skip: !PYTHON }, () => { + // That setup.py hardcodes the CPU index when torch_flavor is rocm on Windows; + // without this rewrite an AMD Windows user would silently get CPU-only torch. + const { argv } = runThroughLauncher( + ['install', 'torch==2.6.0', 'torchvision==0.21.0', '--index-url', 'https://download.pytorch.org/whl/cpu'], + { + ...ROCM_ENV, + MODLY_TORCH_INDEX_URL: 'https://repo.amd.com/rocm/whl-multi-arch/', + MODLY_TORCH_SPECS: JSON.stringify([ + 'torch[device-gfx1200]==2.11.0+rocm7.14.0', + 'torchvision[device-gfx1200]==0.26.0+rocm7.14.0', + ]), + }, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://repo.amd.com/rocm/whl-multi-arch/', + 'torch[device-gfx1200]==2.11.0+rocm7.14.0', + 'torchvision[device-gfx1200]==0.26.0+rocm7.14.0', + ]) +}) + +test('ROCm shim leaves an extension that already chose ROCm alone', { skip: !PYTHON }, () => { + // hunyuan3d-mini's own rocm branch on Linux. Its index is kept; only the + // requirements are normalised to what Modly resolved. + const { argv } = runThroughLauncher( + ['install', 'torch', 'torchvision', '--index-url', 'https://download.pytorch.org/whl/rocm7.2'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + 'torch', 'torchvision', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + ]) +}) + +test('ROCm shim replaces the pinned direct wheel URLs of the ARM64 path', { skip: !PYTHON }, () => { + const { argv } = runThroughLauncher( + [ + 'install', '--retries', '10', + '--extra-index-url', 'https://download.pytorch.org/whl/cu128', + 'https://download-r2.pytorch.org/whl/cu128/torch-2.7.0%2Bcu128-cp311-cp311-manylinux_2_28_aarch64.whl', + 'https://download-r2.pytorch.org/whl/cu128/torchvision-0.22.0-cp311-cp311-manylinux_2_28_aarch64.whl', + ], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', '--retries', '10', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + ]) +}) + +test('ROCm shim handles the --index-url=VALUE spelling', { skip: !PYTHON }, () => { + const { argv } = runThroughLauncher( + ['install', '--index-url=https://download.pytorch.org/whl/cu124', 'torch==2.6.0'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + ]) +}) + +test('ROCm shim leaves non-torch installs untouched', { skip: !PYTHON }, () => { + // The bulk of every setup.py: core deps, rembg, etc. Nothing here is ours to + // rewrite, and an extra --index-url would send them to the ROCm index. + const original = ['install', 'Pillow', 'numpy', 'trimesh', 'rembg', 'onnxruntime'] + const { argv } = runThroughLauncher(original, ROCM_ENV) + assert.deepEqual(argv, original) +}) + +test('ROCm shim keeps PyPI reachable when torch is mixed with other packages', { skip: !PYTHON }, () => { + // The ROCm index mirrors torch's dependency closure only: asking it for + // trimesh returns 403, so a mixed install would hard-fail without this. + const { argv } = runThroughLauncher( + ['install', 'torch==2.6.0', 'trimesh', 'diffusers', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + '--extra-index-url', 'https://pypi.org/simple', + 'torch', 'torchvision', + 'trimesh', 'diffusers', + ]) +}) + +test('ROCm shim does not add PyPI for a torch-only install', { skip: !PYTHON }, () => { + // Keeping PyPI out of a pure torch install avoids pip ever preferring a + // plain CUDA wheel over the ROCm one. + const { argv } = runThroughLauncher( + ['install', '--retries', '10', 'torch==2.6.0', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + ) + assert.ok(!argv.includes('--extra-index-url')) +}) + +test('shim is inert on a CUDA machine', { skip: !PYTHON }, () => { + const original = ['install', 'torch==2.6.0', '--index-url', 'https://download.pytorch.org/whl/cu124'] + const { argv } = runThroughLauncher(original, { + MODLY_TORCH_FLAVOR: 'cuda', + MODLY_TORCH_INDEX_URL: '', + MODLY_TORCH_SPECS: '[]', + }) + assert.deepEqual(argv, original) +}) + +test('--no-cache-dir is still stripped alongside the ROCm rewrite', { skip: !PYTHON }, () => { + const { argv } = runThroughLauncher( + ['install', '--no-cache-dir', 'torch==2.6.0', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + ) + + assert.ok(!argv.includes('--no-cache-dir')) + assert.ok(argv.includes('https://download.pytorch.org/whl/rocm7.2')) +}) diff --git a/electron/main/setup-launcher.ts b/electron/main/setup-launcher.ts new file mode 100644 index 00000000..4207e4a6 --- /dev/null +++ b/electron/main/setup-launcher.ts @@ -0,0 +1,231 @@ +/** + * Python launcher used to run an extension's setup.py. + * + * Extension setup scripts are third-party code we cannot edit, and they install + * PyTorch themselves from an index they hardcode. The launcher wraps them: it + * patches subprocess so every pip invocation passes through a few corrections + * before it runs — dropping CUDA-only indexes on macOS, keeping the shared wheel + * cache alive, and redirecting torch to ROCm wheels on AMD machines. + * + * Kept in its own module so setup-launcher.test.mjs can execute it for real + * against the command shapes the official extensions actually use. + */ + +export const SETUP_LAUNCHER_SOURCE = ` +import json +import os +import re +import runpy +import subprocess +import sys + +setup_py = sys.argv[1] +setup_args = sys.argv[2:] + +_original_run = subprocess.run +_original_check_call = subprocess.check_call +_original_check_output = subprocess.check_output + +_TORCH_REQ_RE = re.compile(r"^(torch|torchvision|torchaudio)(\\[[^\\]]*\\])?\\s*([<>=!~].*)?$", re.I) + +def _is_cuda_torch_index(value): + return isinstance(value, str) and value.startswith("https://download.pytorch.org/whl/cu") + +def _is_torch_requirement(text): + # Matches "torch", "torch==2.6.0", "torch[device-gfx1200]==2.11.0+rocm7.14.0", + # and the pinned direct wheel URLs the ARM64 install path uses. + if _TORCH_REQ_RE.match(text): + return True + if text.startswith(("http://", "https://")): + return any(seg in text for seg in ("/torch-", "/torchvision-", "/torchaudio-")) + return False + +def _mentions_torch(command): + if not isinstance(command, (list, tuple)): + return False + return any(_is_torch_requirement(str(part)) for part in command) + +def _rewrite_command(command): + if sys.platform != "darwin" or not _mentions_torch(command): + return command + if not isinstance(command, (list, tuple)): + return command + + rewritten = [] + changed = False + i = 0 + while i < len(command): + part = command[i] + text = str(part) + if text in ("--index-url", "-i", "--extra-index-url") and i + 1 < len(command) and _is_cuda_torch_index(str(command[i + 1])): + changed = True + i += 2 + continue + if text.startswith("--index-url=") or text.startswith("--extra-index-url="): + value = text.split("=", 1)[1] + if _is_cuda_torch_index(value): + changed = True + i += 1 + continue + rewritten.append(part) + i += 1 + + if changed: + print("[Modly setup compat] Removed CUDA-only PyTorch index on macOS; pip will use macOS wheels.", file=sys.stderr) + return rewritten + return command + +def _is_pip_command(command): + if not isinstance(command, (list, tuple)): + return False + return any("pip" in str(part).lower() for part in command[:3]) + +def _strip_no_cache(command): + # Extension setup scripts often hardcode --no-cache-dir, which forces pip to + # re-download multi-GB wheels on every retry. Modly provides a shared cache + # via PIP_CACHE_DIR, so drop the flag and let pip use it. + if not _is_pip_command(command): + return command + if not any(str(part) == "--no-cache-dir" for part in command): + return command + print("[Modly setup compat] Removed --no-cache-dir so pip reuses the shared wheel cache.", file=sys.stderr) + return [part for part in command if str(part) != "--no-cache-dir"] + +# ROCm redirect. Most extension setup.py scripts predate AMD support and +# hardcode a CUDA index (hunyuan3d-mini even forces the CPU index on Windows), +# so on an AMD machine we swap the whole torch install for the ROCm one Modly +# resolved. An index the extension already pointed at ROCm is left alone. +_ROCM_INDEX = os.environ.get("MODLY_TORCH_INDEX_URL", "") +try: + _ROCM_SPECS = json.loads(os.environ.get("MODLY_TORCH_SPECS", "[]")) +except ValueError: + _ROCM_SPECS = [] + +def _is_pytorch_index(value): + return isinstance(value, str) and "download.pytorch.org/whl/" in value + +def _is_rocm_index(value): + return isinstance(value, str) and "/whl/rocm" in value + +_PIP_VALUE_FLAGS = ( + "--index-url", "-i", "--extra-index-url", "--find-links", "-f", + "--retries", "--timeout", "--cache-dir", "--target", "-t", + "--requirement", "-r", "--constraint", "-c", "--progress-bar", + "--proxy", "--cert", "--client-cert", "--trusted-host", "--log", + "--no-binary", "--only-binary", "--prefix", "--root", "--upgrade-strategy", + "--python-version", "--platform", "--abi", "--implementation", +) + +def _has_non_torch_requirement(command): + # Only look past the subcommand, so the interpreter and pip executable + # paths ahead of it are never mistaken for requirements. + texts = [str(part) for part in command] + start = None + for index, text in enumerate(texts): + if text in ("install", "download", "wheel"): + start = index + 1 + break + if start is None: + return False + + skip_next = False + for text in texts[start:]: + if skip_next: + skip_next = False + continue + if text in _PIP_VALUE_FLAGS: + skip_next = True + continue + if text.startswith("-") or _is_torch_requirement(text): + continue + return True + return False + +def _rewrite_rocm(command): + if os.environ.get("MODLY_TORCH_FLAVOR") != "rocm" or not _ROCM_INDEX or not _ROCM_SPECS: + return command + if not isinstance(command, (list, tuple)): + return command + if not _is_pip_command(command) or not _mentions_torch(command): + return command + + rewritten = [] + insert_at = None + keeps_rocm_index = False + changed = False + i = 0 + while i < len(command): + text = str(command[i]) + if text in ("--index-url", "-i", "--extra-index-url") and i + 1 < len(command): + value = str(command[i + 1]) + if _is_rocm_index(value): + keeps_rocm_index = True + elif _is_pytorch_index(value): + changed = True + i += 2 + continue + rewritten.extend(command[i:i + 2]) + i += 2 + continue + if text.split("=", 1)[0] in ("--index-url", "--extra-index-url") and "=" in text: + value = text.split("=", 1)[1] + if _is_rocm_index(value): + keeps_rocm_index = True + elif _is_pytorch_index(value): + changed = True + i += 1 + continue + elif _is_torch_requirement(text): + if insert_at is None: + insert_at = len(rewritten) + changed = True + i += 1 + continue + rewritten.append(command[i]) + i += 1 + + if not changed: + return command + if insert_at is None: + insert_at = len(rewritten) + injected = list(_ROCM_SPECS) + if not keeps_rocm_index: + index_args = ["--index-url", _ROCM_INDEX] + if _has_non_torch_requirement(command): + # The ROCm index only mirrors torch's own dependency closure + # (numpy, pillow…), not application packages like trimesh or + # diffusers. A pip call that mixes both still needs PyPI reachable. + index_args += ["--extra-index-url", "https://pypi.org/simple"] + injected = index_args + injected + rewritten[insert_at:insert_at] = injected + print("[Modly setup compat] Redirected PyTorch to ROCm wheels: " + " ".join(injected), file=sys.stderr) + return rewritten + +def _transform_command(command): + return _strip_no_cache(_rewrite_rocm(_rewrite_command(command))) + +def _patched_run(*args, **kwargs): + args = list(args) + if args: + args[0] = _transform_command(args[0]) + return _original_run(*args, **kwargs) + +def _patched_check_call(*args, **kwargs): + args = list(args) + if args: + args[0] = _transform_command(args[0]) + return _original_check_call(*args, **kwargs) + +def _patched_check_output(*args, **kwargs): + args = list(args) + if args: + args[0] = _transform_command(args[0]) + return _original_check_output(*args, **kwargs) + +subprocess.run = _patched_run +subprocess.check_call = _patched_check_call +subprocess.check_output = _patched_check_output + +sys.argv = [setup_py] + setup_args +runpy.run_path(setup_py, run_name="__main__") +` From 494dabf11b9d177663640ad2da534b234bbcbda4 Mon Sep 17 00:00:00 2001 From: DrHepa Date: Fri, 21 Aug 2026 07:54:11 +0200 Subject: [PATCH 03/57] feat(models): support multiple Hugging Face sources per node --- README.md | 36 +++ api/routers/model.py | 164 ++++++++++- api/services/generator_registry.py | 27 +- api/services/model_sources.py | 259 ++++++++++++++++++ api/tests/test_generator_registry.py | 56 ++++ api/tests/test_model_router.py | 187 +++++++++++++ api/tests/test_model_sources.py | 103 +++++++ .../main/extension-install-utils.test.mjs | 49 ++++ electron/main/extension-install-utils.ts | 21 +- electron/main/ipc-handlers.ts | 152 ++++++++-- electron/main/model-download-plan.test.mjs | 92 +++++++ electron/main/model-download-plan.ts | 139 ++++++++++ electron/main/model-download-preload.test.mjs | 55 ++++ electron/main/model-downloader.ts | 30 +- electron/main/model-sources.test.mjs | 109 ++++++++ electron/main/model-sources.ts | 190 +++++++++++++ electron/preload/electron-api.ts | 6 +- src/areas/models/ModelsPage.tsx | 44 +-- .../models/components/ExtensionDrawer.tsx | 7 +- .../models/components/extensionShared.tsx | 6 +- src/areas/models/utils.test.mjs | 29 ++ src/areas/models/utils.ts | 24 ++ src/shared/types/electron.d.ts | 6 +- 23 files changed, 1731 insertions(+), 60 deletions(-) create mode 100644 api/services/model_sources.py create mode 100644 api/tests/test_model_router.py create mode 100644 api/tests/test_model_sources.py create mode 100644 electron/main/model-download-plan.test.mjs create mode 100644 electron/main/model-download-plan.ts create mode 100644 electron/main/model-download-preload.test.mjs create mode 100644 electron/main/model-sources.test.mjs create mode 100644 electron/main/model-sources.ts diff --git a/README.md b/README.md index 11184ddc..9414106e 100644 --- a/README.md +++ b/README.md @@ -106,6 +106,42 @@ Modly supports external model and process extensions. Each extension is a GitHub ![Install models](docs/install-models.png) +### Multiple Hugging Face repositories per model node + +A model node whose weights are split across repositories can declare +`model_sources`. Modly validates every source, downloads them sequentially in +one Models-page action, and considers the node installed only when every +declared check exists. + +```json +{ + "id": "generate", + "model_sources": [ + { + "id": "primary", + "provider": "huggingface", + "repo_id": "org/main-model", + "destination": ".", + "checks": ["model.safetensors"] + }, + { + "id": "encoder", + "provider": "huggingface", + "repo_id": "org/encoder", + "revision": "v1.0", + "destination": "auxiliary/encoder", + "include_prefixes": ["config.json", "model.safetensors"], + "checks": ["config.json", "model.safetensors"] + } + ] +} +``` + +`destination`, filters, and checks use safe POSIX paths relative to the node's +model directory. The only supported provider is `huggingface`. Existing nodes +that use `hf_repo`, `download_check`, `hf_include_prefixes`, and +`hf_skip_prefixes` keep their original behavior. + --- ## Workflows diff --git a/api/routers/model.py b/api/routers/model.py index 4f04718b..509fadda 100644 --- a/api/routers/model.py +++ b/api/routers/model.py @@ -8,9 +8,16 @@ from typing import Optional from urllib.error import HTTPError, URLError from urllib.request import Request, urlopen -from fastapi import APIRouter, HTTPException +from fastapi import APIRouter, HTTPException, Request as FastAPIRequest from fastapi.responses import StreamingResponse from services.generator_registry import generator_registry, MODELS_DIR +from services.model_sources import ( + normalize_model_sources, + resolve_download_path, + resolve_model_root, + resolve_source_destination, + validate_source_file_plan, +) router = APIRouter(tags=["model"]) @@ -97,7 +104,7 @@ async def unload_all_models(): return {"unloaded": True} -@router.post("/unload/{model_id}") +@router.post("/unload/{model_id:path}") async def unload_model(model_id: str): """Unloads a model from memory so its files can be safely deleted.""" try: @@ -122,6 +129,159 @@ async def cancel_hf_download(model_id: str): return {"cancelled": True} +@router.post("/hf-download-sources") +async def hf_download_sources(request: FastAPIRequest, model_id: str): + """Download all Hugging Face sources declared for one model node.""" + try: + body = await request.json() + if not isinstance(body, dict): + raise ValueError("Request body must be an object") + sources = normalize_model_sources({"model_sources": body.get("sources")}) + if sources is None: + raise ValueError("sources are required") + model_root = resolve_model_root(MODELS_DIR, model_id) + destinations = { + source["id"]: resolve_source_destination( + MODELS_DIR, model_id, source["destination"] + ) + for source in sources + } + except (TypeError, ValueError) as exc: + raise HTTPException(400, str(exc)) from exc + + authorization = request.headers.get("authorization", "") + hf_token = ( + authorization[7:].strip() + if authorization.lower().startswith("bearer ") + else os.environ.get("HUGGING_FACE_HUB_TOKEN") + or os.environ.get("HF_TOKEN") + or None + ) + control = _new_download_control(model_id) + + async def stream(): + loop = asyncio.get_running_loop() + + def _fmt(data: dict) -> str: + return f"data: {json.dumps(data)}\n\n" + + try: + yield _fmt({"percent": 0, "status": "Listing repository files..."}) + files_by_source: dict[str, list[str]] = {} + + for source in sources: + _check_download_control(control) + + def _list_files(current=source): + from huggingface_hub import list_repo_files + listed = list_repo_files( + current["repo_id"], + revision=current.get("revision"), + token=hf_token, + ) + include = current.get("include_prefixes", []) + skip = current.get("skip_prefixes", []) + return [ + filename for filename in listed + if (not include or any(filename.startswith(prefix) for prefix in include)) + if not any(filename.startswith(prefix) for prefix in skip) + ] + + files = await loop.run_in_executor(None, _list_files) + if not files: + raise RuntimeError( + f'No files found in Hugging Face repo: {source["repo_id"]}' + ) + destination = destinations[source["id"]] + for filename in files: + resolve_download_path(destination, filename) + files_by_source[source["id"]] = files + + validate_source_file_plan(sources, files_by_source) + planned_files = [ + (source, filename) + for source in sources + for filename in files_by_source[source["id"]] + ] + total = len(planned_files) + yield _fmt({"percent": 1, "status": f"Downloading {total} files..."}) + + from huggingface_hub import hf_hub_url + + for index, (source, filename) in enumerate(planned_files): + _check_download_control(control) + display_file = f'{source["id"]}/{filename}' + base_percent = 1 + round(index / total * 94) + yield _fmt({ + "percent": base_percent, + "file": display_file, + "fileIndex": index + 1, + "totalFiles": total, + "status": f"Starting {display_file}", + "bytesDownloaded": 0, + "stalledSeconds": 0, + }) + + queue: asyncio.Queue[dict] = asyncio.Queue() + + def _progress(message: dict) -> None: + message["file"] = display_file + loop.call_soon_threadsafe(queue.put_nowait, message) + + url = hf_hub_url( + repo_id=source["repo_id"], + filename=filename, + revision=source.get("revision"), + ) + future = loop.run_in_executor( + None, + lambda: _download_file_streamed( + url=url, + filename=filename, + dest_dir=str(destinations[source["id"]]), + file_index=index + 1, + total_files=total, + base_percent=base_percent, + progress_cb=_progress, + control=control, + token=hf_token, + ), + ) + while not future.done(): + try: + message = await asyncio.wait_for(queue.get(), timeout=2.0) + except asyncio.TimeoutError: + continue + yield _fmt(message) + + final_size = await future + _check_download_control(control) + yield _fmt({ + "percent": 1 + round((index + 1) / total * 94), + "file": display_file, + "fileIndex": index + 1, + "totalFiles": total, + "status": "Downloaded", + "bytesDownloaded": final_size, + "stalledSeconds": 0, + }) + + yield _fmt({"percent": 100, "status": "done"}) + except DownloadPaused: + yield _fmt({"paused": True, "status": "paused"}) + except DownloadCancelled: + for part in model_root.rglob("*.part"): + part.unlink(missing_ok=True) + yield _fmt({"cancelled": True, "status": "cancelled"}) + except Exception as exc: + yield _fmt({"error": str(exc)}) + finally: + if _download_controls.get(model_id) is control: + _download_controls.pop(model_id, None) + + return StreamingResponse(stream(), media_type="text/event-stream") + + @router.get("/hf-download") async def hf_download( repo_id: str, diff --git a/api/services/generator_registry.py b/api/services/generator_registry.py index ecb6bb6c..348a42cb 100644 --- a/api/services/generator_registry.py +++ b/api/services/generator_registry.py @@ -25,6 +25,7 @@ from services.generators.base import BaseGenerator from services.extension_process import ExtensionProcess, _venv_python +from services.model_sources import model_sources_are_downloaded, normalize_model_sources # ------------------------------------------------------------------ # # Global paths @@ -428,6 +429,9 @@ def _discover_extensions( ext_id = manifest["id"] class_name = manifest["generator_class"] + if "model_sources" in manifest: + raise ValueError("model_sources must be declared on a model node") + if ext_id != ext_dir.name: message = ( f"Extension folder '{ext_dir.name}' declares mismatched " @@ -523,6 +527,7 @@ def _discover_extensions( if nodes: for node in nodes: + model_sources = normalize_model_sources(node) node_manifest = { **manifest, "id": f"{ext_id}/{node['id']}", @@ -537,6 +542,8 @@ def _discover_extensions( "input": node.get("input", "image"), "output": node.get("output", "mesh"), } + if model_sources is not None: + node_manifest["model_sources"] = model_sources full_id = f"{ext_id}/{node['id']}" result[full_id] = (cls_or_None, node_manifest, ext_dir, legacy_context) if subprocess_mode: @@ -699,8 +706,14 @@ def get_active(self) -> BaseGenerator: """Returns the active generator. Downloads and loads if necessary.""" self._assert_not_quarantined(self._active_id) gen = self._generators[self._active_id] + downloaded = self._is_downloaded(self._active_id, gen) + if "model_sources" in self._manifests[self._active_id] and not downloaded: + raise RuntimeError( + "Model sources are incomplete. Download this node's weights " + "from the Modly Models page before generation." + ) if not gen.is_loaded(): - if not gen.is_downloaded(): + if not downloaded: if isinstance(gen, ExtensionProcess): # Let the subprocess handle its own download logic during # load() — some extensions (e.g. mv-adapter) need custom @@ -743,13 +756,21 @@ def switch_model(self, model_id: str) -> None: # Status # ------------------------------------------------------------------ # + def _is_downloaded(self, model_id: str, gen: BaseGenerator) -> bool: + manifest = self._manifests[model_id] + if "model_sources" in manifest: + return model_sources_are_downloaded( + MODELS_DIR, model_id, manifest["model_sources"] + ) + return gen.is_downloaded() + def active_status(self) -> dict: gen = self._generators[self._active_id] manifest = self._manifests[self._active_id] return { "id": self._active_id, "name": manifest.get("name", gen.DISPLAY_NAME), - "downloaded": gen.is_downloaded(), + "downloaded": self._is_downloaded(self._active_id, gen), "loaded": gen.is_loaded(), } @@ -765,7 +786,7 @@ def all_status(self) -> list: "vram_gb": manifest.get("vram_gb", gen.VRAM_GB), "hf_repo": manifest.get("hf_repo", ""), "tags": manifest.get("tags", []), - "downloaded": gen.is_downloaded(), + "downloaded": self._is_downloaded(model_id, gen), "loaded": gen.is_loaded(), "active": model_id == self._active_id, }) diff --git a/api/services/model_sources.py b/api/services/model_sources.py new file mode 100644 index 00000000..5988ae02 --- /dev/null +++ b/api/services/model_sources.py @@ -0,0 +1,259 @@ +"""Validation and readiness helpers for manifest-declared Hugging Face sources.""" + +from __future__ import annotations + +import re +import unicodedata +from pathlib import Path +from typing import Any + + +_SAFE_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]*$") +_WINDOWS_DEVICE = re.compile( + r"^(?:con|prn|aux|nul|com[1-9]|lpt[1-9])(?:\..*)?$", re.IGNORECASE +) +_WINDOWS_UNSAFE = re.compile(r'[<>"|?*\x00-\x1f]') + + +def _portable_segment(value: str, field: str) -> str: + if ( + not value + or value in {".", ".."} + or value.endswith((".", " ")) + or ":" in value + or _WINDOWS_UNSAFE.search(value) + or _WINDOWS_DEVICE.fullmatch(value) + ): + raise ValueError(f'{field} contains unsafe path segment "{value}"') + return value + + +def safe_source_id(value: Any, field: str = "model source id") -> str: + if ( + not isinstance(value, str) + or not value + or value != value.strip() + or _SAFE_ID.fullmatch(value) is None + ): + raise ValueError(f"{field} must be a safe non-empty identifier") + return _portable_segment(value, field) + + +def safe_relative_path(value: Any, field: str, *, allow_dot: bool = False) -> str: + if not isinstance(value, str) or not value or value != value.strip(): + raise ValueError(f"{field} must be a non-empty relative path") + if allow_dot and value == ".": + return value + if value == "." or value.startswith("/") or "\\" in value: + raise ValueError(f"{field} must be a safe relative POSIX path") + for part in value.split("/"): + _portable_segment(part, field) + return value + + +def _safe_prefix(value: Any, field: str) -> str: + path = value[:-1] if isinstance(value, str) and value.endswith("/") else value + safe_relative_path(path, field) + return value + + +def _prefixes(value: Any, field: str) -> list[str] | None: + if not isinstance(value, list): + raise ValueError(f"{field} must be an array") + return [_safe_prefix(entry, f"{field}[{index}]") for index, entry in enumerate(value)] + + +def _safe_repo_id(value: Any, field: str) -> str: + if not isinstance(value, str) or not value or value != value.strip() or "\\" in value: + raise ValueError(f"{field} must be a non-empty Hugging Face repository id") + parts = value.split("/") + if len(parts) > 2 or any( + part in {"", ".", ".."} or _SAFE_ID.fullmatch(part) is None for part in parts + ): + raise ValueError(f"{field} is not a safe Hugging Face repository id") + return value + + +def _safe_revision(value: Any, field: str) -> str | None: + if value is None: + return None + if ( + not isinstance(value, str) + or not value + or value != value.strip() + or value.startswith("/") + or "\\" in value + or "\0" in value + or any(part in {"", ".", ".."} for part in value.split("/")) + ): + raise ValueError(f"{field} must be a safe non-empty revision") + return value + + +def normalize_model_sources(node: dict[str, Any]) -> list[dict[str, Any]] | None: + """Validate only the new contract; legacy fields remain untouched.""" + if "model_sources" not in node: + return None + raw_sources = node["model_sources"] + if not isinstance(raw_sources, list) or not raw_sources: + raise ValueError("model_sources must be a non-empty array") + + aliases: dict[str, str] = {} + sources: list[dict[str, Any]] = [] + for index, raw in enumerate(raw_sources): + field = f"model_sources[{index}]" + if not isinstance(raw, dict): + raise ValueError(f"{field} must be an object") + source_id = safe_source_id(raw.get("id"), f"{field}.id") + alias = unicodedata.normalize("NFC", source_id).casefold() + if alias in aliases: + raise ValueError( + f'model source ids "{aliases[alias]}" and "{source_id}" are not portable-unique' + ) + aliases[alias] = source_id + if raw.get("provider") != "huggingface": + raise ValueError(f'{field}.provider must be "huggingface"') + checks = raw.get("checks") + if not isinstance(checks, list) or not checks: + raise ValueError(f"{field}.checks must be a non-empty array") + + source: dict[str, Any] = { + "id": source_id, + "provider": "huggingface", + "repo_id": _safe_repo_id(raw.get("repo_id"), f"{field}.repo_id"), + "destination": safe_relative_path( + raw.get("destination"), f"{field}.destination", allow_dot=True + ), + "checks": [ + safe_relative_path(check, f"{field}.checks[{check_index}]") + for check_index, check in enumerate(checks) + ], + } + revision = ( + _safe_revision(raw["revision"], f"{field}.revision") + if "revision" in raw + else None + ) + if "revision" in raw and revision is None: + raise ValueError(f"{field}.revision must be a safe non-empty revision") + include = ( + _prefixes(raw["include_prefixes"], f"{field}.include_prefixes") + if "include_prefixes" in raw + else None + ) + skip = ( + _prefixes(raw["skip_prefixes"], f"{field}.skip_prefixes") + if "skip_prefixes" in raw + else None + ) + if revision is not None: + source["revision"] = revision + if include is not None: + source["include_prefixes"] = include + if skip is not None: + source["skip_prefixes"] = skip + sources.append(source) + return sources + + +def _path_has_symlink(root: Path, candidate: Path) -> bool: + root = root.absolute() + candidate = candidate.absolute() + try: + relative = candidate.relative_to(root) + except ValueError: + return True + current = root + if current.exists() and current.is_symlink(): + return True + for part in relative.parts: + current /= part + if current.exists() and current.is_symlink(): + return True + return False + + +def resolve_model_root(models_dir: Path, model_id: str) -> Path: + if not isinstance(model_id, str): + raise ValueError("Model id must be a string") + parts = model_id.split("/") + if len(parts) != 2: + raise ValueError("Model id must identify one extension node") + extension_id = safe_source_id(parts[0], "extension id") + node_id = safe_source_id(parts[1], "model node id") + root = models_dir.absolute() + candidate = root / extension_id / node_id + if _path_has_symlink(root, candidate): + raise ValueError("Model path resolves through a symlink") + try: + candidate.resolve().relative_to(root.resolve()) + except ValueError as exc: + raise ValueError("Model path escapes the models directory") from exc + return candidate + + +def resolve_source_destination(models_dir: Path, model_id: str, destination: str) -> Path: + model_root = resolve_model_root(models_dir, model_id) + safe_destination = safe_relative_path(destination, "destination", allow_dot=True) + candidate = model_root if safe_destination == "." else model_root.joinpath(*safe_destination.split("/")) + if _path_has_symlink(model_root, candidate): + raise ValueError("Source destination resolves through a symlink") + return candidate + + +def resolve_download_path(destination: Path, filename: str) -> Path: + safe_filename = safe_relative_path(filename, "Hugging Face repository file") + candidate = destination.joinpath(*safe_filename.split("/")) + if _path_has_symlink(destination, candidate): + raise ValueError("Download target resolves through a symlink") + return candidate + + +def model_sources_are_downloaded( + models_dir: Path, model_id: str, sources: list[dict[str, Any]] +) -> bool: + try: + model_root = resolve_model_root(models_dir, model_id) + if not model_root.is_dir(): + return False + for source in sources: + destination = resolve_source_destination( + models_dir, model_id, source["destination"] + ) + if not destination.is_dir(): + return False + for check in source["checks"]: + candidate = resolve_download_path(destination, check) + if not candidate.exists() or _path_has_symlink(model_root, candidate): + return False + return bool(sources) + except (KeyError, OSError, TypeError, ValueError): + return False + + +def validate_source_file_plan( + sources: list[dict[str, Any]], files_by_source: dict[str, list[str]] +) -> None: + """Reject cross-source aliases before the first file is written.""" + aliases: dict[str, tuple[str, str]] = {} + for source in sources: + source_id = source["id"] + destination = source["destination"] + for filename in files_by_source[source_id]: + safe_filename = safe_relative_path(filename, f'model source "{source_id}" file') + target = safe_filename if destination == "." else f"{destination}/{safe_filename}" + for value in (target, f"{target}.part"): + alias = unicodedata.normalize("NFC", value).casefold() + for previous_alias, (previous_source, previous_target) in aliases.items(): + if previous_source == source_id: + continue + if ( + alias == previous_alias + or alias.startswith(f"{previous_alias}/") + or previous_alias.startswith(f"{alias}/") + ): + raise ValueError( + "Model sources have a portable target collision: " + f'"{previous_source}:{previous_target}" and "{source_id}:{value}"' + ) + aliases[alias] = (source_id, value) diff --git a/api/tests/test_generator_registry.py b/api/tests/test_generator_registry.py index 7340b72c..ff9d090c 100644 --- a/api/tests/test_generator_registry.py +++ b/api/tests/test_generator_registry.py @@ -155,6 +155,62 @@ def test_legacy_generator_supports_eager_and_lazy_sibling_imports(self) -> None: self.registry.reload() self.assertNotIn(str(extension.resolve()), sys.path) + def test_declared_sources_block_generation_even_when_generator_overrides_readiness(self) -> None: + extension = self._make_extension("multi-source") + manifest = { + "id": "multi-source", + "name": "multi-source", + "type": "model", + "generator_class": "TestGenerator", + "nodes": [{ + "id": "generate", + "model_sources": [ + { + "id": "primary", + "provider": "huggingface", + "repo_id": "org/main", + "destination": ".", + "checks": ["main.bin"], + }, + { + "id": "encoder", + "provider": "huggingface", + "repo_id": "org/encoder", + "destination": "auxiliary/encoder", + "checks": ["encoder.bin"], + }, + ], + }], + } + (extension / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + (extension / "generator.py").write_text( + "\n".join([ + "from services.generators.base import BaseGenerator", + "class TestGenerator(BaseGenerator):", + " def is_downloaded(self): return True", + " def load(self): self._model = object()", + " def generate(self, image_bytes, params, progress_cb=None, cancel_event=None):", + " return self.outputs_dir / 'result.glb'", + ]), + encoding="utf-8", + ) + + self.registry.initialize() + self.registry._active_id = "multi-source/generate" + with self.assertRaisesRegex(RuntimeError, "Model sources are incomplete"): + self.registry.get_active() + self.assertFalse(self.registry.all_status()[0]["downloaded"]) + + model_root = self.models_dir / "multi-source" / "generate" + (model_root / "auxiliary" / "encoder").mkdir(parents=True) + (model_root / "main.bin").write_bytes(b"main") + (model_root / "auxiliary" / "encoder" / "encoder.bin").write_bytes(b"encoder") + self.assertIsNotNone(self.registry.get_active()) + self.assertTrue(self.registry.all_status()[0]["downloaded"]) + (model_root / "main.bin").unlink() + with self.assertRaisesRegex(RuntimeError, "Model sources are incomplete"): + self.registry.get_active() + def test_reload_preserves_legacy_path_owned_by_the_host(self) -> None: extension = self._make_extension("host-owned-path") self._write_manifest(extension, extension_id="host-owned-path") diff --git a/api/tests/test_model_router.py b/api/tests/test_model_router.py new file mode 100644 index 00000000..d0fe1279 --- /dev/null +++ b/api/tests/test_model_router.py @@ -0,0 +1,187 @@ +import asyncio +import json +import sys +import tempfile +import types +import unittest +from pathlib import Path +from unittest.mock import patch + +from starlette.requests import Request + +import routers.model as model_router + + +SOURCES = [ + { + "id": "primary", + "provider": "huggingface", + "repo_id": "org/main", + "destination": ".", + "checks": ["main.bin"], + }, + { + "id": "encoder", + "provider": "huggingface", + "repo_id": "org/encoder", + "destination": "auxiliary/encoder", + "checks": ["encoder.bin"], + }, +] + + +def request_for(sources: list[dict]) -> Request: + body = json.dumps({"sources": sources}).encode() + sent = False + + async def receive(): + nonlocal sent + if sent: + return {"type": "http.disconnect"} + sent = True + return {"type": "http.request", "body": body, "more_body": False} + + return Request({ + "type": "http", + "method": "POST", + "path": "/model/hf-download-sources", + "headers": [(b"authorization", b"Bearer test-token")], + "query_string": b"", + "server": ("test", 80), + "client": ("test", 1), + "scheme": "http", + }, receive) + + +async def collect_events(response) -> list[dict]: + payload = "" + async for chunk in response.body_iterator: + payload += chunk.decode() if isinstance(chunk, bytes) else chunk + return [ + json.loads(block[6:]) + for block in payload.strip().split("\n\n") + if block.startswith("data: ") + ] + + +class MultiSourceRouterTests(unittest.TestCase): + def setUp(self) -> None: + self.tempdir = tempfile.TemporaryDirectory(prefix="modly-model-router-") + self.models_dir = Path(self.tempdir.name) / "models" + self.models_dir.mkdir() + self.old_models_dir = model_router.MODELS_DIR + model_router.MODELS_DIR = self.models_dir + self.old_hf_module = sys.modules.get("huggingface_hub") + + def tearDown(self) -> None: + model_router.MODELS_DIR = self.old_models_dir + model_router._download_controls.clear() + if self.old_hf_module is None: + sys.modules.pop("huggingface_hub", None) + else: + sys.modules["huggingface_hub"] = self.old_hf_module + self.tempdir.cleanup() + + def install_hf_stub(self, files: dict[str, list[str]], calls: list[str]) -> None: + module = types.ModuleType("huggingface_hub") + + def list_repo_files(repo_id, revision=None, token=None): + calls.append(f"list:{repo_id}:{revision}:{token}") + return files[repo_id] + + def hf_hub_url(repo_id, filename, revision=None): + return f"https://example.invalid/{repo_id}/{revision or 'main'}/{filename}" + + module.list_repo_files = list_repo_files + module.hf_hub_url = hf_hub_url + sys.modules["huggingface_hub"] = module + + def test_lists_every_source_before_sequential_download_with_monotonic_progress(self) -> None: + calls: list[str] = [] + controls: list[int] = [] + self.install_hf_stub({"org/main": ["main.bin"], "org/encoder": ["encoder.bin"]}, calls) + + def fake_download(**kwargs): + calls.append(f"download:{kwargs['dest_dir']}:{kwargs['filename']}") + controls.append(id(kwargs["control"])) + target = Path(kwargs["dest_dir"]) / kwargs["filename"] + target.parent.mkdir(parents=True, exist_ok=True) + target.write_bytes(b"data") + kwargs["progress_cb"]({ + "percent": kwargs["base_percent"], + "file": kwargs["filename"], + "fileIndex": kwargs["file_index"], + "totalFiles": kwargs["total_files"], + "status": "Downloading...", + "bytesDownloaded": 4, + "stalledSeconds": 0, + }) + return 4 + + async def run(): + with patch.object(model_router, "_download_file_streamed", fake_download): + response = await model_router.hf_download_sources( + request_for(SOURCES), "pixal3d/generate" + ) + return await collect_events(response) + + events = asyncio.run(run()) + first_download = next(index for index, value in enumerate(calls) if value.startswith("download:")) + self.assertTrue(all(value.startswith("list:") for value in calls[:first_download])) + self.assertEqual(len(set(controls)), 1) + self.assertEqual([event["percent"] for event in events if "percent" in event], sorted( + event["percent"] for event in events if "percent" in event + )) + self.assertEqual(events[-1], {"percent": 100, "status": "done"}) + self.assertTrue((self.models_dir / "pixal3d/generate/main.bin").is_file()) + self.assertTrue((self.models_dir / "pixal3d/generate/auxiliary/encoder/encoder.bin").is_file()) + + def test_pause_cancel_and_resume_reuse_one_model_control(self) -> None: + calls: list[str] = [] + self.install_hf_stub({"org/main": ["main.bin"]}, calls) + source = [SOURCES[0]] + mode = "pause" + + def controlled_download(**kwargs): + target = Path(kwargs["dest_dir"]) / kwargs["filename"] + target.parent.mkdir(parents=True, exist_ok=True) + part = Path(f"{target}.part") + part.write_bytes(b"partial") + if mode == "pause": + kwargs["control"]["pause"].set() + model_router._check_download_control(kwargs["control"]) + if mode == "cancel": + kwargs["control"]["cancel"].set() + model_router._check_download_control(kwargs["control"]) + part.replace(target) + return target.stat().st_size + + async def one_run(): + with patch.object(model_router, "_download_file_streamed", controlled_download): + response = await model_router.hf_download_sources( + request_for(source), "pixal3d/generate" + ) + return await collect_events(response) + + paused = asyncio.run(one_run()) + self.assertTrue(paused[-1]["paused"]) + self.assertTrue((self.models_dir / "pixal3d/generate/main.bin.part").is_file()) + + mode = "cancel" + cancelled = asyncio.run(one_run()) + self.assertTrue(cancelled[-1]["cancelled"]) + self.assertFalse((self.models_dir / "pixal3d/generate/main.bin.part").exists()) + + mode = "resume" + resumed = asyncio.run(one_run()) + self.assertEqual(resumed[-1], {"percent": 100, "status": "done"}) + self.assertTrue((self.models_dir / "pixal3d/generate/main.bin").is_file()) + + def test_composite_model_unload_route_uses_path_converter(self) -> None: + paths = {route.path for route in model_router.router.routes} + self.assertIn("/unload/{model_id:path}", paths) + self.assertEqual(model_router.Request.__module__, "urllib.request") + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_model_sources.py b/api/tests/test_model_sources.py new file mode 100644 index 00000000..72a68c71 --- /dev/null +++ b/api/tests/test_model_sources.py @@ -0,0 +1,103 @@ +import os +import tempfile +import unittest +from pathlib import Path + +from services.model_sources import ( + model_sources_are_downloaded, + normalize_model_sources, + resolve_model_root, + validate_source_file_plan, +) + + +def valid_node() -> dict: + return { + "model_sources": [ + { + "id": "primary", + "provider": "huggingface", + "repo_id": "org/main", + "destination": ".", + "checks": ["pipeline.json"], + }, + { + "id": "encoder", + "provider": "huggingface", + "repo_id": "org/encoder", + "revision": "refs/pr/1", + "destination": "auxiliary/encoder", + "include_prefixes": ["config.json", "weights/"], + "checks": ["config.json", "model.safetensors"], + }, + ] + } + + +class ModelSourcesTests(unittest.TestCase): + def test_validates_new_sources_without_reinterpreting_legacy_fields(self) -> None: + sources = normalize_model_sources(valid_node()) + self.assertEqual([source["id"] for source in sources or []], ["primary", "encoder"]) + self.assertIsNone(normalize_model_sources({ + "hf_repo": "legacy/repo", + "download_check": "../generate/model.safetensors", + "hf_skip_prefixes": ["weights/**"], + })) + + def test_rejects_unsafe_and_non_portable_declarations(self) -> None: + source = valid_node()["model_sources"][0] + for destination in ("../outside", "aux/CON", "aux/name.", "C:/models"): + with self.subTest(destination=destination), self.assertRaises(ValueError): + normalize_model_sources({ + "model_sources": [{**source, "destination": destination}] + }) + with self.assertRaisesRegex(ValueError, "provider"): + normalize_model_sources({ + "model_sources": [{**source, "provider": "url"}] + }) + with self.assertRaisesRegex(ValueError, "portable-unique"): + normalize_model_sources({ + "model_sources": [source, {**source, "id": "PRIMARY"}] + }) + with self.assertRaisesRegex(ValueError, "checks"): + normalize_model_sources({ + "model_sources": [{**source, "checks": []}] + }) + + def test_rejects_portable_cross_source_file_collisions(self) -> None: + sources = normalize_model_sources(valid_node()) or [] + with self.assertRaisesRegex(ValueError, "portable target collision"): + validate_source_file_plan(sources, { + "primary": ["Auxiliary/Encoder/model.safetensors"], + "encoder": ["model.safetensors"], + }) + + def test_requires_all_checks_and_rejects_symlinked_extension_ancestry(self) -> None: + sources = normalize_model_sources(valid_node()) or [] + with tempfile.TemporaryDirectory(prefix="modly-model-sources-") as tmp: + models = Path(tmp) / "models" + model_root = models / "pixal3d" / "generate" + encoder = model_root / "auxiliary" / "encoder" + encoder.mkdir(parents=True) + (model_root / "pipeline.json").write_text("{}", encoding="utf-8") + (encoder / "config.json").write_text("{}", encoding="utf-8") + self.assertFalse(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + (encoder / "model.safetensors").write_bytes(b"x") + self.assertTrue(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + + for child in sorted((models / "pixal3d").rglob("*"), reverse=True): + child.unlink() if child.is_file() else child.rmdir() + (models / "pixal3d").rmdir() + outside = Path(tmp) / "outside" + (outside / "generate").mkdir(parents=True) + try: + os.symlink(outside, models / "pixal3d", target_is_directory=True) + except (NotImplementedError, OSError) as exc: + self.skipTest(f"Symlinks unavailable: {exc}") + with self.assertRaisesRegex(ValueError, "symlink"): + resolve_model_root(models, "pixal3d/generate") + self.assertFalse(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + + +if __name__ == "__main__": + unittest.main() diff --git a/electron/main/extension-install-utils.test.mjs b/electron/main/extension-install-utils.test.mjs index d5d1a389..84139f9a 100644 --- a/electron/main/extension-install-utils.test.mjs +++ b/electron/main/extension-install-utils.test.mjs @@ -53,6 +53,55 @@ test('validateInstallManifest still rejects missing process entry files', () => ) }) +test('validateInstallManifest accepts multi-source nodes and preserves legacy shapes', () => { + const mod = loadModule() + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'multi-model', + generator_class: 'Generator', + nodes: [{ + id: 'generate', + model_sources: [ + { + id: 'primary', provider: 'huggingface', repo_id: 'org/main', + destination: '.', checks: ['pipeline.json'], + }, + { + id: 'encoder', provider: 'huggingface', repo_id: 'org/encoder', + destination: 'auxiliary/encoder', checks: ['model.safetensors'], + }, + ], + }], + }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) + + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'legacy', + generator_class: 'Generator', + nodes: [{ + id: 'projection', + hf_repo: 'org/legacy', + download_check: '../generate/model.safetensors', + hf_skip_prefixes: ['weights/**'], + }], + }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) +}) + +test('validateInstallManifest rejects malformed or process model_sources', () => { + const mod = loadModule() + const source = { + id: 'weights', provider: 'huggingface', repo_id: 'org/model', + destination: '../outside', checks: ['model.safetensors'], + } + assert.throws(() => mod.validateInstallManifest({ + id: 'unsafe', generator_class: 'Generator', + nodes: [{ id: 'generate', model_sources: [source] }], + }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository'), /destination/i) + + assert.throws(() => mod.validateInstallManifest({ + id: 'process', type: 'process', entry: 'processor.js', + nodes: [{ id: 'run', model_sources: [{ ...source, destination: '.' }] }], + }, { hasEntryFile: () => true, hasGeneratorFile: () => false }, 'repository'), /only for model nodes/i) +}) + test('python process setup failures are treated as fatal', () => { const mod = loadModule() diff --git a/electron/main/extension-install-utils.ts b/electron/main/extension-install-utils.ts index 4195cdca..05b965b0 100644 --- a/electron/main/extension-install-utils.ts +++ b/electron/main/extension-install-utils.ts @@ -1,9 +1,16 @@ +import { + normalizeModelSources, + safeModelSourceId, + type ModelSourceNode, +} from './model-sources' + export interface InstallManifest { id?: string type?: 'model' | 'process' entry?: string generator_class?: string - nodes?: Array<{ id?: string }> + model_sources?: unknown + nodes?: Array<{ id?: string; model_sources?: unknown } & ModelSourceNode> } export interface ValidatedInstallManifest { @@ -39,6 +46,18 @@ export function validateInstallManifest( const entryFile = manifest.entry ?? 'processor.js' const nodes = Array.isArray(manifest.nodes) ? manifest.nodes.filter((node) => node?.id) : [] + if (manifest.model_sources !== undefined) { + throw new Error('manifest.json: model_sources must be declared on a model node') + } + for (const node of Array.isArray(manifest.nodes) ? manifest.nodes : []) { + if (node.model_sources === undefined) continue + if (isProcess) { + throw new Error('manifest.json: model_sources is supported only for model nodes') + } + safeModelSourceId(node.id, 'model node id') + normalizeModelSources(node) + } + if (isProcess) { if (!opts.hasEntryFile(entryFile)) { throw new Error(`manifest.json: entry file "${entryFile}" missing from ${sourceLabel}`) diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index a0b2c136..a504d962 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -13,7 +13,15 @@ import { isModelDownloaded, listDownloadedModels, downloadModelFromHF, + downloadModelSourcesFromHF, } from './model-downloader' +import { resolveInstalledModelDownloadPlan } from './model-download-plan' +import { + areModelSourcesDownloaded, + modelHasLocalData, + normalizeModelSources, + resolveModelRoot, +} from './model-sources' import { getSettings, setSettings } from './settings-store' import { checkSetupNeeded, markSetupDone, runFullSetup, getVenvPythonExe, ensureSslPatch } from './python-setup' import { logger } from './logger' @@ -265,7 +273,12 @@ const renameWithRetry = (from: string, to: string, label: string) => renameExtensionWithRetry(from, to, label, logger) export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGetter): void { - const activeDownloads = new Map() + type ActiveDownload = { + progress: { percent: number; file?: string; fileIndex?: number; totalFiles?: number } + done: Promise + finish: () => void + } + const activeDownloads = new Map() // Logging from renderer ipcMain.on('log:error', (_event, message: string) => logger.error(`[Renderer] ${message}`)) ipcMain.handle('log:getPath', () => join(app.getPath('userData'), 'logs', 'modly.log')) @@ -423,7 +436,21 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }) ipcMain.handle('model:delete', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { - const modelDir = join(getSettings(app.getPath('userData')).modelsDir, modelId) + if (activeDownloads.has(modelId)) { + return { success: false, error: 'Cannot remove model weights while their download is active' } + } + let modelDir: string + try { + await resolveInstalledModelDownloadPlan({ + modelId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + modelDir = resolveModelRoot(getSettings(app.getPath('userData')).modelsDir, modelId) + } catch (err) { + return { success: false, error: String(err) } + } // Unload the model and wait for confirmation so file handles are released try { @@ -475,28 +502,74 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe return listDownloadedModels(modelsDir) }) - ipcMain.handle('model:isDownloaded', (_, modelId: string, downloadCheck?: string): boolean => { + ipcMain.handle('model:isDownloaded', async (_, modelId: string): Promise => { const modelsDir = getSettings(app.getPath('userData')).modelsDir - return isModelDownloaded(modelsDir, modelId, downloadCheck) + try { + const plan = await resolveInstalledModelDownloadPlan({ + modelId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + return plan.kind === 'multi-source' + ? areModelSourcesDownloaded(modelsDir, modelId, plan.sources) + : isModelDownloaded(modelsDir, modelId, plan.downloadCheck) + } catch { + return false + } + }) + + ipcMain.handle('model:hasLocalData', async (_, modelId: string): Promise => { + try { + await resolveInstalledModelDownloadPlan({ + modelId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + return modelHasLocalData(getSettings(app.getPath('userData')).modelsDir, modelId) + } catch { + return false + } }) ipcMain.handle('model:activeDownloads', () => - [...activeDownloads.entries()].map(([modelId, progress]) => ({ modelId, ...progress })) + [...activeDownloads.entries()].map(([modelId, active]) => ({ modelId, ...active.progress })) ) ipcMain.handle('model:download', async ( event, - { repoId, modelId, skipPrefixes, includePrefixes }: { repoId: string; modelId: string; skipPrefixes?: string[]; includePrefixes?: string[] }, + modelId: string, ) => { if (activeDownloads.has(modelId)) { return { success: false, error: 'Download already in progress' } } - activeDownloads.set(modelId, { percent: 0 }) + let finish!: () => void + const done = new Promise((resolveDone) => { finish = resolveDone }) + const active: ActiveDownload = { progress: { percent: 0 }, done, finish } + activeDownloads.set(modelId, active) try { - await downloadModelFromHF(repoId, modelId, (progress) => { - activeDownloads.set(modelId, progress) + const plan = await resolveInstalledModelDownloadPlan({ + modelId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + const onProgress = (progress: typeof active.progress) => { + active.progress = progress event.sender.send('model:downloadProgress', { modelId, ...progress }) - }, skipPrefixes, includePrefixes) + } + if (plan.kind === 'multi-source') { + await downloadModelSourcesFromHF(modelId, plan.sources, onProgress) + } else { + await downloadModelFromHF( + plan.repoId, + modelId, + onProgress, + plan.skipPrefixes, + plan.includePrefixes, + ) + } return { success: true } } catch (err: any) { const message = err?.message ?? String(err) @@ -510,7 +583,8 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe } return { success: false, error: String(err) } } finally { - activeDownloads.delete(modelId) + if (activeDownloads.get(modelId) === active) activeDownloads.delete(modelId) + active.finish() } }) @@ -528,17 +602,24 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe ipcMain.handle('model:cancelDownload', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { try { + const active = activeDownloads.get(modelId) await axios.post(`${API_BASE_URL}/model/hf-download/cancel`, null, { params: { model_id: modelId }, timeout: 5000, }) - const modelDir = join(getSettings(app.getPath('userData')).modelsDir, modelId) + if (active) { + await Promise.race([ + active.done, + new Promise((_, reject) => { + setTimeout(() => reject(new Error('Timed out waiting for the download to stop')), 30_000) + }), + ]) + } + const modelDir = resolveModelRoot(getSettings(app.getPath('userData')).modelsDir, modelId) await rmAsync(modelDir, { recursive: true, force: true }) return { success: true } } catch (err) { return { success: false, error: String(err) } - } finally { - activeDownloads.delete(modelId) } }) @@ -835,6 +916,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // extension type type?: 'model' | 'process' entry?: string + model_sources?: unknown // Optional top-level fallbacks — applied to each node if not set on the node params_schema?: unknown[] param_defaults?: Record @@ -851,6 +933,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe download_check?: string hf_skip_prefixes?: string[] hf_include_prefixes?: string[] + model_sources?: unknown }[] } @@ -866,20 +949,30 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe builtin, } - const nodes = (parsed.nodes ?? []).map(n => ({ - id: n.id, - name: n.name ?? n.id, - input: n.input ?? 'image' as const, - inputs: n.inputs, - inputLabels: n.input_labels, - output: n.output ?? 'mesh' as const, - paramsSchema: n.params_schema ?? parsed.params_schema ?? [], - paramDefaults: { ...(parsed.param_defaults ?? {}), ...(n.param_defaults ?? {}) }, - hfRepo: n.hf_repo, - downloadCheck: n.download_check, - hfSkipPrefixes: n.hf_skip_prefixes, - hfIncludePrefixes: n.hf_include_prefixes, - })) + if (parsed.model_sources !== undefined) { + throw new Error('manifest.json: model_sources must be declared on a model node') + } + const nodes = (parsed.nodes ?? []).map(n => { + if (parsed.type === 'process' && n.model_sources !== undefined) { + throw new Error('manifest.json: model_sources is supported only for model nodes') + } + const modelSources = normalizeModelSources(n) + return { + id: n.id, + name: n.name ?? n.id, + input: n.input ?? 'image' as const, + inputs: n.inputs, + inputLabels: n.input_labels, + output: n.output ?? 'mesh' as const, + paramsSchema: n.params_schema ?? parsed.params_schema ?? [], + paramDefaults: { ...(parsed.param_defaults ?? {}), ...(n.param_defaults ?? {}) }, + hfRepo: n.hf_repo, + downloadCheck: n.download_check, + hfSkipPrefixes: n.hf_skip_prefixes, + hfIncludePrefixes: n.hf_include_prefixes, + hasModelSources: modelSources !== undefined, + } + }) if (parsed.type === 'process') { return { ...common, type: 'process' as const, entry: parsed.entry ?? 'processor.js', nodes } @@ -1450,6 +1543,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // Uninstall an extension — built-ins cannot be uninstalled ipcMain.handle('extensions:uninstall', async (_, extensionId: string) => { try { + if ([...activeDownloads.keys()].some((modelId) => modelId.split('/', 1)[0] === extensionId)) { + return { success: false, error: 'Cannot uninstall an extension while its model download is active' } + } // Corrupted folders can carry arbitrary names (manual copies, failed // unzips), so only enforce root confinement for the deletion path. The // strict id pattern still guards the built-in check — a non-conforming diff --git a/electron/main/model-download-plan.test.mjs b/electron/main/model-download-plan.test.mjs new file mode 100644 index 00000000..e5d037a0 --- /dev/null +++ b/electron/main/model-download-plan.test.mjs @@ -0,0 +1,92 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, rmSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-model-plan-module-')), 'model-plan.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/model-download-plan.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +function setupExtension(manifest) { + const root = mkdtempSync(join(tmpdir(), 'modly-action-plan-')) + const user = join(root, 'user') + const builtin = join(root, 'builtin') + const extension = join(user, manifest.id) + mkdirSync(extension, { recursive: true }) + mkdirSync(builtin) + const manifestPath = join(extension, 'manifest.json') + writeFileSync(manifestPath, JSON.stringify(manifest)) + return { root, user, builtin, manifestPath } +} + +test('re-reads the installed manifest for each action and resolves only node-owned sources', async () => { + const { resolveInstalledModelDownloadPlan } = loadModule() + const manifest = { + id: 'pixal3d', + type: 'model', + nodes: [{ + id: 'generate', + model_sources: [{ + id: 'primary', provider: 'huggingface', repo_id: 'org/old', + destination: '.', checks: ['model.safetensors'], + }], + }], + } + const fixture = setupExtension(manifest) + const args = { + modelId: 'pixal3d/generate', + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + } + try { + const first = await resolveInstalledModelDownloadPlan(args) + assert.equal(first.kind, 'multi-source') + assert.equal(first.sources[0].repo_id, 'org/old') + + manifest.nodes[0].model_sources[0].repo_id = 'org/new' + writeFileSync(fixture.manifestPath, JSON.stringify(manifest)) + const second = await resolveInstalledModelDownloadPlan(args) + assert.equal(second.sources[0].repo_id, 'org/new') + } finally { + rmSync(fixture.root, { recursive: true, force: true }) + } +}) + +test('keeps legacy sibling checks and wildcard filters unchanged', async () => { + const { resolveInstalledModelDownloadPlan } = loadModule() + const fixture = setupExtension({ + id: 'triposplat', + type: 'model', + nodes: [{ + id: 'projection', + hf_repo: 'VAST-AI/TripoSplat', + download_check: '../generate/diffusion_models/triposplat_fp16.safetensors', + hf_skip_prefixes: ['weights/**', 'assets/*'], + }], + }) + try { + const plan = await resolveInstalledModelDownloadPlan({ + modelId: 'triposplat/projection', + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + }) + assert.equal(plan.kind, 'legacy') + assert.equal(plan.downloadCheck, '../generate/diffusion_models/triposplat_fp16.safetensors') + assert.deepEqual(plan.skipPrefixes, ['weights/**', 'assets/*']) + } finally { + rmSync(fixture.root, { recursive: true, force: true }) + } +}) diff --git a/electron/main/model-download-plan.ts b/electron/main/model-download-plan.ts new file mode 100644 index 00000000..27d2dba6 --- /dev/null +++ b/electron/main/model-download-plan.ts @@ -0,0 +1,139 @@ +import { existsSync } from 'node:fs' +import { readFile, readdir } from 'node:fs/promises' +import { join } from 'node:path' + +import { + EXT_INCOMPLETE_MARKER, + EXT_REGISTRATION_PENDING_MARKER, + assertSafeExtensionId, + resolveExtensionPathWithinRoot, +} from './extension-path-guard' +import { normalizeModelSources, safeModelSourceId, type ModelSource } from './model-sources' + +interface InstalledNode { + id?: unknown + hf_repo?: unknown + download_check?: unknown + hf_skip_prefixes?: unknown + hf_include_prefixes?: unknown + model_sources?: unknown +} + +interface InstalledManifest { + id?: unknown + type?: unknown + model_sources?: unknown + nodes?: unknown +} + +export type InstalledModelDownloadPlan = { + kind: 'legacy' + modelId: string + extensionId: string + nodeId: string + repoId: string + downloadCheck?: string + skipPrefixes?: string[] + includePrefixes?: string[] +} | { + kind: 'multi-source' + modelId: string + extensionId: string + nodeId: string + sources: ModelSource[] +} + +async function hasPendingRegistration(root: string, extensionId: string): Promise { + try { + const prefix = `${EXT_REGISTRATION_PENDING_MARKER}-${extensionId}-` + return (await readdir(root)).some((name) => ( + name.startsWith(prefix) && /^\d+$/.test(name.slice(prefix.length)) + )) + } catch { + return false + } +} + +function parseManifest(raw: string, extensionId: string, nodeId: string): InstalledModelDownloadPlan { + let parsed: unknown + try { + parsed = JSON.parse(raw) + } catch { + throw new Error(`Extension "${extensionId}" has an invalid manifest.json`) + } + if (typeof parsed !== 'object' || parsed === null || Array.isArray(parsed)) { + throw new Error(`Extension "${extensionId}" has an invalid manifest.json`) + } + const manifest = parsed as InstalledManifest + if (manifest.id !== extensionId) throw new Error(`Installed manifest id does not match extension "${extensionId}"`) + if (manifest.type !== undefined && manifest.type !== 'model') { + throw new Error(`Extension "${extensionId}" is not a model extension`) + } + if (manifest.model_sources !== undefined) { + throw new Error('manifest.json: model_sources must be declared on a model node') + } + if (!Array.isArray(manifest.nodes)) throw new Error(`Extension "${extensionId}" does not declare model nodes`) + + const matches = manifest.nodes.filter((candidate): candidate is InstalledNode => ( + typeof candidate === 'object' + && candidate !== null + && !Array.isArray(candidate) + && (candidate as InstalledNode).id === nodeId + )) + if (matches.length !== 1) { + throw new Error(`Installed manifest must declare model node "${nodeId}" exactly once`) + } + + const node = matches[0] + const modelId = `${extensionId}/${nodeId}` + const sources = normalizeModelSources(node) + if (sources) return { kind: 'multi-source', modelId, extensionId, nodeId, sources } + + if (typeof node.hf_repo !== 'string' || !node.hf_repo) { + throw new Error(`Model node "${modelId}" has no Hugging Face download source`) + } + return { + kind: 'legacy', + modelId, + extensionId, + nodeId, + repoId: node.hf_repo, + downloadCheck: typeof node.download_check === 'string' ? node.download_check : undefined, + skipPrefixes: node.hf_skip_prefixes as string[] | undefined, + includePrefixes: node.hf_include_prefixes as string[] | undefined, + } +} + +/** Re-read the installed manifest for every model action; renderer metadata is never trusted. */ +export async function resolveInstalledModelDownloadPlan(args: { + modelId: unknown + userExtensionsDir: string + builtinExtensionsDir: string + blockedExtensionIds?: ReadonlySet +}): Promise { + if (typeof args.modelId !== 'string') throw new Error('Model id must be a string') + const parts = args.modelId.split('/') + if (parts.length !== 2) throw new Error('Model id must identify one extension node') + const extensionId = assertSafeExtensionId(parts[0]) + const nodeId = safeModelSourceId(parts[1], 'model node id') + if (args.blockedExtensionIds?.has(extensionId)) { + throw new Error(`Extension "${extensionId}" is being installed or repaired`) + } + + const userPath = resolveExtensionPathWithinRoot(args.userExtensionsDir, extensionId) + const builtinPath = resolveExtensionPathWithinRoot(args.builtinExtensionsDir, extensionId) + const extensionPath = existsSync(userPath) ? userPath : existsSync(builtinPath) ? builtinPath : undefined + if (!extensionPath) throw new Error(`Extension "${extensionId}" is not installed`) + const extensionRoot = extensionPath === userPath ? args.userExtensionsDir : args.builtinExtensionsDir + if ( + existsSync(join(extensionPath, EXT_INCOMPLETE_MARKER)) + || existsSync(join(extensionPath, EXT_REGISTRATION_PENDING_MARKER)) + || await hasPendingRegistration(extensionRoot, extensionId) + ) { + throw new Error(`Extension "${extensionId}" has an incomplete installation`) + } + + const manifestPath = join(extensionPath, 'manifest.json') + if (!existsSync(manifestPath)) throw new Error(`Extension "${extensionId}" has no manifest.json`) + return parseManifest(await readFile(manifestPath, 'utf-8'), extensionId, nodeId) +} diff --git a/electron/main/model-download-preload.test.mjs b/electron/main/model-download-preload.test.mjs new file mode 100644 index 00000000..8d9a16b7 --- /dev/null +++ b/electron/main/model-download-preload.test.mjs @@ -0,0 +1,55 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, readFileSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-preload-download-')), 'electron-api.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/preload/electron-api.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +test('renderer model actions send only the model node id', async () => { + const { createElectronApi } = loadModule() + const calls = [] + const ipc = { + invoke: async (...args) => { calls.push(args); return { success: true } }, + send: () => {}, + on: () => {}, + removeAllListeners: () => {}, + } + const api = createElectronApi(ipc, { setZoomFactor: () => {} }) + + await api.model.isDownloaded('pixal3d/generate') + await api.model.hasLocalData('pixal3d/generate') + await api.model.download('pixal3d/generate') + + assert.deepEqual(calls, [ + ['model:isDownloaded', 'pixal3d/generate'], + ['model:hasLocalData', 'pixal3d/generate'], + ['model:download', 'pixal3d/generate'], + ]) +}) + +test('declared partial data is removable and active downloads block destructive actions', () => { + const main = readFileSync(resolve('electron/main/ipc-handlers.ts'), 'utf8') + const page = readFileSync(resolve('src/areas/models/ModelsPage.tsx'), 'utf8') + const drawer = readFileSync(resolve('src/areas/models/components/ExtensionDrawer.tsx'), 'utf8') + + assert.match(main, /model:delete[\s\S]*activeDownloads\.has\(modelId\)/) + assert.match(main, /extensions:uninstall[\s\S]*activeDownloads\.keys\(\)/) + assert.match(page, /window\.electron\.model\.hasLocalData\(fullId\)/) + assert.match(drawer, /localDataIds\.includes\(fullId\) && state\.kind !== 'downloading'/) + assert.match(drawer, /Remove partial model data/) +}) diff --git a/electron/main/model-downloader.ts b/electron/main/model-downloader.ts index 8c5571d9..174b474a 100644 --- a/electron/main/model-downloader.ts +++ b/electron/main/model-downloader.ts @@ -6,6 +6,7 @@ import { existsSync, readdirSync, statSync, readFileSync } from 'fs' import { join } from 'path' import { getSettings } from './settings-store' import { app } from 'electron' +import type { ModelSource } from './model-sources' export interface DownloadProgress { percent: number @@ -121,7 +122,6 @@ export async function downloadModelFromHF( includePrefixes?: string[], ): Promise { const { net } = require('electron') - const STALL_TIMEOUT_MS = 120_000 let url = `${PYTHON_API_URL}/model/hf-download?repo_id=${encodeURIComponent(repoId)}&model_id=${encodeURIComponent(modelId)}` if (skipPrefixes && skipPrefixes.length > 0) { url += `&skip_prefixes=${encodeURIComponent(JSON.stringify(skipPrefixes))}` @@ -136,11 +136,39 @@ export async function downloadModelFromHF( const res = await net.fetch(url) if (!res.ok) throw new Error(`HuggingFace download failed: HTTP ${res.status}`) + await consumeDownloadStream(res, onProgress) +} + +/** Download a validated node-level source plan through one aggregate SSE stream. */ +export async function downloadModelSourcesFromHF( + modelId: string, + sources: ModelSource[], + onProgress: ProgressCallback, +): Promise { + const { net } = require('electron') + const headers: Record = { 'Content-Type': 'application/json' } + const hfToken = getSettings(app.getPath('userData')).hfToken + if (hfToken) headers.Authorization = `Bearer ${hfToken}` + const url = `${PYTHON_API_URL}/model/hf-download-sources?model_id=${encodeURIComponent(modelId)}` + const res = await net.fetch(url, { + method: 'POST', + headers, + body: JSON.stringify({ sources }), + }) + if (!res.ok) throw new Error(`HuggingFace multi-source download failed: HTTP ${res.status}`) + await consumeDownloadStream(res, onProgress) +} + +async function consumeDownloadStream( + res: Response, + onProgress: ProgressCallback, +): Promise { if (!res.body) throw new Error('No response body from HF download stream') const decoder = new TextDecoder() const reader = res.body.getReader() let buffer = '' + const STALL_TIMEOUT_MS = 120_000 async function readWithTimeout() { return await Promise.race([ diff --git a/electron/main/model-sources.test.mjs b/electron/main/model-sources.test.mjs new file mode 100644 index 00000000..b29f0dc5 --- /dev/null +++ b/electron/main/model-sources.test.mjs @@ -0,0 +1,109 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, rmSync, symlinkSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-model-sources-module-')), 'model-sources.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/model-sources.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const validNode = () => ({ + model_sources: [ + { + id: 'primary', + provider: 'huggingface', + repo_id: 'org/main', + destination: '.', + checks: ['pipeline.json'], + }, + { + id: 'encoder', + provider: 'huggingface', + repo_id: 'org/encoder', + revision: 'refs/pr/1', + destination: 'auxiliary/encoder', + include_prefixes: ['config.json', 'weights/'], + skip_prefixes: ['README.md'], + checks: ['config.json', 'model.safetensors'], + }, + ], +}) + +test('validates the new model_sources contract without reinterpreting legacy fields', () => { + const { normalizeModelSources } = loadModule() + const sources = normalizeModelSources(validNode()) + assert.equal(sources.length, 2) + assert.equal(sources[1].destination, 'auxiliary/encoder') + + assert.equal(normalizeModelSources({ + hf_repo: 'legacy/repo', + download_check: '../generate/model.safetensors', + hf_skip_prefixes: ['weights/**'], + }), undefined) +}) + +test('rejects unsafe destinations, unsupported providers, and non-portable source aliases', () => { + const { normalizeModelSources } = loadModule() + const source = validNode().model_sources[0] + for (const destination of ['../outside', 'aux/CON', 'aux/name.', 'C:/models']) { + assert.throws( + () => normalizeModelSources({ model_sources: [{ ...source, destination }] }), + /destination|unsafe/i, + ) + } + assert.throws( + () => normalizeModelSources({ model_sources: [{ ...source, provider: 'url' }] }), + /provider.*huggingface/i, + ) + assert.throws( + () => normalizeModelSources({ model_sources: [source, { ...source, id: 'PRIMARY' }] }), + /portable-unique/i, + ) + assert.throws( + () => normalizeModelSources({ model_sources: [{ ...source, checks: [] }] }), + /checks.*non-empty/i, + ) +}) + +test('requires every declared check and rejects symlinked extension-root ancestry', (t) => { + const { areModelSourcesDownloaded, normalizeModelSources } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-model-readiness-')) + const models = join(root, 'models') + const modelRoot = join(models, 'pixal3d', 'generate') + const sources = normalizeModelSources(validNode()) + mkdirSync(join(modelRoot, 'auxiliary', 'encoder'), { recursive: true }) + writeFileSync(join(modelRoot, 'pipeline.json'), '{}') + writeFileSync(join(modelRoot, 'auxiliary', 'encoder', 'config.json'), '{}') + + try { + assert.equal(areModelSourcesDownloaded(models, 'pixal3d/generate', sources), false) + writeFileSync(join(modelRoot, 'auxiliary', 'encoder', 'model.safetensors'), 'x') + assert.equal(areModelSourcesDownloaded(models, 'pixal3d/generate', sources), true) + + rmSync(join(models, 'pixal3d'), { recursive: true, force: true }) + const outside = join(root, 'outside') + mkdirSync(join(outside, 'generate'), { recursive: true }) + try { + symlinkSync(outside, join(models, 'pixal3d'), 'dir') + } catch (error) { + t.skip(`Symlinks unavailable: ${error}`) + return + } + assert.equal(areModelSourcesDownloaded(models, 'pixal3d/generate', sources), false) + } finally { + rmSync(root, { recursive: true, force: true }) + } +}) diff --git a/electron/main/model-sources.ts b/electron/main/model-sources.ts new file mode 100644 index 00000000..03b82948 --- /dev/null +++ b/electron/main/model-sources.ts @@ -0,0 +1,190 @@ +import { existsSync, lstatSync, readdirSync } from 'node:fs' +import { isAbsolute, relative, resolve } from 'node:path' + +export interface ModelSource { + id: string + provider: 'huggingface' + repo_id: string + revision?: string + destination: string + include_prefixes?: string[] + skip_prefixes?: string[] + checks: string[] +} + +export interface ModelSourceNode { + model_sources?: unknown +} + +const SAFE_ID = /^[A-Za-z0-9][A-Za-z0-9._-]*$/ +const WINDOWS_DEVICE = /^(?:con|prn|aux|nul|com[1-9]|lpt[1-9])(?:\..*)?$/i +const WINDOWS_UNSAFE = /[<>"|?*\u0000-\u001f]/ + +function portableSegment(value: string, field: string): string { + if ( + !value + || value === '.' + || value === '..' + || value.endsWith('.') + || value.endsWith(' ') + || value.includes(':') + || WINDOWS_UNSAFE.test(value) + || WINDOWS_DEVICE.test(value) + ) { + throw new Error(`${field} contains unsafe path segment "${value}"`) + } + return value +} + +export function safeModelSourceId(value: unknown, field = 'model source id'): string { + if (typeof value !== 'string' || !value || value !== value.trim() || !SAFE_ID.test(value)) { + throw new Error(`${field} must be a safe non-empty identifier`) + } + return portableSegment(value, field) +} + +export function safeModelRelativePath(value: unknown, field: string, allowDot = false): string { + if (typeof value !== 'string' || !value || value !== value.trim()) { + throw new Error(`${field} must be a non-empty relative path`) + } + if (allowDot && value === '.') return value + if (value === '.' || value.startsWith('/') || value.includes('\\') || isAbsolute(value)) { + throw new Error(`${field} must be a safe relative POSIX path`) + } + for (const part of value.split('/')) portableSegment(part, field) + return value +} + +function safePrefix(value: unknown, field: string): string { + if (typeof value !== 'string') return safeModelRelativePath(value, field) + const path = value.endsWith('/') ? value.slice(0, -1) : value + safeModelRelativePath(path, field) + return value +} + +function optionalPrefixes(value: unknown, field: string): string[] | undefined { + if (value === undefined) return undefined + if (!Array.isArray(value)) throw new Error(`${field} must be an array`) + return value.map((entry, index) => safePrefix(entry, `${field}[${index}]`)) +} + +function safeRepoId(value: unknown, field: string): string { + if (typeof value !== 'string' || !value || value !== value.trim() || value.includes('\\')) { + throw new Error(`${field} must be a non-empty Hugging Face repository id`) + } + const parts = value.split('/') + if (parts.length > 2 || parts.some((part) => !SAFE_ID.test(part) || part === '.' || part === '..')) { + throw new Error(`${field} is not a safe Hugging Face repository id`) + } + return value +} + +function safeRevision(value: unknown, field: string): string | undefined { + if (value === undefined) return undefined + if (typeof value !== 'string' || !value || value !== value.trim() || value.startsWith('/') || value.includes('\\') || value.includes('\0')) { + throw new Error(`${field} must be a safe non-empty revision`) + } + if (value.split('/').some((part) => !part || part === '.' || part === '..')) { + throw new Error(`${field} must be a safe revision`) + } + return value +} + +export function normalizeModelSources(node: ModelSourceNode): ModelSource[] | undefined { + if (!Object.prototype.hasOwnProperty.call(node, 'model_sources')) return undefined + if (!Array.isArray(node.model_sources) || node.model_sources.length === 0) { + throw new Error('model_sources must be a non-empty array') + } + + const seen = new Map() + return node.model_sources.map((raw, index) => { + const field = `model_sources[${index}]` + if (typeof raw !== 'object' || raw === null || Array.isArray(raw)) { + throw new Error(`${field} must be an object`) + } + const value = raw as Record + const id = safeModelSourceId(value.id, `${field}.id`) + const alias = id.normalize('NFC').toLowerCase() + const previous = seen.get(alias) + if (previous) throw new Error(`model source ids "${previous}" and "${id}" are not portable-unique`) + seen.set(alias, id) + if (value.provider !== 'huggingface') throw new Error(`${field}.provider must be "huggingface"`) + if (!Array.isArray(value.checks) || value.checks.length === 0) { + throw new Error(`${field}.checks must be a non-empty array`) + } + + const source: ModelSource = { + id, + provider: 'huggingface', + repo_id: safeRepoId(value.repo_id, `${field}.repo_id`), + destination: safeModelRelativePath(value.destination, `${field}.destination`, true), + checks: value.checks.map((check, checkIndex) => ( + safeModelRelativePath(check, `${field}.checks[${checkIndex}]`) + )), + } + const revision = safeRevision(value.revision, `${field}.revision`) + const include = optionalPrefixes(value.include_prefixes, `${field}.include_prefixes`) + const skip = optionalPrefixes(value.skip_prefixes, `${field}.skip_prefixes`) + if (revision !== undefined) source.revision = revision + if (include !== undefined) source.include_prefixes = include + if (skip !== undefined) source.skip_prefixes = skip + return source + }) +} + +function pathHasSymlink(root: string, candidate: string): boolean { + const rootPath = resolve(root) + const rel = relative(rootPath, resolve(candidate)) + if (rel === '..' || rel.startsWith('../') || rel.startsWith('..\\') || isAbsolute(rel)) return true + let current = rootPath + try { + if (existsSync(current) && lstatSync(current).isSymbolicLink()) return true + for (const part of rel.split(/[/\\]/).filter(Boolean)) { + current = resolve(current, part) + if (existsSync(current) && lstatSync(current).isSymbolicLink()) return true + } + } catch { + return true + } + return false +} + +export function resolveModelRoot(modelsDir: string, modelId: string): string { + if (typeof modelId !== 'string') throw new Error('Model id must be a string') + const parts = modelId.split('/') + if (parts.length !== 2) throw new Error('Model id must identify one extension node') + const extensionId = safeModelSourceId(parts[0], 'extension id') + const nodeId = safeModelSourceId(parts[1], 'model node id') + const root = resolve(modelsDir) + const modelRoot = resolve(root, extensionId, nodeId) + if (pathHasSymlink(root, modelRoot)) throw new Error('Model path resolves through a symlink') + return modelRoot +} + +export function areModelSourcesDownloaded(modelsDir: string, modelId: string, sources: ModelSource[]): boolean { + try { + const modelRoot = resolveModelRoot(modelsDir, modelId) + if (!existsSync(modelRoot)) return false + return sources.every((source) => { + const destination = source.destination === '.' + ? modelRoot + : resolve(modelRoot, ...source.destination.split('/')) + if (!existsSync(destination) || pathHasSymlink(modelRoot, destination)) return false + return source.checks.every((check) => { + const candidate = resolve(destination, ...check.split('/')) + return existsSync(candidate) && !pathHasSymlink(modelRoot, candidate) + }) + }) + } catch { + return false + } +} + +export function modelHasLocalData(modelsDir: string, modelId: string): boolean { + try { + const modelRoot = resolveModelRoot(modelsDir, modelId) + return existsSync(modelRoot) && readdirSync(modelRoot).length > 0 + } catch { + return false + } +} diff --git a/electron/preload/electron-api.ts b/electron/preload/electron-api.ts index af211f0d..c929204c 100644 --- a/electron/preload/electron-api.ts +++ b/electron/preload/electron-api.ts @@ -111,9 +111,9 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra model: { export: (args: { outputUrl: string; format: string }) => ipcRenderer.invoke('model:export', args), listDownloaded: () => ipcRenderer.invoke('model:listDownloaded'), - isDownloaded: (modelId: string, downloadCheck?: string) => ipcRenderer.invoke('model:isDownloaded', modelId, downloadCheck), - download: (repoId: string, modelId: string, skipPrefixes?: string[], includePrefixes?: string[]) => - ipcRenderer.invoke('model:download', { repoId, modelId, skipPrefixes, includePrefixes }), + isDownloaded: (modelId: string) => ipcRenderer.invoke('model:isDownloaded', modelId), + hasLocalData: (modelId: string) => ipcRenderer.invoke('model:hasLocalData', modelId), + download: (modelId: string) => ipcRenderer.invoke('model:download', modelId), pauseDownload: (modelId: string) => ipcRenderer.invoke('model:pauseDownload', modelId), cancelDownload: (modelId: string) => ipcRenderer.invoke('model:cancelDownload', modelId), delete: (modelId: string) => ipcRenderer.invoke('model:delete', modelId), diff --git a/src/areas/models/ModelsPage.tsx b/src/areas/models/ModelsPage.tsx index 895d4e8f..66f8e5d4 100644 --- a/src/areas/models/ModelsPage.tsx +++ b/src/areas/models/ModelsPage.tsx @@ -2,11 +2,11 @@ import { useEffect, useMemo, useRef, useState } from 'react' import { createPortal } from 'react-dom' import { useExtensionsStore } from '@shared/stores/extensionsStore' import type { AnyExtension, ModelExtension } from '@shared/types/electron.d' -import { formatModelName } from './utils' +import { deleteModelsThenUninstallExtension, formatModelName } from './utils' import { ExtensionCard } from './components/ExtensionCard' import type { ExtensionNode } from './components/ExtensionCard' import { ExtensionDrawer } from './components/ExtensionDrawer' -import { ICONS } from './components/extensionShared' +import { ICONS, nodeHasManagedWeights } from './components/extensionShared' // ─── Filters & sorts ────────────────────────────────────────────────────────── @@ -49,6 +49,7 @@ export default function ModelsPage(): JSX.Element { // Model weight state (needed for node install status + uninstall cleanup) const [installedVariantIds, setInstalledVariantIds] = useState([]) + const [localDataIds, setLocalDataIds] = useState([]) const [downloading, setDownloading] = useState { @@ -161,11 +168,11 @@ export default function ModelsPage(): JSX.Element { // ── Node install / download controls ────────────────────────────────────── function handleInstallNode(node: ExtensionNode, fullId: string) { - if (!node.hfRepo) return + if (!nodeHasManagedWeights(node)) return setDownloading((prev) => ({ ...prev, [fullId]: { ...(prev[fullId] ?? { percent: 0 }), paused: false, status: 'Starting…' } })) - window.electron.model.download(node.hfRepo!, fullId, node.hfSkipPrefixes, node.hfIncludePrefixes).then((result: { success: boolean; paused?: boolean; cancelled?: boolean }) => { + window.electron.model.download(fullId).then((result) => { if (!result.success && !result.paused && !result.cancelled) { - setGhErr('Download failed') + setGhErr(result.error ?? 'Download failed') setDownloading((prev) => { const n = { ...prev }; delete n[fullId]; return n }) } }) @@ -174,7 +181,7 @@ export default function ModelsPage(): JSX.Element { function handleInstallAll(ext: AnyExtension) { if (ext.type !== 'model') return for (const node of ext.nodes) { - if (!node.hfRepo) continue + if (!nodeHasManagedWeights(node)) continue const fullId = `${ext.id}/${node.id}` if (installedVariantIds.includes(fullId) || downloading[fullId]) continue handleInstallNode(node, fullId) @@ -188,7 +195,9 @@ export default function ModelsPage(): JSX.Element { async function handleCancelDownload(fullId: string) { setDownloading((prev) => { const n = { ...prev }; delete n[fullId]; return n }) - await window.electron.model.cancelDownload(fullId) + const result = await window.electron.model.cancelDownload(fullId) + if (!result.success) setGhErr(result.error ?? 'Could not cancel download') + await refreshInstalledIds(useExtensionsStore.getState().modelExtensions) } async function handleUninstallNode(fullId: string) { @@ -228,8 +237,8 @@ export default function ModelsPage(): JSX.Element { function openUninstallModal(extId: string) { const ext = allExtensions.find((e) => e.id === extId) if (ext?.type === 'model') { - const installedModels = ext.nodes.filter((n) => installedVariantIds.includes(`${extId}/${n.id}`)) - setModelsToDelete(new Set(installedModels.map((n) => `${extId}/${n.id}`))) + const localModels = ext.nodes.filter((n) => localDataIds.includes(`${extId}/${n.id}`)) + setModelsToDelete(new Set(localModels.map((n) => `${extId}/${n.id}`))) } else { setModelsToDelete(new Set()) } @@ -237,10 +246,12 @@ export default function ModelsPage(): JSX.Element { } async function handleUninstallExtension(extId: string) { - for (const modelId of modelsToDelete) { - await window.electron.model.delete(modelId) - } - const result = await uninstallExt(extId) + const result = await deleteModelsThenUninstallExtension( + extId, + modelsToDelete, + (modelId) => window.electron.model.delete(modelId), + uninstallExt, + ) if (!result.success) { // Keep the dialog open so the failure is visible (locked folder, etc.) setUninstallError(result.error ?? 'Could not delete the extension folder.') @@ -617,6 +628,7 @@ export default function ModelsPage(): JSX.Element { { const ext = allExtensions.find((e) => e.id === uninstallTarget) const installedModels = ext?.type === 'model' - ? ext.nodes.filter((n) => installedVariantIds.includes(`${uninstallTarget}/${n.id}`)) + ? ext.nodes.filter((n) => localDataIds.includes(`${uninstallTarget}/${n.id}`)) : [] return createPortal( diff --git a/src/areas/models/components/ExtensionDrawer.tsx b/src/areas/models/components/ExtensionDrawer.tsx index 54adab9c..2f154c71 100644 --- a/src/areas/models/components/ExtensionDrawer.tsx +++ b/src/areas/models/components/ExtensionDrawer.tsx @@ -16,6 +16,7 @@ import { finishExtensionRepair, isExtensionRepairable } from '../utils' interface Props { ext: AnyExtension installedIds: string[] + localDataIds: string[] downloading: DownloadMap loadError?: string disabled?: boolean @@ -31,7 +32,7 @@ interface Props { } export function ExtensionDrawer({ - ext, installedIds, downloading, loadError, disabled, + ext, installedIds, localDataIds, downloading, loadError, disabled, onInstall, onInstallAll, onPauseDownload, onCancelDownload, onUninstallNode, onUninstall, onRepaired, onSynced, onClose, }: Props): JSX.Element { @@ -203,11 +204,11 @@ export function ExtensionDrawer({ onResume={() => onInstall(node, fullId)} onCancel={() => onCancelDownload(fullId)} /> - {state.kind === 'installed' && ( + {localDataIds.includes(fullId) && state.kind !== 'downloading' && ( + )} + + + {pointLights.length === 0 && ( +

No point lights yet.

+ )} + + {pointLights.map((pl) => ( +
+
+ onPointLightsChange(pointLights.map((p) => p.id === pl.id ? { ...p, color: c } : p))} + /> + Point + {pl.intensity.toFixed(1)} + +
+ onPointLightsChange(pointLights.map((p) => p.id === pl.id ? { ...p, intensity: parseFloat(e.target.value) } : p))} + className="w-full h-1.5 accent-violet-500 cursor-pointer" /> +
+ ))} + + - onPointLightsChange(pointLights.map((p) => p.id === pl.id ? { ...p, intensity: parseFloat(e.target.value) } : p))} className="w-full h-1.5 accent-violet-500 cursor-pointer" /> From 9bdaa90a340cd8b32c01204574b1ace1d059d1cf Mon Sep 17 00:00:00 2001 From: iammojogo-sudo Date: Sat, 12 Sep 2026 13:07:25 -0400 Subject: [PATCH 22/57] fix: clear gizmo mode on selection change; add rotate/scale gizmos for point lights - Fix gizmo re-activation bug: clear gizmoMode on any selection change (mesh or point light), not just when both are deselected - Show all transform tools (move/rotate/scale) when a point light is selected - Add RotateGizmo and ScaleGizmo to PointLightMarker for consistency --- src/areas/generate/GeneratePage.tsx | 56 ++++++++++------------ src/areas/generate/components/Viewer3D.tsx | 14 +++++- 2 files changed, 39 insertions(+), 31 deletions(-) diff --git a/src/areas/generate/GeneratePage.tsx b/src/areas/generate/GeneratePage.tsx index 6ce5e894..658893b8 100644 --- a/src/areas/generate/GeneratePage.tsx +++ b/src/areas/generate/GeneratePage.tsx @@ -193,7 +193,7 @@ function LightPopover({ onClose, pointLights, onPointLightsChange, - onSelectPointLight, + onSelectPointLight: _onSelectPointLight, }: { settings: LightSettings onChange: (s: LightSettings) => void @@ -692,11 +692,11 @@ export default function GeneratePage(): JSX.Element { const el = document.activeElement as HTMLElement | null if (el && (el instanceof HTMLInputElement || el instanceof HTMLTextAreaElement || el.isContentEditable)) return if (e.key === 'Escape') { setGizmoMode((m) => (m ? null : m)); return } - if (!hasModel || (!meshSelected && !selectedPointLightId)) return + if (!meshSelected && !selectedPointLightId) return const k = e.key.toLowerCase() if (k === 'w') setGizmoMode('translate') - else if (k === 'r' && meshSelected) setGizmoMode('rotate') - else if (k === 's' && meshSelected) setGizmoMode('scale') + else if (k === 'r') setGizmoMode('rotate') + else if (k === 's') setGizmoMode('scale') } window.addEventListener('keydown', handler) return () => window.removeEventListener('keydown', handler) @@ -1167,32 +1167,28 @@ export default function GeneratePage(): JSX.Element { - {meshSelected && ( - <> - setGizmoMode((m) => (m === 'rotate' ? null : 'rotate'))} - > - - - - - - setGizmoMode((m) => (m === 'scale' ? null : 'scale'))} - > - - - - - - - - - )} + setGizmoMode((m) => (m === 'rotate' ? null : 'rotate'))} + > + + + + + + setGizmoMode((m) => (m === 'scale' ? null : 'scale'))} + > + + + + + + + )} diff --git a/src/areas/generate/components/Viewer3D.tsx b/src/areas/generate/components/Viewer3D.tsx index ebd17323..234143bc 100644 --- a/src/areas/generate/components/Viewer3D.tsx +++ b/src/areas/generate/components/Viewer3D.tsx @@ -852,7 +852,7 @@ function PointLightMarker({ { e.stopPropagation(); onSelect() }} + onClick={(e) => { e.stopPropagation(); onSelect() }} > @@ -863,6 +863,18 @@ function PointLightMarker({ onDragEnd={() => onPositionChange([group.position.x, group.position.y, group.position.z])} /> )} + {group && isSelected && gizmoMode === 'rotate' && ( + onPositionChange([group.position.x, group.position.y, group.position.z])} + /> + )} + {group && isSelected && gizmoMode === 'scale' && ( + onPositionChange([group.position.x, group.position.y, group.position.z])} + /> + )} ) } From 7f5a505ed0a50af7e410d839069e57244ffe2ca1 Mon Sep 17 00:00:00 2001 From: kevin9327 <5299031+kevin9327@users.noreply.github.com> Date: Sun, 13 Sep 2026 06:42:19 +0900 Subject: [PATCH 23/57] fix(process-runner): rebuild a cached runner when its workspace moves Process-extension runners are cached per extension id and the cache ignored the arguments of every call after the first. The workspace folder is baked into each runner at construction, so after the workspace is moved in Settings (which updates paths at runtime, without a restart) workflow process nodes such as Mesh Optimizer kept writing their output into the previous workspace, where the viewer can no longer find it. Remember the arguments each runner was built with and replace the runner when they differ; identical arguments keep reusing the warm worker as before. Co-Authored-By: Claude Opus 5 --- electron/main/process-runner.test.mjs | 114 ++++++++++++++++++++++++++ electron/main/process-runner.ts | 17 +++- 2 files changed, 129 insertions(+), 2 deletions(-) create mode 100644 electron/main/process-runner.test.mjs diff --git a/electron/main/process-runner.test.mjs b/electron/main/process-runner.test.mjs new file mode 100644 index 00000000..6ce5656a --- /dev/null +++ b/electron/main/process-runner.test.mjs @@ -0,0 +1,114 @@ +/** + * Process-extension runners are cached per extension id and reused across + * workflow runs. The cache must not hand back a runner that was built for + * different arguments: the workspace folder is baked into each runner, so after + * the user moves the workspace in Settings (which updates the backend at + * runtime, without a restart) a stale runner keeps writing node output into the + * old folder, where the viewer can no longer find it. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-runner-test-')), 'process-runner.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/process-runner.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +// A JS process extension that reports the workspace it was given, plus how many +// runs this worker has served (module state survives only while it is reused). +function makeJsExtension(root) { + const extDir = join(root, 'js-ext') + mkdirSync(extDir, { recursive: true }) + writeFileSync(join(extDir, 'processor.js'), [ + 'let runs = 0', + 'module.exports = async (input, params, context) => {', + ' runs += 1', + ' return { filePath: context.workspaceDir, text: String(runs) }', + '}', + '', + ].join('\n')) + return extDir +} + +// A "Python" process extension driven through the same stdin/stdout protocol. +// Node stands in for the interpreter so the test needs no Python install. +function makeStdioExtension(root) { + const extDir = join(root, 'py-ext') + mkdirSync(extDir, { recursive: true }) + writeFileSync(join(extDir, 'processor.cjs'), [ + "let raw = ''", + "process.stdin.on('data', (chunk) => { raw += chunk })", + "process.stdin.on('end', () => {", + ' const data = JSON.parse(raw)', + " process.stdout.write(JSON.stringify({ type: 'done', result: { filePath: data.workspaceDir } }) + '\\n')", + '})', + '', + ].join('\n')) + return extDir +} + +test('a JS process runner follows the workspace after it moves', async () => { + const { getProcessRunner, terminateAllProcessRunners } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-runner-js-')) + const extDir = makeJsExtension(root) + const oldWorkspace = join(root, 'workspace-old') + const newWorkspace = join(root, 'workspace-new') + try { + const before = await getProcessRunner('js-ext', extDir, 'processor.js', oldWorkspace, root).run({}, {}) + assert.equal(before.filePath, oldWorkspace) + + const after = await getProcessRunner('js-ext', extDir, 'processor.js', newWorkspace, root).run({}, {}) + assert.equal(after.filePath, newWorkspace) + } finally { + terminateAllProcessRunners() + } +}) + +test('a Python process runner follows the workspace after it moves', async () => { + const { getPythonProcessRunner, terminateAllProcessRunners } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-runner-py-')) + const extDir = makeStdioExtension(root) + const oldWorkspace = join(root, 'workspace-old') + const newWorkspace = join(root, 'workspace-new') + try { + const before = await getPythonProcessRunner('py-ext', process.execPath, extDir, 'processor.cjs', oldWorkspace, root).run({}, {}) + assert.equal(before.filePath, oldWorkspace) + + const after = await getPythonProcessRunner('py-ext', process.execPath, extDir, 'processor.cjs', newWorkspace, root).run({}, {}) + assert.equal(after.filePath, newWorkspace) + } finally { + terminateAllProcessRunners() + } +}) + +test('unchanged arguments keep reusing the same warm runner', async () => { + const { getProcessRunner, terminateAllProcessRunners } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-runner-reuse-')) + const extDir = makeJsExtension(root) + const workspace = join(root, 'workspace') + try { + const first = getProcessRunner('js-ext', extDir, 'processor.js', workspace, root) + assert.equal((await first.run({}, {})).text, '1') + + const second = getProcessRunner('js-ext', extDir, 'processor.js', workspace, root) + assert.equal(second, first) + // Same worker thread: its module state carried over instead of reloading. + assert.equal((await second.run({}, {})).text, '2') + } finally { + terminateAllProcessRunners() + } +}) diff --git a/electron/main/process-runner.ts b/electron/main/process-runner.ts index 758f62d2..f29e5dae 100644 --- a/electron/main/process-runner.ts +++ b/electron/main/process-runner.ts @@ -270,6 +270,19 @@ export function getExtPythonExe(extDir: string): string | null { // ─── Registry (one runner per extension id, reused across calls) ────────────── const registry = new Map() +// The arguments each cached runner was built with. A runner bakes them in at +// construction, so a call with different ones — e.g. after the workspace is +// moved in Settings, which updates paths without a restart — must not get the +// old runner back, or node output keeps landing in the previous folder. +const registryArgs = new Map() + +function canReuseRunner(extensionId: string, args: string[]): boolean { + const key = JSON.stringify(args) + if (registry.has(extensionId) && registryArgs.get(extensionId) === key) return true + terminateProcessRunner(extensionId) + registryArgs.set(extensionId, key) + return false +} export function getProcessRunner( extensionId: string, @@ -278,7 +291,7 @@ export function getProcessRunner( workspaceDir: string, tempDir: string, ): ProcessRunner { - if (!registry.has(extensionId)) { + if (!canReuseRunner(extensionId, [extDir, entry, workspaceDir, tempDir])) { registry.set(extensionId, new ProcessRunner(extDir, entry, workspaceDir, tempDir)) } return registry.get(extensionId)! as ProcessRunner @@ -292,7 +305,7 @@ export function getPythonProcessRunner( workspaceDir: string, tempDir: string, ): PythonProcessRunner { - if (!registry.has(extensionId)) { + if (!canReuseRunner(extensionId, [pythonExe, extDir, entry, workspaceDir, tempDir])) { registry.set(extensionId, new PythonProcessRunner(pythonExe, extDir, entry, workspaceDir, tempDir)) } return registry.get(extensionId)! as PythonProcessRunner From 2b56f867b0cccf5ad7359acb7940c9ac7535b0d3 Mon Sep 17 00:00:00 2001 From: kevin9327 <5299031+kevin9327@users.noreply.github.com> Date: Sun, 13 Sep 2026 06:59:02 +0900 Subject: [PATCH 24/57] fix(process-runner): settle a JS process run when its worker dies A JS process extension runs in a worker thread kept warm between runs, and run() only listened for the worker's 'done'/'error' messages. If the worker died mid-run (an uncaught error outside the awaited processor call, running out of memory on a large mesh, process.exit) neither message ever arrived: the workflow waited on that node forever, and because the dead worker stayed cached, every later run of the node hung the same way until the app was restarted. Listen for the worker's 'error' and 'exit' events for the duration of a run, reject with the cause, and drop the dead worker so the next run starts a fresh one. Errors thrown by the processor itself are still reported through its 'error' message and keep the warm worker. Co-Authored-By: Claude Opus 5 --- .../main/process-runner-worker-exit.test.mjs | 83 +++++++++++++++++++ electron/main/process-runner.ts | 31 ++++++- 2 files changed, 112 insertions(+), 2 deletions(-) create mode 100644 electron/main/process-runner-worker-exit.test.mjs diff --git a/electron/main/process-runner-worker-exit.test.mjs b/electron/main/process-runner-worker-exit.test.mjs new file mode 100644 index 00000000..068ccc7e --- /dev/null +++ b/electron/main/process-runner-worker-exit.test.mjs @@ -0,0 +1,83 @@ +/** + * A JS process extension runs in a worker thread that is kept warm between + * runs. If that worker dies mid-run -- an uncaught error outside the awaited + * processor call, running out of memory on a large mesh, process.exit() -- it + * never posts 'done' or 'error'. The run must settle with an error instead of + * leaving the workflow waiting forever, and the next run must get a fresh + * worker rather than posting into the dead one. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-worker-exit-test-')), 'process-runner.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/process-runner.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +// Behaves according to params.mode; `runs` counts runs served by this worker, +// so a fresh worker starts again from 1. +function makeRunner() { + const { ProcessRunner } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-worker-exit-')) + const extDir = join(root, 'ext') + mkdirSync(extDir, { recursive: true }) + writeFileSync(join(extDir, 'processor.js'), [ + 'let runs = 0', + 'module.exports = async (input, params) => {', + ' runs += 1', + " if (params.mode === 'throw') throw new Error('bad input')", + " if (params.mode === 'crash') {", + " setTimeout(() => { throw new Error('worker blew up') }, 0)", + ' return new Promise(() => {})', + ' }', + " if (params.mode === 'exit') process.exit(3)", + ' return { text: String(runs) }', + '}', + '', + ].join('\n')) + return new ProcessRunner(extDir, 'processor.js', join(root, 'workspace'), root) +} + +test('a run whose worker crashes rejects instead of hanging', { timeout: 5000 }, async () => { + const runner = makeRunner() + try { + await assert.rejects(runner.run({}, { mode: 'crash' }), /worker blew up/) + } finally { + runner.terminate() + } +}) + +test('after its worker exits, the runner starts a fresh one for the next run', { timeout: 5000 }, async () => { + const runner = makeRunner() + try { + await assert.rejects(runner.run({}, { mode: 'exit' }), /exited with code 3/) + assert.deepEqual(await runner.run({}, { mode: 'ok' }), { text: '1' }) + } finally { + runner.terminate() + } +}) + +test('an error thrown by the processor still rejects with its message and keeps the warm worker', { timeout: 5000 }, async () => { + const runner = makeRunner() + try { + await assert.rejects(runner.run({}, { mode: 'throw' }), { message: 'Error: bad input' }) + // Same worker thread: its run counter carried over instead of restarting. + assert.deepEqual(await runner.run({}, { mode: 'ok' }), { text: '2' }) + } finally { + runner.terminate() + } +}) diff --git a/electron/main/process-runner.ts b/electron/main/process-runner.ts index 758f62d2..d12a7eec 100644 --- a/electron/main/process-runner.ts +++ b/electron/main/process-runner.ts @@ -127,25 +127,52 @@ export class ProcessRunner implements IProcessRunner { const worker = this.worker! return new Promise((resolve, reject) => { + const settle = () => { + worker.off('message', handler) + worker.off('error', onError) + worker.off('exit', onExit) + } const handler = (msg: { type: string; result?: ProcessResult; message?: string; percent?: number; label?: string }) => { if (msg.type === 'progress') { onProgress?.(msg.percent ?? 0, msg.label ?? '') } else if (msg.type === 'log') { onLog?.(msg.message ?? '') } else if (msg.type === 'done') { - worker.off('message', handler) + settle() resolve(msg.result ?? {}) } else if (msg.type === 'error') { - worker.off('message', handler) + settle() reject(new Error(msg.message)) } } + // A worker that dies mid-run (an uncaught error, out of memory, + // process.exit) never posts 'done' or 'error'. Settle the run instead of + // waiting forever, and drop the dead worker so the next run starts a + // fresh one rather than posting into it. + const onError = (err: Error) => { + settle() + this.discardWorker(worker) + reject(err) + } + const onExit = (code: number) => { + settle() + this.discardWorker(worker) + reject(new Error(`Process extension worker exited with code ${code}`)) + } worker.on('message', handler) + worker.on('error', onError) + worker.on('exit', onExit) worker.postMessage({ action: 'run', input, params }) }) } + private discardWorker(worker: Worker): void { + if (this.worker !== worker) return + this.worker = null + this.ready = false + } + terminate(): void { this.worker?.terminate() this.worker = null From f2f88bc39345066712a3bdea7b2faccb3a13ee2b Mon Sep 17 00:00:00 2001 From: DrHepa <162889656+DrHepa@users.noreply.github.com> Date: Sun, 13 Sep 2026 11:35:45 +0200 Subject: [PATCH 25/57] fix(mesh): validate registry params and map operation errors --- api/routers/optimize.py | 5 +- api/services/mesh_ops/registry.py | 34 ++++++++++++ api/tests/test_mesh_ops_registry.py | 71 ++++++++++++++++++++++++ api/tests/test_optimize_mesh_ops.py | 84 ++++++++++++++++++++++++++++- 4 files changed, 192 insertions(+), 2 deletions(-) diff --git a/api/routers/optimize.py b/api/routers/optimize.py index 44178a09..303def73 100644 --- a/api/routers/optimize.py +++ b/api/routers/optimize.py @@ -14,6 +14,7 @@ from services.generator_registry import WORKSPACE_DIR from services.mesh_ops import ( MeshOpContext, + MeshOpExecutionError, MeshOpNotFoundError, MeshOpResult, MeshOpUnavailableError, @@ -88,7 +89,9 @@ def _run_operation( return mesh_ops_registry.run(operation_id, input_path, params, context) except MeshOpNotFoundError as exc: raise HTTPException(404, f"Unknown mesh operation: {operation_id}") from exc - except MeshOpUnavailableError as exc: + except FileNotFoundError as exc: + raise HTTPException(404, str(exc)) from exc + except (MeshOpUnavailableError, MeshOpExecutionError) as exc: raise HTTPException(503, str(exc)) from exc except (TypeError, ValueError) as exc: raise HTTPException(400, str(exc)) from exc diff --git a/api/services/mesh_ops/registry.py b/api/services/mesh_ops/registry.py index 02f3f901..67928e5e 100644 --- a/api/services/mesh_ops/registry.py +++ b/api/services/mesh_ops/registry.py @@ -2,6 +2,7 @@ import re from copy import deepcopy +from math import isfinite from pathlib import Path from typing import Any, Iterable, Mapping, Optional @@ -11,6 +12,29 @@ _OP_ID = re.compile(r"^[a-z][a-z0-9_-]*$") +def _apply_numeric_bounds( + parameter_id: str, + value: Any, + schema: Mapping[str, Any], +) -> Any: + """Clamp a supplied numeric value to the bounds declared by its schema.""" + minimum = schema.get("min") + maximum = schema.get("max") + if minimum is None and maximum is None: + return value + + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise TypeError(f"Parameter {parameter_id!r} must be numeric") + if isinstance(value, float) and not isfinite(value): + raise ValueError(f"Parameter {parameter_id!r} must be finite") + + if minimum is not None: + value = max(minimum, value) + if maximum is not None: + value = min(maximum, value) + return value + + class MeshOpsRegistry: """Stores mesh operations and provides one invocation path for every caller.""" @@ -55,4 +79,14 @@ def run( if params: resolved_params.update(params) + for schema in operation.params_schema: + parameter_id = schema.get("id") + if parameter_id not in resolved_params: + continue + resolved_params[parameter_id] = _apply_numeric_bounds( + parameter_id, + resolved_params[parameter_id], + schema, + ) + return operation.fn(path, resolved_params, context) diff --git a/api/tests/test_mesh_ops_registry.py b/api/tests/test_mesh_ops_registry.py index 7274c9dd..31a777a2 100644 --- a/api/tests/test_mesh_ops_registry.py +++ b/api/tests/test_mesh_ops_registry.py @@ -62,6 +62,77 @@ def operation(input_path, params, context): 3, ) + def test_run_clamps_supplied_values_to_schema_bounds(self) -> None: + calls = [] + + def operation(input_path, params, context): + calls.append(params) + return MeshOpResult(input_path) + + registry = MeshOpsRegistry( + [ + MeshOp( + id="bounded", + label="Bounded", + params_schema=( + { + "id": "count", + "type": "int", + "default": 3, + "min": 1, + "max": 5, + }, + { + "id": "strength", + "type": "float", + "default": 0.5, + "min": 0.1, + "max": 1.0, + }, + ), + fn=operation, + category="test", + ) + ] + ) + + with tempfile.TemporaryDirectory() as directory: + input_path = Path(directory) / "mesh.glb" + input_path.touch() + context = MeshOpContext(Path(directory), Path(directory)) + registry.run( + "bounded", + input_path, + {"count": 99, "strength": -2.0}, + context, + ) + + self.assertEqual(calls[0], {"count": 5, "strength": 0.1}) + + def test_run_rejects_non_numeric_values_for_bounded_params(self) -> None: + operation = MeshOp( + id="bounded", + label="Bounded", + params_schema=( + {"id": "amount", "type": "int", "min": 1, "max": 5}, + ), + fn=lambda path, params, context: MeshOpResult(path), + category="test", + ) + registry = MeshOpsRegistry([operation]) + + with tempfile.TemporaryDirectory() as directory: + input_path = Path(directory) / "mesh.glb" + input_path.touch() + context = MeshOpContext(Path(directory), Path(directory)) + with self.assertRaisesRegex(TypeError, "'amount' must be numeric"): + registry.run( + "bounded", + input_path, + {"amount": "a lot"}, + context, + ) + def test_invalid_duplicate_and_unknown_ids_are_rejected(self) -> None: operation = MeshOp( id="valid", diff --git a/api/tests/test_optimize_mesh_ops.py b/api/tests/test_optimize_mesh_ops.py index 6e46dd9f..a1f0c5fe 100644 --- a/api/tests/test_optimize_mesh_ops.py +++ b/api/tests/test_optimize_mesh_ops.py @@ -6,7 +6,13 @@ from fastapi import HTTPException from routers import optimize -from services.mesh_ops import MeshOpNotFoundError, MeshOpResult +from services.mesh_ops import ( + MeshOp, + MeshOpExecutionError, + MeshOpNotFoundError, + MeshOpResult, + MeshOpsRegistry, +) class _FakeRegistry: @@ -120,6 +126,82 @@ def run(self, operation_id, input_path, params, context): self.assertEqual(raised.exception.status_code, 404) + def test_generic_run_route_enforces_registry_bounds(self) -> None: + calls = [] + + def operation(input_path, params, context): + calls.append(params) + return MeshOpResult(input_path) + + registry = MeshOpsRegistry( + [ + MeshOp( + id="bounded", + label="Bounded", + params_schema=( + { + "id": "amount", + "type": "int", + "default": 3, + "min": 1, + "max": 5, + }, + ), + fn=operation, + category="test", + ) + ] + ) + + with tempfile.TemporaryDirectory() as directory: + workspace = Path(directory) + input_path = workspace / "input.glb" + input_path.touch() + with ( + patch.object(optimize, "WORKSPACE_DIR", workspace), + patch.object(optimize, "mesh_ops_registry", registry), + ): + optimize.run_mesh_operation( + "bounded", + optimize.MeshOpRequest( + path="input.glb", + params={"amount": 99}, + ), + ) + + self.assertEqual(calls, [{"amount": 5}]) + + def test_missing_input_during_execution_is_a_404(self) -> None: + class MissingInputRegistry: + def run(self, operation_id, input_path, params, context): + raise FileNotFoundError(f"Input mesh not found: {input_path}") + + self._assert_operation_error(MissingInputRegistry(), 404) + + def test_backend_execution_failure_is_a_503(self) -> None: + class FailedBackendRegistry: + def run(self, operation_id, input_path, params, context): + raise MeshOpExecutionError("mesh-optimizer exited with code 1") + + self._assert_operation_error(FailedBackendRegistry(), 503) + + def _assert_operation_error(self, registry, expected_status: int) -> None: + with tempfile.TemporaryDirectory() as directory: + workspace = Path(directory) + input_path = workspace / "input.glb" + input_path.touch() + with ( + patch.object(optimize, "WORKSPACE_DIR", workspace), + patch.object(optimize, "mesh_ops_registry", registry), + self.assertRaises(HTTPException) as raised, + ): + optimize.run_mesh_operation( + "decimate", + optimize.MeshOpRequest(path="input.glb"), + ) + + self.assertEqual(raised.exception.status_code, expected_status) + if __name__ == "__main__": unittest.main() From b350595f5da99afacd50f0f73f084e98a89c3fb7 Mon Sep 17 00:00:00 2001 From: Lightning Pixel Date: Sun, 13 Sep 2026 13:29:08 +0200 Subject: [PATCH 26/57] fix(mesh-ops): validate boolean and select params in generic operation endpoint The generic POST /optimize/op/{op_name} endpoint only clamped numeric parameters against their schema bounds, leaving boolean and select-type params (e.g. smooth's mode) unvalidated. An invalid select value would silently fall through to a default branch instead of raising an error. _validate_param now checks booleans are actual bools and select values are one of the schema's declared options, delegating to the existing numeric bounds check otherwise. --- api/services/mesh_ops/registry.py | 26 +++++++++++++++++++++++++- 1 file changed, 25 insertions(+), 1 deletion(-) diff --git a/api/services/mesh_ops/registry.py b/api/services/mesh_ops/registry.py index 67928e5e..5d755340 100644 --- a/api/services/mesh_ops/registry.py +++ b/api/services/mesh_ops/registry.py @@ -35,6 +35,30 @@ def _apply_numeric_bounds( return value +def _validate_param( + parameter_id: str, + value: Any, + schema: Mapping[str, Any], +) -> Any: + """Validate (and where applicable, clamp) a supplied param against its schema.""" + param_type = schema.get("type") + + if param_type == "boolean": + if not isinstance(value, bool): + raise TypeError(f"Parameter {parameter_id!r} must be a boolean") + return value + + if param_type == "select": + allowed = {option["value"] for option in schema.get("options", [])} + if allowed and value not in allowed: + raise ValueError( + f"Parameter {parameter_id!r} must be one of {sorted(allowed)}" + ) + return value + + return _apply_numeric_bounds(parameter_id, value, schema) + + class MeshOpsRegistry: """Stores mesh operations and provides one invocation path for every caller.""" @@ -83,7 +107,7 @@ def run( parameter_id = schema.get("id") if parameter_id not in resolved_params: continue - resolved_params[parameter_id] = _apply_numeric_bounds( + resolved_params[parameter_id] = _validate_param( parameter_id, resolved_params[parameter_id], schema, From 82c460595d5a54d5f30f71f3e670a34f4632bcf5 Mon Sep 17 00:00:00 2001 From: Lightning Pixel Date: Sun, 13 Sep 2026 23:04:22 +0200 Subject: [PATCH 27/57] fix(point-lights): wire panel selection, fix gizmo carry-over, add selection outline, stop persisting lights - Selecting a point light from the light panel list now selects it in the 3D viewer too (was wired but unused). - Selecting a mesh/point light directly (without deselecting first) now clears the active gizmo tool, instead of only clearing it when both selections become empty. - Selected point light markers now get a violet outline that hugs the bulb icon shape, matching the mesh selection outline. - Point lights are no longer persisted across app restarts; the scene starts empty on launch. --- src/areas/generate/GeneratePage.tsx | 36 ++++++++++-- src/areas/generate/components/Viewer3D.tsx | 68 +++++++++++++--------- src/shared/stores/appStore.ts | 1 - 3 files changed, 71 insertions(+), 34 deletions(-) diff --git a/src/areas/generate/GeneratePage.tsx b/src/areas/generate/GeneratePage.tsx index 658893b8..04a059c3 100644 --- a/src/areas/generate/GeneratePage.tsx +++ b/src/areas/generate/GeneratePage.tsx @@ -193,13 +193,15 @@ function LightPopover({ onClose, pointLights, onPointLightsChange, - onSelectPointLight: _onSelectPointLight, + selectedPointLightId, + onSelectPointLight, }: { settings: LightSettings onChange: (s: LightSettings) => void onClose: () => void pointLights: PointLight[] onPointLightsChange: (lights: PointLight[]) => void + selectedPointLightId: string | null onSelectPointLight: (id: string | null) => void }) { function lightRow( @@ -292,7 +294,15 @@ function LightPopover({ )} {pointLights.map((pl) => ( -
+
onSelectPointLight(pl.id)} + className={`flex flex-col gap-1.5 p-2 rounded-lg bg-zinc-800/40 border cursor-pointer transition-colors ${ + pl.id === selectedPointLightId + ? 'border-violet-500' + : 'border-zinc-700/40 hover:border-zinc-600' + }`} + >
Point {pl.intensity.toFixed(1)}
@@ -1201,7 +1227,7 @@ export default function GeneratePage(): JSX.Element { gizmoUndoRef={gizmoUndoRef} pointLights={pointLights} selectedPointLightId={selectedPointLightId} - onSelectPointLight={setSelectedPointLightId} + onSelectPointLight={handleSelectPointLight} onPointLightsChange={setPointLights} /> diff --git a/src/areas/generate/components/Viewer3D.tsx b/src/areas/generate/components/Viewer3D.tsx index 234143bc..a2621cc6 100644 --- a/src/areas/generate/components/Viewer3D.tsx +++ b/src/areas/generate/components/Viewer3D.tsx @@ -64,7 +64,9 @@ function createCheckerTexture(): THREE.CanvasTexture { return tex } -function makeLightBulbTexture(color: string): THREE.CanvasTexture { +const SELECTION_OUTLINE_COLOR = '#8b5cf6' + +function makeLightBulbTexture(color: string, isSelected: boolean): THREE.CanvasTexture { const size = 64 const canvas = document.createElement('canvas') canvas.width = canvas.height = size @@ -73,35 +75,45 @@ function makeLightBulbTexture(color: string): THREE.CanvasTexture { const cx = size / 2 const cy = size / 2 - // Rays - ctx.strokeStyle = color - ctx.lineWidth = 2 - ctx.lineCap = 'round' - for (let i = 0; i < 6; i++) { - const a = (i / 6) * Math.PI * 2 - Math.PI / 2 - const r1 = 20 - const r2 = 27 + // Draws the bulb glyph (rays + circle + base). `pad` grows every part by a + // few pixels — used to lay down an oversized violet silhouette behind the + // normal-sized icon, so the outline hugs the actual glyph shape instead of + // being a plain circle around it. + const drawGlyph = (fillColor: string, pad: number) => { + ctx.strokeStyle = fillColor + ctx.lineWidth = 2 + pad * 2 + ctx.lineCap = 'round' + for (let i = 0; i < 6; i++) { + const a = (i / 6) * Math.PI * 2 - Math.PI / 2 + const r1 = 20 - pad + const r2 = 27 + pad + ctx.beginPath() + ctx.moveTo(cx + Math.cos(a) * r1, cy + Math.sin(a) * r1) + ctx.lineTo(cx + Math.cos(a) * r2, cy + Math.sin(a) * r2) + ctx.stroke() + } + ctx.beginPath() - ctx.moveTo(cx + Math.cos(a) * r1, cy + Math.sin(a) * r1) - ctx.lineTo(cx + Math.cos(a) * r2, cy + Math.sin(a) * r2) - ctx.stroke() - } + ctx.arc(cx, cy - 1, 12 + pad, 0, Math.PI * 2) + ctx.fillStyle = fillColor + ctx.fill() + if (pad === 0) { + ctx.strokeStyle = '#ffffff' + ctx.lineWidth = 1.5 + ctx.stroke() + } - // Bulb circle - ctx.beginPath() - ctx.arc(cx, cy - 1, 12, 0, Math.PI * 2) - ctx.fillStyle = color - ctx.fill() - ctx.strokeStyle = '#ffffff' - ctx.lineWidth = 1.5 - ctx.stroke() + ctx.fillStyle = fillColor + ctx.fillRect(cx - 4 - pad, cy + 11 - pad, 8 + pad * 2, 8 + pad * 2) + if (pad === 0) { + ctx.strokeStyle = '#ffffff' + ctx.lineWidth = 1 + ctx.strokeRect(cx - 4, cy + 11, 8, 8) + } + } - // Base (small rectangle below bulb) - ctx.fillStyle = color - ctx.fillRect(cx - 4, cy + 11, 8, 8) - ctx.strokeStyle = '#ffffff' - ctx.lineWidth = 1 - ctx.strokeRect(cx - 4, cy + 11, 8, 8) + if (isSelected) drawGlyph(SELECTION_OUTLINE_COLOR, 2.5) + drawGlyph(color, 0) return new THREE.CanvasTexture(canvas) } @@ -844,7 +856,7 @@ function PointLightMarker({ onPositionChange: (pos: [number, number, number]) => void }) { const [group, setGroup] = useState(null) - const iconTexture = useMemo(() => makeLightBulbTexture(light.color), [light.color]) + const iconTexture = useMemo(() => makeLightBulbTexture(light.color, isSelected), [light.color, isSelected]) return ( <> diff --git a/src/shared/stores/appStore.ts b/src/shared/stores/appStore.ts index a8cfc9c5..bb5ddd0e 100644 --- a/src/shared/stores/appStore.ts +++ b/src/shared/stores/appStore.ts @@ -315,7 +315,6 @@ export const useAppStore = create()( useAtkinsonFont: state.useAtkinsonFont, uiScale: state.uiScale, lightSettings: state.lightSettings, - pointLights: state.pointLights, }), } ) From aca9a79b319777ea3ab8b599526f5f5b9e2257d6 Mon Sep 17 00:00:00 2001 From: weng haishi <74546450+wenghaishi@users.noreply.github.com> Date: Mon, 14 Sep 2026 16:54:01 +0800 Subject: [PATCH 28/57] feat: add "Open in OrcaSlicer" export action MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit One-click hand-off from a generated model to OrcaSlicer via its orcaslicer://open?file= deeplink, in the Export dropdown. - Backend: GET /export/slicer/{fmt}/{token}/model.{fmt} converts the workspace GLB to STL on the fly — bakes scene-graph transforms, reorients Y-up->Z-up, normalizes print size. Path-only URL ending in the filename (no query string), since OrcaSlicer derives the import format from the URL's final segment; ancestry-based path containment. - Electron: slicer:open IPC opens the deeplink and reports failure so the UI can fall back when OrcaSlicer isn't installed. - Frontend: "Open in OrcaSlicer" item in the Export dropdown, shown for sliceable workspace meshes; pure deeplink builder with unit tests. Tests: api/tests/test_export_router.py and orcaSlicerLink.test.ts. Co-Authored-By: Claude Opus 4.8 --- api/routers/export.py | 107 ++++++++++++++++++ api/tests/test_export_router.py | 130 ++++++++++++++++++++++ electron/main/ipc-handlers.ts | 15 +++ electron/preload/electron-api.ts | 6 + package.json | 2 +- src/areas/generate/GeneratePage.tsx | 38 +++++++ src/areas/generate/orcaSlicerLink.test.ts | 51 +++++++++ src/areas/generate/orcaSlicerLink.ts | 44 ++++++++ src/shared/types/electron.d.ts | 3 + 9 files changed, 395 insertions(+), 1 deletion(-) create mode 100644 api/tests/test_export_router.py create mode 100644 src/areas/generate/orcaSlicerLink.test.ts create mode 100644 src/areas/generate/orcaSlicerLink.ts diff --git a/api/routers/export.py b/api/routers/export.py index 2a2f2bf3..f40a9045 100644 --- a/api/routers/export.py +++ b/api/routers/export.py @@ -1,4 +1,7 @@ +import base64 +import binascii import io +import math import trimesh from fastapi import APIRouter, HTTPException @@ -10,6 +13,110 @@ SUPPORTED = {"glb", "stl", "obj", "ply"} +# Formats OrcaSlicer's importer accepts (see the orcaslicer://open contract). +# GLB is deliberately excluded — OrcaSlicer cannot import glTF/GLB, so a .glb +# deeplink downloads but silently fails to slice. +SLICER_FORMATS = {"stl", "obj"} +SLICER_MEDIA_TYPES = {"stl": "model/stl", "obj": "text/plain"} + +# Image-to-3D output has no inherent physical scale (a single photo carries no +# real-world size), and AI generators emit roughly unit-sized meshes — which +# import into a slicer as an invisible ~1 mm speck. Normalise the longest +# bounding-box edge to a sane, obviously-printable default; the user rescales +# in OrcaSlicer as needed. +DEFAULT_PRINT_LONGEST_MM = 50.0 + + +def _to_single_mesh(loaded: object) -> "trimesh.Trimesh": + """Flatten a loaded GLB into one Trimesh, baking scene-graph node transforms. + + ``trimesh.util.concatenate(scene.geometry.values())`` would DROP the node + transforms and misassemble a multi-node scene, so flatten at the scene level + where the graph transforms are applied. + """ + if isinstance(loaded, trimesh.Trimesh): + return loaded + if isinstance(loaded, trimesh.Scene): + if len(loaded.geometry) == 0: + raise HTTPException(422, "Mesh contains no geometry") + # Bake the scene-graph node transforms into a single mesh. The spelling + # varies across trimesh versions — to_mesh()/to_geometry() are the modern + # APIs (4.6+); dump(concatenate=True) is the pre-removal fallback for 4.5. + for flatten in (lambda s: s.to_mesh(), lambda s: s.to_geometry(), lambda s: s.dump(concatenate=True)): + try: + result = flatten(loaded) + except (AttributeError, TypeError): + continue + if isinstance(result, trimesh.Trimesh): + return result + if isinstance(result, (list, tuple)) and result: + return trimesh.util.concatenate(result) + # Fallback: concatenate the geometry as-is (may ignore node transforms). + return trimesh.util.concatenate(list(loaded.geometry.values())) + raise HTTPException(422, "Unsupported mesh contents") + + +def _scale_to_print_size(mesh: "trimesh.Trimesh", longest_mm: float = DEFAULT_PRINT_LONGEST_MM) -> None: + """Uniformly scale ``mesh`` in place so its longest bbox edge is ``longest_mm``.""" + extents = mesh.extents + longest = float(max(extents)) if extents is not None and len(extents) else 0.0 + if longest > 1e-9 and math.isfinite(longest): + mesh.apply_scale(longest_mm / longest) + + +@router.get("/slicer/{fmt}/{token}/{filename}") +def export_for_slicer(fmt: str, token: str, filename: str): + """Serve a generated GLB converted to a slicer-importable mesh, at a URL + shaped for OrcaSlicer's ``orcaslicer://open?file=`` deeplink. + + The URL is intentionally path-only and ends in the real filename+extension + (e.g. ``/export/slicer/stl//model.stl``). OrcaSlicer + downloads the URL and derives the import filename — and therefore the mesh + format — from the URL's FINAL path segment, so a query string (``?path=...``) + would corrupt the parsed extension and the model would silently fail to + import. ``token`` is the url-safe-base64 of the workspace-relative source + path; ``filename`` (e.g. ``model.stl``) is what OrcaSlicer names the download. + """ + fmt = fmt.lower() + if fmt not in SLICER_FORMATS: + raise HTTPException(400, f"Unsupported slicer format: {fmt}. Supported: {', '.join(sorted(SLICER_FORMATS))}") + if not filename.lower().endswith(f".{fmt}"): + raise HTTPException(400, "Filename must end with the requested format extension") + + try: + padded = token + "=" * (-len(token) % 4) + rel_path = base64.urlsafe_b64decode(padded.encode("ascii")).decode("utf-8") + except (binascii.Error, UnicodeDecodeError, ValueError): + raise HTTPException(400, "Malformed source token") + + # Containment check via ancestry, not string prefix: `startswith` would let a + # sibling like `-other/...` slip through, and `..` escapes resolve + # outside the workspace and fail this check. + workspace = WORKSPACE_DIR.resolve() + full_path = (workspace / rel_path).resolve() + if full_path != workspace and workspace not in full_path.parents: + raise HTTPException(400, "Invalid path") + if not full_path.is_file(): + raise HTTPException(404, f"File not found: {rel_path}") + + mesh = _to_single_mesh(trimesh.load(str(full_path))) + # glTF/GLB is Y-up; OrcaSlicer's world is Z-up. Rotate +90° about X so the + # model imports standing upright instead of on its side. (Modly's own viewer + # rests generated meshes on the Y=0 plane, confirming Y is the up axis.) + mesh.apply_transform(trimesh.transformations.rotation_matrix(math.pi / 2, [1, 0, 0])) + _scale_to_print_size(mesh) + + data = mesh.export(file_type=fmt) + if isinstance(data, str): + data = data.encode("utf-8") + return Response( + content=data, + media_type=SLICER_MEDIA_TYPES.get(fmt, "application/octet-stream"), + # Fixed name (not the client-supplied segment) — keeps arbitrary input out + # of the response header. OrcaSlicer names the file from the URL anyway. + headers={"Content-Disposition": f'attachment; filename="model.{fmt}"'}, + ) + @router.get("/{fmt}") def export_mesh(fmt: str, path: str): diff --git a/api/tests/test_export_router.py b/api/tests/test_export_router.py new file mode 100644 index 00000000..bc07188f --- /dev/null +++ b/api/tests/test_export_router.py @@ -0,0 +1,130 @@ +import base64 +import io +import tempfile +import unittest +from pathlib import Path + +from fastapi import HTTPException + +# The export router imports trimesh at module load; skip the whole suite (rather +# than breaking `unittest discover`) in minimal environments without it. +try: + import numpy as np + import trimesh + + import routers.export as export_router + + HAVE_TRIMESH = True +except Exception: # noqa: BLE001 + HAVE_TRIMESH = False + + +def _token(rel_path: str) -> str: + return base64.urlsafe_b64encode(rel_path.encode("utf-8")).decode("ascii").rstrip("=") + + +def _load_stl(resp) -> "trimesh.Trimesh": + return trimesh.load(io.BytesIO(resp.body), file_type="stl") + + +@unittest.skipUnless(HAVE_TRIMESH, "trimesh not installed") +class ExportForSlicerTests(unittest.TestCase): + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self.workspace = Path(self._tmp.name).resolve() + self._orig_workspace = export_router.WORKSPACE_DIR + export_router.WORKSPACE_DIR = self.workspace + # A box that is tallest along Y (glTF up-axis). Exported to GLB, it + # reloads as a Scene so the flatten path is exercised too. + box = trimesh.creation.box(extents=[10.0, 30.0, 10.0]) + self.rel = "Workflows/hero.glb" + (self.workspace / "Workflows").mkdir(parents=True, exist_ok=True) + box.export(str(self.workspace / self.rel)) + + def tearDown(self) -> None: + export_router.WORKSPACE_DIR = self._orig_workspace + self._tmp.cleanup() + + def test_converts_glb_to_stl_with_download_filename(self) -> None: + resp = export_router.export_for_slicer("stl", _token(self.rel), "model.stl") + self.assertEqual(resp.media_type, "model/stl") + self.assertIn('filename="model.stl"', resp.headers["content-disposition"]) + mesh = _load_stl(resp) + self.assertGreater(len(mesh.faces), 0) + + def test_reorients_y_up_to_z_up(self) -> None: + # The box is tallest in Y; after the Y->Z rotation it must be tallest in + # Z so it imports standing upright on the slicer bed. + resp = export_router.export_for_slicer("stl", _token(self.rel), "model.stl") + ex = _load_stl(resp).extents + self.assertEqual(int(np.argmax(ex)), 2, f"expected Z to be the tallest axis, got extents {ex}") + + def test_normalizes_longest_edge_to_default_print_size(self) -> None: + resp = export_router.export_for_slicer("stl", _token(self.rel), "model.stl") + longest = float(max(_load_stl(resp).extents)) + self.assertAlmostEqual(longest, export_router.DEFAULT_PRINT_LONGEST_MM, places=3) + + def test_rejects_unsupported_format(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("glb", _token(self.rel), "model.glb") + self.assertEqual(ctx.exception.status_code, 400) + + def test_rejects_filename_extension_mismatch(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token(self.rel), "model.obj") + self.assertEqual(ctx.exception.status_code, 400) + + def test_rejects_malformed_token(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", "!!!not-base64!!!", "model.stl") + self.assertEqual(ctx.exception.status_code, 400) + + def test_rejects_path_traversal(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token("../escape.glb"), "model.stl") + self.assertEqual(ctx.exception.status_code, 400) + + def test_rejects_sibling_prefix_escape(self) -> None: + # A sibling dir whose name starts with the workspace dir name must not be + # reachable — the old str.startswith containment guard would allow it. + sibling = self.workspace.parent / (self.workspace.name + "-secret") + sibling.mkdir(parents=True, exist_ok=True) + (sibling / "x.glb").write_bytes(b"nope") + rel = f"../{self.workspace.name}-secret/x.glb" + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token(rel), "model.stl") + self.assertEqual(ctx.exception.status_code, 400) + + def test_missing_file_is_404(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token("Workflows/nope.glb"), "model.stl") + self.assertEqual(ctx.exception.status_code, 404) + + +@unittest.skipUnless(HAVE_TRIMESH, "trimesh not installed") +class FlattenAndScaleHelperTests(unittest.TestCase): + def test_flatten_bakes_scene_node_transforms(self) -> None: + # Two boxes placed at different positions via scene-graph transforms. + # util.concatenate(geometry.values()) would ignore the transforms; the + # scene-level flatten must reflect them in the combined bounds. + scene = trimesh.Scene() + scene.add_geometry(trimesh.creation.box(extents=[2, 2, 2]), transform=trimesh.transformations.translation_matrix([0, 0, 0])) + scene.add_geometry(trimesh.creation.box(extents=[2, 2, 2]), transform=trimesh.transformations.translation_matrix([100, 0, 0])) + mesh = export_router._to_single_mesh(scene) + self.assertIsInstance(mesh, trimesh.Trimesh) + # Combined X extent spans both boxes: ~101 (from -1 to 101). + self.assertGreater(mesh.extents[0], 100.0) + + def test_scale_to_print_size(self) -> None: + mesh = trimesh.creation.box(extents=[1.0, 2.0, 4.0]) + export_router._scale_to_print_size(mesh, longest_mm=80.0) + self.assertAlmostEqual(float(max(mesh.extents)), 80.0, places=3) + + def test_scale_ignores_degenerate_mesh(self) -> None: + # A single point cloud has zero extent; scaling must not divide by zero. + mesh = trimesh.Trimesh(vertices=[[0, 0, 0]], faces=[]) + export_router._scale_to_print_size(mesh) # must not raise + + +if __name__ == "__main__": + unittest.main() diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 005f1f78..f097a559 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -595,6 +595,21 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // Shell ipcMain.handle('shell:openExternal', (_, url: string) => shell.openExternal(url)) + // Open a model in OrcaSlicer via its orcaslicer://open?file= deeplink. + // Returns success/error so the renderer can surface a fallback (e.g. when + // OrcaSlicer is not installed and no app is registered for the scheme). + ipcMain.handle('slicer:open', async (_, url: string): Promise<{ success: boolean; error?: string }> => { + if (typeof url !== 'string' || !url.startsWith('orcaslicer://')) { + return { success: false, error: 'slicer:open requires an orcaslicer:// URL' } + } + try { + await shell.openExternal(url) + return { success: true } + } catch (err) { + return { success: false, error: err instanceof Error ? err.message : String(err) } + } + }) + // App info // System memory (used/available/total bytes). // On macOS, matches Activity Monitor's "Memory Used": diff --git a/electron/preload/electron-api.ts b/electron/preload/electron-api.ts index 89fa69ec..5d482f80 100644 --- a/electron/preload/electron-api.ts +++ b/electron/preload/electron-api.ts @@ -43,6 +43,12 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra // Shell utilities shell: { openExternal: (url: string) => ipcRenderer.invoke('shell:openExternal', url) }, + // Slicer integration — open a model in OrcaSlicer via its deeplink + slicer: { + open: (url: string): Promise<{ success: boolean; error?: string }> => + ipcRenderer.invoke('slicer:open', url) as Promise<{ success: boolean; error?: string }>, + }, + // System info system: { memory: (): Promise<{ total: number; used: number; available: number }> => diff --git a/package.json b/package.json index df135498..fdedfaad 100644 --- a/package.json +++ b/package.json @@ -13,7 +13,7 @@ "prepare-resources": "node scripts/download-python-embed.js", "test": "npm run test:py && npm run test:node", "test:py": "node scripts/run-pytests.mjs", - "test:node": "node --test --experimental-strip-types --experimental-loader ./scripts/node-ts-extensionless-loader.mjs src/shared/types/assetLibrary.test.ts src/areas/generate/assetLibraryProjection.test.ts src/areas/generate/assetLibraryService.test.ts src/areas/generate/assetLibraryUi.test.ts electron/main/artifact-registry-service.test.ts electron/main/extension-path-guard.test.ts electron/preload/artifact-registry-preload.test.ts && node --test electron/main/*.test.mjs src/**/*.test.mjs", + "test:node": "node --test --experimental-strip-types --experimental-loader ./scripts/node-ts-extensionless-loader.mjs src/shared/types/assetLibrary.test.ts src/areas/generate/assetLibraryProjection.test.ts src/areas/generate/assetLibraryService.test.ts src/areas/generate/assetLibraryUi.test.ts src/areas/generate/orcaSlicerLink.test.ts electron/main/artifact-registry-service.test.ts electron/main/extension-path-guard.test.ts electron/preload/artifact-registry-preload.test.ts && node --test electron/main/*.test.mjs src/**/*.test.mjs", "package": "cross-env CSC_IDENTITY_AUTO_DISCOVERY=false npm run build && npm run prepare-resources && electron-builder", "package:mac": "cross-env CSC_IDENTITY_AUTO_DISCOVERY=false npm run build && npm run prepare-resources && electron-builder --mac --arm64", "lint": "eslint ." diff --git a/src/areas/generate/GeneratePage.tsx b/src/areas/generate/GeneratePage.tsx index a6dbe4f3..f970993c 100644 --- a/src/areas/generate/GeneratePage.tsx +++ b/src/areas/generate/GeneratePage.tsx @@ -8,6 +8,7 @@ import GenerationHUD from './components/GenerationHUD' import Viewer3D from './components/Viewer3D' import WorkflowPanel from './components/WorkflowPanel' import { getDefaultAssetLibraryService } from './assetLibraryService' +import { buildOrcaSlicerDeepLink, canOpenInOrcaSlicer } from './orcaSlicerLink' import { resolveAssetLibraryOpenTarget, type ProjectedAssetLibraryEntry } from './assetLibraryProjection' import { ASSET_LIBRARY_SORT_OPTIONS, @@ -41,9 +42,13 @@ const EXPORT_FORMATS = [ function ExportDropdown({ onExport, onClose, + onOpenInSlicer, + canOpenInSlicer, }: { onExport: (f: 'glb' | 'obj' | 'stl' | 'ply') => void onClose: () => void + onOpenInSlicer: () => void + canOpenInSlicer: boolean }) { return (
@@ -57,6 +62,23 @@ function ExportDropdown({ {desc} ))} + {canOpenInSlicer && ( + <> +
+ + + )}
) } @@ -611,6 +633,7 @@ export default function GeneratePage(): JSX.Element { }, [undoMesh, redoMesh]) const hasModel = currentJob?.status === 'done' && !!currentJob.outputUrl + const showOpenInSlicer = hasModel && canOpenInOrcaSlicer(currentJob?.outputUrl) // Drop the active transform tool when the mesh is deselected, so it doesn't // silently re-activate on the next selection. @@ -660,6 +683,19 @@ export default function GeneratePage(): JSX.Element { link.click() } + async function handleOpenInOrcaSlicer() { + if (!currentJob?.outputUrl) return + try { + const link = buildOrcaSlicerDeepLink(apiUrl, currentJob.outputUrl) + const result = await window.electron.slicer.open(link) + if (!result.success) { + showError(result.error ?? 'Could not open OrcaSlicer. Make sure it is installed.') + } + } catch (err) { + showError(err instanceof Error ? err.message : 'Could not open OrcaSlicer.') + } + } + function getOptimizePath(url: string): string { if (url.startsWith('/workspace/')) { return url.slice('/workspace/'.length) @@ -971,6 +1007,8 @@ export default function GeneratePage(): JSX.Element { void} onClose={() => setOpenPanel(null)} + onOpenInSlicer={() => { void handleOpenInOrcaSlicer() }} + canOpenInSlicer={showOpenInSlicer} /> )}
diff --git a/src/areas/generate/orcaSlicerLink.test.ts b/src/areas/generate/orcaSlicerLink.test.ts new file mode 100644 index 00000000..2f7e47d3 --- /dev/null +++ b/src/areas/generate/orcaSlicerLink.test.ts @@ -0,0 +1,51 @@ +import assert from 'node:assert/strict' +import test from 'node:test' + +import { + SLICER_FORMAT, + buildOrcaSlicerDeepLink, + canOpenInOrcaSlicer, + encodeWorkspacePathToken, +} from './orcaSlicerLink.ts' + +test('builds an orcaslicer://open deeplink whose file= is a percent-encoded, query-less URL ending in model.stl', () => { + const link = buildOrcaSlicerDeepLink('http://localhost:8765', '/workspace/Workflows/checkpoints/hero.glb') + assert.ok(link.startsWith('orcaslicer://open?file=')) + const modelUrl = decodeURIComponent(link.slice('orcaslicer://open?file='.length)) + // OrcaSlicer derives the import format from the URL's final path segment, so + // it must end in the real extension and carry no query string. + assert.ok(!modelUrl.includes('?'), 'model URL must not contain a query string') + assert.ok(modelUrl.endsWith('/model.stl'), 'model URL must end in model.stl') + assert.equal( + modelUrl, + `http://localhost:8765/export/slicer/stl/${encodeWorkspacePathToken('Workflows/checkpoints/hero.glb')}/model.stl`, + ) +}) + +test('token round-trips a workspace path through url-safe base64 (matches the API decode)', () => { + const path = 'Workflows/checkpoints/hero model (v2).glb' + const token = encodeWorkspacePathToken(path) + assert.ok(!/[+/=]/.test(token), 'token must be url-safe with no padding') + // Decode the way the Python API does: restore padding, then urlsafe-decode. + const padded = token + '='.repeat((4 - (token.length % 4)) % 4) + const decoded = Buffer.from(padded.replace(/-/g, '+').replace(/_/g, '/'), 'base64').toString('utf-8') + assert.equal(decoded, path) +}) + +test('strips a trailing slash from the api origin', () => { + const link = buildOrcaSlicerDeepLink('http://localhost:8765/', '/workspace/a.glb') + const modelUrl = decodeURIComponent(link.slice('orcaslicer://open?file='.length)) + assert.equal(modelUrl, `http://localhost:8765/export/slicer/stl/${encodeWorkspacePathToken('a.glb')}/model.stl`) +}) + +test('canOpenInOrcaSlicer accepts workspace meshes and rejects splats, imports, and empty', () => { + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/hero.glb'), true) + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/scan.ply'), false) + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/scan.splat'), false) + assert.equal(canOpenInOrcaSlicer('/optimize/serve-file?path=/tmp/x.glb'), false) + assert.equal(canOpenInOrcaSlicer(undefined), false) +}) + +test('SLICER_FORMAT is a format OrcaSlicer can import', () => { + assert.equal(SLICER_FORMAT, 'stl') +}) diff --git a/src/areas/generate/orcaSlicerLink.ts b/src/areas/generate/orcaSlicerLink.ts new file mode 100644 index 00000000..2bade12e --- /dev/null +++ b/src/areas/generate/orcaSlicerLink.ts @@ -0,0 +1,44 @@ +// Builds the OrcaSlicer deeplink for a generated mesh. +// +// OrcaSlicer registers the `orcaslicer://open?file=` scheme; its handler +// downloads the http(s) URL in `file=` and imports it, deriving the filename — +// and therefore the mesh format — from the URL's FINAL path segment. That means +// the served URL must be path-only and end in a real `model.` with NO +// query string, and the whole thing must be percent-encoded. OrcaSlicer cannot +// import GLB, so we point at the backend's slicer-export route which converts to +// STL on the fly. + +/** Format handed to OrcaSlicer. STL is universal and OrcaSlicer auto-repairs it. */ +export const SLICER_FORMAT = 'stl' + +/** URL-safe base64 (no padding) of a UTF-8 string — matches the API's token decode. */ +export function encodeWorkspacePathToken(workspacePath: string): string { + const bytes = new TextEncoder().encode(workspacePath) + let binary = '' + for (const b of bytes) binary += String.fromCharCode(b) + return btoa(binary).replace(/\+/g, '-').replace(/\//g, '_').replace(/=+$/, '') +} + +/** + * Whether a generation output can be opened in OrcaSlicer: it must be a mesh + * served from the workspace (Gaussian splats and non-workspace imports are not + * sliceable through this route). + */ +export function canOpenInOrcaSlicer(outputUrl: string | undefined): boolean { + if (!outputUrl) return false + return outputUrl.startsWith('/workspace/') && !/\.(ply|splat)$/i.test(outputUrl) +} + +/** + * Build the `orcaslicer://open?file=...` deeplink for a generated mesh. + * + * @param apiUrl Modly backend origin, e.g. `http://localhost:8765` + * @param outputUrl workspace URL of the mesh, e.g. `/workspace/Foo/hero.glb` + */ +export function buildOrcaSlicerDeepLink(apiUrl: string, outputUrl: string): string { + const workspacePath = outputUrl.replace(/^\/workspace\//, '') + const token = encodeWorkspacePathToken(workspacePath) + const base = apiUrl.replace(/\/+$/, '') + const modelUrl = `${base}/export/slicer/${SLICER_FORMAT}/${token}/model.${SLICER_FORMAT}` + return `orcaslicer://open?file=${encodeURIComponent(modelUrl)}` +} diff --git a/src/shared/types/electron.d.ts b/src/shared/types/electron.d.ts index 1a4d6fde..b674ec00 100644 --- a/src/shared/types/electron.d.ts +++ b/src/shared/types/electron.d.ts @@ -154,6 +154,9 @@ declare global { shell: { openExternal: (url: string) => Promise } + slicer: { + open: (url: string) => Promise<{ success: boolean; error?: string }> + } system: { memory: () => Promise<{ total: number; used: number; available: number }> } From 63b731c505a4d9ded15a75fe4ddb0a26fa093847 Mon Sep 17 00:00:00 2001 From: DrHepa <162889656+DrHepa@users.noreply.github.com> Date: Mon, 14 Sep 2026 12:17:51 +0200 Subject: [PATCH 29/57] feat(models): support extension-scoped shared weight groups --- README.md | 55 ++++ api/routers/model.py | 12 +- api/runner.py | 26 +- api/services/extension_process.py | 13 +- api/services/generator_registry.py | 78 +++++- api/services/generators/base.py | 3 + api/services/model_sources.py | 144 +++++++++- api/tests/test_extension_process.py | 18 ++ api/tests/test_generator_registry.py | 63 +++++ api/tests/test_model_router.py | 23 ++ api/tests/test_model_sources.py | 57 ++++ api/tests/test_runner.py | 11 + .../main/extension-install-utils.test.mjs | 57 ++++ electron/main/extension-install-utils.ts | 36 ++- electron/main/ipc-handlers.ts | 261 ++++++++++++++++-- electron/main/model-download-plan.test.mjs | 91 ++++++ electron/main/model-download-plan.ts | 129 ++++++++- electron/main/model-download-preload.test.mjs | 12 +- electron/main/model-sources.test.mjs | 60 ++++ electron/main/model-sources.ts | 153 +++++++++- electron/preload/electron-api.ts | 3 + src/areas/models/ModelsPage.tsx | 69 ++++- .../models/components/ExtensionDrawer.tsx | 58 +++- .../models/components/extensionShared.tsx | 2 +- src/shared/types/electron.d.ts | 18 ++ 25 files changed, 1366 insertions(+), 86 deletions(-) diff --git a/README.md b/README.md index b162cf23..d7bca2e6 100644 --- a/README.md +++ b/README.md @@ -148,6 +148,61 @@ supported provider is `huggingface`. Existing nodes that use `hf_repo`, `download_check`, `hf_include_prefixes`, and `hf_skip_prefixes` keep their original behavior. +### Shared weights inside one model extension + +Multi-node model extensions can declare extension-scoped `weight_groups` and +reference them from any sibling node. Shared files are downloaded once under +`//_shared/`, while node-specific +`model_sources` stay under the node's existing model directory. + +```json +{ + "id": "pixal3d", + "type": "model", + "weight_groups": [ + { + "id": "pixal3d-base", + "model_sources": [ + { + "id": "base", + "provider": "huggingface", + "repo_id": "TencentARC/Pixal3D", + "revision": "", + "destination": ".", + "checks": ["pipeline.json"] + } + ] + } + ], + "nodes": [ + { + "id": "generate", + "weight_groups": ["pixal3d-base"] + }, + { + "id": "worldsculpt", + "weight_groups": ["pixal3d-base"], + "model_sources": [ + { + "id": "adapter", + "provider": "huggingface", + "repo_id": "AlayaLab/WorldSculpt", + "revision": "", + "destination": ".", + "checks": ["model.safetensors"] + } + ] + } + ] +} +``` + +At runtime, `MODEL_DIR` remains the selected node's private directory. +Subprocess extensions also receive `MODEL_ID`, `MODEL_NODE_ID`, and a JSON +`SHARED_MODEL_DIRS` map. Direct generators receive the same resolved mapping in +`shared_model_dirs`. Removing private node data never removes a shared group; +shared-group removal is a separate action that identifies every affected node. + --- ## Workflows diff --git a/api/routers/model.py b/api/routers/model.py index 0b40d155..3fd937c0 100644 --- a/api/routers/model.py +++ b/api/routers/model.py @@ -14,8 +14,8 @@ from services.model_sources import ( normalize_model_sources, resolve_download_path, - resolve_model_root, - resolve_source_destination, + resolve_source_destination_at_root, + resolve_weight_storage_root, validate_source_file_plan, ) @@ -131,7 +131,7 @@ async def cancel_hf_download(model_id: str): @router.post("/hf-download-sources") async def hf_download_sources(request: FastAPIRequest, model_id: str): - """Download all Hugging Face sources declared for one model node.""" + """Download sources into one validated node or extension-shared target.""" try: body = await request.json() if not isinstance(body, dict): @@ -140,10 +140,10 @@ async def hf_download_sources(request: FastAPIRequest, model_id: str): if raw_sources is None: raise ValueError("sources are required") sources = normalize_model_sources({"model_sources": raw_sources}) - model_root = resolve_model_root(MODELS_DIR, model_id) + model_root = resolve_weight_storage_root(MODELS_DIR, model_id) destinations = { - source["id"]: resolve_source_destination( - MODELS_DIR, model_id, source["destination"] + source["id"]: resolve_source_destination_at_root( + model_root, source["destination"] ) for source in sources } diff --git a/api/runner.py b/api/runner.py index 0dd21392..489d12a0 100644 --- a/api/runner.py +++ b/api/runner.py @@ -34,6 +34,15 @@ # MODEL_DIR is set by ExtensionProcess to match its own model_dir (composite node id path). # Falls back to MODELS_DIR/manifest_id for standalone/legacy use. _MODEL_DIR_OVERRIDE = os.environ.get("MODEL_DIR", "") +_MODEL_ID_OVERRIDE = os.environ.get("MODEL_ID", "") +_MODEL_NODE_ID_OVERRIDE = os.environ.get("MODEL_NODE_ID", "") +try: + _SHARED_MODEL_DIRS = { + str(group_id): Path(path) + for group_id, path in json.loads(os.environ.get("SHARED_MODEL_DIRS", "{}")).items() + } +except (AttributeError, TypeError, ValueError, json.JSONDecodeError): + _SHARED_MODEL_DIRS = {} # Inject Modly's api/ so generator.py can do: # from services.generators.base import BaseGenerator, ... @@ -87,8 +96,12 @@ def load_generator(manifest: dict): return getattr(mod, manifest["generator_class"]) -def _select_node(manifest: dict, model_dir_override: str) -> dict: +def _select_node( + manifest: dict, model_dir_override: str, node_id_override: str = "" +) -> dict: nodes = manifest.get("nodes") or [] + if nodes and node_id_override: + return next((n for n in nodes if n.get("id") == node_id_override), nodes[0]) if nodes and model_dir_override: node_id = Path(model_dir_override).name return next((n for n in nodes if n.get("id") == node_id), nodes[0]) @@ -157,7 +170,7 @@ def _apply_manifest_metadata(gen, manifest: dict, node: dict) -> None: def main() -> None: manifest = json.loads((EXT_DIR / "manifest.json").read_text(encoding="utf-8")) - model_id = manifest["id"] + model_id = _MODEL_ID_OVERRIDE or manifest["id"] try: GenClass = load_generator(manifest) @@ -167,11 +180,9 @@ def main() -> None: "traceback": traceback.format_exc()}) return - # Support both flat manifest (legacy) and nodes[] format. - # Use MODEL_DIR to find the correct node for multi-node extensions: - # MODEL_DIR is set by ExtensionProcess to MODELS_DIR/ext_id/node_id, - # so its last component matches the node id. - node = _select_node(manifest, _MODEL_DIR_OVERRIDE) + # Support both flat manifest (legacy) and nodes[] format. The host passes an + # explicit node id; MODEL_DIR name inference remains only as a legacy fallback. + node = _select_node(manifest, _MODEL_DIR_OVERRIDE, _MODEL_NODE_ID_OVERRIDE) # Announce readiness and send params_schema so ExtensionProcess # can serve it without needing to query the subprocess later. @@ -184,6 +195,7 @@ def main() -> None: # Falls back to MODELS_DIR/manifest_id for legacy / standalone use. model_dir = Path(_MODEL_DIR_OVERRIDE) if _MODEL_DIR_OVERRIDE else MODELS_DIR / model_id gen = GenClass(model_dir, WORKSPACE_DIR) + gen.shared_model_dirs = dict(_SHARED_MODEL_DIRS) _apply_manifest_metadata(gen, manifest, node) # Active cancel events keyed by request id diff --git a/api/services/extension_process.py b/api/services/extension_process.py index 67565d36..ab9431d6 100644 --- a/api/services/extension_process.py +++ b/api/services/extension_process.py @@ -45,6 +45,7 @@ def __init__(self, ext_dir: Path, manifest: dict) -> None: self.manifest = manifest self.model_dir = None # set by registry after init self.outputs_dir = None # set by registry after init + self.shared_model_dirs: dict[str, Path] = {} self._proc: Optional[subprocess.Popen] = None self._queue: queue.Queue = queue.Queue() @@ -84,11 +85,17 @@ def _build_env(self) -> dict: # Setting it inside generator.py is too late, since generator.py # itself imports torch before calling select_device(). env.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1") - # Pass the exact model_dir so runner.py doesn't have to re-derive it - # from manifest["id"] (which is the ext_id, not the composite node id). - # runner.py extracts the node id from MODEL_DIR's trailing path component. + # Keep capability identity separate from storage identity. MODEL_DIR + # retains its node-private meaning; shared roots are passed explicitly. if self.model_dir is not None: env["MODEL_DIR"] = str(self.model_dir) + env["MODEL_ID"] = self.MODEL_ID + env["MODEL_NODE_ID"] = self.manifest.get( + "node_id", self.MODEL_ID.split("/", 1)[-1] + ) + env["SHARED_MODEL_DIRS"] = json.dumps( + {group_id: str(path) for group_id, path in self.shared_model_dirs.items()} + ) # Extension venvs are based on python-embed which ships without a CA bundle. # Only set SSL_CERT_FILE if not already provided (preserves corporate/custom certs). if "SSL_CERT_FILE" not in env: diff --git a/api/services/generator_registry.py b/api/services/generator_registry.py index 348a42cb..4f642e97 100644 --- a/api/services/generator_registry.py +++ b/api/services/generator_registry.py @@ -25,7 +25,15 @@ from services.generators.base import BaseGenerator from services.extension_process import ExtensionProcess, _venv_python -from services.model_sources import model_sources_are_downloaded, normalize_model_sources +from services.model_sources import ( + model_sources_are_downloaded, + normalize_model_sources, + normalize_weight_group_references, + normalize_weight_groups, + resolve_weight_group_root, + safe_source_id, + weight_group_sources_are_downloaded, +) # ------------------------------------------------------------------ # # Global paths @@ -432,6 +440,8 @@ def _discover_extensions( if "model_sources" in manifest: raise ValueError("model_sources must be declared on a model node") + weight_groups = normalize_weight_groups(manifest) + if ext_id != ext_dir.name: message = ( f"Extension folder '{ext_dir.name}' declares mismatched " @@ -452,6 +462,30 @@ def _discover_extensions( node for node in raw_nodes if isinstance(node, dict) and node.get("id") ] + group_by_id = {group["id"]: group for group in weight_groups or []} + uses_shared_weights = weight_groups is not None or any( + "weight_groups" in node for node in nodes + ) + if uses_shared_weights: + for node in nodes: + raw_node_id = node.get("id") + if ( + weight_groups is not None + and isinstance(raw_node_id, str) + and raw_node_id.casefold() == "_shared" + ): + raise ValueError('model node id "_shared" is reserved') + node_id = safe_source_id(raw_node_id, "model node id") + normalize_weight_group_references( + node, + weight_groups, + field_name=f"nodes[{node_id}].weight_groups", + ) + if "weight_groups" in node and "hf_repo" in node: + raise ValueError( + f'model node "{node_id}" must use model_sources for private ' + "weights when weight_groups are declared" + ) # Markers left while setup or runtime registration is unfinished: # the folder is not ready to be loaded. The readable manifest lets @@ -528,6 +562,11 @@ def _discover_extensions( if nodes: for node in nodes: model_sources = normalize_model_sources(node) + group_ids = normalize_weight_group_references( + node, + weight_groups, + field_name=f"nodes[{node['id']}].weight_groups", + ) or [] node_manifest = { **manifest, "id": f"{ext_id}/{node['id']}", @@ -541,6 +580,7 @@ def _discover_extensions( "params_schema": node.get("params_schema", manifest.get("params_schema", [])), "input": node.get("input", "image"), "output": node.get("output", "mesh"), + "weight_groups": [group_by_id[group_id] for group_id in group_ids], } if model_sources is not None: node_manifest["model_sources"] = model_sources @@ -630,6 +670,13 @@ def initialize( gen.download_check = manifest.get("download_check", "") gen._params_schema = manifest.get("params_schema", []) + gen.shared_model_dirs = { + group["id"]: resolve_weight_group_root( + MODELS_DIR, manifest.get("ext_id", model_id.split("/", 1)[0]), group["id"] + ) + for group in manifest.get("weight_groups", []) + } + self._generators[model_id] = gen self._manifests[model_id] = manifest self._errors.pop(model_id, None) @@ -707,9 +754,12 @@ def get_active(self) -> BaseGenerator: self._assert_not_quarantined(self._active_id) gen = self._generators[self._active_id] downloaded = self._is_downloaded(self._active_id, gen) - if "model_sources" in self._manifests[self._active_id] and not downloaded: + if ( + "model_sources" in self._manifests[self._active_id] + or self._manifests[self._active_id].get("weight_groups") + ) and not downloaded: raise RuntimeError( - "Model sources are incomplete. Download this node's weights " + "Model sources are incomplete. Download this node's shared and private weights " "from the Modly Models page before generation." ) if not gen.is_loaded(): @@ -758,10 +808,21 @@ def switch_model(self, model_id: str) -> None: def _is_downloaded(self, model_id: str, gen: BaseGenerator) -> bool: manifest = self._manifests[model_id] + private_ready = True if "model_sources" in manifest: - return model_sources_are_downloaded( + private_ready = model_sources_are_downloaded( MODELS_DIR, model_id, manifest["model_sources"] ) + shared_ready = all( + weight_group_sources_are_downloaded( + MODELS_DIR, + manifest.get("ext_id", model_id.split("/", 1)[0]), + group, + ) + for group in manifest.get("weight_groups", []) + ) + if "model_sources" in manifest or manifest.get("weight_groups"): + return private_ready and shared_ready return gen.is_downloaded() def active_status(self) -> dict: @@ -812,6 +873,15 @@ def update_paths(self, models_dir: Optional[Path], workspace_dir: Optional[Path] _self_module.MODELS_DIR = models_dir for model_id, gen in self._generators.items(): gen.model_dir = models_dir / model_id + manifest = self._manifests[model_id] + gen.shared_model_dirs = { + group["id"]: resolve_weight_group_root( + models_dir, + manifest.get("ext_id", model_id.split("/", 1)[0]), + group["id"], + ) + for group in manifest.get("weight_groups", []) + } if workspace_dir is not None: workspace_dir.mkdir(parents=True, exist_ok=True) diff --git a/api/services/generators/base.py b/api/services/generators/base.py index fd62ceef..c9344538 100644 --- a/api/services/generators/base.py +++ b/api/services/generators/base.py @@ -90,6 +90,9 @@ def __init__(self, model_dir: Path, outputs_dir: Path) -> None: self.hf_skip_prefixes: list = [] self.download_check: str = "" # relative path to check in model_dir self._params_schema: list = [] # params declared in the manifest + # Host-resolved extension-scoped shared weight roots, keyed by group id. + # Model identity and the private model_dir remain unchanged. + self.shared_model_dirs: dict[str, Path] = {} # ------------------------------------------------------------------ # # Model lifecycle diff --git a/api/services/model_sources.py b/api/services/model_sources.py index 592d1342..81c6c41a 100644 --- a/api/services/model_sources.py +++ b/api/services/model_sources.py @@ -90,18 +90,20 @@ def _safe_revision(value: Any, field: str) -> str | None: return value -def normalize_model_sources(node: dict[str, Any]) -> list[dict[str, Any]] | None: +def normalize_model_sources( + node: dict[str, Any], *, field_name: str = "model_sources" +) -> list[dict[str, Any]] | None: """Validate only the new contract; legacy fields remain untouched.""" if "model_sources" not in node: return None raw_sources = node["model_sources"] if not isinstance(raw_sources, list) or not raw_sources: - raise ValueError("model_sources must be a non-empty array") + raise ValueError(f"{field_name} must be a non-empty array") aliases: dict[str, str] = {} sources: list[dict[str, Any]] = [] for index, raw in enumerate(raw_sources): - field = f"model_sources[{index}]" + field = f"{field_name}[{index}]" if not isinstance(raw, dict): raise ValueError(f"{field} must be an object") source_id = safe_source_id(raw.get("id"), f"{field}.id") @@ -156,6 +158,74 @@ def normalize_model_sources(node: dict[str, Any]) -> list[dict[str, Any]] | None return sources +def normalize_weight_groups(manifest: dict[str, Any]) -> list[dict[str, Any]] | None: + if "weight_groups" not in manifest: + return None + raw_groups = manifest["weight_groups"] + if not isinstance(raw_groups, list) or not raw_groups: + raise ValueError("weight_groups must be a non-empty array") + + aliases: dict[str, str] = {} + groups: list[dict[str, Any]] = [] + for index, raw in enumerate(raw_groups): + field = f"weight_groups[{index}]" + if not isinstance(raw, dict): + raise ValueError(f"{field} must be an object") + raw_group_id = raw.get("id") + if isinstance(raw_group_id, str) and raw_group_id.casefold() == "_shared": + raise ValueError(f'{field}.id uses the reserved identifier "_shared"') + group_id = safe_source_id(raw_group_id, f"{field}.id") + alias = unicodedata.normalize("NFC", group_id).casefold() + if alias in aliases: + raise ValueError( + f'weight group ids "{aliases[alias]}" and "{group_id}" ' + "are not portable-unique" + ) + aliases[alias] = group_id + sources = normalize_model_sources( + {"model_sources": raw.get("model_sources")}, + field_name=f"{field}.model_sources", + ) + groups.append({"id": group_id, "model_sources": sources}) + return groups + + +def normalize_weight_group_references( + node: dict[str, Any], + groups: list[dict[str, Any]] | None, + *, + field_name: str = "weight_groups", +) -> list[str] | None: + if "weight_groups" not in node: + return None + raw_refs = node["weight_groups"] + if not isinstance(raw_refs, list) or not raw_refs: + raise ValueError(f"{field_name} must be a non-empty array of weight group ids") + + available = { + unicodedata.normalize("NFC", group["id"]).casefold(): group["id"] + for group in groups or [] + } + aliases: dict[str, str] = {} + refs: list[str] = [] + for index, raw in enumerate(raw_refs): + group_id = safe_source_id(raw, f"{field_name}[{index}]") + alias = unicodedata.normalize("NFC", group_id).casefold() + if alias in aliases: + raise ValueError( + f'weight group references "{aliases[alias]}" and "{group_id}" ' + "are not portable-unique" + ) + aliases[alias] = group_id + canonical = available.get(alias) + if canonical is None: + raise ValueError( + f'{field_name}[{index}] references unknown weight group "{group_id}"' + ) + refs.append(canonical) + return refs + + def _path_has_symlink(root: Path, candidate: Path) -> bool: root = root.absolute() candidate = candidate.absolute() @@ -180,6 +250,8 @@ def resolve_model_root(models_dir: Path, model_id: str) -> Path: if len(parts) != 2: raise ValueError("Model id must identify one extension node") extension_id = safe_source_id(parts[0], "extension id") + if parts[1].casefold() == "_shared": + raise ValueError('Model node id "_shared" is reserved') node_id = safe_source_id(parts[1], "model node id") root = models_dir.absolute() candidate = root / extension_id / node_id @@ -192,6 +264,35 @@ def resolve_model_root(models_dir: Path, model_id: str) -> Path: return candidate +def resolve_weight_group_root(models_dir: Path, extension_id: str, group_id: str) -> Path: + safe_extension_id = safe_source_id(extension_id, "extension id") + if isinstance(group_id, str) and group_id.casefold() == "_shared": + raise ValueError('Weight group id "_shared" is reserved') + safe_group_id = safe_source_id(group_id, "weight group id") + root = models_dir.absolute() + candidate = root / safe_extension_id / "_shared" / safe_group_id + if _path_has_symlink(root, candidate): + raise ValueError("Weight group path resolves through a symlink") + try: + candidate.resolve().relative_to(root.resolve()) + except ValueError as exc: + raise ValueError("Weight group path escapes the models directory") from exc + return candidate + + +def resolve_weight_storage_root(models_dir: Path, target_id: str) -> Path: + if not isinstance(target_id, str): + raise ValueError("Weight target id must be a string") + parts = target_id.split("/") + if len(parts) == 2: + return resolve_model_root(models_dir, target_id) + if len(parts) == 3 and parts[1] == "_shared": + return resolve_weight_group_root(models_dir, parts[0], parts[2]) + raise ValueError( + "Weight target id must identify one model node or extension weight group" + ) + + def resolve_source_destination(models_dir: Path, model_id: str, destination: str) -> Path: model_root = resolve_model_root(models_dir, model_id) safe_destination = safe_relative_path(destination, "destination", allow_dot=True) @@ -201,6 +302,18 @@ def resolve_source_destination(models_dir: Path, model_id: str, destination: str return candidate +def resolve_source_destination_at_root(model_root: Path, destination: str) -> Path: + safe_destination = safe_relative_path(destination, "destination", allow_dot=True) + candidate = ( + model_root + if safe_destination == "." + else model_root.joinpath(*safe_destination.split("/")) + ) + if _path_has_symlink(model_root, candidate): + raise ValueError("Source destination resolves through a symlink") + return candidate + + def resolve_download_path(destination: Path, filename: str) -> Path: safe_filename = safe_relative_path(filename, "Hugging Face repository file") candidate = destination.joinpath(*safe_filename.split("/")) @@ -214,11 +327,20 @@ def model_sources_are_downloaded( ) -> bool: try: model_root = resolve_model_root(models_dir, model_id) + return model_sources_are_downloaded_at_root(model_root, sources) + except (KeyError, OSError, TypeError, ValueError): + return False + + +def model_sources_are_downloaded_at_root( + model_root: Path, sources: list[dict[str, Any]] +) -> bool: + try: if not model_root.is_dir(): return False for source in sources: - destination = resolve_source_destination( - models_dir, model_id, source["destination"] + destination = resolve_source_destination_at_root( + model_root, source["destination"] ) if not destination.is_dir(): return False @@ -235,6 +357,18 @@ def model_sources_are_downloaded( return False +def weight_group_sources_are_downloaded( + models_dir: Path, extension_id: str, group: dict[str, Any] +) -> bool: + try: + return model_sources_are_downloaded_at_root( + resolve_weight_group_root(models_dir, extension_id, group["id"]), + group["model_sources"], + ) + except (KeyError, OSError, TypeError, ValueError): + return False + + def validate_source_file_plan( sources: list[dict[str, Any]], files_by_source: dict[str, list[str]] ) -> None: diff --git a/api/tests/test_extension_process.py b/api/tests/test_extension_process.py index 348e293f..e4791f43 100644 --- a/api/tests/test_extension_process.py +++ b/api/tests/test_extension_process.py @@ -1,4 +1,5 @@ import io +import json import platform import queue import unittest @@ -102,6 +103,23 @@ def test_sets_worker_model_dir_when_known(self) -> None: env = proc._build_env() self.assertEqual(env.get("MODEL_DIR"), str(Path("/tmp/models/ext/node"))) + def test_sets_explicit_node_identity_and_shared_weight_dirs(self) -> None: + proc = ExtensionProcess( + ext_dir=Path("/tmp/extensions/ext"), + manifest={"id": "ext/quality", "node_id": "quality"}, + ) + proc.model_dir = Path("/tmp/models/ext/quality") + proc.shared_model_dirs = {"base": Path("/tmp/models/ext/_shared/base")} + + env = proc._build_env() + + self.assertEqual(env["MODEL_ID"], "ext/quality") + self.assertEqual(env["MODEL_NODE_ID"], "quality") + self.assertEqual( + json.loads(env["SHARED_MODEL_DIRS"]), + {"base": "/tmp/models/ext/_shared/base"}, + ) + class MissingModuleExtractionTests(unittest.TestCase): def test_extracts_module_name_from_message(self) -> None: diff --git a/api/tests/test_generator_registry.py b/api/tests/test_generator_registry.py index ff9d090c..71ffb72a 100644 --- a/api/tests/test_generator_registry.py +++ b/api/tests/test_generator_registry.py @@ -211,6 +211,69 @@ def test_declared_sources_block_generation_even_when_generator_overrides_readine with self.assertRaisesRegex(RuntimeError, "Model sources are incomplete"): self.registry.get_active() + def test_shared_groups_gate_all_dependents_and_keep_private_dirs_separate(self) -> None: + extension = self._make_extension("shared-model") + manifest = { + "id": "shared-model", + "name": "shared-model", + "type": "model", + "generator_class": "TestGenerator", + "weight_groups": [{ + "id": "base", + "model_sources": [{ + "id": "base", + "provider": "huggingface", + "repo_id": "org/base", + "destination": ".", + "checks": ["base.bin"], + }], + }], + "nodes": [ + {"id": "generate", "weight_groups": ["base"]}, + { + "id": "adapter", + "weight_groups": ["base"], + "model_sources": [{ + "id": "adapter", + "provider": "huggingface", + "repo_id": "org/adapter", + "destination": ".", + "checks": ["adapter.bin"], + }], + }, + ], + } + (extension / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + (extension / "generator.py").write_text( + "\n".join([ + "from services.generators.base import BaseGenerator", + "class TestGenerator(BaseGenerator):", + " def load(self): self._model = object()", + " def generate(self, image_bytes, params, progress_cb=None, cancel_event=None):", + " return self.outputs_dir / 'result.glb'", + ]), + encoding="utf-8", + ) + + self.registry.initialize() + base_root = self.models_dir / "shared-model" / "_shared" / "base" + generate = self.registry.get_generator("shared-model/generate") + adapter = self.registry.get_generator("shared-model/adapter") + self.assertEqual(generate.shared_model_dirs, {"base": base_root}) + self.assertEqual(adapter.shared_model_dirs, {"base": base_root}) + self.assertFalse(self.registry._is_downloaded("shared-model/generate", generate)) + self.assertFalse(self.registry._is_downloaded("shared-model/adapter", adapter)) + + base_root.mkdir(parents=True) + (base_root / "base.bin").write_bytes(b"base") + self.assertTrue(self.registry._is_downloaded("shared-model/generate", generate)) + self.assertFalse(self.registry._is_downloaded("shared-model/adapter", adapter)) + + private_root = self.models_dir / "shared-model" / "adapter" + private_root.mkdir(parents=True) + (private_root / "adapter.bin").write_bytes(b"adapter") + self.assertTrue(self.registry._is_downloaded("shared-model/adapter", adapter)) + def test_reload_preserves_legacy_path_owned_by_the_host(self) -> None: extension = self._make_extension("host-owned-path") self._write_manifest(extension, extension_id="host-owned-path") diff --git a/api/tests/test_model_router.py b/api/tests/test_model_router.py index 3abaef0b..9bd9f71a 100644 --- a/api/tests/test_model_router.py +++ b/api/tests/test_model_router.py @@ -177,6 +177,29 @@ async def one_run(): self.assertEqual(resumed[-1], {"percent": 100, "status": "done"}) self.assertTrue((self.models_dir / "pixal3d/generate/main.bin").is_file()) + def test_shared_target_downloads_under_extension_reserved_root(self) -> None: + calls: list[str] = [] + self.install_hf_stub({"org/main": ["main.bin"]}, calls) + + def fake_download(**kwargs): + target = Path(kwargs["dest_dir"]) / kwargs["filename"] + target.parent.mkdir(parents=True, exist_ok=True) + target.write_bytes(b"shared") + return target.stat().st_size + + async def run(): + with patch.object(model_router, "_download_file_streamed", fake_download): + response = await model_router.hf_download_sources( + request_for([SOURCES[0]]), "pixal3d/_shared/base" + ) + return await collect_events(response) + + events = asyncio.run(run()) + self.assertEqual(events[-1], {"percent": 100, "status": "done"}) + self.assertTrue( + (self.models_dir / "pixal3d/_shared/base/main.bin").is_file() + ) + def test_rejects_a_check_filtered_out_of_the_source_plan(self) -> None: calls: list[str] = [] self.install_hf_stub({"org/main": ["other.bin"]}, calls) diff --git a/api/tests/test_model_sources.py b/api/tests/test_model_sources.py index cdab245a..3cbe661e 100644 --- a/api/tests/test_model_sources.py +++ b/api/tests/test_model_sources.py @@ -6,8 +6,13 @@ from services.model_sources import ( model_sources_are_downloaded, normalize_model_sources, + normalize_weight_group_references, + normalize_weight_groups, resolve_model_root, + resolve_weight_group_root, + resolve_weight_storage_root, validate_source_file_plan, + weight_group_sources_are_downloaded, ) @@ -112,6 +117,58 @@ def test_requires_all_checks_and_rejects_symlinked_extension_ancestry(self) -> N resolve_model_root(models, "pixal3d/generate") self.assertFalse(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + def test_validates_group_references_and_rejects_portable_aliases(self) -> None: + groups = normalize_weight_groups({ + "weight_groups": [{ + "id": "Base-Weights", + "model_sources": valid_node()["model_sources"], + }] + }) + self.assertEqual( + normalize_weight_group_references( + {"weight_groups": ["base-weights"]}, groups + ), + ["Base-Weights"], + ) + with self.assertRaisesRegex(ValueError, "unknown weight group"): + normalize_weight_group_references( + {"weight_groups": ["missing"]}, groups + ) + with self.assertRaisesRegex(ValueError, "portable-unique"): + normalize_weight_groups({ + "weight_groups": [ + {"id": "base", "model_sources": valid_node()["model_sources"]}, + {"id": "BASE", "model_sources": valid_node()["model_sources"]}, + ] + }) + + def test_shared_group_uses_reserved_extension_storage_root(self) -> None: + group = (normalize_weight_groups({ + "weight_groups": [{ + "id": "base", + "model_sources": [{ + "id": "primary", + "provider": "huggingface", + "repo_id": "org/base", + "destination": ".", + "checks": ["model.bin"], + }], + }] + }) or [])[0] + with tempfile.TemporaryDirectory(prefix="modly-shared-sources-") as tmp: + models = Path(tmp) / "models" + group_root = models / "demo" / "_shared" / "base" + self.assertEqual(resolve_weight_group_root(models, "demo", "base"), group_root) + self.assertEqual( + resolve_weight_storage_root(models, "demo/_shared/base"), group_root + ) + with self.assertRaisesRegex(ValueError, "reserved"): + resolve_model_root(models, "demo/_shared") + self.assertFalse(weight_group_sources_are_downloaded(models, "demo", group)) + group_root.mkdir(parents=True) + (group_root / "model.bin").write_bytes(b"weights") + self.assertTrue(weight_group_sources_are_downloaded(models, "demo", group)) + if __name__ == "__main__": unittest.main() diff --git a/api/tests/test_runner.py b/api/tests/test_runner.py index 8fce3d31..a8faeeaf 100644 --- a/api/tests/test_runner.py +++ b/api/tests/test_runner.py @@ -32,6 +32,17 @@ def test_select_node_uses_model_dir_override(self) -> None: self.assertEqual(node["id"], "quality") + def test_select_node_prefers_explicit_node_id_over_storage_path(self) -> None: + manifest = {"nodes": [{"id": "fast"}, {"id": "quality"}]} + + node = _select_node( + manifest, + str(Path("/tmp/ext/_shared/base")), + "quality", + ) + + self.assertEqual(node["id"], "quality") + def test_ready_schema_falls_back_to_selected_node_schema(self) -> None: class GenClass: @classmethod diff --git a/electron/main/extension-install-utils.test.mjs b/electron/main/extension-install-utils.test.mjs index 84139f9a..121b9353 100644 --- a/electron/main/extension-install-utils.test.mjs +++ b/electron/main/extension-install-utils.test.mjs @@ -83,6 +83,13 @@ test('validateInstallManifest accepts multi-source nodes and preserves legacy sh hf_skip_prefixes: ['weights/**'], }], }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) + + // Nodes that do not opt into managed sources keep the pre-existing validation path. + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'legacy-unmanaged', + generator_class: 'Generator', + nodes: [{ id: 'legacy node' }], + }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) }) test('validateInstallManifest rejects malformed or process model_sources', () => { @@ -102,6 +109,56 @@ test('validateInstallManifest rejects malformed or process model_sources', () => }, { hasEntryFile: () => true, hasGeneratorFile: () => false }, 'repository'), /only for model nodes/i) }) +test('validateInstallManifest accepts shared groups with private sources', () => { + const mod = loadModule() + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'shared-model', + generator_class: 'Generator', + weight_groups: [{ + id: 'base', + model_sources: [{ + id: 'base', provider: 'huggingface', repo_id: 'org/base', + destination: '.', checks: ['base.bin'], + }], + }], + nodes: [ + { id: 'base-node', weight_groups: ['base'] }, + { + id: 'adapter-node', + weight_groups: ['base'], + model_sources: [{ + id: 'adapter', provider: 'huggingface', repo_id: 'org/adapter', + destination: '.', checks: ['adapter.bin'], + }], + }, + ], + }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) +}) + +test('validateInstallManifest rejects unsafe shared-weight contracts', () => { + const mod = loadModule() + const group = { + id: 'base', + model_sources: [{ + id: 'base', provider: 'huggingface', repo_id: 'org/base', + destination: '.', checks: ['base.bin'], + }], + } + const files = { hasEntryFile: () => true, hasGeneratorFile: () => true } + assert.throws(() => mod.validateInstallManifest({ + id: 'unknown', generator_class: 'Generator', + weight_groups: [group], nodes: [{ id: 'generate', weight_groups: ['missing'] }], + }, files, 'repository'), /unknown weight group/i) + assert.throws(() => mod.validateInstallManifest({ + id: 'reserved', generator_class: 'Generator', + weight_groups: [group], nodes: [{ id: '_shared', weight_groups: ['base'] }], + }, files, 'repository'), /reserved/i) + assert.throws(() => mod.validateInstallManifest({ + id: 'process', type: 'process', entry: 'processor.py', weight_groups: [group], + nodes: [{ id: 'run' }], + }, files, 'repository'), /only for model extensions/i) +}) + test('python process setup failures are treated as fatal', () => { const mod = loadModule() diff --git a/electron/main/extension-install-utils.ts b/electron/main/extension-install-utils.ts index 05b965b0..d50e4acc 100644 --- a/electron/main/extension-install-utils.ts +++ b/electron/main/extension-install-utils.ts @@ -1,7 +1,9 @@ import { normalizeModelSources, + normalizeWeightGroupReferences, + normalizeWeightGroups, safeModelSourceId, - type ModelSourceNode, + type ModelWeightNode, } from './model-sources' export interface InstallManifest { @@ -10,7 +12,13 @@ export interface InstallManifest { entry?: string generator_class?: string model_sources?: unknown - nodes?: Array<{ id?: string; model_sources?: unknown } & ModelSourceNode> + weight_groups?: unknown + nodes?: Array<{ + id?: string + hf_repo?: unknown + model_sources?: unknown + weight_groups?: unknown + } & ModelWeightNode> } export interface ValidatedInstallManifest { @@ -49,13 +57,27 @@ export function validateInstallManifest( if (manifest.model_sources !== undefined) { throw new Error('manifest.json: model_sources must be declared on a model node') } + if (isProcess && manifest.weight_groups !== undefined) { + throw new Error('manifest.json: weight_groups is supported only for model extensions') + } + const weightGroups = normalizeWeightGroups(manifest) for (const node of Array.isArray(manifest.nodes) ? manifest.nodes : []) { - if (node.model_sources === undefined) continue - if (isProcess) { - throw new Error('manifest.json: model_sources is supported only for model nodes') + const usesSharedWeights = weightGroups !== undefined || node.weight_groups !== undefined + if (usesSharedWeights && typeof node.id === 'string' && node.id.toLowerCase() === '_shared') { + throw new Error('manifest.json: model node id "_shared" is reserved') + } + if (isProcess && (node.model_sources !== undefined || node.weight_groups !== undefined)) { + throw new Error('manifest.json: model_sources and weight_groups are supported only for model nodes') + } + if (!usesSharedWeights && node.model_sources === undefined) continue + const nodeId = safeModelSourceId(node.id, 'model node id') + if (node.model_sources !== undefined) normalizeModelSources(node) + normalizeWeightGroupReferences(node, weightGroups, `nodes[${nodeId}].weight_groups`) + if (node.weight_groups !== undefined && node.hf_repo !== undefined) { + throw new Error( + `manifest.json: model node "${nodeId}" must use model_sources for private weights when weight_groups are declared`, + ) } - safeModelSourceId(node.id, 'model node id') - normalizeModelSources(node) } if (isProcess) { diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 60f3cf2b..245cc4f2 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -14,14 +14,27 @@ import { listDownloadedModels, downloadModelFromHF, downloadModelSourcesFromHF, + type DownloadProgress, } from './model-downloader' -import { resolveInstalledModelDownloadPlan } from './model-download-plan' +import { + resolveInstalledExtensionSharedWeightGroups, + resolveInstalledModelDownloadPlan, +} from './model-download-plan' import { areModelSourcesDownloaded, + areModelSourcesDownloadedAtRoot, + areWeightGroupSourcesDownloaded, modelHasLocalData, normalizeModelSources, + normalizeWeightGroupReferences, + normalizeWeightGroups, removePartialDownloadArtifacts, + resolveExtensionModelRoot, resolveModelRoot, + resolveWeightGroupRoot, + resolveWeightStorageRoot, + safeModelSourceId, + weightStorageHasLocalData, } from './model-sources' import { getSettings, setSettings } from './settings-store' import { checkSetupNeeded, markSetupDone, runFullSetup, getVenvPythonExe, ensureSslPatch } from './python-setup' @@ -149,11 +162,14 @@ const renameWithRetry = (from: string, to: string, label: string) => export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGetter): void { type ActiveDownload = { - progress: { percent: number; file?: string; fileIndex?: number; totalFiles?: number } + progress: DownloadProgress done: Promise finish: () => void + targetRoots: string[] + currentTargetId?: string } const activeDownloads = new Map() + const activeWeightTargets = new Map() // Logging from renderer ipcMain.on('log:error', (_event, message: string) => logger.error(`[Renderer] ${message}`)) ipcMain.handle('log:getPath', () => join(app.getPath('userData'), 'logs', 'modly.log')) @@ -409,9 +425,14 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe builtinExtensionsDir: getBuiltinExtensionsDir(), blockedExtensionIds: activeExtensionInstalls, }) - return plan.kind === 'multi-source' - ? areModelSourcesDownloaded(modelsDir, modelId, plan.sources) - : isModelDownloaded(modelsDir, modelId, plan.downloadCheck) + if (plan.kind === 'multi-source') { + const privateReady = plan.sources.length === 0 + || areModelSourcesDownloaded(modelsDir, modelId, plan.sources) + return privateReady && plan.sharedGroups.every((group) => ( + areWeightGroupSourcesDownloaded(modelsDir, plan.extensionId, group) + )) + } + return isModelDownloaded(modelsDir, modelId, plan.downloadCheck) } catch { return false } @@ -431,6 +452,102 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe } }) + ipcMain.handle('model:sharedGroups', async (_, extensionId: string) => { + try { + const groups = await resolveInstalledExtensionSharedWeightGroups({ + extensionId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + const modelsDir = getSettings(app.getPath('userData')).modelsDir + return groups.map((group) => ({ + id: group.id, + targetId: group.targetId, + dependentModelIds: group.dependentModelIds, + downloaded: areWeightGroupSourcesDownloaded(modelsDir, extensionId, group), + hasLocalData: weightStorageHasLocalData(modelsDir, group.targetId), + })) + } catch { + return [] + } + }) + + ipcMain.handle('model:deleteSharedGroup', async ( + _, + extensionId: string, + groupId: string, + ): Promise<{ success: boolean; error?: string }> => { + try { + const groups = await resolveInstalledExtensionSharedWeightGroups({ + extensionId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + const group = groups.find((candidate) => candidate.id === groupId) + if (!group) return { success: false, error: `Unknown shared weight group: ${groupId}` } + const groupRoot = resolveWeightGroupRoot( + getSettings(app.getPath('userData')).modelsDir, + extensionId, + group.id, + ) + if (activeWeightTargets.has(groupRoot)) { + return { success: false, error: 'Cannot remove shared weights while their download is active' } + } + await Promise.all(group.dependentModelIds.map(async (dependentModelId) => { + try { + await axios.post( + `${API_BASE_URL}/model/unload/${encodeURIComponent(dependentModelId)}`, + {}, + { timeout: 10_000 }, + ) + } catch { /* an unloaded or unavailable model does not block file removal */ } + })) + await new Promise(resolve => setTimeout(resolve, 1_500)) + const removed = await rmWithRetry(groupRoot, 'shared-model-delete') + if (removed.ok) return { success: true } + return { + success: false, + error: removed.locked + ? 'Shared model files are still locked. Close any programs using them and try again.' + : String(removed.error), + } + } catch (err) { + return { success: false, error: String(err) } + } + }) + + ipcMain.handle('model:deleteExtensionWeights', async ( + _, + extensionId: string, + ): Promise<{ success: boolean; error?: string }> => { + try { + const safeExtensionId = assertSafeExtensionId(extensionId) + if ([...activeDownloads.keys()].some((modelId) => modelId.split('/', 1)[0] === safeExtensionId)) { + return { success: false, error: 'Cannot remove extension weights while a download is active' } + } + const extensionRoot = resolveExtensionModelRoot( + getSettings(app.getPath('userData')).modelsDir, + safeExtensionId, + ) + try { + await axios.post(`${API_BASE_URL}/model/unload-all`, {}, { timeout: 10_000 }) + await new Promise(resolve => setTimeout(resolve, 1_500)) + } catch { /* still attempt deletion when the API is unavailable */ } + const removed = await rmWithRetry(extensionRoot, 'extension-model-delete') + if (removed.ok) return { success: true } + return { + success: false, + error: removed.locked + ? 'Extension model files are still locked. Close any programs using them and try again.' + : String(removed.error), + } + } catch (err) { + return { success: false, error: String(err) } + } + }) + ipcMain.handle('model:activeDownloads', () => [...activeDownloads.entries()].map(([modelId, active]) => ({ modelId, ...active.progress })) ) @@ -442,24 +559,81 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe if (activeDownloads.has(modelId)) { return { success: false, error: 'Download already in progress' } } - let finish!: () => void - const done = new Promise((resolveDone) => { finish = resolveDone }) - const active: ActiveDownload = { progress: { percent: 0 }, done, finish } - activeDownloads.set(modelId, active) + let plan: Awaited> try { - const plan = await resolveInstalledModelDownloadPlan({ + plan = await resolveInstalledModelDownloadPlan({ modelId, userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, builtinExtensionsDir: getBuiltinExtensionsDir(), blockedExtensionIds: activeExtensionInstalls, }) + } catch (err) { + return { success: false, error: String(err) } + } + if (activeDownloads.has(modelId)) { + return { success: false, error: 'Download already in progress' } + } + + const modelsDir = getSettings(app.getPath('userData')).modelsDir + const managedTargets = plan.kind === 'multi-source' + ? [ + ...plan.sharedGroups.map((group) => ({ + targetId: group.targetId, + label: `Shared · ${group.id}`, + sources: group.sources, + })), + ...(plan.sources.length > 0 ? [{ + targetId: modelId, + label: 'Node-specific', + sources: plan.sources, + }] : []), + ].filter((target) => !areModelSourcesDownloadedAtRoot( + resolveWeightStorageRoot(modelsDir, target.targetId), + target.sources, + )) + : [] + const targetRoots = plan.kind === 'multi-source' + ? managedTargets.map((target) => resolveWeightStorageRoot(modelsDir, target.targetId)) + : [resolveModelRoot(modelsDir, modelId)] + const conflict = targetRoots.find((root) => activeWeightTargets.has(root)) + if (conflict) { + return { + success: false, + error: `Weights are already being downloaded by ${activeWeightTargets.get(conflict)}`, + } + } + + let finish!: () => void + const done = new Promise((resolveDone) => { finish = resolveDone }) + const active: ActiveDownload = { progress: { percent: 0 }, done, finish, targetRoots } + activeDownloads.set(modelId, active) + for (const root of targetRoots) activeWeightTargets.set(root, modelId) + try { const onProgress = (progress: typeof active.progress) => { active.progress = progress event.sender.send('model:downloadProgress', { modelId, ...progress }) } if (plan.kind === 'multi-source') { - await downloadModelSourcesFromHF(modelId, plan.sources, onProgress) + if (managedTargets.length === 0) { + onProgress({ percent: 100 }) + } + for (const [index, target] of managedTargets.entries()) { + active.currentTargetId = target.targetId + await downloadModelSourcesFromHF(target.targetId, target.sources, (progress) => { + const aggregatePercent = Math.min( + 99, + Math.round(((index + progress.percent / 100) / managedTargets.length) * 100), + ) + onProgress({ + ...progress, + percent: aggregatePercent, + status: progress.status ? `${target.label} · ${progress.status}` : target.label, + }) + }) + } + if (managedTargets.length > 0) onProgress({ percent: 100, status: 'done' }) } else { + active.currentTargetId = modelId await downloadModelFromHF( plan.repoId, modelId, @@ -482,14 +656,19 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe return { success: false, error: String(err) } } finally { if (activeDownloads.get(modelId) === active) activeDownloads.delete(modelId) + for (const root of targetRoots) { + if (activeWeightTargets.get(root) === modelId) activeWeightTargets.delete(root) + } active.finish() } }) ipcMain.handle('model:pauseDownload', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { try { + const targetId = activeDownloads.get(modelId)?.currentTargetId + if (!targetId) return { success: false, error: 'No active download target' } await axios.post(`${API_BASE_URL}/model/hf-download/pause`, null, { - params: { model_id: modelId }, + params: { model_id: targetId }, timeout: 5000, }) return { success: true } @@ -501,10 +680,12 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe ipcMain.handle('model:cancelDownload', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { try { const active = activeDownloads.get(modelId) - await axios.post(`${API_BASE_URL}/model/hf-download/cancel`, null, { - params: { model_id: modelId }, - timeout: 5000, - }) + if (active?.currentTargetId) { + await axios.post(`${API_BASE_URL}/model/hf-download/cancel`, null, { + params: { model_id: active.currentTargetId }, + timeout: 5000, + }) + } if (active) { await Promise.race([ active.done, @@ -513,11 +694,13 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }), ]) } - const modelDir = resolveModelRoot(getSettings(app.getPath('userData')).modelsDir, modelId) // Only remove in-progress `.part` files — a model can now have multiple sources // sharing this directory, and any source that already finished downloading // must survive cancelling the ones still in flight. - await removePartialDownloadArtifacts(modelDir) + await Promise.all((active?.targetRoots ?? [resolveModelRoot( + getSettings(app.getPath('userData')).modelsDir, + modelId, + )]).map((root) => removePartialDownloadArtifacts(root))) return { success: true } } catch (err) { return { success: false, error: String(err) } @@ -818,6 +1001,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe type?: 'model' | 'process' entry?: string model_sources?: unknown + weight_groups?: unknown // Optional top-level fallbacks — applied to each node if not set on the node params_schema?: unknown[] param_defaults?: Record @@ -835,6 +1019,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe hf_skip_prefixes?: string[] hf_include_prefixes?: string[] model_sources?: unknown + weight_groups?: unknown }[] } @@ -853,13 +1038,34 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe if (parsed.model_sources !== undefined) { throw new Error('manifest.json: model_sources must be declared on a model node') } + if (parsed.type === 'process' && parsed.weight_groups !== undefined) { + throw new Error('manifest.json: weight_groups is supported only for model extensions') + } + const weightGroups = normalizeWeightGroups(parsed) const nodes = (parsed.nodes ?? []).map(n => { - if (parsed.type === 'process' && n.model_sources !== undefined) { - throw new Error('manifest.json: model_sources is supported only for model nodes') + const usesManagedWeights = weightGroups !== undefined + || n.model_sources !== undefined + || n.weight_groups !== undefined + if (weightGroups !== undefined && typeof n.id === 'string' && n.id.toLowerCase() === '_shared') { + throw new Error('manifest.json: model node id "_shared" is reserved') + } + const nodeId = usesManagedWeights ? safeModelSourceId(n.id, 'model node id') : n.id + if (parsed.type === 'process' && (n.model_sources !== undefined || n.weight_groups !== undefined)) { + throw new Error('manifest.json: model_sources and weight_groups are supported only for model nodes') } const modelSources = normalizeModelSources(n) + const groupRefs = normalizeWeightGroupReferences( + n, + weightGroups, + `nodes[${n.id}].weight_groups`, + ) + if (groupRefs && n.hf_repo !== undefined) { + throw new Error( + `manifest.json: model node "${nodeId}" must use model_sources for private weights when weight_groups are declared`, + ) + } return { - id: n.id, + id: nodeId, name: n.name ?? n.id, input: n.input ?? 'image' as const, inputs: n.inputs, @@ -872,6 +1078,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe hfSkipPrefixes: n.hf_skip_prefixes, hfIncludePrefixes: n.hf_include_prefixes, hasModelSources: modelSources !== undefined, + weightGroups: groupRefs, } }) @@ -879,7 +1086,17 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe return { ...common, type: 'process' as const, entry: parsed.entry ?? 'processor.js', nodes } } - return { ...common, type: 'model' as const, nodes } + return { + ...common, + type: 'model' as const, + nodes, + weightGroups: (weightGroups ?? []).map((group) => ({ + id: group.id, + dependentNodeIds: nodes + .filter((node) => node.weightGroups?.includes(group.id)) + .map((node) => node.id), + })), + } } async function reloadAndValidateModelExtension( diff --git a/electron/main/model-download-plan.test.mjs b/electron/main/model-download-plan.test.mjs index e5d037a0..28aae53c 100644 --- a/electron/main/model-download-plan.test.mjs +++ b/electron/main/model-download-plan.test.mjs @@ -90,3 +90,94 @@ test('keeps legacy sibling checks and wildcard filters unchanged', async () => { rmSync(fixture.root, { recursive: true, force: true }) } }) + +test('composes shared base groups with private node sources and dependency metadata', async () => { + const { + resolveInstalledExtensionSharedWeightGroups, + resolveInstalledModelDownloadPlan, + } = loadModule() + const fixture = setupExtension({ + id: 'pixal3d', + type: 'model', + weight_groups: [{ + id: 'base', + model_sources: [{ + id: 'base-model', provider: 'huggingface', repo_id: 'org/base', + destination: '.', checks: ['pipeline.json'], + }], + }], + nodes: [ + { id: 'generate', weight_groups: ['base'] }, + { + id: 'worldsculpt', + weight_groups: ['base'], + model_sources: [{ + id: 'adapter', provider: 'huggingface', repo_id: 'org/adapter', + destination: '.', checks: ['adapter.bin'], + }], + }, + ], + }) + try { + const plan = await resolveInstalledModelDownloadPlan({ + modelId: 'pixal3d/worldsculpt', + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + }) + assert.equal(plan.kind, 'multi-source') + assert.equal(plan.sources[0].repo_id, 'org/adapter') + assert.equal(plan.sharedGroups[0].targetId, 'pixal3d/_shared/base') + assert.deepEqual(plan.sharedGroups[0].dependentModelIds, [ + 'pixal3d/generate', + 'pixal3d/worldsculpt', + ]) + + const groups = await resolveInstalledExtensionSharedWeightGroups({ + extensionId: 'pixal3d', + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + }) + assert.equal(groups.length, 1) + assert.equal(groups[0].sources[0].repo_id, 'org/base') + } finally { + rmSync(fixture.root, { recursive: true, force: true }) + } +}) + +test('rejects unknown groups, reserved node ids, and legacy private aliases', async () => { + const { resolveInstalledModelDownloadPlan } = loadModule() + for (const manifest of [ + { + id: 'unknown-group', type: 'model', nodes: [{ id: 'generate', weight_groups: ['missing'] }], + }, + { + id: 'reserved-node', type: 'model', nodes: [{ id: '_shared', hf_repo: 'org/model' }], + }, + { + id: 'legacy-private', type: 'model', + weight_groups: [{ + id: 'base', + model_sources: [{ + id: 'base', provider: 'huggingface', repo_id: 'org/base', + destination: '.', checks: ['base.bin'], + }], + }], + nodes: [{ id: 'generate', weight_groups: ['base'], hf_repo: 'org/private' }], + }, + ]) { + const fixture = setupExtension(manifest) + const nodeId = manifest.nodes[0].id + try { + await assert.rejects( + resolveInstalledModelDownloadPlan({ + modelId: `${manifest.id}/${nodeId}`, + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + }), + /unknown weight group|reserved|must use model_sources/i, + ) + } finally { + rmSync(fixture.root, { recursive: true, force: true }) + } + } +}) diff --git a/electron/main/model-download-plan.ts b/electron/main/model-download-plan.ts index 27d2dba6..8b17f694 100644 --- a/electron/main/model-download-plan.ts +++ b/electron/main/model-download-plan.ts @@ -8,7 +8,15 @@ import { assertSafeExtensionId, resolveExtensionPathWithinRoot, } from './extension-path-guard' -import { normalizeModelSources, safeModelSourceId, type ModelSource } from './model-sources' +import { + normalizeModelSources, + normalizeWeightGroupReferences, + normalizeWeightGroups, + safeModelSourceId, + weightGroupTargetId, + type ModelSource, + type ModelWeightGroup, +} from './model-sources' interface InstalledNode { id?: unknown @@ -17,15 +25,22 @@ interface InstalledNode { hf_skip_prefixes?: unknown hf_include_prefixes?: unknown model_sources?: unknown + weight_groups?: unknown } interface InstalledManifest { id?: unknown type?: unknown model_sources?: unknown + weight_groups?: unknown nodes?: unknown } +export interface InstalledSharedWeightGroup extends ModelWeightGroup { + targetId: string + dependentModelIds: string[] +} + export type InstalledModelDownloadPlan = { kind: 'legacy' modelId: string @@ -41,6 +56,7 @@ export type InstalledModelDownloadPlan = { extensionId: string nodeId: string sources: ModelSource[] + sharedGroups: InstalledSharedWeightGroup[] } async function hasPendingRegistration(root: string, extensionId: string): Promise { @@ -54,6 +70,37 @@ async function hasPendingRegistration(root: string, extensionId: string): Promis } } +function installedSharedGroups( + manifest: InstalledManifest, + extensionId: string, + nodes: InstalledNode[], +): InstalledSharedWeightGroup[] { + const groups = normalizeWeightGroups(manifest) + if (groups === undefined) return [] + const groupDependents = new Map() + for (const candidate of nodes) { + if (typeof candidate.id === 'string' && candidate.id.toLowerCase() === '_shared') { + throw new Error('manifest.json: model node id "_shared" is reserved') + } + const candidateId = safeModelSourceId(candidate.id, 'model node id') + const refs = normalizeWeightGroupReferences( + candidate, + groups, + `nodes[${candidateId}].weight_groups`, + ) ?? [] + for (const groupId of refs) { + const dependents = groupDependents.get(groupId) ?? [] + dependents.push(`${extensionId}/${candidateId}`) + groupDependents.set(groupId, dependents) + } + } + return groups.map((group) => ({ + ...group, + targetId: weightGroupTargetId(extensionId, group.id), + dependentModelIds: groupDependents.get(group.id) ?? [], + })) +} + function parseManifest(raw: string, extensionId: string, nodeId: string): InstalledModelDownloadPlan { let parsed: unknown try { @@ -74,12 +121,12 @@ function parseManifest(raw: string, extensionId: string, nodeId: string): Instal } if (!Array.isArray(manifest.nodes)) throw new Error(`Extension "${extensionId}" does not declare model nodes`) - const matches = manifest.nodes.filter((candidate): candidate is InstalledNode => ( - typeof candidate === 'object' - && candidate !== null - && !Array.isArray(candidate) - && (candidate as InstalledNode).id === nodeId + const nodes = manifest.nodes.filter((candidate): candidate is InstalledNode => ( + typeof candidate === 'object' && candidate !== null && !Array.isArray(candidate) )) + const allSharedGroups = installedSharedGroups(manifest, extensionId, nodes) + + const matches = nodes.filter((candidate) => candidate.id === nodeId) if (matches.length !== 1) { throw new Error(`Installed manifest must declare model node "${nodeId}" exactly once`) } @@ -87,7 +134,27 @@ function parseManifest(raw: string, extensionId: string, nodeId: string): Instal const node = matches[0] const modelId = `${extensionId}/${nodeId}` const sources = normalizeModelSources(node) - if (sources) return { kind: 'multi-source', modelId, extensionId, nodeId, sources } + const refs = normalizeWeightGroupReferences( + node, + allSharedGroups, + `nodes[${nodeId}].weight_groups`, + ) ?? [] + const sharedGroups = refs.map((groupId) => ( + allSharedGroups.find((candidate) => candidate.id === groupId)! + )) + if (sources || sharedGroups.length > 0) { + if (sharedGroups.length > 0 && node.hf_repo !== undefined) { + throw new Error(`Model node "${modelId}" must use model_sources for private weights when weight_groups are declared`) + } + return { + kind: 'multi-source', + modelId, + extensionId, + nodeId, + sources: sources ?? [], + sharedGroups, + } + } if (typeof node.hf_repo !== 'string' || !node.hf_repo) { throw new Error(`Model node "${modelId}" has no Hugging Face download source`) @@ -104,6 +171,53 @@ function parseManifest(raw: string, extensionId: string, nodeId: string): Instal } } + +export async function resolveInstalledExtensionSharedWeightGroups(args: { + extensionId: unknown + userExtensionsDir: string + builtinExtensionsDir: string + blockedExtensionIds?: ReadonlySet +}): Promise { + const extensionId = assertSafeExtensionId(args.extensionId) + if (args.blockedExtensionIds?.has(extensionId)) { + throw new Error(`Extension "${extensionId}" is being installed or repaired`) + } + const userPath = resolveExtensionPathWithinRoot(args.userExtensionsDir, extensionId) + const builtinPath = resolveExtensionPathWithinRoot(args.builtinExtensionsDir, extensionId) + const extensionPath = existsSync(userPath) ? userPath : existsSync(builtinPath) ? builtinPath : undefined + if (!extensionPath) throw new Error(`Extension "${extensionId}" is not installed`) + const extensionRoot = extensionPath === userPath ? args.userExtensionsDir : args.builtinExtensionsDir + if ( + existsSync(join(extensionPath, EXT_INCOMPLETE_MARKER)) + || existsSync(join(extensionPath, EXT_REGISTRATION_PENDING_MARKER)) + || await hasPendingRegistration(extensionRoot, extensionId) + ) { + throw new Error(`Extension "${extensionId}" has an incomplete installation`) + } + + const manifestPath = join(extensionPath, 'manifest.json') + if (!existsSync(manifestPath)) throw new Error(`Extension "${extensionId}" has no manifest.json`) + let parsed: unknown + try { + parsed = JSON.parse(await readFile(manifestPath, 'utf-8')) + } catch { + throw new Error(`Extension "${extensionId}" has an invalid manifest.json`) + } + if (typeof parsed !== 'object' || parsed === null || Array.isArray(parsed)) { + throw new Error(`Extension "${extensionId}" has an invalid manifest.json`) + } + const manifest = parsed as InstalledManifest + if (manifest.id !== extensionId) throw new Error(`Installed manifest id does not match extension "${extensionId}"`) + if (manifest.type !== undefined && manifest.type !== 'model') { + throw new Error(`Extension "${extensionId}" is not a model extension`) + } + if (!Array.isArray(manifest.nodes)) throw new Error(`Extension "${extensionId}" does not declare model nodes`) + const nodes = manifest.nodes.filter((candidate): candidate is InstalledNode => ( + typeof candidate === 'object' && candidate !== null && !Array.isArray(candidate) + )) + return installedSharedGroups(manifest, extensionId, nodes) +} + /** Re-read the installed manifest for every model action; renderer metadata is never trusted. */ export async function resolveInstalledModelDownloadPlan(args: { modelId: unknown @@ -115,6 +229,7 @@ export async function resolveInstalledModelDownloadPlan(args: { const parts = args.modelId.split('/') if (parts.length !== 2) throw new Error('Model id must identify one extension node') const extensionId = assertSafeExtensionId(parts[0]) + if (parts[1].toLowerCase() === '_shared') throw new Error('Model node id "_shared" is reserved') const nodeId = safeModelSourceId(parts[1], 'model node id') if (args.blockedExtensionIds?.has(extensionId)) { throw new Error(`Extension "${extensionId}" is being installed or repaired`) diff --git a/electron/main/model-download-preload.test.mjs b/electron/main/model-download-preload.test.mjs index 8d9a16b7..aa414479 100644 --- a/electron/main/model-download-preload.test.mjs +++ b/electron/main/model-download-preload.test.mjs @@ -20,7 +20,7 @@ function loadModule() { return require(outfile) } -test('renderer model actions send only the model node id', async () => { +test('renderer model actions keep node and shared-weight identities explicit', async () => { const { createElectronApi } = loadModule() const calls = [] const ipc = { @@ -33,12 +33,18 @@ test('renderer model actions send only the model node id', async () => { await api.model.isDownloaded('pixal3d/generate') await api.model.hasLocalData('pixal3d/generate') + await api.model.sharedGroups('pixal3d') await api.model.download('pixal3d/generate') + await api.model.deleteSharedGroup('pixal3d', 'base') + await api.model.deleteExtensionWeights('pixal3d') assert.deepEqual(calls, [ ['model:isDownloaded', 'pixal3d/generate'], ['model:hasLocalData', 'pixal3d/generate'], + ['model:sharedGroups', 'pixal3d'], ['model:download', 'pixal3d/generate'], + ['model:deleteSharedGroup', 'pixal3d', 'base'], + ['model:deleteExtensionWeights', 'pixal3d'], ]) }) @@ -48,8 +54,12 @@ test('declared partial data is removable and active downloads block destructive const drawer = readFileSync(resolve('src/areas/models/components/ExtensionDrawer.tsx'), 'utf8') assert.match(main, /model:delete[\s\S]*activeDownloads\.has\(modelId\)/) + assert.match(main, /model:deleteSharedGroup[\s\S]*activeWeightTargets\.has\(groupRoot\)/) + assert.match(main, /model:deleteExtensionWeights[\s\S]*resolveExtensionModelRoot/) assert.match(main, /extensions:uninstall[\s\S]*activeDownloads\.keys\(\)/) assert.match(page, /window\.electron\.model\.hasLocalData\(fullId\)/) + assert.match(page, /deleteExtensionWeights\(extId\)/) assert.match(drawer, /localDataIds\.includes\(fullId\) && state\.kind !== 'downloading'/) assert.match(drawer, /Remove partial model data/) + assert.match(drawer, /following nodes will become unavailable/) }) diff --git a/electron/main/model-sources.test.mjs b/electron/main/model-sources.test.mjs index da8a54fc..15dfc5a8 100644 --- a/electron/main/model-sources.test.mjs +++ b/electron/main/model-sources.test.mjs @@ -113,3 +113,63 @@ test('requires every declared check and rejects symlinked extension-root ancestr rmSync(root, { recursive: true, force: true }) } }) + +test('validates extension-scoped groups and canonicalizes sibling references', () => { + const { normalizeWeightGroups, normalizeWeightGroupReferences } = loadModule() + const groups = normalizeWeightGroups({ + weight_groups: [{ + id: 'Base-Weights', + model_sources: validNode().model_sources, + }], + }) + assert.deepEqual( + normalizeWeightGroupReferences({ weight_groups: ['base-weights'] }, groups), + ['Base-Weights'], + ) + assert.throws( + () => normalizeWeightGroupReferences({ weight_groups: ['missing'] }, groups), + /unknown weight group/i, + ) + assert.throws( + () => normalizeWeightGroups({ + weight_groups: [ + { id: 'base', model_sources: validNode().model_sources }, + { id: 'BASE', model_sources: validNode().model_sources }, + ], + }), + /portable-unique/i, + ) +}) + +test('stores and checks shared weights under the reserved extension root', () => { + const { + areWeightGroupSourcesDownloaded, + normalizeWeightGroups, + resolveModelRoot, + resolveWeightGroupRoot, + resolveWeightStorageRoot, + } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-shared-readiness-')) + const models = join(root, 'models') + const [group] = normalizeWeightGroups({ + weight_groups: [{ + id: 'base', + model_sources: [{ + id: 'primary', provider: 'huggingface', repo_id: 'org/base', + destination: '.', checks: ['model.bin'], + }], + }], + }) + const groupRoot = join(models, 'demo', '_shared', 'base') + try { + assert.equal(resolveWeightGroupRoot(models, 'demo', 'base'), groupRoot) + assert.equal(resolveWeightStorageRoot(models, 'demo/_shared/base'), groupRoot) + assert.throws(() => resolveModelRoot(models, 'demo/_shared'), /reserved/i) + assert.equal(areWeightGroupSourcesDownloaded(models, 'demo', group), false) + mkdirSync(groupRoot, { recursive: true }) + writeFileSync(join(groupRoot, 'model.bin'), 'weights') + assert.equal(areWeightGroupSourcesDownloaded(models, 'demo', group), true) + } finally { + rmSync(root, { recursive: true, force: true }) + } +}) diff --git a/electron/main/model-sources.ts b/electron/main/model-sources.ts index dac76930..cbc69e42 100644 --- a/electron/main/model-sources.ts +++ b/electron/main/model-sources.ts @@ -17,6 +17,19 @@ export interface ModelSourceNode { model_sources?: unknown } +export interface ModelWeightGroup { + id: string + sources: ModelSource[] +} + +export interface ModelWeightManifest { + weight_groups?: unknown +} + +export interface ModelWeightNode extends ModelSourceNode { + weight_groups?: unknown +} + const SAFE_ID = /^[A-Za-z0-9][A-Za-z0-9._-]*$/ const WINDOWS_DEVICE = /^(?:con|prn|aux|nul|com[1-9]|lpt[1-9])(?:\..*)?$/i const WINDOWS_UNSAFE = /[<>"|?*\u0000-\u001f]/ @@ -91,15 +104,18 @@ function safeRevision(value: unknown, field: string): string | undefined { return value } -export function normalizeModelSources(node: ModelSourceNode): ModelSource[] | undefined { +export function normalizeModelSources( + node: ModelSourceNode, + fieldName = 'model_sources', +): ModelSource[] | undefined { if (!Object.prototype.hasOwnProperty.call(node, 'model_sources')) return undefined if (!Array.isArray(node.model_sources) || node.model_sources.length === 0) { - throw new Error('model_sources must be a non-empty array') + throw new Error(`${fieldName} must be a non-empty array`) } const seen = new Map() return node.model_sources.map((raw, index) => { - const field = `model_sources[${index}]` + const field = `${fieldName}[${index}]` if (typeof raw !== 'object' || raw === null || Array.isArray(raw)) { throw new Error(`${field} must be an object`) } @@ -133,6 +149,58 @@ export function normalizeModelSources(node: ModelSourceNode): ModelSource[] | un }) } +export function normalizeWeightGroups(manifest: ModelWeightManifest): ModelWeightGroup[] | undefined { + if (!Object.prototype.hasOwnProperty.call(manifest, 'weight_groups')) return undefined + if (!Array.isArray(manifest.weight_groups) || manifest.weight_groups.length === 0) { + throw new Error('weight_groups must be a non-empty array') + } + + const seen = new Map() + return manifest.weight_groups.map((raw, index) => { + const field = `weight_groups[${index}]` + if (typeof raw !== 'object' || raw === null || Array.isArray(raw)) { + throw new Error(`${field} must be an object`) + } + const value = raw as Record + if (typeof value.id === 'string' && value.id.toLowerCase() === '_shared') { + throw new Error(`${field}.id uses the reserved identifier "_shared"`) + } + const id = safeModelSourceId(value.id, `${field}.id`) + const alias = id.normalize('NFC').toLowerCase() + const previous = seen.get(alias) + if (previous) throw new Error(`weight group ids "${previous}" and "${id}" are not portable-unique`) + seen.set(alias, id) + const sources = normalizeModelSources( + { model_sources: value.model_sources }, + `${field}.model_sources`, + ) + return { id, sources: sources! } + }) +} + +export function normalizeWeightGroupReferences( + node: ModelWeightNode, + groups: ModelWeightGroup[] | undefined, + fieldName = 'weight_groups', +): string[] | undefined { + if (!Object.prototype.hasOwnProperty.call(node, 'weight_groups')) return undefined + if (!Array.isArray(node.weight_groups) || node.weight_groups.length === 0) { + throw new Error(`${fieldName} must be a non-empty array of weight group ids`) + } + const available = new Map((groups ?? []).map((group) => [group.id.normalize('NFC').toLowerCase(), group.id])) + const seen = new Map() + return node.weight_groups.map((raw, index) => { + const id = safeModelSourceId(raw, `${fieldName}[${index}]`) + const alias = id.normalize('NFC').toLowerCase() + const previous = seen.get(alias) + if (previous) throw new Error(`weight group references "${previous}" and "${id}" are not portable-unique`) + seen.set(alias, id) + const canonical = available.get(alias) + if (!canonical) throw new Error(`${fieldName}[${index}] references unknown weight group "${id}"`) + return canonical + }) +} + function pathHasSymlink(root: string, candidate: string): boolean { const rootPath = resolve(root) const rel = relative(rootPath, resolve(candidate)) @@ -155,17 +223,55 @@ export function resolveModelRoot(modelsDir: string, modelId: string): string { const parts = modelId.split('/') if (parts.length !== 2) throw new Error('Model id must identify one extension node') const extensionId = safeModelSourceId(parts[0], 'extension id') + if (parts[1].toLowerCase() === '_shared') throw new Error('Model node id "_shared" is reserved') const nodeId = safeModelSourceId(parts[1], 'model node id') const root = resolve(modelsDir) - const modelRoot = resolve(root, extensionId, nodeId) + const extensionRoot = resolveExtensionModelRoot(modelsDir, extensionId) + const modelRoot = resolve(extensionRoot, nodeId) if (pathHasSymlink(root, modelRoot)) throw new Error('Model path resolves through a symlink') return modelRoot } -export function areModelSourcesDownloaded(modelsDir: string, modelId: string, sources: ModelSource[]): boolean { +export function resolveExtensionModelRoot(modelsDir: string, extensionId: string): string { + const safeExtensionId = safeModelSourceId(extensionId, 'extension id') + const root = resolve(modelsDir) + const extensionRoot = resolve(root, safeExtensionId) + if (pathHasSymlink(root, extensionRoot)) throw new Error('Extension model path resolves through a symlink') + return extensionRoot +} + +export function resolveWeightGroupRoot(modelsDir: string, extensionId: string, groupId: string): string { + const safeExtensionId = safeModelSourceId(extensionId, 'extension id') + if (typeof groupId === 'string' && groupId.toLowerCase() === '_shared') { + throw new Error('Weight group id "_shared" is reserved') + } + const safeGroupId = safeModelSourceId(groupId, 'weight group id') + const root = resolve(modelsDir) + const extensionRoot = resolveExtensionModelRoot(modelsDir, safeExtensionId) + const groupRoot = resolve(extensionRoot, '_shared', safeGroupId) + if (pathHasSymlink(root, groupRoot)) throw new Error('Weight group path resolves through a symlink') + return groupRoot +} + +export function weightGroupTargetId(extensionId: string, groupId: string): string { + const safeExtensionId = safeModelSourceId(extensionId, 'extension id') + const safeGroupId = safeModelSourceId(groupId, 'weight group id') + return `${safeExtensionId}/_shared/${safeGroupId}` +} + +export function resolveWeightStorageRoot(modelsDir: string, targetId: string): string { + if (typeof targetId !== 'string') throw new Error('Weight target id must be a string') + const parts = targetId.split('/') + if (parts.length === 2) return resolveModelRoot(modelsDir, targetId) + if (parts.length === 3 && parts[1] === '_shared') { + return resolveWeightGroupRoot(modelsDir, parts[0], parts[2]) + } + throw new Error('Weight target id must identify one model node or extension weight group') +} + +export function areModelSourcesDownloadedAtRoot(modelRoot: string, sources: ModelSource[]): boolean { try { - const modelRoot = resolveModelRoot(modelsDir, modelId) - if (!existsSync(modelRoot)) return false + if (!existsSync(modelRoot) || pathHasSymlink(modelRoot, modelRoot)) return false return sources.every((source) => { const destination = source.destination === '.' ? modelRoot @@ -187,6 +293,30 @@ export function areModelSourcesDownloaded(modelsDir: string, modelId: string, so } } +export function areModelSourcesDownloaded(modelsDir: string, modelId: string, sources: ModelSource[]): boolean { + try { + const modelRoot = resolveModelRoot(modelsDir, modelId) + return areModelSourcesDownloadedAtRoot(modelRoot, sources) + } catch { + return false + } +} + +export function areWeightGroupSourcesDownloaded( + modelsDir: string, + extensionId: string, + group: ModelWeightGroup, +): boolean { + try { + return areModelSourcesDownloadedAtRoot( + resolveWeightGroupRoot(modelsDir, extensionId, group.id), + group.sources, + ) + } catch { + return false + } +} + export function modelHasLocalData(modelsDir: string, modelId: string): boolean { try { const modelRoot = resolveModelRoot(modelsDir, modelId) @@ -196,6 +326,15 @@ export function modelHasLocalData(modelsDir: string, modelId: string): boolean { } } +export function weightStorageHasLocalData(modelsDir: string, targetId: string): boolean { + try { + const root = resolveWeightStorageRoot(modelsDir, targetId) + return existsSync(root) && readdirSync(root).length > 0 + } catch { + return false + } +} + // Mirrors the backend's cancel cleanup (api/routers/model.py): only the in-progress // `.part` files are removed, so completed sources already on disk survive a cancel. export async function removePartialDownloadArtifacts(modelRoot: string): Promise { diff --git a/electron/preload/electron-api.ts b/electron/preload/electron-api.ts index dae65e69..52eae5a2 100644 --- a/electron/preload/electron-api.ts +++ b/electron/preload/electron-api.ts @@ -119,10 +119,13 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra listDownloaded: () => ipcRenderer.invoke('model:listDownloaded'), isDownloaded: (modelId: string) => ipcRenderer.invoke('model:isDownloaded', modelId), hasLocalData: (modelId: string) => ipcRenderer.invoke('model:hasLocalData', modelId), + sharedGroups: (extensionId: string) => ipcRenderer.invoke('model:sharedGroups', extensionId), download: (modelId: string) => ipcRenderer.invoke('model:download', modelId), pauseDownload: (modelId: string) => ipcRenderer.invoke('model:pauseDownload', modelId), cancelDownload: (modelId: string) => ipcRenderer.invoke('model:cancelDownload', modelId), delete: (modelId: string) => ipcRenderer.invoke('model:delete', modelId), + deleteSharedGroup: (extensionId: string, groupId: string) => ipcRenderer.invoke('model:deleteSharedGroup', extensionId, groupId), + deleteExtensionWeights: (extensionId: string) => ipcRenderer.invoke('model:deleteExtensionWeights', extensionId), unloadAll: () => ipcRenderer.invoke('model:unloadAll'), showInFolder: (modelId: string) => ipcRenderer.invoke('model:showInFolder', modelId), activeDownloads: (): Promise<{ modelId: string; percent: number; file?: string; fileIndex?: number; totalFiles?: number }[]> => diff --git a/src/areas/models/ModelsPage.tsx b/src/areas/models/ModelsPage.tsx index 66f8e5d4..407b7cef 100644 --- a/src/areas/models/ModelsPage.tsx +++ b/src/areas/models/ModelsPage.tsx @@ -1,7 +1,7 @@ import { useEffect, useMemo, useRef, useState } from 'react' import { createPortal } from 'react-dom' import { useExtensionsStore } from '@shared/stores/extensionsStore' -import type { AnyExtension, ModelExtension } from '@shared/types/electron.d' +import type { AnyExtension, ModelExtension, SharedWeightGroupState } from '@shared/types/electron.d' import { deleteModelsThenUninstallExtension, formatModelName } from './utils' import { ExtensionCard } from './components/ExtensionCard' import type { ExtensionNode } from './components/ExtensionCard' @@ -50,6 +50,7 @@ export default function ModelsPage(): JSX.Element { // Model weight state (needed for node install status + uninstall cleanup) const [installedVariantIds, setInstalledVariantIds] = useState([]) const [localDataIds, setLocalDataIds] = useState([]) + const [sharedGroupStates, setSharedGroupStates] = useState>({}) const [downloading, setDownloading] = useState = {} for (const ext of exts) { + sharedStates[ext.id] = await window.electron.model.sharedGroups(ext.id) for (const node of ext.nodes) { if (!nodeHasManagedWeights(node)) continue const fullId = `${ext.id}/${node.id}` @@ -100,6 +103,7 @@ export default function ModelsPage(): JSX.Element { } setInstalledVariantIds(ids) setLocalDataIds(localIds) + setSharedGroupStates(sharedStates) } useEffect(() => { @@ -167,24 +171,23 @@ export default function ModelsPage(): JSX.Element { // ── Node install / download controls ────────────────────────────────────── - function handleInstallNode(node: ExtensionNode, fullId: string) { + async function handleInstallNode(node: ExtensionNode, fullId: string) { if (!nodeHasManagedWeights(node)) return setDownloading((prev) => ({ ...prev, [fullId]: { ...(prev[fullId] ?? { percent: 0 }), paused: false, status: 'Starting…' } })) - window.electron.model.download(fullId).then((result) => { - if (!result.success && !result.paused && !result.cancelled) { - setGhErr(result.error ?? 'Download failed') - setDownloading((prev) => { const n = { ...prev }; delete n[fullId]; return n }) - } - }) + const result = await window.electron.model.download(fullId) + if (!result.success && !result.paused && !result.cancelled) { + setGhErr(result.error ?? 'Download failed') + setDownloading((prev) => { const n = { ...prev }; delete n[fullId]; return n }) + } } - function handleInstallAll(ext: AnyExtension) { + async function handleInstallAll(ext: AnyExtension) { if (ext.type !== 'model') return for (const node of ext.nodes) { if (!nodeHasManagedWeights(node)) continue const fullId = `${ext.id}/${node.id}` if (installedVariantIds.includes(fullId) || downloading[fullId]) continue - handleInstallNode(node, fullId) + await handleInstallNode(node, fullId) } } @@ -205,6 +208,12 @@ export default function ModelsPage(): JSX.Element { refreshInstalledIds(useExtensionsStore.getState().modelExtensions) } + async function handleDeleteSharedGroup(extensionId: string, groupId: string) { + const result = await window.electron.model.deleteSharedGroup(extensionId, groupId) + await refreshInstalledIds(useExtensionsStore.getState().modelExtensions) + return result + } + // ── GitHub extension install ─────────────────────────────────────────────── async function handleGHInstall() { @@ -238,7 +247,10 @@ export default function ModelsPage(): JSX.Element { const ext = allExtensions.find((e) => e.id === extId) if (ext?.type === 'model') { const localModels = ext.nodes.filter((n) => localDataIds.includes(`${extId}/${n.id}`)) - setModelsToDelete(new Set(localModels.map((n) => `${extId}/${n.id}`))) + const hasSharedLocalData = (sharedGroupStates[extId] ?? []).some((group) => group.hasLocalData) + setModelsToDelete(ext.weightGroups?.length && (localModels.length > 0 || hasSharedLocalData) + ? new Set([`${extId}/*`]) + : new Set(localModels.map((n) => `${extId}/${n.id}`))) } else { setModelsToDelete(new Set()) } @@ -249,7 +261,9 @@ export default function ModelsPage(): JSX.Element { const result = await deleteModelsThenUninstallExtension( extId, modelsToDelete, - (modelId) => window.electron.model.delete(modelId), + (modelId) => modelId === `${extId}/*` + ? window.electron.model.deleteExtensionWeights(extId) + : window.electron.model.delete(modelId), uninstallExt, ) if (!result.success) { @@ -630,6 +644,7 @@ export default function ModelsPage(): JSX.Element { installedIds={installedVariantIds} localDataIds={localDataIds} downloading={downloading} + sharedGroups={sharedGroupStates[selectedExt.id] ?? []} loadError={extLoadError(selectedExt)} disabled={isBusy} onInstall={handleInstallNode} @@ -637,6 +652,7 @@ export default function ModelsPage(): JSX.Element { onPauseDownload={handlePauseDownload} onCancelDownload={handleCancelDownload} onUninstallNode={handleUninstallNode} + onDeleteSharedGroup={handleDeleteSharedGroup} onUninstall={(extId) => openUninstallModal(extId)} onRepaired={() => reloadExtensions()} onSynced={() => reloadExtensions()} @@ -650,6 +666,10 @@ export default function ModelsPage(): JSX.Element { const installedModels = ext?.type === 'model' ? ext.nodes.filter((n) => localDataIds.includes(`${uninstallTarget}/${n.id}`)) : [] + const sharedExtensionHasData = ext?.type === 'model' && Boolean(ext.weightGroups?.length) && ( + installedModels.length > 0 + || (sharedGroupStates[uninstallTarget] ?? []).some((group) => group.hasLocalData) + ) return createPortal(
- {installedModels.length > 0 && ( + {sharedExtensionHasData ? ( +
+

+ Also delete downloaded model weights: +

+ +
+ ) : installedModels.length > 0 && (

Also delete downloaded model weights: diff --git a/src/areas/models/components/ExtensionDrawer.tsx b/src/areas/models/components/ExtensionDrawer.tsx index 2f154c71..38cec758 100644 --- a/src/areas/models/components/ExtensionDrawer.tsx +++ b/src/areas/models/components/ExtensionDrawer.tsx @@ -1,5 +1,5 @@ import { useEffect, useState } from 'react' -import type { AnyExtension, ExtensionNode } from '@shared/types/electron.d' +import type { AnyExtension, ExtensionNode, SharedWeightGroupState } from '@shared/types/electron.d' import { useNavStore } from '@shared/stores/navStore' import { DownloadMap, @@ -18,6 +18,7 @@ interface Props { installedIds: string[] localDataIds: string[] downloading: DownloadMap + sharedGroups: SharedWeightGroupState[] loadError?: string disabled?: boolean onInstall: (node: ExtensionNode, fullId: string) => void @@ -25,6 +26,7 @@ interface Props { onPauseDownload: (fullId: string) => void onCancelDownload: (fullId: string) => void onUninstallNode: (fullId: string) => void + onDeleteSharedGroup: (extensionId: string, groupId: string) => Promise<{ success: boolean; error?: string }> onUninstall: (extId: string) => void onRepaired: () => void | Promise onSynced: () => void @@ -32,9 +34,9 @@ interface Props { } export function ExtensionDrawer({ - ext, installedIds, localDataIds, downloading, loadError, disabled, + ext, installedIds, localDataIds, downloading, sharedGroups, loadError, disabled, onInstall, onInstallAll, onPauseDownload, onCancelDownload, - onUninstallNode, onUninstall, onRepaired, onSynced, onClose, + onUninstallNode, onDeleteSharedGroup, onUninstall, onRepaired, onSynced, onClose, }: Props): JSX.Element { const navigate = useNavStore((s) => s.navigate) const [repairing, setRepairing] = useState(false) @@ -85,6 +87,16 @@ export function ExtensionDrawer({ } } + async function handleDeleteSharedGroup(group: SharedWeightGroupState) { + const dependents = group.dependentModelIds.join(', ') + if (!window.confirm( + `Remove shared weights "${group.id}"? The following nodes will become unavailable: ${dependents}`, + )) return + setSyncError(null) + const result = await onDeleteSharedGroup(ext.id, group.id) + if (!result.success) setSyncError(result.error ?? 'Could not remove shared model weights.') + } + const error = syncError ?? repairError ?? loadError return ( @@ -169,6 +181,46 @@ export function ExtensionDrawer({

{/* Nodes */} + {isModel && sharedGroups.length > 0 && ( +
+
+ Shared weights +
+
+ {sharedGroups.map((group) => ( +
+
+
+
{group.id}
+
+ Shared by {group.dependentModelIds.length} node{group.dependentModelIds.length === 1 ? '' : 's'} +
+
+
+ + {group.downloaded ? 'Shared · Installed' : 'Shared · Required'} + + {group.hasLocalData && ( + + )} +
+
+
+ ))} +
+
+ )} +
{isModel ? `Nodes · ${done}/${total} installed` : `Actions · ${total}`} diff --git a/src/areas/models/components/extensionShared.tsx b/src/areas/models/components/extensionShared.tsx index 8bf865e3..49c7fb58 100644 --- a/src/areas/models/components/extensionShared.tsx +++ b/src/areas/models/components/extensionShared.tsx @@ -23,7 +23,7 @@ export type NodeUiState = | { kind: 'installed' } export function nodeHasManagedWeights(node: ExtensionNode): boolean { - return Boolean(node.hfRepo || node.hasModelSources) + return Boolean(node.hfRepo || node.hasModelSources || node.weightGroups?.length) } export function getNodeState( diff --git a/src/shared/types/electron.d.ts b/src/shared/types/electron.d.ts index 840c5e76..95eae548 100644 --- a/src/shared/types/electron.d.ts +++ b/src/shared/types/electron.d.ts @@ -25,6 +25,12 @@ export interface ExtensionNode { hfSkipPrefixes?: string[] hfIncludePrefixes?: string[] hasModelSources?: boolean + weightGroups?: string[] +} + +export interface SharedWeightGroup { + id: string + dependentNodeIds: string[] } export interface ModelExtension { @@ -39,6 +45,7 @@ export interface ModelExtension { source?: string localPath?: string nodes: ExtensionNode[] + weightGroups?: SharedWeightGroup[] /** Folder exists but is not a loadable extension — see manifestError */ corrupted?: boolean /** Why the folder is corrupted: manifest gone, manifest unparseable, or install never completed */ @@ -210,10 +217,13 @@ declare global { activeDownloads: () => Promise<{ modelId: string; percent: number; file?: string; fileIndex?: number; totalFiles?: number }[]> isDownloaded: (modelId: string) => Promise hasLocalData: (modelId: string) => Promise + sharedGroups: (extensionId: string) => Promise download: (modelId: string) => Promise<{ success: boolean; error?: string; paused?: boolean; cancelled?: boolean }> pauseDownload: (modelId: string) => Promise<{ success: boolean; error?: string }> cancelDownload: (modelId: string) => Promise<{ success: boolean; error?: string }> delete: (modelId: string) => Promise<{ success: boolean; error?: string }> + deleteSharedGroup: (extensionId: string, groupId: string) => Promise<{ success: boolean; error?: string }> + deleteExtensionWeights: (extensionId: string) => Promise<{ success: boolean; error?: string }> unloadAll: () => Promise<{ success: boolean; error?: string }> showInFolder: (modelId: string) => Promise onProgress: (cb: (data: { @@ -322,3 +332,11 @@ declare global { } } } + +export interface SharedWeightGroupState { + id: string + targetId: string + dependentModelIds: string[] + downloaded: boolean + hasLocalData: boolean +} From 0e97260880ec483882caa44817a2077421b5efdd Mon Sep 17 00:00:00 2001 From: DrHepa <162889656+DrHepa@users.noreply.github.com> Date: Mon, 14 Sep 2026 18:57:47 +0200 Subject: [PATCH 30/57] fix(models): harden shared weight lifecycle and runtime contracts Reserve physical roots across downloads, cancellation and removal; preserve paused targets and require confirmed runtime shutdown before deleting files. Stop install-all on interrupted work and refresh dependent readiness after partial success. Fix live model paths and explicit generator identity, reject portable node aliases, and exercise the regressions on Windows/Linux with Python 3.11 and 3.12. --- .github/workflows/model-weights.yml | 40 ++++ README.md | 8 +- api/routers/model.py | 23 ++- api/runner.py | 2 + api/services/generator_registry.py | 7 + api/services/generators/base.py | 1 + api/services/model_sources.py | 15 ++ api/tests/test_generator_registry.py | 10 + api/tests/test_model_router.py | 50 ++++- api/tests/test_model_sources.py | 6 + api/tests/test_runner.py | 17 ++ .../main/extension-install-utils.test.mjs | 13 ++ electron/main/extension-install-utils.ts | 4 + electron/main/ipc-handlers.ts | 159 +++++++++------ electron/main/model-download-plan.ts | 4 + electron/main/model-download-preload.test.mjs | 7 +- electron/main/model-sources.ts | 13 ++ electron/main/model-weight-ipc.test.mjs | 190 ++++++++++++++++++ electron/main/model-weight-operations.ts | 35 ++++ electron/preload/electron-api.ts | 2 + src/areas/models/ModelsPage.tsx | 53 +++-- src/areas/models/utils.test.mjs | 36 ++++ src/areas/models/utils.ts | 32 +++ src/shared/types/electron.d.ts | 2 + 24 files changed, 632 insertions(+), 97 deletions(-) create mode 100644 .github/workflows/model-weights.yml create mode 100644 electron/main/model-weight-ipc.test.mjs create mode 100644 electron/main/model-weight-operations.ts diff --git a/.github/workflows/model-weights.yml b/.github/workflows/model-weights.yml new file mode 100644 index 00000000..7de918a3 --- /dev/null +++ b/.github/workflows/model-weights.yml @@ -0,0 +1,40 @@ +name: Model weight regressions + +on: + pull_request: + branches: [dev, main] + paths: + - 'api/**' + - 'electron/**' + - 'src/areas/models/**' + - 'src/shared/types/electron.d.ts' + - 'package*.json' + - '.github/workflows/model-weights.yml' + +permissions: + contents: read + +jobs: + model-weights: + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest, windows-latest] + python: ['3.11', '3.12'] + runs-on: ${{ matrix.os }} + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-node@v4 + with: + node-version: '22' + cache: npm + - uses: actions/setup-python@v5 + with: + python-version: ${{ matrix.python }} + - run: npm ci --ignore-scripts --no-audit --no-fund + - run: python -m pip install fastapi httpx + - name: Download, deletion, manifests and install queue + run: node --test electron/main/model-weight-ipc.test.mjs electron/main/model-sources.test.mjs electron/main/model-download-plan.test.mjs electron/main/model-download-preload.test.mjs electron/main/extension-install-utils.test.mjs src/areas/models/utils.test.mjs + - name: Python runtime and download regressions + working-directory: api + run: python -m unittest tests.test_model_sources tests.test_model_router tests.test_generator_registry tests.test_extension_process tests.test_runner diff --git a/README.md b/README.md index d7bca2e6..4d5aa6b7 100644 --- a/README.md +++ b/README.md @@ -199,8 +199,12 @@ reference them from any sibling node. Shared files are downloaded once under At runtime, `MODEL_DIR` remains the selected node's private directory. Subprocess extensions also receive `MODEL_ID`, `MODEL_NODE_ID`, and a JSON -`SHARED_MODEL_DIRS` map. Direct generators receive the same resolved mapping in -`shared_model_dirs`. Removing private node data never removes a shared group; +`SHARED_MODEL_DIRS` map in their environment. Both direct and subprocess generator +instances receive `MODEL_ID`, `MODEL_NODE_ID`, and the resolved mapping in +`shared_model_dirs` before `load()`. Direct generators use these instance attributes, +not process-global environment variables, to distinguish sibling nodes. +Shared groups are installed through their dependent nodes; the drawer exposes +shared-group status and explicit removal. Removing private node data never removes a shared group; shared-group removal is a separate action that identifies every affected node. --- diff --git a/api/routers/model.py b/api/routers/model.py index 3fd937c0..8f2dde3d 100644 --- a/api/routers/model.py +++ b/api/routers/model.py @@ -10,7 +10,9 @@ from urllib.request import Request, urlopen from fastapi import APIRouter, HTTPException, Request as FastAPIRequest from fastapi.responses import StreamingResponse -from services.generator_registry import generator_registry, MODELS_DIR +from services.generator_registry import generator_registry +import services.generator_registry as registry_module +from services.extension_process import ExtensionProcess from services.model_sources import ( normalize_model_sources, resolve_download_path, @@ -109,10 +111,19 @@ async def unload_model(model_id: str): """Unloads a model from memory so its files can be safely deleted.""" try: gen = generator_registry.get_generator(model_id) + except ValueError as exc: + if model_id in generator_registry._generators: + raise HTTPException(409, str(exc)) from exc + return {"unloaded": True} # No runtime registered for these files. + # unload() on ExtensionProcess deliberately swallows IPC errors; deletion + # needs a confirmed process exit so no worker can retain file handles. + if isinstance(gen, ExtensionProcess): + gen.stop() + else: gen.unload() - return {"unloaded": True} - except ValueError: - return {"unloaded": True} # already not loaded, that's fine + if gen.is_loaded(): + raise HTTPException(409, "Model is still loaded; weights were preserved") + return {"unloaded": True} @router.post("/hf-download/pause") @@ -140,7 +151,7 @@ async def hf_download_sources(request: FastAPIRequest, model_id: str): if raw_sources is None: raise ValueError("sources are required") sources = normalize_model_sources({"model_sources": raw_sources}) - model_root = resolve_weight_storage_root(MODELS_DIR, model_id) + model_root = resolve_weight_storage_root(registry_module.MODELS_DIR, model_id) destinations = { source["id"]: resolve_source_destination_at_root( model_root, source["destination"] @@ -305,7 +316,7 @@ async def hf_download( """ import json as _json import os - dest_dir = str(MODELS_DIR / model_id) + dest_dir = str(registry_module.MODELS_DIR / model_id) # Prefer skip_prefixes passed directly from the client (authoritative, no registry dep) if skip_prefixes: try: diff --git a/api/runner.py b/api/runner.py index 489d12a0..040f4590 100644 --- a/api/runner.py +++ b/api/runner.py @@ -195,6 +195,8 @@ def main() -> None: # Falls back to MODELS_DIR/manifest_id for legacy / standalone use. model_dir = Path(_MODEL_DIR_OVERRIDE) if _MODEL_DIR_OVERRIDE else MODELS_DIR / model_id gen = GenClass(model_dir, WORKSPACE_DIR) + gen.MODEL_ID = model_id + gen.MODEL_NODE_ID = node.get("id", "") gen.shared_model_dirs = dict(_SHARED_MODEL_DIRS) _apply_manifest_metadata(gen, manifest, node) diff --git a/api/services/generator_registry.py b/api/services/generator_registry.py index 4f642e97..1ee4644b 100644 --- a/api/services/generator_registry.py +++ b/api/services/generator_registry.py @@ -30,6 +30,7 @@ normalize_model_sources, normalize_weight_group_references, normalize_weight_groups, + validate_model_node_ids, resolve_weight_group_root, safe_source_id, weight_group_sources_are_downloaded, @@ -466,6 +467,8 @@ def _discover_extensions( uses_shared_weights = weight_groups is not None or any( "weight_groups" in node for node in nodes ) + if uses_shared_weights or any("model_sources" in node for node in nodes): + validate_model_node_ids(raw_nodes) if uses_shared_weights: for node in nodes: raw_node_id = node.get("id") @@ -670,6 +673,8 @@ def initialize( gen.download_check = manifest.get("download_check", "") gen._params_schema = manifest.get("params_schema", []) + gen.MODEL_ID = model_id + gen.MODEL_NODE_ID = manifest.get("node_id", "") gen.shared_model_dirs = { group["id"]: resolve_weight_group_root( MODELS_DIR, manifest.get("ext_id", model_id.split("/", 1)[0]), group["id"] @@ -895,6 +900,8 @@ def unload_all(self) -> None: gen.stop() else: gen.unload() + if gen.is_loaded(): + raise RuntimeError("Model is still loaded; weights were preserved") # Singleton diff --git a/api/services/generators/base.py b/api/services/generators/base.py index c9344538..de1a9f79 100644 --- a/api/services/generators/base.py +++ b/api/services/generators/base.py @@ -78,6 +78,7 @@ class BaseGenerator(ABC): # Metadata — override in each subclass # ------------------------------------------------------------------ # MODEL_ID: str = "" + MODEL_NODE_ID: str = "" DISPLAY_NAME: str = "" VRAM_GB: int = 0 # Minimum recommended VRAM (in GB) diff --git a/api/services/model_sources.py b/api/services/model_sources.py index 81c6c41a..638e1754 100644 --- a/api/services/model_sources.py +++ b/api/services/model_sources.py @@ -158,6 +158,21 @@ def normalize_model_sources( return sources +def validate_model_node_ids(nodes: list[dict[str, Any]]) -> None: + """Managed node roots must be unique on case-insensitive filesystems too.""" + seen: set[str] = set() + for node in nodes: + if not isinstance(node, dict): + raise ValueError("model node must be an object") + if isinstance(node.get("id"), str) and node["id"].casefold() == "_shared": + raise ValueError('model node id "_shared" is reserved') + node_id = safe_source_id(node.get("id"), "model node id") + alias = node_id.casefold() + if alias in seen: + raise ValueError(f'model node id "{node_id}" is not portable-unique') + seen.add(alias) + + def normalize_weight_groups(manifest: dict[str, Any]) -> list[dict[str, Any]] | None: if "weight_groups" not in manifest: return None diff --git a/api/tests/test_generator_registry.py b/api/tests/test_generator_registry.py index 71ffb72a..c8782514 100644 --- a/api/tests/test_generator_registry.py +++ b/api/tests/test_generator_registry.py @@ -261,6 +261,10 @@ def test_shared_groups_gate_all_dependents_and_keep_private_dirs_separate(self) adapter = self.registry.get_generator("shared-model/adapter") self.assertEqual(generate.shared_model_dirs, {"base": base_root}) self.assertEqual(adapter.shared_model_dirs, {"base": base_root}) + self.assertEqual(generate.MODEL_ID, "shared-model/generate") + self.assertEqual(generate.MODEL_NODE_ID, "generate") + self.assertEqual(adapter.MODEL_ID, "shared-model/adapter") + self.assertEqual(adapter.MODEL_NODE_ID, "adapter") self.assertFalse(self.registry._is_downloaded("shared-model/generate", generate)) self.assertFalse(self.registry._is_downloaded("shared-model/adapter", adapter)) @@ -273,6 +277,12 @@ def test_shared_groups_gate_all_dependents_and_keep_private_dirs_separate(self) private_root.mkdir(parents=True) (private_root / "adapter.bin").write_bytes(b"adapter") self.assertTrue(self.registry._is_downloaded("shared-model/adapter", adapter)) + relocated = self.root / "relocated-models" + self.registry.update_paths(relocated, None) + self.assertEqual(adapter.model_dir, relocated / "shared-model/adapter") + self.assertEqual(adapter.shared_model_dirs, {"base": relocated / "shared-model/_shared/base"}) + self.assertEqual(adapter.MODEL_NODE_ID, "adapter") + self.assertFalse(self.registry._is_downloaded("shared-model/adapter", adapter)) def test_reload_preserves_legacy_path_owned_by_the_host(self) -> None: extension = self._make_extension("host-owned-path") diff --git a/api/tests/test_model_router.py b/api/tests/test_model_router.py index 9bd9f71a..5f2c3506 100644 --- a/api/tests/test_model_router.py +++ b/api/tests/test_model_router.py @@ -69,12 +69,12 @@ def setUp(self) -> None: self.tempdir = tempfile.TemporaryDirectory(prefix="modly-model-router-") self.models_dir = Path(self.tempdir.name) / "models" self.models_dir.mkdir() - self.old_models_dir = model_router.MODELS_DIR - model_router.MODELS_DIR = self.models_dir + self.old_models_dir = model_router.registry_module.MODELS_DIR + model_router.registry_module.MODELS_DIR = self.models_dir self.old_hf_module = sys.modules.get("huggingface_hub") def tearDown(self) -> None: - model_router.MODELS_DIR = self.old_models_dir + model_router.registry_module.MODELS_DIR = self.old_models_dir model_router._download_controls.clear() if self.old_hf_module is None: sys.modules.pop("huggingface_hub", None) @@ -214,6 +214,50 @@ async def run(): self.assertIn("excluded from its download plan", events[-1]["error"]) self.assertFalse((self.models_dir / "pixal3d/generate/other.bin").exists()) + def test_download_uses_live_storage_after_settings_update(self): + calls = [] + self.install_hf_stub({"org/main": ["main.bin"]}, calls) + relocated = self.models_dir.parent / "relocated" + registry = model_router.registry_module.GeneratorRegistry() + registry.update_paths(relocated, None) + + def fake_download(**kwargs): + target = Path(kwargs["dest_dir"]) / kwargs["filename"] + target.parent.mkdir(parents=True, exist_ok=True) + target.write_bytes(b"complete") + return target.stat().st_size + + async def run(): + with patch.object(model_router, "_download_file_streamed", fake_download): + for target_id in ("pixal3d/generate", "pixal3d/_shared/base"): + response = await model_router.hf_download_sources(request_for([SOURCES[0]]), target_id) + await collect_events(response) + self.assertTrue((relocated / target_id / "main.bin").is_file()) + self.assertFalse((self.models_dir / target_id).exists()) + asyncio.run(run()) + + def test_removal_unload_requires_confirmed_process_stop(self): + from unittest.mock import Mock + from services.extension_process import ExtensionProcess + gen = Mock(spec=ExtensionProcess) + with patch.object(model_router.generator_registry, "get_generator", return_value=gen): + self.assertEqual(asyncio.run(model_router.unload_model("demo/a")), {"unloaded": True}) + gen.stop.assert_called_once() + gen.unload.assert_not_called() + gen.stop.side_effect = RuntimeError("still running") + with self.assertRaisesRegex(RuntimeError, "still running"): + asyncio.run(model_router.unload_model("demo/a")) + + def test_removal_rejects_direct_generator_that_remains_loaded(self): + from unittest.mock import Mock + from fastapi import HTTPException + gen = Mock() + gen.is_loaded.return_value = True + with patch.object(model_router.generator_registry, "get_generator", return_value=gen): + with self.assertRaises(HTTPException) as error: + asyncio.run(model_router.unload_model("demo/a")) + self.assertEqual(error.exception.status_code, 409) + def test_composite_model_unload_route_uses_path_converter(self) -> None: paths = {route.path for route in model_router.router.routes} self.assertIn("/unload/{model_id:path}", paths) diff --git a/api/tests/test_model_sources.py b/api/tests/test_model_sources.py index 3cbe661e..3ca3b5fe 100644 --- a/api/tests/test_model_sources.py +++ b/api/tests/test_model_sources.py @@ -12,6 +12,7 @@ resolve_weight_group_root, resolve_weight_storage_root, validate_source_file_plan, + validate_model_node_ids, weight_group_sources_are_downloaded, ) @@ -40,6 +41,11 @@ def valid_node() -> dict: class ModelSourcesTests(unittest.TestCase): + def test_managed_node_ids_reject_case_aliases(self): + for ids in (("Fast", "fast"), ("fast", "fast")): + with self.subTest(ids=ids), self.assertRaisesRegex(ValueError, "portable-unique"): + validate_model_node_ids([{"id": node_id} for node_id in ids]) + def test_validates_new_sources_without_reinterpreting_legacy_fields(self) -> None: sources = normalize_model_sources(valid_node()) self.assertEqual([source["id"] for source in sources or []], ["primary", "encoder"]) diff --git a/api/tests/test_runner.py b/api/tests/test_runner.py index a8faeeaf..3560bd4b 100644 --- a/api/tests/test_runner.py +++ b/api/tests/test_runner.py @@ -250,6 +250,23 @@ def run(self, actions: list) -> list: return [json.loads(line) for line in out.getvalue().splitlines() if line.strip()] +class RuntimeIdentityTests(unittest.TestCase): + def test_runner_exposes_selected_identity_independently_of_storage(self): + from unittest.mock import patch + driver = _RunnerDriver(_FAKE_TEXGEN_GENERATOR, "FakeTexGen") + manifest = {"id": "demo-ext", "generator_class": "FakeTexGen", + "nodes": [{"id": "a"}, {"id": "b"}]} + (driver.ext_dir / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + with patch.object(runner, "_MODEL_ID_OVERRIDE", "demo-ext/b"), \ + patch.object(runner, "_MODEL_NODE_ID_OVERRIDE", "b"), \ + patch.object(runner, "_MODEL_DIR_OVERRIDE", str(driver.ext_dir / "unrelated-storage")): + driver.run([]) + gen = driver.generator_module.INSTANCES[0] + self.assertEqual(gen.MODEL_ID, "demo-ext/b") + self.assertEqual(gen.MODEL_NODE_ID, "b") + self.assertEqual(gen.model_dir.name, "unrelated-storage") + + class GeneratorLoadedStateTests(unittest.TestCase): def test_reports_loaded_state(self) -> None: gen = type("Gen", (), {"is_loaded": lambda self: True})() diff --git a/electron/main/extension-install-utils.test.mjs b/electron/main/extension-install-utils.test.mjs index 121b9353..96f4f1e0 100644 --- a/electron/main/extension-install-utils.test.mjs +++ b/electron/main/extension-install-utils.test.mjs @@ -423,3 +423,16 @@ test('incompleteInstallRecoveryAction chooses restore, removal, or no-op', () => backupExists: true, }), 'none') }) + +test('managed model node ids reject portable aliases before installation', () => { + const { validateInstallManifest } = loadModule() + const opts = { hasGeneratorFile: () => true, hasEntryFile: () => true } + for (const ids of [['Fast', 'fast'], ['fast', 'fast']]) { + const manifest = { + id: 'demo', type: 'model', generator_class: 'Generator', + weight_groups: [{ id: 'base', model_sources: [{ id: 'main', provider: 'huggingface', repo_id: 'org/base', destination: '.', checks: ['weights.bin'] }] }], + nodes: ids.map((id) => ({ id, weight_groups: ['base'] })), + } + assert.throws(() => validateInstallManifest(manifest, opts, 'test'), /portable-unique/) + } +}) diff --git a/electron/main/extension-install-utils.ts b/electron/main/extension-install-utils.ts index d50e4acc..4be2085b 100644 --- a/electron/main/extension-install-utils.ts +++ b/electron/main/extension-install-utils.ts @@ -2,6 +2,7 @@ import { normalizeModelSources, normalizeWeightGroupReferences, normalizeWeightGroups, + validateModelNodeIds, safeModelSourceId, type ModelWeightNode, } from './model-sources' @@ -61,6 +62,9 @@ export function validateInstallManifest( throw new Error('manifest.json: weight_groups is supported only for model extensions') } const weightGroups = normalizeWeightGroups(manifest) + if (weightGroups || nodes.some((node) => node.model_sources !== undefined || node.weight_groups !== undefined)) { + validateModelNodeIds(manifest.nodes ?? []) + } for (const node of Array.isArray(manifest.nodes) ? manifest.nodes : []) { const usesSharedWeights = weightGroups !== undefined || node.weight_groups !== undefined if (usesSharedWeights && typeof node.id === 'string' && node.id.toLowerCase() === '_shared') { diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 245cc4f2..5b86e3f1 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -28,6 +28,7 @@ import { normalizeModelSources, normalizeWeightGroupReferences, normalizeWeightGroups, + validateModelNodeIds, removePartialDownloadArtifacts, resolveExtensionModelRoot, resolveModelRoot, @@ -80,6 +81,7 @@ import { } from './extension-install-recovery' import { registerWorkspaceAssetLibraryIpcHandlers } from './artifact-registry-service' import { updatesSupported } from './updater' +import { ModelWeightOperations } from './model-weight-operations' type WindowGetter = () => BrowserWindow | null const pExecFile = promisify(execFile) @@ -167,9 +169,24 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe finish: () => void targetRoots: string[] currentTargetId?: string + stopRequested?: 'pause' | 'cancel' } const activeDownloads = new Map() - const activeWeightTargets = new Map() + const weightOperations = new ModelWeightOperations() + // Paused/error sessions retain their original paths even after settings change. + const interruptedTargets = new Map() + const notifyWeightChange = () => { + getWindow()?.webContents.send('model:weightsChanged') + } + async function unloadForRemoval(modelIds?: string[]) { + const urls = modelIds + ? modelIds.map((id) => `${API_BASE_URL}/model/unload/${encodeURIComponent(id)}`) + : [`${API_BASE_URL}/model/unload-all`] + for (const url of urls) { + const response = await axios.post(url, {}, { timeout: 40_000 }) + if (response.data?.unloaded !== true) throw new Error('Model unload was not confirmed; weights were preserved') + } + } // Logging from renderer ipcMain.on('log:error', (_event, message: string) => logger.error(`[Renderer] ${message}`)) ipcMain.handle('log:getPath', () => join(app.getPath('userData'), 'logs', 'modly.log')) @@ -366,23 +383,16 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe return { success: false, error: String(err) } } - // Unload the model and wait for confirmation so file handles are released try { - await axios.post(`${API_BASE_URL}/model/unload/${encodeURIComponent(modelId)}`, {}, { timeout: 10_000 }) - // Give the OS a moment to release file locks (Windows holds handles briefly after close) - await new Promise(resolve => setTimeout(resolve, 1_500)) - } catch { - // Unload failed (model may not be loaded) — still attempt deletion - } - - // Retry removal — Windows may return EBUSY/EPERM if handles linger - const removed = await rmWithRetry(modelDir, 'model-delete') - if (removed.ok) return { success: true } - return { - success: false, - error: removed.locked - ? 'Model files are still locked after several attempts. Close any programs using the model and try again.' - : String(removed.error), + const removed = await weightOperations.remove( + [modelDir], () => unloadForRemoval([modelId]), () => rmWithRetry(modelDir, 'model-delete'), + ) + notifyWeightChange() + return removed.ok ? { success: true } : { + success: false, error: removed.locked ? 'Model files are still locked. Try again after closing the model.' : String(removed.error), + } + } catch (err) { + return { success: false, error: String(err) } } }) @@ -492,20 +502,12 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe extensionId, group.id, ) - if (activeWeightTargets.has(groupRoot)) { - return { success: false, error: 'Cannot remove shared weights while their download is active' } - } - await Promise.all(group.dependentModelIds.map(async (dependentModelId) => { - try { - await axios.post( - `${API_BASE_URL}/model/unload/${encodeURIComponent(dependentModelId)}`, - {}, - { timeout: 10_000 }, - ) - } catch { /* an unloaded or unavailable model does not block file removal */ } - })) - await new Promise(resolve => setTimeout(resolve, 1_500)) - const removed = await rmWithRetry(groupRoot, 'shared-model-delete') + const removed = await weightOperations.remove( + [groupRoot], + () => unloadForRemoval(group.dependentModelIds), + () => rmWithRetry(groupRoot, 'shared-model-delete'), + ) + notifyWeightChange() if (removed.ok) return { success: true } return { success: false, @@ -531,11 +533,10 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe getSettings(app.getPath('userData')).modelsDir, safeExtensionId, ) - try { - await axios.post(`${API_BASE_URL}/model/unload-all`, {}, { timeout: 10_000 }) - await new Promise(resolve => setTimeout(resolve, 1_500)) - } catch { /* still attempt deletion when the API is unavailable */ } - const removed = await rmWithRetry(extensionRoot, 'extension-model-delete') + const removed = await weightOperations.remove( + [extensionRoot], () => unloadForRemoval(), () => rmWithRetry(extensionRoot, 'extension-model-delete'), + ) + notifyWeightChange() if (removed.ok) return { success: true } return { success: false, @@ -575,7 +576,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe } const modelsDir = getSettings(app.getPath('userData')).modelsDir - const managedTargets = plan.kind === 'multi-source' + const allTargets = plan.kind === 'multi-source' ? [ ...plan.sharedGroups.map((group) => ({ targetId: group.targetId, @@ -587,27 +588,26 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe label: 'Node-specific', sources: plan.sources, }] : []), - ].filter((target) => !areModelSourcesDownloadedAtRoot( - resolveWeightStorageRoot(modelsDir, target.targetId), - target.sources, - )) + ] : [] const targetRoots = plan.kind === 'multi-source' - ? managedTargets.map((target) => resolveWeightStorageRoot(modelsDir, target.targetId)) + ? allTargets.map((target) => resolveWeightStorageRoot(modelsDir, target.targetId)) : [resolveModelRoot(modelsDir, modelId)] - const conflict = targetRoots.find((root) => activeWeightTargets.has(root)) - if (conflict) { - return { - success: false, - error: `Weights are already being downloaded by ${activeWeightTargets.get(conflict)}`, - } + let release: () => void + try { + release = weightOperations.acquire(`downloading ${modelId}`, targetRoots) + } catch (err) { + return { success: false, error: String(err) } } + const managedTargets = allTargets.filter((target) => !areModelSourcesDownloadedAtRoot( + resolveWeightStorageRoot(modelsDir, target.targetId), target.sources, + )) let finish!: () => void const done = new Promise((resolveDone) => { finish = resolveDone }) const active: ActiveDownload = { progress: { percent: 0 }, done, finish, targetRoots } activeDownloads.set(modelId, active) - for (const root of targetRoots) activeWeightTargets.set(root, modelId) + interruptedTargets.set(modelId, targetRoots) try { const onProgress = (progress: typeof active.progress) => { active.progress = progress @@ -618,6 +618,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe onProgress({ percent: 100 }) } for (const [index, target] of managedTargets.entries()) { + if (active.stopRequested) throw new Error(`Model download ${active.stopRequested === 'pause' ? 'paused' : 'cancelled'}`) active.currentTargetId = target.targetId await downloadModelSourcesFromHF(target.targetId, target.sources, (progress) => { const aggregatePercent = Math.min( @@ -630,7 +631,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe status: progress.status ? `${target.label} · ${progress.status}` : target.label, }) }) + notifyWeightChange() } + if (active.stopRequested) throw new Error(`Model download ${active.stopRequested === 'pause' ? 'paused' : 'cancelled'}`) if (managedTargets.length > 0) onProgress({ percent: 100, status: 'done' }) } else { active.currentTargetId = modelId @@ -642,6 +645,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe plan.includePrefixes, ) } + interruptedTargets.delete(modelId) return { success: true } } catch (err: any) { const message = err?.message ?? String(err) @@ -656,17 +660,18 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe return { success: false, error: String(err) } } finally { if (activeDownloads.get(modelId) === active) activeDownloads.delete(modelId) - for (const root of targetRoots) { - if (activeWeightTargets.get(root) === modelId) activeWeightTargets.delete(root) - } + release() active.finish() + notifyWeightChange() } }) ipcMain.handle('model:pauseDownload', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { try { - const targetId = activeDownloads.get(modelId)?.currentTargetId - if (!targetId) return { success: false, error: 'No active download target' } + const active = activeDownloads.get(modelId) + const targetId = active?.currentTargetId + if (!active || !targetId) return { success: false, error: 'No active download target' } + active.stopRequested = 'pause' await axios.post(`${API_BASE_URL}/model/hf-download/pause`, null, { params: { model_id: targetId }, timeout: 5000, @@ -680,6 +685,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe ipcMain.handle('model:cancelDownload', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { try { const active = activeDownloads.get(modelId) + if (active) active.stopRequested = 'cancel' if (active?.currentTargetId) { await axios.post(`${API_BASE_URL}/model/hf-download/cancel`, null, { params: { model_id: active.currentTargetId }, @@ -687,20 +693,30 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }) } if (active) { - await Promise.race([ - active.done, - new Promise((_, reject) => { - setTimeout(() => reject(new Error('Timed out waiting for the download to stop')), 30_000) - }), - ]) + let timer: ReturnType | undefined + try { + await Promise.race([ + active.done, + new Promise((_, reject) => { + timer = setTimeout(() => reject(new Error('Timed out waiting for the download to stop')), 30_000) + }), + ]) + } finally { + clearTimeout(timer) + } + } + const roots = active?.targetRoots ?? interruptedTargets.get(modelId) ?? [resolveModelRoot( + getSettings(app.getPath('userData')).modelsDir, modelId, + )] + // Another sibling may have resumed these targets after our session stopped. + const release = weightOperations.acquire(`cancelling ${modelId}`, roots) + try { + await Promise.all(roots.map((root) => removePartialDownloadArtifacts(root))) + interruptedTargets.delete(modelId) + } finally { + release() + notifyWeightChange() } - // Only remove in-progress `.part` files — a model can now have multiple sources - // sharing this directory, and any source that already finished downloading - // must survive cancelling the ones still in flight. - await Promise.all((active?.targetRoots ?? [resolveModelRoot( - getSettings(app.getPath('userData')).modelsDir, - modelId, - )]).map((root) => removePartialDownloadArtifacts(root))) return { success: true } } catch (err) { return { success: false, error: String(err) } @@ -794,6 +810,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }) ipcMain.handle('settings:set', async (_event, patch: { modelsDir?: string; workspaceDir?: string; extensionsDir?: string; hfToken?: string }) => { + if (patch.modelsDir !== undefined && weightOperations.busy) { + throw new Error('Cannot change model storage while model weights are busy') + } const updated = setSettings(app.getPath('userData'), patch) // Keep main-process env in sync so child processes spawned after token change inherit it if (patch.hfToken !== undefined) { @@ -1042,6 +1061,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe throw new Error('manifest.json: weight_groups is supported only for model extensions') } const weightGroups = normalizeWeightGroups(parsed) + if (weightGroups || parsed.nodes?.some((node) => node.model_sources !== undefined || node.weight_groups !== undefined)) { + validateModelNodeIds(parsed.nodes ?? []) + } const nodes = (parsed.nodes ?? []).map(n => { const usesManagedWeights = weightGroups !== undefined || n.model_sources !== undefined @@ -2000,6 +2022,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // Update FastAPI paths at runtime (without restarting) ipcMain.handle('api:updatePaths', async (_event, patch: { modelsDir?: string; workspaceDir?: string; extensionsDir?: string }) => { try { + if (patch.modelsDir !== undefined && weightOperations.busy) { + throw new Error('Cannot change model storage while model weights are busy') + } await axios.post(`${API_BASE_URL}/settings/paths`, { models_dir: patch.modelsDir, workspace_dir: patch.workspaceDir, diff --git a/electron/main/model-download-plan.ts b/electron/main/model-download-plan.ts index 8b17f694..fe786420 100644 --- a/electron/main/model-download-plan.ts +++ b/electron/main/model-download-plan.ts @@ -12,6 +12,7 @@ import { normalizeModelSources, normalizeWeightGroupReferences, normalizeWeightGroups, + validateModelNodeIds, safeModelSourceId, weightGroupTargetId, type ModelSource, @@ -76,6 +77,9 @@ function installedSharedGroups( nodes: InstalledNode[], ): InstalledSharedWeightGroup[] { const groups = normalizeWeightGroups(manifest) + if (groups || nodes.some((node) => node.model_sources !== undefined || node.weight_groups !== undefined)) { + validateModelNodeIds(manifest.nodes as InstalledNode[]) + } if (groups === undefined) return [] const groupDependents = new Map() for (const candidate of nodes) { diff --git a/electron/main/model-download-preload.test.mjs b/electron/main/model-download-preload.test.mjs index aa414479..1fc9ebd3 100644 --- a/electron/main/model-download-preload.test.mjs +++ b/electron/main/model-download-preload.test.mjs @@ -48,15 +48,10 @@ test('renderer model actions keep node and shared-weight identities explicit', a ]) }) -test('declared partial data is removable and active downloads block destructive actions', () => { - const main = readFileSync(resolve('electron/main/ipc-handlers.ts'), 'utf8') +test('UI exposes partial-data removal and dependent-node warnings', () => { const page = readFileSync(resolve('src/areas/models/ModelsPage.tsx'), 'utf8') const drawer = readFileSync(resolve('src/areas/models/components/ExtensionDrawer.tsx'), 'utf8') - assert.match(main, /model:delete[\s\S]*activeDownloads\.has\(modelId\)/) - assert.match(main, /model:deleteSharedGroup[\s\S]*activeWeightTargets\.has\(groupRoot\)/) - assert.match(main, /model:deleteExtensionWeights[\s\S]*resolveExtensionModelRoot/) - assert.match(main, /extensions:uninstall[\s\S]*activeDownloads\.keys\(\)/) assert.match(page, /window\.electron\.model\.hasLocalData\(fullId\)/) assert.match(page, /deleteExtensionWeights\(extId\)/) assert.match(drawer, /localDataIds\.includes\(fullId\) && state\.kind !== 'downloading'/) diff --git a/electron/main/model-sources.ts b/electron/main/model-sources.ts index cbc69e42..0cd3274d 100644 --- a/electron/main/model-sources.ts +++ b/electron/main/model-sources.ts @@ -149,6 +149,19 @@ export function normalizeModelSources( }) } +export function validateModelNodeIds(nodes: Array<{ id?: unknown }>): void { + const seen = new Set() + for (const node of nodes) { + if (typeof node?.id === 'string' && node.id.toLowerCase() === '_shared') { + throw new Error('model node id "_shared" is reserved') + } + const id = safeModelSourceId(node?.id, 'model node id') + const alias = id.toLowerCase() + if (seen.has(alias)) throw new Error(`model node id "${id}" is not portable-unique`) + seen.add(alias) + } +} + export function normalizeWeightGroups(manifest: ModelWeightManifest): ModelWeightGroup[] | undefined { if (!Object.prototype.hasOwnProperty.call(manifest, 'weight_groups')) return undefined if (!Array.isArray(manifest.weight_groups) || manifest.weight_groups.length === 0) { diff --git a/electron/main/model-weight-ipc.test.mjs b/electron/main/model-weight-ipc.test.mjs new file mode 100644 index 00000000..d109a2d0 --- /dev/null +++ b/electron/main/model-weight-ipc.test.mjs @@ -0,0 +1,190 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync, transformSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, writeFileSync, readFileSync, existsSync, rmSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' +import vm from 'node:vm' + +const require = createRequire(import.meta.url) +function moduleFromCode(code, dependencies = require) { + const module = { exports: {} } + vm.runInNewContext(code, { module, exports: module.exports, require: dependencies, process, console, Buffer, setTimeout, clearTimeout }) + return module.exports +} +function loadModule(path) { + return moduleFromCode(buildSync({ entryPoints: [resolve(path)], bundle: true, platform: 'node', format: 'cjs', write: false }).outputFiles[0].text) +} +const realSources = loadModule('electron/main/model-sources.ts') +const realPlan = loadModule('electron/main/model-download-plan.ts') +const realGuard = loadModule('electron/main/extension-path-guard.ts') +const realOperations = loadModule('electron/main/model-weight-operations.ts') +const ipcCode = transformSync(readFileSync('electron/main/ipc-handlers.ts', 'utf8'), { loader: 'ts', format: 'cjs' }).code +const source = { id: 'weights', provider: 'huggingface', repo_id: 'org/base', destination: '.', checks: ['weights.bin'] } +function deferred() { + let resolve + const promise = new Promise((r) => { resolve = r }) + return { promise, resolve } +} +function fixture(t) { + const root = mkdtempSync(join(tmpdir(), 'modly-weight-ipc-')) + t.after(() => rmSync(root, { recursive: true, force: true })) + const settings = { modelsDir: join(root, 'models'), extensionsDir: join(root, 'extensions') } + const extension = join(settings.extensionsDir, 'demo') + mkdirSync(extension, { recursive: true }) + writeFileSync(join(extension, 'manifest.json'), JSON.stringify({ + id: 'demo', type: 'model', weight_groups: [{ id: 'base', model_sources: [source] }], + nodes: [{ id: 'a', weight_groups: ['base'] }, { id: 'b', weight_groups: ['base'], model_sources: [{ ...source, repo_id: 'org/adapter' }] }], + })) + const handlers = new Map(), events = [], removed = [], calls = [] + const hooks = { + unload: async () => ({ data: { unloaded: true } }), + download: async (id) => { + const dir = realSources.resolveWeightStorageRoot(settings.modelsDir, id) + mkdirSync(dir, { recursive: true }) + writeFileSync(join(dir, 'weights.bin'), 'complete') + }, + } + const stub = new Proxy({}, { get: () => () => {} }) + const deps = (name) => { + if (name === 'electron') return { ipcMain: { handle: (id, fn) => handlers.set(id, fn), on: () => {} }, app: { getPath: () => root, on: () => {} } } + if (name === 'axios') return { post: (...args) => hooks.unload(...args) } + if (name === './model-sources') return realSources + if (name === './model-download-plan') return realPlan + if (name === './extension-path-guard') return realGuard + if (name === './model-weight-operations') return realOperations + if (name === './settings-store') return { getSettings: () => settings, setSettings: (_, patch) => Object.assign(settings, patch) } + if (name === './builtin-sync') return { getBuiltinExtensionsDir: () => join(root, 'builtin') } + if (name === './python-bridge') return { API_BASE_URL: 'http://test' } + if (name === './model-downloader') return { + downloadModelSourcesFromHF: async (...args) => { calls.push(args[0]); await hooks.download(...args) }, + } + if (name === './extension-install-recovery') return { rmWithRetry: async (dir) => { removed.push(dir); rmSync(dir, { recursive: true, force: true }); return { ok: true } } } + if (name.startsWith('./') || name === 'electron-updater') return stub + return require(name) + } + moduleFromCode(ipcCode, deps).setupIpcHandlers({}, () => ({ webContents: { send: (...event) => events.push(event) } })) + const event = { sender: { send: (...args) => events.push(args) } } + const invoke = (name, ...args) => handlers.get(`model:${name}`)(event, ...args) + return { settings, hooks, invoke, removed, events, calls, handlers, root } +} + +for (const action of ['deleteSharedGroup', 'deleteExtensionWeights']) { + test(`${action} reserves its root before awaiting unload and releases it afterwards`, async (t) => { + const f = fixture(t), entered = deferred(), finish = deferred() + f.hooks.unload = async () => { entered.resolve(); await finish.promise; return { data: { unloaded: true } } } + const deleting = f.invoke(action, 'demo', 'base') + await entered.promise + const blocked = await f.invoke('download', 'demo/a') + assert.equal(blocked.success, false) + assert.match(blocked.error, /busy/) + assert.equal(f.calls.length, 0) + finish.resolve() + assert.equal((await deleting).success, true) + assert.equal((await f.invoke('download', 'demo/a')).success, true) + }) + test(`${action} preserves files when unload fails or is not confirmed`, async (t) => { + const f = fixture(t) + for (const unload of [async () => { throw new Error('timeout') }, async () => ({ data: {} })]) { + f.hooks.unload = unload + assert.equal((await f.invoke(action, 'demo', 'base')).success, false) + assert.equal(f.removed.length, 0) + } + // Failure must release the lease as well. + assert.equal((await f.invoke('download', 'demo/a')).success, true) + }) +} + +test('active download blocks deletion, including a complete base needed by a private adapter', async (t) => { + const f = fixture(t) + await f.invoke('download', 'demo/a') + const entered = deferred(), finish = deferred() + f.hooks.download = async () => { entered.resolve(); await finish.promise } + const downloading = f.invoke('download', 'demo/b') + await entered.promise + assert.equal((await f.invoke('deleteSharedGroup', 'demo', 'base')).success, false) + assert.equal(f.removed.length, 0) + finish.resolve() + await downloading +}) + +test('paused shared cancellation cleans the original target and preserves completed files', async (t) => { + const f = fixture(t) + const groupDir = join(f.settings.modelsDir, 'demo', '_shared', 'base') + f.hooks.download = async () => { + mkdirSync(groupDir, { recursive: true }) + writeFileSync(join(groupDir, 'weights.bin.part'), 'partial') + writeFileSync(join(groupDir, 'finished.bin'), 'complete') + throw new Error('Model download paused') + } + assert.equal((await f.invoke('download', 'demo/a')).paused, true) + f.settings.modelsDir = join(f.root, 'new-models') + assert.equal((await f.invoke('cancelDownload', 'demo/a')).success, true) + assert.equal(existsSync(join(groupDir, 'weights.bin.part')), false) + assert.equal(existsSync(join(groupDir, 'finished.bin')), true) +}) + +test('paused cleanup cannot delete a sibling download partials', async (t) => { + const f = fixture(t) + f.hooks.download = async () => { throw new Error('Model download paused') } + await f.invoke('download', 'demo/a') + const entered = deferred(), finish = deferred() + f.hooks.download = async () => { entered.resolve(); await finish.promise } + const sibling = f.invoke('download', 'demo/b') + await entered.promise + const result = await f.invoke('cancelDownload', 'demo/a') + assert.equal(result.success, false) + assert.match(result.error, /busy/) + finish.resolve() + await sibling +}) + +test('private failure notifies readiness changes and preserves the usable shared base', async (t) => { + const f = fixture(t), normalDownload = f.hooks.download + f.hooks.download = async (id) => { + if (id === 'demo/b') throw new Error('adapter failed') + await normalDownload(id) + } + const result = await f.invoke('download', 'demo/b') + assert.equal(result.success, false) + assert.equal(await f.invoke('isDownloaded', 'demo/a'), true) + assert.equal(await f.invoke('isDownloaded', 'demo/b'), false) + assert.ok(f.events.filter(([name]) => name === 'model:weightsChanged').length >= 2) +}) + +test('physical locks reject case aliases and parent roots without blocking unrelated roots', () => { + const locks = new realOperations.ModelWeightOperations() + const release = locks.acquire('node', [resolve('Models/Demo/Fast')]) + assert.throws(() => locks.acquire('alias', [resolve('models/demo/fast')]), /busy/) + assert.throws(() => locks.acquire('extension', [resolve('models/demo')]), /busy/) + locks.acquire('sibling', [resolve('models/demo/other')])() + release() + assert.equal(locks.busy, false) +}) + +for (const action of ['pauseDownload', 'cancelDownload']) { + test(`${action} at a target boundary does not start the private source`, async (t) => { + const f = fixture(t), entered = deferred(), finish = deferred() + f.hooks.download = async () => { entered.resolve(); await finish.promise } + const downloading = f.invoke('download', 'demo/b') + await entered.promise + const stopping = f.invoke(action, 'demo/b') + finish.resolve() + assert.equal((await stopping).success, true) + const result = await downloading + assert.equal(result[action === 'pauseDownload' ? 'paused' : 'cancelled'], true) + assert.deepEqual(f.calls, ['demo/_shared/base']) + }) +} + +test('model path changes are rejected while weights are reserved', async (t) => { + const f = fixture(t), entered = deferred(), finish = deferred() + f.hooks.download = async () => { entered.resolve(); await finish.promise } + const downloading = f.invoke('download', 'demo/a') + await entered.promise + await assert.rejects(f.handlers.get('settings:set')(null, { modelsDir: join(f.root, 'new') }), /busy/) + assert.equal((await f.handlers.get('api:updatePaths')(null, { modelsDir: join(f.root, 'new') })).success, false) + finish.resolve() + await downloading +}) diff --git a/electron/main/model-weight-operations.ts b/electron/main/model-weight-operations.ts new file mode 100644 index 00000000..7ec92181 --- /dev/null +++ b/electron/main/model-weight-operations.ts @@ -0,0 +1,35 @@ +import { resolve, sep } from 'node:path' + +/** Atomic in-process reservations, including ancestor/descendant conflicts. + * Fold case conservatively so Windows aliases cannot obtain separate leases. + */ +export class ModelWeightOperations { + private leases = new Map() + + get busy(): boolean { return this.leases.size > 0 } + + acquire(owner: string, roots: string[]): () => void { + const canonical = roots.map((root) => resolve(root).normalize('NFC').toLowerCase()) + const contains = (parent: string, child: string) => ( + child === parent || child.startsWith(parent.endsWith(sep) ? parent : parent + sep) + ) + for (const lease of this.leases.values()) { + if (canonical.some((root) => lease.roots.some((other) => contains(root, other) || contains(other, root)))) { + throw new Error(`Model weights are busy: ${lease.owner}`) + } + } + const token = Symbol(owner) + this.leases.set(token, { owner, roots: canonical }) + return () => { this.leases.delete(token) } + } + + async remove(roots: string[], unload: () => Promise, remove: () => Promise): Promise { + const release = this.acquire('removing model weights', roots) + try { + await unload() + return await remove() + } finally { + release() + } + } +} diff --git a/electron/preload/electron-api.ts b/electron/preload/electron-api.ts index 52eae5a2..42d24835 100644 --- a/electron/preload/electron-api.ts +++ b/electron/preload/electron-api.ts @@ -158,6 +158,8 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra })) }, offProgress: () => ipcRenderer.removeAllListeners('model:downloadProgress'), + onWeightsChanged: (cb: () => void) => ipcRenderer.on('model:weightsChanged', cb), + offWeightsChanged: () => ipcRenderer.removeAllListeners('model:weightsChanged'), }, // App metadata diff --git a/src/areas/models/ModelsPage.tsx b/src/areas/models/ModelsPage.tsx index 407b7cef..f0bee1e5 100644 --- a/src/areas/models/ModelsPage.tsx +++ b/src/areas/models/ModelsPage.tsx @@ -2,7 +2,7 @@ import { useEffect, useMemo, useRef, useState } from 'react' import { createPortal } from 'react-dom' import { useExtensionsStore } from '@shared/stores/extensionsStore' import type { AnyExtension, ModelExtension, SharedWeightGroupState } from '@shared/types/electron.d' -import { deleteModelsThenUninstallExtension, formatModelName } from './utils' +import { deleteModelsThenUninstallExtension, formatModelName, installModelAndRefresh, installModelQueue } from './utils' import { ExtensionCard } from './components/ExtensionCard' import type { ExtensionNode } from './components/ExtensionCard' import { ExtensionDrawer } from './components/ExtensionDrawer' @@ -81,10 +81,14 @@ export default function ModelsPage(): JSX.Element { const [ghUrl, setGhUrl] = useState('') const [ghErr, setGhErr] = useState(null) + const installedRefreshRevision = useRef(0) + const installQueues = useRef(new Set()) + // ── Init ────────────────────────────────────────────────────────────────── // Check each model node individually via filesystem IPC — reliable regardless of API state async function refreshInstalledIds(exts: ModelExtension[]) { + const revision = ++installedRefreshRevision.current const ids: string[] = [] const localIds: string[] = [] const sharedStates: Record = {} @@ -101,6 +105,7 @@ export default function ModelsPage(): JSX.Element { if (hasLocalData) localIds.push(fullId) } } + if (revision !== installedRefreshRevision.current) return setInstalledVariantIds(ids) setLocalDataIds(localIds) setSharedGroupStates(sharedStates) @@ -119,6 +124,9 @@ export default function ModelsPage(): JSX.Element { } refreshInstalledIds(exts) }) + window.electron.model.onWeightsChanged(() => { + void refreshInstalledIds(useExtensionsStore.getState().modelExtensions) + }) window.electron.model.onProgress(({ modelId: id, percent, file, fileIndex, totalFiles, status, bytesDownloaded, totalBytes, stalledSeconds, paused, cancelled }) => { if (cancelled) { setDownloading((prev) => { const n = { ...prev }; delete n[id]; return n }) @@ -148,7 +156,10 @@ export default function ModelsPage(): JSX.Element { }) } }) - return () => window.electron.model.offProgress() + return () => { + window.electron.model.offProgress() + window.electron.model.offWeightsChanged() + } // eslint-disable-next-line react-hooks/exhaustive-deps -- register the progress listener once on mount }, []) @@ -172,22 +183,38 @@ export default function ModelsPage(): JSX.Element { // ── Node install / download controls ────────────────────────────────────── async function handleInstallNode(node: ExtensionNode, fullId: string) { - if (!nodeHasManagedWeights(node)) return + if (!nodeHasManagedWeights(node)) return { success: true } setDownloading((prev) => ({ ...prev, [fullId]: { ...(prev[fullId] ?? { percent: 0 }), paused: false, status: 'Starting…' } })) - const result = await window.electron.model.download(fullId) - if (!result.success && !result.paused && !result.cancelled) { - setGhErr(result.error ?? 'Download failed') - setDownloading((prev) => { const n = { ...prev }; delete n[fullId]; return n }) + try { + const result = await installModelAndRefresh( + () => window.electron.model.download(fullId), + () => refreshInstalledIds(useExtensionsStore.getState().modelExtensions), + ) + if (!result.success && !result.paused && !result.cancelled) { + setGhErr(result.error ?? 'Download failed') + } + if (!result.paused) setDownloading((prev) => { const next = { ...prev }; delete next[fullId]; return next }) + return result + } catch (err) { + const error = String(err) + setGhErr(error) + setDownloading((prev) => { const next = { ...prev }; delete next[fullId]; return next }) + return { success: false, error } } } async function handleInstallAll(ext: AnyExtension) { - if (ext.type !== 'model') return - for (const node of ext.nodes) { - if (!nodeHasManagedWeights(node)) continue - const fullId = `${ext.id}/${node.id}` - if (installedVariantIds.includes(fullId) || downloading[fullId]) continue - await handleInstallNode(node, fullId) + if (ext.type !== 'model' || installQueues.current.has(ext.id)) return + installQueues.current.add(ext.id) + try { + const nodes = new Map(ext.nodes.filter(nodeHasManagedWeights).map((node) => [`${ext.id}/${node.id}`, node])) + await installModelQueue( + nodes.keys(), + (id) => window.electron.model.isDownloaded(id), + (id) => handleInstallNode(nodes.get(id)!, id), + ) + } finally { + installQueues.current.delete(ext.id) } } diff --git a/src/areas/models/utils.test.mjs b/src/areas/models/utils.test.mjs index 2dec654f..531cfe92 100644 --- a/src/areas/models/utils.test.mjs +++ b/src/areas/models/utils.test.mjs @@ -114,3 +114,39 @@ test('failed selected-weight deletion aborts extension uninstall and preserves i assert.equal(uninstallCalls, 0) assert.deepEqual(result, { success: false, error: 'Model weights are locked.' }) }) + +for (const result of [{ success: false, paused: true }, { success: false, cancelled: true }, { success: false, error: 'failed' }]) { + test(`install queue stops after ${JSON.stringify(result)} without reactivating a shared sibling`, async () => { + const { installModelQueue } = loadModule() + const installed = [] + const actual = await installModelQueue(['a', 'b'], async () => false, async (id) => { + installed.push(id) + return result + }) + assert.deepEqual(installed, ['a']) + assert.deepEqual(actual, result) + }) +} + +test('install queue rechecks readiness so completed shared siblings are not downloaded twice', async () => { + const { installModelQueue } = loadModule() + const ready = new Set(), calls = [] + await installModelQueue(['a', 'b'], async (id) => ready.has(id), async (id) => { + calls.push(id) + ready.add('a'); ready.add('b') + return { success: true } + }) + assert.deepEqual(calls, ['a']) +}) + +test('partial success is refreshed on failure, pause, cancel and rejected IPC', async () => { + const { installModelAndRefresh } = loadModule() + for (const result of [{ success: false, error: 'adapter' }, { success: false, paused: true }, { success: false, cancelled: true }]) { + let refreshed = false + await installModelAndRefresh(async () => result, async () => { refreshed = true }) + assert.equal(refreshed, true) + } + let refreshed = false + await assert.rejects(installModelAndRefresh(async () => { throw new Error('IPC') }, async () => { refreshed = true }), /IPC/) + assert.equal(refreshed, true) +}) diff --git a/src/areas/models/utils.ts b/src/areas/models/utils.ts index 2b1aac53..33c3cf4d 100644 --- a/src/areas/models/utils.ts +++ b/src/areas/models/utils.ts @@ -45,3 +45,35 @@ export async function deleteModelsThenUninstallExtension( return uninstallExtension(extensionId) } + +export interface ModelInstallResult { + success: boolean + error?: string + paused?: boolean + cancelled?: boolean +} + +export async function installModelAndRefresh( + install: () => Promise, + refresh: () => Promise, +): Promise { + try { + return await install() + } finally { + // A failed private source can leave a completed base usable by siblings. + await refresh() + } +} + +export async function installModelQueue( + modelIds: Iterable, + isReady: (id: string) => Promise, + install: (id: string) => Promise, +): Promise { + for (const id of modelIds) { + if (await isReady(id)) continue + const result = await install(id) + if (!result.success) return result + } + return { success: true } +} diff --git a/src/shared/types/electron.d.ts b/src/shared/types/electron.d.ts index 95eae548..ddf02fe1 100644 --- a/src/shared/types/electron.d.ts +++ b/src/shared/types/electron.d.ts @@ -240,6 +240,8 @@ declare global { cancelled?: boolean }) => void) => void offProgress: () => void + onWeightsChanged: (cb: () => void) => void + offWeightsChanged: () => void } app: { info: () => Promise<{ From 2dd0ce6160403f1fb60b746a6f668b80abaefd76 Mon Sep 17 00:00:00 2001 From: Lorchie Date: Fri, 18 Sep 2026 10:55:50 +0200 Subject: [PATCH 31/57] feat(models): install and remove weight variants per node A model node that publishes the same weights in several variants (quantizations, precisions) can declare `weight_variants` next to `hf_repo`. Each variant is downloaded and deleted on its own from the Extensions drawer, and the params_schema select that picks one labels the variants that are missing from disk. - manifest validation mirrored in electron/main/model-sources.ts and api/services/model_sources.py: safe relative paths, no overlap between two variants, download_check outside every variant, no combination with model_sources, and the selecting param must exist and offer every id. - every install downloads the shared files first (all variants excluded), then the requested variant, or the default one; files already complete on disk are skipped, so a second variant only fetches its own weights. - new IPC: model:installedWeightVariants (null when unreadable) and model:deleteWeightVariant, which unloads the model before removal and touches only the files of that variant. - generation refuses a variant that is not installed (GeneratorRegistry.assert_weight_variant_installed) and picking one in a node opens the extension drawer. --- README.md | 63 +++++++ api/routers/generation.py | 1 + api/services/generator_registry.py | 24 ++- api/services/model_sources.py | 169 +++++++++++++++++ api/tests/test_generation_router.py | 20 ++ api/tests/test_generator_registry.py | 55 ++++++ api/tests/test_model_sources.py | 68 +++++++ .../main/extension-install-utils.test.mjs | 26 +++ electron/main/extension-install-utils.ts | 11 +- electron/main/ipc-handlers.ts | 162 +++++++++++------ electron/main/model-download-plan.test.mjs | 52 ++++++ electron/main/model-download-plan.ts | 49 ++++- electron/main/model-sources.test.mjs | 84 +++++++++ electron/main/model-sources.ts | 172 +++++++++++++++++- electron/preload/electron-api.ts | 6 +- .../generate/components/WorkflowPanel.tsx | 10 +- src/areas/models/ModelsPage.tsx | 51 +++--- src/areas/models/components/ExtensionCard.tsx | 7 +- .../models/components/ExtensionDrawer.tsx | 98 ++++++++-- .../models/components/extensionShared.tsx | 1 + src/areas/workflows/mockExtensions.ts | 4 +- src/areas/workflows/nodes/ExtensionNode.tsx | 14 +- src/shared/stores/extensionsStore.ts | 23 +++ src/shared/stores/navStore.ts | 8 +- src/shared/types/electron.d.ts | 15 +- src/shared/utils/weightVariants.test.mjs | 63 +++++++ src/shared/utils/weightVariants.ts | 35 ++++ 27 files changed, 1187 insertions(+), 104 deletions(-) create mode 100644 src/shared/utils/weightVariants.test.mjs create mode 100644 src/shared/utils/weightVariants.ts diff --git a/README.md b/README.md index b162cf23..c7a9d7b7 100644 --- a/README.md +++ b/README.md @@ -148,6 +148,69 @@ supported provider is `huggingface`. Existing nodes that use `hf_repo`, `download_check`, `hf_include_prefixes`, and `hf_skip_prefixes` keep their original behavior. +### Separately installable weight variants + +A model node that publishes the same weights in several variants (quantizations, +precisions…) can declare `weight_variants` next to `hf_repo`. The Extensions page +lists every variant under the node, and each one is downloaded or deleted on its own. + +```json +{ + "id": "generate", + "hf_repo": "org/model-gguf", + "download_check": "pipeline.json", + "hf_include_prefixes": ["pipeline.json", "encoder/", "dit/"], + "params_schema": [ + { + "id": "quant", + "label": "Quantization", + "type": "select", + "default": "Q5_K_M", + "options": [ + { "value": "Q4_K_M", "label": "Q4_K_M" }, + { "value": "Q5_K_M", "label": "Q5_K_M" } + ] + } + ], + "weight_variants": { + "param": "quant", + "default": "Q5_K_M", + "options": [ + { + "id": "Q4_K_M", + "label": "Q4_K_M", + "size_gb": 2.4, + "vram_gb": 6, + "include_prefixes": ["dit/model_Q4_K_M.gguf"], + "checks": ["dit/model_Q4_K_M.gguf"] + }, + { + "id": "Q5_K_M", + "include_prefixes": ["dit/model_Q5_K_M.gguf"], + "checks": ["dit/model_Q5_K_M.gguf"] + } + ] + } +} +``` + +- `param` names the `params_schema` select whose values are the variant ids. That + param must exist on the node (or on the extension, as its fallback), and when it + declares `options` they must cover every variant id. +- `size_gb` (download size) and `vram_gb` (approximate VRAM the variant needs) are + optional positive numbers, shown next to the variant when present. +- Every install downloads the shared files (`hf_include_prefixes`, with every + variant's files excluded automatically) plus one variant: the one asked for, or the + `default` one — the first option when `default` is omitted. Files already complete + on disk are skipped, so adding a second variant only fetches that variant. +- A variant is installed when all of its `checks` exist; the node is installed once + its `download_check` and at least one variant are present. +- Generation fails with an explicit message when the selected variant is not + installed, and the node's selector labels those options `(not installed)`. +- `include_prefixes` and `checks` are safe POSIX paths relative to the node's model + directory. Prefixes of two variants cannot overlap, `download_check` stays outside + every variant, and `weight_variants` cannot be combined with `model_sources`. + --- ## Workflows diff --git a/api/routers/generation.py b/api/routers/generation.py index 8481deb4..b4566036 100644 --- a/api/routers/generation.py +++ b/api/routers/generation.py @@ -185,6 +185,7 @@ def progress_cb(pct: int, step: str = "") -> None: try: loop = asyncio.get_running_loop() + generator_registry.assert_weight_variant_installed(params) # Check if the model needs to be loaded BEFORE calling get_active(), # because get_active() loads the model in a blocking manner. diff --git a/api/services/generator_registry.py b/api/services/generator_registry.py index 348a42cb..453a98d1 100644 --- a/api/services/generator_registry.py +++ b/api/services/generator_registry.py @@ -25,7 +25,12 @@ from services.generators.base import BaseGenerator from services.extension_process import ExtensionProcess, _venv_python -from services.model_sources import model_sources_are_downloaded, normalize_model_sources +from services.model_sources import ( + missing_weight_variant, + model_sources_are_downloaded, + normalize_model_sources, + normalize_weight_variants, +) # ------------------------------------------------------------------ # # Global paths @@ -528,6 +533,9 @@ def _discover_extensions( if nodes: for node in nodes: model_sources = normalize_model_sources(node) + weight_variants = normalize_weight_variants( + node, node.get("params_schema", manifest.get("params_schema", [])) + ) node_manifest = { **manifest, "id": f"{ext_id}/{node['id']}", @@ -544,6 +552,8 @@ def _discover_extensions( } if model_sources is not None: node_manifest["model_sources"] = model_sources + if weight_variants is not None: + node_manifest["weight_variants"] = weight_variants full_id = f"{ext_id}/{node['id']}" result[full_id] = (cls_or_None, node_manifest, ext_dir, legacy_context) if subprocess_mode: @@ -724,6 +734,18 @@ def get_active(self) -> BaseGenerator: gen.load() return gen + def assert_weight_variant_installed(self, params: dict) -> None: + """Refuse generation when the weight variant selected by params is not installed.""" + manifest = self._manifests.get(self._active_id, {}) + option = missing_weight_variant( + MODELS_DIR, self._active_id, manifest.get("weight_variants"), params + ) + if option is not None: + raise RuntimeError( + f'{option["label"]} weights for {self._active_id} are not installed. ' + "Install them from the Extensions page, or select an installed variant." + ) + def get_generator(self, model_id: str) -> BaseGenerator: self._assert_not_quarantined(model_id) if model_id not in self._generators: diff --git a/api/services/model_sources.py b/api/services/model_sources.py index 592d1342..1dee2711 100644 --- a/api/services/model_sources.py +++ b/api/services/model_sources.py @@ -2,6 +2,7 @@ from __future__ import annotations +import math import re import unicodedata from pathlib import Path @@ -270,3 +271,171 @@ def validate_source_file_plan( f'"{previous_source}:{previous_target}" and "{source_id}:{value}"' ) aliases[alias] = (source_id, value) + + +def _prefixes_overlap(left: list[str], right: list[str]) -> bool: + return any( + a.lower().startswith(b.lower()) or b.lower().startswith(a.lower()) + for a in left + for b in right + ) + + +def _assert_param_offers_variants(param: str, params_schema: Any, ids: list[str]) -> None: + """Variant ids must be selectable, so the param they key has to exist and offer them.""" + entry = next( + (p for p in params_schema if isinstance(p, dict) and p.get("id") == param), + None, + ) if isinstance(params_schema, list) else None + if entry is None: + raise ValueError(f'weight_variants.param must name a params_schema entry ("{param}")') + options = entry.get("options") + if not isinstance(options, list): + return + values = { + str(option.get("value")) if isinstance(option, dict) else str(option) + for option in options + } + missing = [variant_id for variant_id in ids if variant_id not in values] + if missing: + raise ValueError( + f'the "{param}" param must offer every weight variant id ' + f'(missing: {", ".join(missing)})' + ) + + +def normalize_weight_variants( + node: dict[str, Any], params_schema: Any = None +) -> dict[str, Any] | None: + """Validate a node's separately installable weight variants (e.g. quantizations).""" + if "weight_variants" not in node: + return None + if "model_sources" in node: + raise ValueError("weight_variants cannot be combined with model_sources") + if not isinstance(node.get("hf_repo"), str) or not node["hf_repo"]: + raise ValueError("weight_variants requires hf_repo on the same node") + raw = node["weight_variants"] + if not isinstance(raw, dict): + raise ValueError("weight_variants must be an object") + param = safe_source_id(raw.get("param"), "weight_variants.param") + raw_options = raw.get("options") + if not isinstance(raw_options, list) or not raw_options: + raise ValueError("weight_variants.options must be a non-empty array") + + aliases: dict[str, str] = {} + options: list[dict[str, Any]] = [] + for index, raw_option in enumerate(raw_options): + field = f"weight_variants.options[{index}]" + if not isinstance(raw_option, dict): + raise ValueError(f"{field} must be an object") + variant_id = safe_source_id(raw_option.get("id"), f"{field}.id") + alias = unicodedata.normalize("NFC", variant_id).casefold() + if alias in aliases: + raise ValueError( + f'weight variant ids "{aliases[alias]}" and "{variant_id}" are not portable-unique' + ) + aliases[alias] = variant_id + label = raw_option.get("label", variant_id) + if not isinstance(label, str) or not label.strip(): + raise ValueError(f"{field}.label must be a non-empty string") + size_gb = raw_option.get("size_gb") + if size_gb is not None and ( + isinstance(size_gb, bool) + or not isinstance(size_gb, (int, float)) + or not math.isfinite(size_gb) + or size_gb <= 0 + ): + raise ValueError(f"{field}.size_gb must be a positive number") + vram_gb = raw_option.get("vram_gb") + if vram_gb is not None and ( + isinstance(vram_gb, bool) + or not isinstance(vram_gb, (int, float)) + or not math.isfinite(vram_gb) + or vram_gb <= 0 + ): + raise ValueError(f"{field}.vram_gb must be a positive number") + include = _prefixes(raw_option.get("include_prefixes"), f"{field}.include_prefixes") + if not include: + raise ValueError(f"{field}.include_prefixes must be a non-empty array") + checks = raw_option.get("checks") + if not isinstance(checks, list) or not checks: + raise ValueError(f"{field}.checks must be a non-empty array") + safe_checks: list[str] = [] + for check_index, check in enumerate(checks): + path = safe_relative_path(check, f"{field}.checks[{check_index}]") + if not any(path.startswith(prefix) for prefix in include): + raise ValueError( + f"{field}.checks[{check_index}] is not covered by its include_prefixes" + ) + safe_checks.append(path) + option: dict[str, Any] = { + "id": variant_id, + "label": label, + "include_prefixes": include, + "checks": safe_checks, + } + if size_gb is not None: + option["size_gb"] = size_gb + if vram_gb is not None: + option["vram_gb"] = vram_gb + options.append(option) + + for index, option in enumerate(options): + for other in options[index + 1:]: + if _prefixes_overlap(option["include_prefixes"], other["include_prefixes"]): + raise ValueError( + f'weight variants "{option["id"]}" and "{other["id"]}" share files' + ) + download_check = node.get("download_check") + if isinstance(download_check, str) and any( + _prefixes_overlap([download_check], option["include_prefixes"]) for option in options + ): + raise ValueError("download_check must name a file outside every weight variant") + default = raw.get("default", options[0]["id"]) + if not any(option["id"] == default for option in options): + raise ValueError("weight_variants.default must name one of its options") + _assert_param_offers_variants( + param, + node.get("params_schema") if params_schema is None else params_schema, + [option["id"] for option in options], + ) + return {"param": param, "default": default, "options": options} + + +def installed_weight_variants( + models_dir: Path, model_id: str, variants: dict[str, Any] +) -> list[str]: + try: + model_root = resolve_model_root(models_dir, model_id) + except ValueError: + return [] + + def _present(check: str) -> bool: + try: + candidate = resolve_download_path(model_root, check) + return candidate.is_file() and candidate.stat().st_size > 0 + except (OSError, ValueError): + return False + + return [ + option["id"] + for option in variants["options"] + if all(_present(check) for check in option["checks"]) + ] + + +def missing_weight_variant( + models_dir: Path, + model_id: str, + variants: dict[str, Any] | None, + params: dict[str, Any], +) -> dict[str, Any] | None: + """The variant selected by params when its files are not installed, else None.""" + if not variants: + return None + selected = params.get(variants["param"]) + selected_id = variants["default"] if selected is None else str(selected) + option = next((o for o in variants["options"] if o["id"] == selected_id), None) + if option is None or option["id"] in installed_weight_variants(models_dir, model_id, variants): + return None + return option diff --git a/api/tests/test_generation_router.py b/api/tests/test_generation_router.py index 20fdda94..18a57253 100644 --- a/api/tests/test_generation_router.py +++ b/api/tests/test_generation_router.py @@ -42,6 +42,9 @@ def active_status(self) -> dict: # Report loaded so _run_generation skips the download/load thread. return {"loaded": True, "name": "fake", "downloaded": True} + def assert_weight_variant_installed(self, params: dict) -> None: + pass + def get_active(self) -> _FakeGenerator: return self._gen @@ -101,6 +104,23 @@ def test_output_lands_under_the_current_workspace(self) -> None: self.assertEqual(job.status, "done") self.assertEqual(job.output_url, "/workspace/MyColl/model.glb") + def test_missing_weight_variant_fails_the_job_before_generation(self) -> None: + class _MissingVariantRegistry(_FakeRegistry): + def assert_weight_variant_installed(self, params: dict) -> None: + raise RuntimeError(f'{params["gguf_quant"]} weights for trellis2/generate are not installed.') + + gen = _FakeGenerator() + generation.generator_registry = _MissingVariantRegistry(gen) + job_id = "job-missing-variant" + generation._jobs[job_id] = JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = threading.Event() + asyncio.run(generation._run_generation(job_id, b"img", {"gguf_quant": "Q6_K"}, "MyColl")) + + job = generation._jobs[job_id] + self.assertEqual(job.status, "error") + self.assertIn("Q6_K weights for trellis2/generate are not installed", job.error) + self.assertIsNone(gen.outputs_dir) + class GenerateFromImageWorkspaceTests(unittest.TestCase): """The request path must survive the relocation too, not just the worker: diff --git a/api/tests/test_generator_registry.py b/api/tests/test_generator_registry.py index ff9d090c..86736b39 100644 --- a/api/tests/test_generator_registry.py +++ b/api/tests/test_generator_registry.py @@ -211,6 +211,61 @@ def test_declared_sources_block_generation_even_when_generator_overrides_readine with self.assertRaisesRegex(RuntimeError, "Model sources are incomplete"): self.registry.get_active() + def test_selected_weight_variant_must_be_installed_before_generation(self) -> None: + def variant(quant: str) -> dict: + return { + "id": quant, + "include_prefixes": [f"dit_{quant}.gguf"], + "checks": [f"dit_{quant}.gguf"], + } + + extension = self._make_extension("quantized") + manifest = { + "id": "quantized", + "name": "quantized", + "type": "model", + "generator_class": "TestGenerator", + "params_schema": [ + {"id": "quant", "type": "select", "options": [{"value": "Q4"}, {"value": "Q5"}]} + ], + "nodes": [{ + "id": "generate", + "hf_repo": "org/model-gguf", + "download_check": "pipeline.json", + "weight_variants": { + "param": "quant", + "default": "Q5", + "options": [variant("Q4"), variant("Q5")], + }, + }], + } + (extension / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + (extension / "generator.py").write_text( + "\n".join([ + "from services.generators.base import BaseGenerator", + "class TestGenerator(BaseGenerator):", + " def is_downloaded(self): return True", + " def load(self): self._model = object()", + " def generate(self, image_bytes, params, progress_cb=None, cancel_event=None):", + " return self.outputs_dir / 'result.glb'", + ]), + encoding="utf-8", + ) + + self.registry.initialize() + self.registry._active_id = "quantized/generate" + manifest_variants = self.registry.get_manifest("quantized/generate")["weight_variants"] + self.assertEqual([option["id"] for option in manifest_variants["options"]], ["Q4", "Q5"]) + + model_root = self.models_dir / "quantized" / "generate" + model_root.mkdir(parents=True) + (model_root / "dit_Q5.gguf").write_bytes(b"q5") + self.registry.assert_weight_variant_installed({}) + self.registry.assert_weight_variant_installed({"quant": "Q5"}) + self.registry.assert_weight_variant_installed({"quant": "fp16"}) + with self.assertRaisesRegex(RuntimeError, "Q4 weights for quantized/generate are not installed"): + self.registry.assert_weight_variant_installed({"quant": "Q4"}) + def test_reload_preserves_legacy_path_owned_by_the_host(self) -> None: extension = self._make_extension("host-owned-path") self._write_manifest(extension, extension_id="host-owned-path") diff --git a/api/tests/test_model_sources.py b/api/tests/test_model_sources.py index cdab245a..6bed7b73 100644 --- a/api/tests/test_model_sources.py +++ b/api/tests/test_model_sources.py @@ -4,8 +4,11 @@ from pathlib import Path from services.model_sources import ( + installed_weight_variants, + missing_weight_variant, model_sources_are_downloaded, normalize_model_sources, + normalize_weight_variants, resolve_model_root, validate_source_file_plan, ) @@ -112,6 +115,71 @@ def test_requires_all_checks_and_rejects_symlinked_extension_ancestry(self) -> N resolve_model_root(models, "pixal3d/generate") self.assertFalse(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + def test_weight_variants_validate_and_report_the_missing_selection(self) -> None: + def variant(quant: str) -> dict: + return { + "id": quant, + "include_prefixes": [f"dit/model_{quant}.gguf"], + "checks": [f"dit/model_{quant}.gguf"], + } + + node = { + "hf_repo": "org/model-gguf", + "download_check": "pipeline.json", + "params_schema": [ + {"id": "quant", "type": "select", "options": [{"value": "Q4"}, {"value": "Q5"}]} + ], + "weight_variants": {"param": "quant", "options": [variant("Q4"), variant("Q5")]}, + } + variants = normalize_weight_variants(node) or {} + self.assertEqual(variants["default"], "Q4") + self.assertIsNone(normalize_weight_variants({"hf_repo": "org/model"})) + + def with_options(options: list, **extra) -> dict: + return {**node, "weight_variants": {"param": "quant", "options": options, **extra}} + + overlapping = {**variant("Q5"), "include_prefixes": ["dit/"]} + for broken, message in ( + ({**node, "hf_repo": ""}, "requires hf_repo"), + ({**node, "model_sources": []}, "cannot be combined"), + ({**node, "download_check": "dit/model_Q4.gguf"}, "download_check"), + (with_options([variant("Q4"), overlapping]), "share files"), + (with_options([variant("Q4")], default="Q8"), "default"), + (with_options([{**variant("Q4"), "checks": ["x.gguf"]}]), "not covered"), + (with_options([{**variant("Q4"), "size_gb": True}]), "size_gb"), + (with_options([{**variant("Q4"), "vram_gb": 0}]), "vram_gb"), + (with_options([{**variant("Q4"), "vram_gb": "6"}]), "vram_gb"), + ({**node, "params_schema": []}, "must name a params_schema entry"), + ( + {**node, "params_schema": [{"id": "quant", "options": [{"value": "Q5"}]}]}, + "must offer every weight variant id", + ), + ): + with self.subTest(message=message), self.assertRaisesRegex(ValueError, message): + normalize_weight_variants(broken) + + # A param that declares no options is left to the node to interpret. + self.assertIsNotNone( + normalize_weight_variants({**node, "params_schema": [{"id": "quant", "type": "string"}]}) + ) + self.assertNotIn("vram_gb", variants["options"][0]) + with_vram = normalize_weight_variants(with_options([{**variant("Q4"), "vram_gb": 6.5}])) or {} + self.assertEqual(with_vram["options"][0]["vram_gb"], 6.5) + + with tempfile.TemporaryDirectory(prefix="modly-weight-variants-") as tmp: + models = Path(tmp) / "models" + dit = models / "trellis" / "generate" / "dit" + dit.mkdir(parents=True) + (dit / "model_Q5.gguf").write_bytes(b"q5") + (dit / "model_Q4.gguf").write_bytes(b"") + self.assertEqual(installed_weight_variants(models, "trellis/generate", variants), ["Q5"]) + missing = missing_weight_variant(models, "trellis/generate", variants, {}) + self.assertEqual((missing or {}).get("id"), "Q4") + for params in ({"quant": "Q5"}, {"quant": "fp16"}): + with self.subTest(params=params): + self.assertIsNone(missing_weight_variant(models, "trellis/generate", variants, params)) + self.assertIsNone(missing_weight_variant(models, "trellis/generate", None, {"quant": "Q4"})) + if __name__ == "__main__": unittest.main() diff --git a/electron/main/extension-install-utils.test.mjs b/electron/main/extension-install-utils.test.mjs index 84139f9a..de49d929 100644 --- a/electron/main/extension-install-utils.test.mjs +++ b/electron/main/extension-install-utils.test.mjs @@ -102,6 +102,32 @@ test('validateInstallManifest rejects malformed or process model_sources', () => }, { hasEntryFile: () => true, hasGeneratorFile: () => false }, 'repository'), /only for model nodes/i) }) +test('validateInstallManifest validates weight variants and keeps them off process nodes', () => { + const mod = loadModule() + const files = { hasEntryFile: () => true, hasGeneratorFile: () => true } + const weightVariants = { + param: 'quant', + options: [{ id: 'Q4', include_prefixes: ['dit/model_Q4.gguf'], checks: ['dit/model_Q4.gguf'] }], + } + const paramsSchema = [{ id: 'quant', type: 'select', options: [{ value: 'Q4' }] }] + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'quantized', generator_class: 'Generator', params_schema: paramsSchema, + nodes: [{ id: 'generate', hf_repo: 'org/model', weight_variants: weightVariants }], + }, files, 'repository')) + assert.throws(() => mod.validateInstallManifest({ + id: 'quantized', generator_class: 'Generator', + nodes: [{ id: 'generate', hf_repo: 'org/model', weight_variants: weightVariants }], + }, files, 'repository'), /must name a params_schema entry/) + assert.throws(() => mod.validateInstallManifest({ + id: 'quantized', generator_class: 'Generator', params_schema: paramsSchema, + nodes: [{ id: 'generate', weight_variants: weightVariants }], + }, files, 'repository'), /requires hf_repo/) + assert.throws(() => mod.validateInstallManifest({ + id: 'proc', type: 'process', entry: 'processor.js', + nodes: [{ id: 'run', hf_repo: 'org/model', weight_variants: weightVariants }], + }, files, 'repository'), /weight_variants is supported only for model nodes/) +}) + test('python process setup failures are treated as fatal', () => { const mod = loadModule() diff --git a/electron/main/extension-install-utils.ts b/electron/main/extension-install-utils.ts index 05b965b0..6fc14e7f 100644 --- a/electron/main/extension-install-utils.ts +++ b/electron/main/extension-install-utils.ts @@ -1,7 +1,9 @@ import { normalizeModelSources, + normalizeWeightVariants, safeModelSourceId, type ModelSourceNode, + type WeightVariantNode, } from './model-sources' export interface InstallManifest { @@ -10,7 +12,8 @@ export interface InstallManifest { entry?: string generator_class?: string model_sources?: unknown - nodes?: Array<{ id?: string; model_sources?: unknown } & ModelSourceNode> + params_schema?: unknown + nodes?: Array<{ id?: string; model_sources?: unknown } & ModelSourceNode & WeightVariantNode> } export interface ValidatedInstallManifest { @@ -50,12 +53,14 @@ export function validateInstallManifest( throw new Error('manifest.json: model_sources must be declared on a model node') } for (const node of Array.isArray(manifest.nodes) ? manifest.nodes : []) { - if (node.model_sources === undefined) continue + if (node.model_sources === undefined && node.weight_variants === undefined) continue if (isProcess) { - throw new Error('manifest.json: model_sources is supported only for model nodes') + const field = node.model_sources !== undefined ? 'model_sources' : 'weight_variants' + throw new Error(`manifest.json: ${field} is supported only for model nodes`) } safeModelSourceId(node.id, 'model node id') normalizeModelSources(node) + normalizeWeightVariants(node, node.params_schema ?? manifest.params_schema) } if (isProcess) { diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 60f3cf2b..deb9b3e3 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -14,12 +14,16 @@ import { listDownloadedModels, downloadModelFromHF, downloadModelSourcesFromHF, + type DownloadProgress, } from './model-downloader' -import { resolveInstalledModelDownloadPlan } from './model-download-plan' +import { legacyDownloadSteps, resolveInstalledModelDownloadPlan } from './model-download-plan' import { areModelSourcesDownloaded, + installedWeightVariants, + listWeightVariantFiles, modelHasLocalData, normalizeModelSources, + normalizeWeightVariants, removePartialDownloadArtifacts, resolveModelRoot, } from './model-sources' @@ -149,11 +153,29 @@ const renameWithRetry = (from: string, to: string, label: string) => export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGetter): void { type ActiveDownload = { - progress: { percent: number; file?: string; fileIndex?: number; totalFiles?: number } + progress: DownloadProgress & { variantId?: string } done: Promise finish: () => void } const activeDownloads = new Map() + const resolveModelPlan = (modelId: unknown) => resolveInstalledModelDownloadPlan({ + modelId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + const LOCKED_MODEL_FILES_ERROR = 'Model files are still locked after several attempts. Close any programs using the model and try again.' + + // Unload and wait for confirmation so file handles are released before removal. + async function unloadModelBeforeRemoval(modelId: string): Promise { + try { + await axios.post(`${API_BASE_URL}/model/unload/${encodeURIComponent(modelId)}`, {}, { timeout: 10_000 }) + // Give the OS a moment to release file locks (Windows holds handles briefly after close) + await new Promise(resolve => setTimeout(resolve, 1_500)) + } catch { + // Unload failed (model may not be loaded) — still attempt deletion + } + } // Logging from renderer ipcMain.on('log:error', (_event, message: string) => logger.error(`[Renderer] ${message}`)) ipcMain.handle('log:getPath', () => join(app.getPath('userData'), 'logs', 'modly.log')) @@ -339,35 +361,45 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe } let modelDir: string try { - await resolveInstalledModelDownloadPlan({ - modelId, - userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, - builtinExtensionsDir: getBuiltinExtensionsDir(), - blockedExtensionIds: activeExtensionInstalls, - }) + await resolveModelPlan(modelId) modelDir = resolveModelRoot(getSettings(app.getPath('userData')).modelsDir, modelId) } catch (err) { return { success: false, error: String(err) } } - // Unload the model and wait for confirmation so file handles are released - try { - await axios.post(`${API_BASE_URL}/model/unload/${encodeURIComponent(modelId)}`, {}, { timeout: 10_000 }) - // Give the OS a moment to release file locks (Windows holds handles briefly after close) - await new Promise(resolve => setTimeout(resolve, 1_500)) - } catch { - // Unload failed (model may not be loaded) — still attempt deletion - } + await unloadModelBeforeRemoval(modelId) // Retry removal — Windows may return EBUSY/EPERM if handles linger const removed = await rmWithRetry(modelDir, 'model-delete') if (removed.ok) return { success: true } - return { - success: false, - error: removed.locked - ? 'Model files are still locked after several attempts. Close any programs using the model and try again.' - : String(removed.error), + return { success: false, error: removed.locked ? LOCKED_MODEL_FILES_ERROR : String(removed.error) } + }) + + ipcMain.handle('model:deleteWeightVariant', async (_, modelId: string, variantId: string): Promise<{ success: boolean; error?: string }> => { + if (activeDownloads.has(modelId)) { + return { success: false, error: 'Cannot remove model weights while their download is active' } } + let files: string[] + try { + const plan = await resolveModelPlan(modelId) + const variant = plan.kind === 'legacy' + ? plan.weightVariants?.options.find((option) => option.id === variantId) + : undefined + if (!variant) throw new Error(`Model node "${modelId}" has no weight variant "${String(variantId)}"`) + files = await listWeightVariantFiles(getSettings(app.getPath('userData')).modelsDir, modelId, variant) + } catch (err) { + return { success: false, error: String(err) } + } + if (files.length === 0) return { success: true } + + await unloadModelBeforeRemoval(modelId) + for (const file of files) { + const removed = await rmWithRetry(file, 'model-variant-delete') + if (!removed.ok) { + return { success: false, error: removed.locked ? LOCKED_MODEL_FILES_ERROR : String(removed.error) } + } + } + return { success: true } }) ipcMain.handle('model:showInFolder', (_, modelId: string) => { @@ -403,28 +435,31 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe ipcMain.handle('model:isDownloaded', async (_, modelId: string): Promise => { const modelsDir = getSettings(app.getPath('userData')).modelsDir try { - const plan = await resolveInstalledModelDownloadPlan({ - modelId, - userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, - builtinExtensionsDir: getBuiltinExtensionsDir(), - blockedExtensionIds: activeExtensionInstalls, - }) - return plan.kind === 'multi-source' - ? areModelSourcesDownloaded(modelsDir, modelId, plan.sources) - : isModelDownloaded(modelsDir, modelId, plan.downloadCheck) + const plan = await resolveModelPlan(modelId) + if (plan.kind === 'multi-source') return areModelSourcesDownloaded(modelsDir, modelId, plan.sources) + return isModelDownloaded(modelsDir, modelId, plan.downloadCheck) + && (!plan.weightVariants || installedWeightVariants(modelsDir, modelId, plan.weightVariants).length > 0) } catch { return false } }) + // null means "unknown" (unreadable plan, or a node without variants) — the renderer + // must not read an empty array as "no variant installed". + ipcMain.handle('model:installedWeightVariants', async (_, modelId: string): Promise => { + try { + const plan = await resolveModelPlan(modelId) + return plan.kind === 'legacy' && plan.weightVariants + ? installedWeightVariants(getSettings(app.getPath('userData')).modelsDir, modelId, plan.weightVariants) + : null + } catch { + return null + } + }) + ipcMain.handle('model:hasLocalData', async (_, modelId: string): Promise => { try { - await resolveInstalledModelDownloadPlan({ - modelId, - userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, - builtinExtensionsDir: getBuiltinExtensionsDir(), - blockedExtensionIds: activeExtensionInstalls, - }) + await resolveModelPlan(modelId) return modelHasLocalData(getSettings(app.getPath('userData')).modelsDir, modelId) } catch { return false @@ -438,45 +473,47 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe ipcMain.handle('model:download', async ( event, modelId: string, + requestedVariantId?: string | null, ) => { if (activeDownloads.has(modelId)) { return { success: false, error: 'Download already in progress' } } + const variantId = requestedVariantId ?? undefined let finish!: () => void const done = new Promise((resolveDone) => { finish = resolveDone }) - const active: ActiveDownload = { progress: { percent: 0 }, done, finish } + const active: ActiveDownload = { progress: { percent: 0, variantId }, done, finish } activeDownloads.set(modelId, active) try { - const plan = await resolveInstalledModelDownloadPlan({ - modelId, - userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, - builtinExtensionsDir: getBuiltinExtensionsDir(), - blockedExtensionIds: activeExtensionInstalls, - }) - const onProgress = (progress: typeof active.progress) => { - active.progress = progress - event.sender.send('model:downloadProgress', { modelId, ...progress }) + const plan = await resolveModelPlan(modelId) + const onProgress = (progress: DownloadProgress) => { + active.progress = { ...progress, variantId } + event.sender.send('model:downloadProgress', { modelId, variantId, ...progress }) } if (plan.kind === 'multi-source') { + if (variantId !== undefined) throw new Error(`Model node "${modelId}" does not declare weight variants`) await downloadModelSourcesFromHF(modelId, plan.sources, onProgress) } else { - await downloadModelFromHF( - plan.repoId, - modelId, - onProgress, - plan.skipPrefixes, - plan.includePrefixes, - ) + // Shared files and the default variant are separate passes sharing one 0–100 bar. + const steps = legacyDownloadSteps(plan, variantId) + for (const [index, step] of steps.entries()) { + await downloadModelFromHF( + plan.repoId, + modelId, + (progress) => onProgress({ ...progress, percent: Math.round((index * 100 + progress.percent) / steps.length) }), + step.skipPrefixes, + step.includePrefixes, + ) + } } return { success: true } } catch (err: any) { const message = err?.message ?? String(err) if (message.includes('paused')) { - event.sender.send('model:downloadProgress', { modelId, percent: 0, status: 'paused', paused: true }) + event.sender.send('model:downloadProgress', { modelId, variantId, percent: 0, status: 'paused', paused: true }) return { success: false, paused: true } } if (message.includes('cancelled')) { - event.sender.send('model:downloadProgress', { modelId, percent: 0, status: 'cancelled', cancelled: true }) + event.sender.send('model:downloadProgress', { modelId, variantId, percent: 0, status: 'cancelled', cancelled: true }) return { success: false, cancelled: true } } return { success: false, error: String(err) } @@ -835,6 +872,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe hf_skip_prefixes?: string[] hf_include_prefixes?: string[] model_sources?: unknown + weight_variants?: unknown }[] } @@ -857,7 +895,11 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe if (parsed.type === 'process' && n.model_sources !== undefined) { throw new Error('manifest.json: model_sources is supported only for model nodes') } + if (parsed.type === 'process' && n.weight_variants !== undefined) { + throw new Error('manifest.json: weight_variants is supported only for model nodes') + } const modelSources = normalizeModelSources(n) + const weightVariants = normalizeWeightVariants(n, n.params_schema ?? parsed.params_schema) return { id: n.id, name: n.name ?? n.id, @@ -872,6 +914,16 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe hfSkipPrefixes: n.hf_skip_prefixes, hfIncludePrefixes: n.hf_include_prefixes, hasModelSources: modelSources !== undefined, + weightVariants: weightVariants && { + param: weightVariants.param, + default: weightVariants.default, + options: weightVariants.options.map((option) => ({ + id: option.id, + label: option.label, + sizeGb: option.size_gb, + vramGb: option.vram_gb, + })), + }, } }) diff --git a/electron/main/model-download-plan.test.mjs b/electron/main/model-download-plan.test.mjs index e5d037a0..520aa780 100644 --- a/electron/main/model-download-plan.test.mjs +++ b/electron/main/model-download-plan.test.mjs @@ -90,3 +90,55 @@ test('keeps legacy sibling checks and wildcard filters unchanged', async () => { rmSync(fixture.root, { recursive: true, force: true }) } }) + +test('downloads a weight-variant node as shared files first, then one variant at a time', async () => { + const { legacyDownloadSteps, resolveInstalledModelDownloadPlan } = loadModule() + const fixture = setupExtension({ + id: 'trellis', + type: 'model', + nodes: [{ + id: 'generate', + hf_repo: 'org/model-gguf', + download_check: 'pipeline.json', + hf_include_prefixes: ['pipeline.json', 'dit/'], + hf_skip_prefixes: ['README.md'], + params_schema: [{ id: 'quant', type: 'select', options: [{ value: 'Q4' }, { value: 'Q5' }] }], + weight_variants: { + param: 'quant', + default: 'Q5', + options: ['Q4', 'Q5'].map((quant) => ({ + id: quant, + include_prefixes: [`dit/model_${quant}.gguf`], + checks: [`dit/model_${quant}.gguf`], + })), + }, + }], + }) + try { + const plan = await resolveInstalledModelDownloadPlan({ + modelId: 'trellis/generate', + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + }) + assert.equal(plan.kind, 'legacy') + assert.deepEqual(legacyDownloadSteps(plan), [ + { includePrefixes: ['pipeline.json', 'dit/'], skipPrefixes: ['README.md', 'dit/model_Q4.gguf', 'dit/model_Q5.gguf'] }, + { includePrefixes: ['dit/model_Q5.gguf'], skipPrefixes: ['README.md'] }, + ]) + // Picking one variant still fetches the shared files first, so a node installed + // variant-first is never left without its pipeline files. + assert.deepEqual(legacyDownloadSteps(plan, 'Q4'), [ + { includePrefixes: ['pipeline.json', 'dit/'], skipPrefixes: ['README.md', 'dit/model_Q4.gguf', 'dit/model_Q5.gguf'] }, + { includePrefixes: ['dit/model_Q4.gguf'], skipPrefixes: ['README.md'] }, + ]) + assert.throws(() => legacyDownloadSteps(plan, 'Q8'), /no weight variant "Q8"/) + + const withoutVariants = { ...plan, weightVariants: undefined } + assert.deepEqual(legacyDownloadSteps(withoutVariants), [ + { includePrefixes: ['pipeline.json', 'dit/'], skipPrefixes: ['README.md'] }, + ]) + assert.throws(() => legacyDownloadSteps(withoutVariants, 'Q4'), /no weight variant/) + } finally { + rmSync(fixture.root, { recursive: true, force: true }) + } +}) diff --git a/electron/main/model-download-plan.ts b/electron/main/model-download-plan.ts index 27d2dba6..e4a4905e 100644 --- a/electron/main/model-download-plan.ts +++ b/electron/main/model-download-plan.ts @@ -8,7 +8,13 @@ import { assertSafeExtensionId, resolveExtensionPathWithinRoot, } from './extension-path-guard' -import { normalizeModelSources, safeModelSourceId, type ModelSource } from './model-sources' +import { + normalizeModelSources, + normalizeWeightVariants, + safeModelSourceId, + type ModelSource, + type WeightVariants, +} from './model-sources' interface InstalledNode { id?: unknown @@ -17,12 +23,15 @@ interface InstalledNode { hf_skip_prefixes?: unknown hf_include_prefixes?: unknown model_sources?: unknown + weight_variants?: unknown + params_schema?: unknown } interface InstalledManifest { id?: unknown type?: unknown model_sources?: unknown + params_schema?: unknown nodes?: unknown } @@ -35,6 +44,7 @@ export type InstalledModelDownloadPlan = { downloadCheck?: string skipPrefixes?: string[] includePrefixes?: string[] + weightVariants?: WeightVariants } | { kind: 'multi-source' modelId: string @@ -86,6 +96,7 @@ function parseManifest(raw: string, extensionId: string, nodeId: string): Instal const node = matches[0] const modelId = `${extensionId}/${nodeId}` + const weightVariants = normalizeWeightVariants(node, node.params_schema ?? manifest.params_schema) const sources = normalizeModelSources(node) if (sources) return { kind: 'multi-source', modelId, extensionId, nodeId, sources } @@ -101,7 +112,43 @@ function parseManifest(raw: string, extensionId: string, nodeId: string): Instal downloadCheck: typeof node.download_check === 'string' ? node.download_check : undefined, skipPrefixes: node.hf_skip_prefixes as string[] | undefined, includePrefixes: node.hf_include_prefixes as string[] | undefined, + ...(weightVariants ? { weightVariants } : {}), + } +} + +export interface LegacyDownloadStep { + includePrefixes?: string[] + skipPrefixes?: string[] +} + +/** + * Hugging Face filter passes for one download action. A node that declares weight + * variants always fetches its shared files first (every variant excluded), then the + * requested variant — its default one when no variant id is given. The shared pass + * skips files already complete on disk, so it stays cheap on a resume or a second + * variant. + */ +export function legacyDownloadSteps( + plan: Extract, + variantId?: unknown, +): LegacyDownloadStep[] { + const variants = plan.weightVariants + if (!variants) { + if (variantId !== undefined) { + throw new Error(`Model node "${plan.modelId}" has no weight variant "${String(variantId)}"`) + } + return [{ includePrefixes: plan.includePrefixes, skipPrefixes: plan.skipPrefixes }] } + const selected = variantId === undefined ? variants.default : variantId + const variant = variants.options.find((option) => option.id === selected) + if (!variant) throw new Error(`Model node "${plan.modelId}" has no weight variant "${String(variantId)}"`) + return [ + { + includePrefixes: plan.includePrefixes, + skipPrefixes: [...(plan.skipPrefixes ?? []), ...variants.options.flatMap((option) => option.include_prefixes)], + }, + { includePrefixes: variant.include_prefixes, skipPrefixes: plan.skipPrefixes }, + ] } /** Re-read the installed manifest for every model action; renderer metadata is never trusted. */ diff --git a/electron/main/model-sources.test.mjs b/electron/main/model-sources.test.mjs index da8a54fc..63183ce3 100644 --- a/electron/main/model-sources.test.mjs +++ b/electron/main/model-sources.test.mjs @@ -113,3 +113,87 @@ test('requires every declared check and rejects symlinked extension-root ancestr rmSync(root, { recursive: true, force: true }) } }) + +const quantNode = () => ({ + hf_repo: 'org/model-gguf', + download_check: 'pipeline.json', + params_schema: [{ + id: 'quant', + label: 'Quantization', + type: 'select', + default: 'Q5', + options: [{ value: 'Q4', label: 'Q4' }, { value: 'Q5', label: 'Q5' }], + }], + weight_variants: { + param: 'quant', + default: 'Q5', + options: ['Q4', 'Q5'].map((quant) => ({ + id: quant, + include_prefixes: [`dit/model_${quant}.gguf`], + checks: [`dit/model_${quant}.gguf`], + })), + }, +}) + +test('validates weight variants and rejects ambiguous or unsafe declarations', () => { + const { normalizeWeightVariants } = loadModule() + const variants = normalizeWeightVariants(quantNode()) + assert.equal(variants.param, 'quant') + assert.equal(variants.default, 'Q5') + assert.deepEqual(variants.options.map((option) => option.label), ['Q4', 'Q5']) + assert.equal(normalizeWeightVariants({ hf_repo: 'org/model' }), undefined) + + const node = quantNode() + const [q4, q5] = node.weight_variants.options + const withOptions = (options, extra = {}) => ({ ...node, weight_variants: { ...node.weight_variants, options, ...extra } }) + const cases = [ + [{ ...node, hf_repo: undefined }, /requires hf_repo/], + [{ ...node, model_sources: [] }, /cannot be combined/], + [{ ...node, download_check: 'dit/model_Q4.gguf' }, /download_check/], + [withOptions([q4, q5], { default: 'Q8' }), /default/], + [withOptions([q4, { ...q5, include_prefixes: ['dit/'] }]), /share files/], + [withOptions([q4, { ...q5, id: 'q4' }]), /portable-unique/], + [withOptions([{ ...q4, checks: ['other.gguf'] }]), /not covered/], + [withOptions([{ ...q4, include_prefixes: ['../outside'] }]), /unsafe/], + [withOptions([{ ...q4, size_gb: -1 }]), /size_gb/], + [withOptions([{ ...q4, vram_gb: 0 }]), /vram_gb/], + [withOptions([{ ...q4, vram_gb: '6' }]), /vram_gb/], + [{ ...node, params_schema: undefined }, /must name a params_schema entry/], + [{ ...node, params_schema: [{ id: 'steps', type: 'int' }] }, /must name a params_schema entry/], + [ + { ...node, params_schema: [{ id: 'quant', type: 'select', options: [{ value: 'Q5' }] }] }, + /must offer every weight variant id \(missing: Q4\)/, + ], + ] + for (const [candidate, error] of cases) assert.throws(() => normalizeWeightVariants(candidate), error) + + // A param that declares no options (or an unknown shape) is left to the node. + assert.doesNotThrow(() => normalizeWeightVariants({ ...node, params_schema: [{ id: 'quant', type: 'string' }] })) + assert.equal('vram_gb' in variants.options[0], false) + assert.equal(normalizeWeightVariants(withOptions([{ ...q4, vram_gb: 6.5 }, q5])).options[0].vram_gb, 6.5) +}) + +test('reports installed variants and lists only the files of the variant being removed', async () => { + const { installedWeightVariants, listWeightVariantFiles, normalizeWeightVariants } = loadModule() + const variants = normalizeWeightVariants(quantNode()) + const root = mkdtempSync(join(tmpdir(), 'modly-weight-variants-')) + const models = join(root, 'models') + const nodeRoot = join(models, 'trellis', 'generate') + mkdirSync(join(nodeRoot, 'dit'), { recursive: true }) + writeFileSync(join(nodeRoot, 'pipeline.json'), '{}') + writeFileSync(join(nodeRoot, 'dit', 'model_Q5.gguf'), 'q5') + writeFileSync(join(nodeRoot, 'dit', 'model_Q4.gguf'), '') + writeFileSync(join(nodeRoot, 'dit', 'model_Q4.gguf.part'), 'partial') + + try { + assert.deepEqual(installedWeightVariants(models, 'trellis/generate', variants), ['Q5']) + assert.deepEqual(installedWeightVariants(models, '../escape', variants), []) + const files = await listWeightVariantFiles(models, 'trellis/generate', variants.options[0]) + assert.deepEqual( + files.map((file) => file.slice(nodeRoot.length + 1).replaceAll('\\', '/')).sort(), + ['dit/model_Q4.gguf', 'dit/model_Q4.gguf.part'], + ) + } finally { + rmSync(root, { recursive: true, force: true }) + } +}) diff --git a/electron/main/model-sources.ts b/electron/main/model-sources.ts index dac76930..291c05f3 100644 --- a/electron/main/model-sources.ts +++ b/electron/main/model-sources.ts @@ -1,6 +1,6 @@ import { existsSync, lstatSync, readdirSync, statSync } from 'node:fs' import { readdir, rm } from 'node:fs/promises' -import { isAbsolute, relative, resolve } from 'node:path' +import { isAbsolute, relative, resolve, sep } from 'node:path' export interface ModelSource { id: string @@ -17,6 +17,31 @@ export interface ModelSourceNode { model_sources?: unknown } +/** One separately installable set of files inside a node's model directory (e.g. a quantization). */ +export interface WeightVariant { + id: string + label: string + size_gb?: number + vram_gb?: number + include_prefixes: string[] + checks: string[] +} + +export interface WeightVariants { + /** params_schema id whose value selects the variant at generation time */ + param: string + default: string + options: WeightVariant[] +} + +export interface WeightVariantNode { + weight_variants?: unknown + hf_repo?: unknown + download_check?: unknown + model_sources?: unknown + params_schema?: unknown +} + const SAFE_ID = /^[A-Za-z0-9][A-Za-z0-9._-]*$/ const WINDOWS_DEVICE = /^(?:con|prn|aux|nul|com[1-9]|lpt[1-9])(?:\..*)?$/i const WINDOWS_UNSAFE = /[<>"|?*\u0000-\u001f]/ @@ -133,6 +158,114 @@ export function normalizeModelSources(node: ModelSourceNode): ModelSource[] | un }) } +function prefixesOverlap(left: string[], right: string[]): boolean { + return left.some((a) => right.some((b) => { + const lowerA = a.toLowerCase() + const lowerB = b.toLowerCase() + return lowerA.startsWith(lowerB) || lowerB.startsWith(lowerA) + })) +} + +/** Variant ids must be selectable, so the param they key has to exist and offer them. */ +function assertParamOffersVariants(param: string, paramsSchema: unknown, ids: string[]): void { + const entry = Array.isArray(paramsSchema) + ? paramsSchema.find((p) => typeof p === 'object' && p !== null && (p as { id?: unknown }).id === param) + : undefined + if (!entry) throw new Error(`weight_variants.param must name a params_schema entry ("${param}")`) + const options = (entry as { options?: unknown }).options + if (!Array.isArray(options)) return + const values = new Set(options.map((option) => + typeof option === 'object' && option !== null ? String((option as { value?: unknown }).value) : String(option))) + const missing = ids.filter((id) => !values.has(id)) + if (missing.length > 0) { + throw new Error(`the "${param}" param must offer every weight variant id (missing: ${missing.join(', ')})`) + } +} + +export function normalizeWeightVariants( + node: WeightVariantNode, + paramsSchema?: unknown, +): WeightVariants | undefined { + if (!Object.prototype.hasOwnProperty.call(node, 'weight_variants')) return undefined + if (Object.prototype.hasOwnProperty.call(node, 'model_sources')) { + throw new Error('weight_variants cannot be combined with model_sources') + } + if (typeof node.hf_repo !== 'string' || !node.hf_repo) { + throw new Error('weight_variants requires hf_repo on the same node') + } + const raw = node.weight_variants + if (typeof raw !== 'object' || raw === null || Array.isArray(raw)) { + throw new Error('weight_variants must be an object') + } + const value = raw as Record + const param = safeModelSourceId(value.param, 'weight_variants.param') + if (!Array.isArray(value.options) || value.options.length === 0) { + throw new Error('weight_variants.options must be a non-empty array') + } + + const seen = new Map() + const options = value.options.map((rawOption, index): WeightVariant => { + const field = `weight_variants.options[${index}]` + if (typeof rawOption !== 'object' || rawOption === null || Array.isArray(rawOption)) { + throw new Error(`${field} must be an object`) + } + const option = rawOption as Record + const id = safeModelSourceId(option.id, `${field}.id`) + const alias = id.normalize('NFC').toLowerCase() + const previous = seen.get(alias) + if (previous) throw new Error(`weight variant ids "${previous}" and "${id}" are not portable-unique`) + seen.set(alias, id) + const label = option.label ?? id + if (typeof label !== 'string' || !label.trim()) throw new Error(`${field}.label must be a non-empty string`) + const sizeGb = option.size_gb + if (sizeGb !== undefined && (typeof sizeGb !== 'number' || !Number.isFinite(sizeGb) || sizeGb <= 0)) { + throw new Error(`${field}.size_gb must be a positive number`) + } + const vramGb = option.vram_gb + if (vramGb !== undefined && (typeof vramGb !== 'number' || !Number.isFinite(vramGb) || vramGb <= 0)) { + throw new Error(`${field}.vram_gb must be a positive number`) + } + const include = optionalPrefixes(option.include_prefixes, `${field}.include_prefixes`) + if (!include?.length) throw new Error(`${field}.include_prefixes must be a non-empty array`) + if (!Array.isArray(option.checks) || option.checks.length === 0) { + throw new Error(`${field}.checks must be a non-empty array`) + } + const checks = option.checks.map((check, checkIndex) => { + const path = safeModelRelativePath(check, `${field}.checks[${checkIndex}]`) + if (!include.some((prefix) => path.startsWith(prefix))) { + throw new Error(`${field}.checks[${checkIndex}] is not covered by its include_prefixes`) + } + return path + }) + return { + id, + label, + ...(sizeGb === undefined ? {} : { size_gb: sizeGb }), + ...(vramGb === undefined ? {} : { vram_gb: vramGb }), + include_prefixes: include, + checks, + } + }) + + options.forEach((variant, index) => { + for (const other of options.slice(index + 1)) { + if (prefixesOverlap(variant.include_prefixes, other.include_prefixes)) { + throw new Error(`weight variants "${variant.id}" and "${other.id}" share files`) + } + } + }) + const downloadCheck = node.download_check + if (typeof downloadCheck === 'string' && options.some((variant) => prefixesOverlap([downloadCheck], variant.include_prefixes))) { + throw new Error('download_check must name a file outside every weight variant') + } + const defaultId = value.default ?? options[0].id + if (!options.some((variant) => variant.id === defaultId)) { + throw new Error('weight_variants.default must name one of its options') + } + assertParamOffersVariants(param, paramsSchema ?? node.params_schema, options.map((variant) => variant.id)) + return { param, default: defaultId as string, options } +} + function pathHasSymlink(root: string, candidate: string): boolean { const rootPath = resolve(root) const rel = relative(rootPath, resolve(candidate)) @@ -196,6 +329,43 @@ export function modelHasLocalData(modelsDir: string, modelId: string): boolean { } } +function isDownloadedFile(root: string, relativePath: string): boolean { + const candidate = resolve(root, ...relativePath.split('/')) + if (!existsSync(candidate) || pathHasSymlink(root, candidate)) return false + try { + const stat = statSync(candidate) + return stat.isFile() && stat.size > 0 + } catch { + return false + } +} + +export function installedWeightVariants(modelsDir: string, modelId: string, variants: WeightVariants): string[] { + try { + const modelRoot = resolveModelRoot(modelsDir, modelId) + return variants.options + .filter((variant) => variant.checks.every((check) => isDownloadedFile(modelRoot, check))) + .map((variant) => variant.id) + } catch { + return [] + } +} + +/** Files of one variant on disk (in-progress `.part` files included), never through a symlink. */ +export async function listWeightVariantFiles(modelsDir: string, modelId: string, variant: WeightVariant): Promise { + const modelRoot = resolveModelRoot(modelsDir, modelId) + if (!existsSync(modelRoot)) return [] + const entries = await readdir(modelRoot, { recursive: true, withFileTypes: true }) + return entries + .filter((entry) => entry.isFile()) + .map((entry) => resolve(entry.parentPath ?? modelRoot, entry.name)) + .filter((path) => { + const relativePath = relative(modelRoot, path).split(sep).join('/') + return variant.include_prefixes.some((prefix) => relativePath.startsWith(prefix)) + && !pathHasSymlink(modelRoot, path) + }) +} + // Mirrors the backend's cancel cleanup (api/routers/model.py): only the in-progress // `.part` files are removed, so completed sources already on disk survive a cancel. export async function removePartialDownloadArtifacts(modelRoot: string): Promise { diff --git a/electron/preload/electron-api.ts b/electron/preload/electron-api.ts index dae65e69..69c033f4 100644 --- a/electron/preload/electron-api.ts +++ b/electron/preload/electron-api.ts @@ -119,10 +119,14 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra listDownloaded: () => ipcRenderer.invoke('model:listDownloaded'), isDownloaded: (modelId: string) => ipcRenderer.invoke('model:isDownloaded', modelId), hasLocalData: (modelId: string) => ipcRenderer.invoke('model:hasLocalData', modelId), - download: (modelId: string) => ipcRenderer.invoke('model:download', modelId), + download: (modelId: string, variantId?: string) => (variantId === undefined + ? ipcRenderer.invoke('model:download', modelId) + : ipcRenderer.invoke('model:download', modelId, variantId)), pauseDownload: (modelId: string) => ipcRenderer.invoke('model:pauseDownload', modelId), cancelDownload: (modelId: string) => ipcRenderer.invoke('model:cancelDownload', modelId), delete: (modelId: string) => ipcRenderer.invoke('model:delete', modelId), + installedWeightVariants: (modelId: string) => ipcRenderer.invoke('model:installedWeightVariants', modelId), + deleteWeightVariant: (modelId: string, variantId: string) => ipcRenderer.invoke('model:deleteWeightVariant', modelId, variantId), unloadAll: () => ipcRenderer.invoke('model:unloadAll'), showInFolder: (modelId: string) => ipcRenderer.invoke('model:showInFolder', modelId), activeDownloads: (): Promise<{ modelId: string; percent: number; file?: string; fileIndex?: number; totalFiles?: number }[]> => diff --git a/src/areas/generate/components/WorkflowPanel.tsx b/src/areas/generate/components/WorkflowPanel.tsx index 2995ddea..c15034b5 100644 --- a/src/areas/generate/components/WorkflowPanel.tsx +++ b/src/areas/generate/components/WorkflowPanel.tsx @@ -17,6 +17,7 @@ import type { WorkflowExtension } from '@areas/workflows/mockExtensions' import type { Workflow, WFNode, WFEdge, ParamSchema } from '@shared/types/electron.d' import { PICKER_LABELS, openParamPicker, resolvePickerIntent } from '@shared/utils/paramPicker' import { PickerIcon } from '@shared/components/ui' +import { isMissingWeightVariant, withWeightVariantAvailability } from '@shared/utils/weightVariants' import ChatPanel from './ChatPanel' type PanelMode = 'basic' | 'chat' @@ -383,6 +384,8 @@ function WaitParamRow({ nodeId }: { nodeId: string }) { function ExtensionParamRow({ nodeId, ext, nodes, onPatch }: { nodeId: string; ext: WorkflowExtension; nodes: FlowNode[]; onPatch: PatchFn }) { const [expanded, setExpanded] = useState(true) + const installedVariants = useExtensionsStore((s) => s.installedWeightVariants[ext.id]) + const openExtension = useNavStore((s) => s.openExtension) const node = nodes.find((n) => n.id === nodeId) const data = node?.data as { enabled: boolean; params: Record } | undefined const enabled = data?.enabled ?? true @@ -430,8 +433,11 @@ function ExtensionParamRow({ nodeId, ext, nodes, onPatch }: { nodeId: string; ex
- onPatch(nodeId, { params: { ...(data?.params ?? {}), [param.id]: v } })} /> + { + onPatch(nodeId, { params: { ...(data?.params ?? {}), [param.id]: v } }) + if (isMissingWeightVariant(param.id, v, ext.weightVariants, installedVariants)) openExtension(ext.extensionId) + }} />
) diff --git a/src/areas/models/ModelsPage.tsx b/src/areas/models/ModelsPage.tsx index 66f8e5d4..c4eaba0e 100644 --- a/src/areas/models/ModelsPage.tsx +++ b/src/areas/models/ModelsPage.tsx @@ -1,12 +1,13 @@ import { useEffect, useMemo, useRef, useState } from 'react' import { createPortal } from 'react-dom' import { useExtensionsStore } from '@shared/stores/extensionsStore' +import { useNavStore } from '@shared/stores/navStore' import type { AnyExtension, ModelExtension } from '@shared/types/electron.d' import { deleteModelsThenUninstallExtension, formatModelName } from './utils' import { ExtensionCard } from './components/ExtensionCard' import type { ExtensionNode } from './components/ExtensionCard' import { ExtensionDrawer } from './components/ExtensionDrawer' -import { ICONS, nodeHasManagedWeights } from './components/extensionShared' +import { ICONS, nodeHasManagedWeights, type DownloadMap } from './components/extensionShared' // ─── Filters & sorts ────────────────────────────────────────────────────────── @@ -48,19 +49,9 @@ export default function ModelsPage(): JSX.Element { ) // Model weight state (needed for node install status + uninstall cleanup) - const [installedVariantIds, setInstalledVariantIds] = useState([]) + const [installedNodeIds, setInstalledNodeIds] = useState([]) const [localDataIds, setLocalDataIds] = useState([]) - const [downloading, setDownloading] = useState>({}) + const [downloading, setDownloading] = useState({}) // Uninstall modal state const [uninstallTarget, setUninstallTarget] = useState(null) @@ -75,6 +66,15 @@ export default function ModelsPage(): JSX.Element { const [selectedId, setSelectedId] = useState(null) const searchRef = useRef(null) + // Another page asked to open an extension (e.g. a weight variant picked but not installed) + const extensionToOpen = useNavStore((s) => s.extensionToOpen) + const clearExtensionToOpen = useNavStore((s) => s.clearExtensionToOpen) + useEffect(() => { + if (!extensionToOpen) return + setSelectedId(extensionToOpen) + clearExtensionToOpen() + }, [extensionToOpen, clearExtensionToOpen]) + // GitHub extension install form const [showGHForm, setShowGHForm] = useState(false) const [ghUrl, setGhUrl] = useState('') @@ -98,8 +98,9 @@ export default function ModelsPage(): JSX.Element { if (hasLocalData) localIds.push(fullId) } } - setInstalledVariantIds(ids) + setInstalledNodeIds(ids) setLocalDataIds(localIds) + await useExtensionsStore.getState().refreshInstalledWeightVariants() } useEffect(() => { @@ -115,7 +116,7 @@ export default function ModelsPage(): JSX.Element { } refreshInstalledIds(exts) }) - window.electron.model.onProgress(({ modelId: id, percent, file, fileIndex, totalFiles, status, bytesDownloaded, totalBytes, stalledSeconds, paused, cancelled }) => { + window.electron.model.onProgress(({ modelId: id, variantId, percent, file, fileIndex, totalFiles, status, bytesDownloaded, totalBytes, stalledSeconds, paused, cancelled }) => { if (cancelled) { setDownloading((prev) => { const n = { ...prev }; delete n[id]; return n }) return @@ -134,6 +135,7 @@ export default function ModelsPage(): JSX.Element { totalBytes: totalBytes ?? current?.totalBytes, stalledSeconds: stalledSeconds ?? current?.stalledSeconds, paused, + variantId: variantId ?? current?.variantId, }, } }) @@ -167,10 +169,10 @@ export default function ModelsPage(): JSX.Element { // ── Node install / download controls ────────────────────────────────────── - function handleInstallNode(node: ExtensionNode, fullId: string) { + function handleInstallNode(node: ExtensionNode, fullId: string, variantId?: string) { if (!nodeHasManagedWeights(node)) return - setDownloading((prev) => ({ ...prev, [fullId]: { ...(prev[fullId] ?? { percent: 0 }), paused: false, status: 'Starting…' } })) - window.electron.model.download(fullId).then((result) => { + setDownloading((prev) => ({ ...prev, [fullId]: { ...(prev[fullId] ?? { percent: 0 }), variantId, paused: false, status: 'Starting…' } })) + window.electron.model.download(fullId, variantId).then((result) => { if (!result.success && !result.paused && !result.cancelled) { setGhErr(result.error ?? 'Download failed') setDownloading((prev) => { const n = { ...prev }; delete n[fullId]; return n }) @@ -183,7 +185,7 @@ export default function ModelsPage(): JSX.Element { for (const node of ext.nodes) { if (!nodeHasManagedWeights(node)) continue const fullId = `${ext.id}/${node.id}` - if (installedVariantIds.includes(fullId) || downloading[fullId]) continue + if (installedNodeIds.includes(fullId) || downloading[fullId]) continue handleInstallNode(node, fullId) } } @@ -205,6 +207,12 @@ export default function ModelsPage(): JSX.Element { refreshInstalledIds(useExtensionsStore.getState().modelExtensions) } + async function handleDeleteWeightVariant(fullId: string, variantId: string) { + const result = await window.electron.model.deleteWeightVariant(fullId, variantId) + await refreshInstalledIds(useExtensionsStore.getState().modelExtensions) + return result + } + // ── GitHub extension install ─────────────────────────────────────────────── async function handleGHInstall() { @@ -322,7 +330,7 @@ export default function ModelsPage(): JSX.Element { } const cardHandlers = { - installedIds: installedVariantIds, + installedIds: installedNodeIds, downloading, disabled: isBusy, onInstall: handleInstallNode, @@ -627,7 +635,7 @@ export default function ModelsPage(): JSX.Element { {selectedExt && ( openUninstallModal(extId)} onRepaired={() => reloadExtensions()} onSynced={() => reloadExtensions()} diff --git a/src/areas/models/components/ExtensionCard.tsx b/src/areas/models/components/ExtensionCard.tsx index 7b0da394..5ceb729a 100644 --- a/src/areas/models/components/ExtensionCard.tsx +++ b/src/areas/models/components/ExtensionCard.tsx @@ -125,7 +125,12 @@ export function ExtensionCard({ className="flex items-center justify-between gap-2.5 px-2.5 py-1.5 rounded-lg bg-white/[0.02] border border-zinc-800" >
- {node.name} + + {node.name} + {isModel && node.weightVariants && ( + · {node.weightVariants.options.length} variants + )} +
{isModel && ( diff --git a/src/areas/models/components/ExtensionDrawer.tsx b/src/areas/models/components/ExtensionDrawer.tsx index 2f154c71..288f12b1 100644 --- a/src/areas/models/components/ExtensionDrawer.tsx +++ b/src/areas/models/components/ExtensionDrawer.tsx @@ -1,11 +1,13 @@ import { useEffect, useState } from 'react' import type { AnyExtension, ExtensionNode } from '@shared/types/electron.d' +import { useExtensionsStore } from '@shared/stores/extensionsStore' import { useNavStore } from '@shared/stores/navStore' import { DownloadMap, ICONS, IOBadge, NodeInstallControl, + NodeUiState, TypePill, extInstallSummary, formatBytes, @@ -20,11 +22,12 @@ interface Props { downloading: DownloadMap loadError?: string disabled?: boolean - onInstall: (node: ExtensionNode, fullId: string) => void + onInstall: (node: ExtensionNode, fullId: string, variantId?: string) => void onInstallAll: (ext: AnyExtension) => void onPauseDownload: (fullId: string) => void onCancelDownload: (fullId: string) => void onUninstallNode: (fullId: string) => void + onDeleteWeightVariant: (fullId: string, variantId: string) => Promise<{ success: boolean; error?: string }> onUninstall: (extId: string) => void onRepaired: () => void | Promise onSynced: () => void @@ -34,13 +37,15 @@ interface Props { export function ExtensionDrawer({ ext, installedIds, localDataIds, downloading, loadError, disabled, onInstall, onInstallAll, onPauseDownload, onCancelDownload, - onUninstallNode, onUninstall, onRepaired, onSynced, onClose, + onUninstallNode, onDeleteWeightVariant, onUninstall, onRepaired, onSynced, onClose, }: Props): JSX.Element { const navigate = useNavStore((s) => s.navigate) + const installedWeightVariants = useExtensionsStore((s) => s.installedWeightVariants) const [repairing, setRepairing] = useState(false) const [repairError, setRepairError] = useState(null) const [syncing, setSyncing] = useState(false) const [syncError, setSyncError] = useState(null) + const [variantError, setVariantError] = useState(null) const isModel = ext.type === 'model' // Built-ins are corrupted-flagged too (builtin-sync repairs them on restart), @@ -85,7 +90,13 @@ export function ExtensionDrawer({ } } - const error = syncError ?? repairError ?? loadError + async function handleDeleteWeightVariant(fullId: string, variantId: string) { + setVariantError(null) + const result = await onDeleteWeightVariant(fullId, variantId) + if (!result.success) setVariantError(result.error ?? 'Could not remove these weights') + } + + const error = variantError ?? syncError ?? repairError ?? loadError return ( <> @@ -178,12 +189,19 @@ export function ExtensionDrawer({ const fullId = `${ext.id}/${node.id}` const state = getNodeState(ext.id, node, installedIds, downloading) const dl = state.kind === 'downloading' ? state.dl : null + const variants = isModel ? node.weightVariants : undefined + // A download started without a variant (Install all) fetches the default one. + const dlVariantId = dl && variants ? (dl.variantId ?? variants.default) : undefined + const dlVariant = variants?.options.find((option) => option.id === dlVariantId) + const installedLabels = variants?.options + .filter((option) => installedWeightVariants[fullId]?.includes(option.id)) + .map((option) => option.label) ?? [] const sub = state.kind === 'ready' ? 'Available on the node graph' - : state.kind === 'installed' ? 'Installed' + : state.kind === 'installed' ? (installedLabels.length > 0 ? `${installedLabels.join(', ')} installed` : 'Installed') : state.kind === 'available' ? 'Not installed' : dl?.paused ? 'Download paused' - : `Downloading… ${dl?.percent ?? 0}%` + : `Downloading${dlVariant ? ` ${dlVariant.label}` : ''}… ${dl?.percent ?? 0}%` return (
@@ -196,14 +214,16 @@ export function ExtensionDrawer({ {isModel && (
- onInstall(node, fullId)} - onPause={() => onPauseDownload(fullId)} - onResume={() => onInstall(node, fullId)} - onCancel={() => onCancelDownload(fullId)} - /> + {!variants && ( + onInstall(node, fullId)} + onPause={() => onPauseDownload(fullId)} + onResume={() => onInstall(node, fullId)} + onCancel={() => onCancelDownload(fullId)} + /> + )} {localDataIds.includes(fullId) && state.kind !== 'downloading' && (
+ {/* Weight variants — each one installs and deletes on its own */} + {variants && ( +
+ {variants.options.map((option) => { + const installed = installedWeightVariants[fullId]?.includes(option.id) ?? false + const variantDl = dlVariantId === option.id ? dl : null + const variantState: NodeUiState = variantDl + ? { kind: 'downloading', dl: variantDl } + : installed ? { kind: 'installed' } : { kind: 'available' } + return ( +
+
+ {option.label} + {option.sizeGb !== undefined && ( + {option.sizeGb} GB + )} + {option.vramGb !== undefined && ( + ~{option.vramGb} GB VRAM + )} + {option.id === variants.default && ( + default + )} +
+
+ onInstall(node, fullId, option.id)} + onPause={() => onPauseDownload(fullId)} + onResume={() => onInstall(node, fullId, option.id)} + onCancel={() => onCancelDownload(fullId)} + /> + {installed && !dl && ( + + )} +
+
+ ) + })} +
+ )} + {/* Download detail */} {dl && (
diff --git a/src/areas/models/components/extensionShared.tsx b/src/areas/models/components/extensionShared.tsx index 8bf865e3..e5682c55 100644 --- a/src/areas/models/components/extensionShared.tsx +++ b/src/areas/models/components/extensionShared.tsx @@ -12,6 +12,7 @@ export interface DownloadInfo { totalBytes?: number stalledSeconds?: number paused?: boolean + variantId?: string // set when the download targets one weight variant of the node } export type DownloadMap = Record diff --git a/src/areas/workflows/mockExtensions.ts b/src/areas/workflows/mockExtensions.ts index 2bbc8cc6..1eab87bb 100644 --- a/src/areas/workflows/mockExtensions.ts +++ b/src/areas/workflows/mockExtensions.ts @@ -1,6 +1,6 @@ import type { ModelExtension, ProcessExtension } from '@shared/stores/extensionsStore' export type { ParamSchema } from '@shared/types/electron.d' -import type { ParamSchema } from '@shared/types/electron.d' +import type { ParamSchema, WeightVariantsInfo } from '@shared/types/electron.d' export interface WorkflowExtension { id: string // "ext_id/node_id" @@ -17,6 +17,7 @@ export interface WorkflowExtension { params: ParamSchema[] builtin: boolean type: 'model' | 'process' + weightVariants?: WeightVariantsInfo } function applyParamDefaults( @@ -77,6 +78,7 @@ export function buildAllWorkflowExtensions( params: applyParamDefaults(node.paramsSchema as ParamSchema[], node.paramDefaults), builtin: ext.builtin, type: 'model', + weightVariants: node.weightVariants, }) } } diff --git a/src/areas/workflows/nodes/ExtensionNode.tsx b/src/areas/workflows/nodes/ExtensionNode.tsx index abe2f9f6..ec5f83d9 100644 --- a/src/areas/workflows/nodes/ExtensionNode.tsx +++ b/src/areas/workflows/nodes/ExtensionNode.tsx @@ -1,11 +1,13 @@ import { useCallback, useEffect, useRef, useLayoutEffect, useState } from 'react' import { Handle, Position, useReactFlow } from '@xyflow/react' import { useExtensionsStore } from '@shared/stores/extensionsStore' +import { useNavStore } from '@shared/stores/navStore' import { buildAllWorkflowExtensions } from '../mockExtensions' import type { ParamSchema } from '../mockExtensions' import type { WFNodeData } from '@shared/types/electron.d' import { PICKER_LABELS, openParamPicker, resolvePickerIntent } from '@shared/utils/paramPicker' import { PickerIcon } from '@shared/components/ui' +import { isMissingWeightVariant, withWeightVariantAvailability } from '@shared/utils/weightVariants' import { useWorkflowRunStore } from '../workflowRunStore' import BaseNode from './BaseNode' @@ -167,8 +169,10 @@ export default function ExtensionNode({ id, data, selected }: { id: string; data const [handleTops, setHandleTops] = useState([]) const { modelExtensions, processExtensions } = useExtensionsStore() + const installedVariants = useExtensionsStore((s) => (data.extensionId ? s.installedWeightVariants[data.extensionId] : undefined)) const allExtensions = buildAllWorkflowExtensions(modelExtensions, processExtensions) const ext = allExtensions.find((e) => e.id === data.extensionId) + const openExtension = useNavStore((s) => s.openExtension) const inputs = ext?.inputs // defined → multi-input mode const isMulti = inputs && inputs.length > 1 @@ -313,7 +317,15 @@ export default function ExtensionNode({ id, data, selected }: { id: string; data
- patchParam(param.id, v)} resolvedParams={resolvedParams} /> + { + patchParam(param.id, v) + if (ext && isMissingWeightVariant(param.id, v, ext.weightVariants, installedVariants)) openExtension(ext.extensionId) + }} + resolvedParams={resolvedParams} + />
) diff --git a/src/shared/stores/extensionsStore.ts b/src/shared/stores/extensionsStore.ts index 91657a0a..ada611f2 100644 --- a/src/shared/stores/extensionsStore.ts +++ b/src/shared/stores/extensionsStore.ts @@ -24,8 +24,11 @@ interface ExtensionsStore { installProgress: InstallProgress | null installError: string | null loadErrors: Record + /** Installed weight variant ids, keyed by "ext_id/node_id", for nodes that declare variants */ + installedWeightVariants: Record loadExtensions: () => Promise + refreshInstalledWeightVariants: () => Promise installFromGitHub: (url: string) => Promise<{ success: boolean; error?: string }> installFromLocal: () => Promise<{ success: boolean; error?: string; cancelled?: boolean; needsRepair?: boolean }> uninstall: (extensionId: string) => Promise<{ success: boolean; error?: string }> @@ -54,6 +57,7 @@ export const useExtensionsStore = create((set, get) => ({ installProgress: null, installError: null, loadErrors: {}, + installedWeightVariants: {}, // ── Load list ────────────────────────────────────────────────────────────── @@ -66,11 +70,30 @@ export const useExtensionsStore = create((set, get) => ({ ...extensions, loading: false, }) + await get().refreshInstalledWeightVariants() } catch { set({ loading: false }) } }, + async refreshInstalledWeightVariants() { + const entries = await Promise.all( + get().modelExtensions.flatMap((ext) => ext.nodes + .filter((node) => node.weightVariants) + .map(async (node) => { + const fullId = `${ext.id}/${node.id}` + return [fullId, await window.electron.model.installedWeightVariants(fullId)] as const + })), + ) + // A node whose state could not be read stays absent: undefined reads as "unknown", + // which the UI keeps neutral, while [] would claim no variant is installed. + set({ + installedWeightVariants: Object.fromEntries( + entries.filter((entry): entry is readonly [string, string[]] => entry[1] !== null), + ), + }) + }, + // ── Install from GitHub ──────────────────────────────────────────────────── async installFromGitHub(url: string) { diff --git a/src/shared/stores/navStore.ts b/src/shared/stores/navStore.ts index c6854424..62afa2f5 100644 --- a/src/shared/stores/navStore.ts +++ b/src/shared/stores/navStore.ts @@ -4,10 +4,16 @@ export type Page = 'generate' | 'workflows' | 'models' | 'settings' interface NavState { currentPage: Page + extensionToOpen: string | null navigate: (page: Page) => void + openExtension: (extensionId: string) => void + clearExtensionToOpen: () => void } export const useNavStore = create((set) => ({ currentPage: 'generate', - navigate: (page) => set({ currentPage: page }) + extensionToOpen: null, + navigate: (page) => set({ currentPage: page }), + openExtension: (extensionId) => set({ currentPage: 'models', extensionToOpen: extensionId }), + clearExtensionToOpen: () => set({ extensionToOpen: null }), })) diff --git a/src/shared/types/electron.d.ts b/src/shared/types/electron.d.ts index 840c5e76..26f9bf3b 100644 --- a/src/shared/types/electron.d.ts +++ b/src/shared/types/electron.d.ts @@ -25,6 +25,13 @@ export interface ExtensionNode { hfSkipPrefixes?: string[] hfIncludePrefixes?: string[] hasModelSources?: boolean + weightVariants?: WeightVariantsInfo +} + +export interface WeightVariantsInfo { + param: string // params_schema id whose value selects the variant + default: string + options: { id: string; label: string; sizeGb?: number; vramGb?: number }[] } export interface ModelExtension { @@ -207,17 +214,21 @@ declare global { model: { export: (args: { outputUrl: string; format: string }) => Promise<{ success: boolean; error?: string }> listDownloaded: () => Promise<{ id: string; name: string; size_gb: number }[]> - activeDownloads: () => Promise<{ modelId: string; percent: number; file?: string; fileIndex?: number; totalFiles?: number }[]> + activeDownloads: () => Promise<{ modelId: string; variantId?: string; percent: number; file?: string; fileIndex?: number; totalFiles?: number }[]> isDownloaded: (modelId: string) => Promise hasLocalData: (modelId: string) => Promise - download: (modelId: string) => Promise<{ success: boolean; error?: string; paused?: boolean; cancelled?: boolean }> + download: (modelId: string, variantId?: string) => Promise<{ success: boolean; error?: string; paused?: boolean; cancelled?: boolean }> pauseDownload: (modelId: string) => Promise<{ success: boolean; error?: string }> cancelDownload: (modelId: string) => Promise<{ success: boolean; error?: string }> delete: (modelId: string) => Promise<{ success: boolean; error?: string }> + /** Installed variant ids, or null when the node's install state could not be read */ + installedWeightVariants: (modelId: string) => Promise + deleteWeightVariant: (modelId: string, variantId: string) => Promise<{ success: boolean; error?: string }> unloadAll: () => Promise<{ success: boolean; error?: string }> showInFolder: (modelId: string) => Promise onProgress: (cb: (data: { modelId: string + variantId?: string percent: number file?: string fileIndex?: number diff --git a/src/shared/utils/weightVariants.test.mjs b/src/shared/utils/weightVariants.test.mjs new file mode 100644 index 00000000..1659391b --- /dev/null +++ b/src/shared/utils/weightVariants.test.mjs @@ -0,0 +1,63 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-weight-variants-test-')), 'weightVariants.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('src/shared/utils/weightVariants.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { withWeightVariantAvailability, isMissingWeightVariant } = loadModule() + +const variants = { + param: 'gguf_quant', + default: 'Q5_K_M', + options: [{ id: 'Q4_K_M', label: 'Q4_K_M' }, { id: 'Q5_K_M', label: 'Q5_K_M' }], +} + +const quantParam = { + id: 'gguf_quant', + label: 'Quantization', + type: 'select', + default: 'Q5_K_M', + options: [ + { value: 'Q4_K_M', label: 'Q4_K_M' }, + { value: 'Q5_K_M', label: 'Q5_K_M' }, + { value: 'auto', label: 'Auto' }, + ], +} + +test('labels declared variants that are not installed and keeps their values', () => { + const marked = withWeightVariantAvailability(quantParam, variants, ['Q5_K_M']) + assert.deepEqual(marked.options.map((option) => option.label), ['Q4_K_M (not installed)', 'Q5_K_M', 'Auto']) + assert.deepEqual(marked.options.map((option) => option.value), ['Q4_K_M', 'Q5_K_M', 'auto']) +}) + +test('returns the param untouched when availability is unknown or the param selects nothing', () => { + assert.equal(withWeightVariantAvailability(quantParam, variants, undefined), quantParam) + assert.equal(withWeightVariantAvailability(quantParam, undefined, []), quantParam) + const steps = { id: 'steps', label: 'Steps', type: 'int', default: 25 } + assert.equal(withWeightVariantAvailability(steps, variants, []), steps) +}) + +test('flags a selected variant only when availability is known and it is not installed', () => { + assert.equal(isMissingWeightVariant('gguf_quant', 'Q4_K_M', variants, ['Q5_K_M']), true) + assert.equal(isMissingWeightVariant('gguf_quant', 'Q5_K_M', variants, ['Q5_K_M']), false) + assert.equal(isMissingWeightVariant('gguf_quant', 'auto', variants, []), false) + assert.equal(isMissingWeightVariant('steps', 'Q4_K_M', variants, []), false) + assert.equal(isMissingWeightVariant('gguf_quant', 'Q4_K_M', variants, undefined), false) + assert.equal(isMissingWeightVariant('gguf_quant', 'Q4_K_M', undefined, []), false) +}) diff --git a/src/shared/utils/weightVariants.ts b/src/shared/utils/weightVariants.ts new file mode 100644 index 00000000..252fd6d5 --- /dev/null +++ b/src/shared/utils/weightVariants.ts @@ -0,0 +1,35 @@ +import type { ParamSchema, WeightVariantsInfo } from '@shared/types/electron.d' + +/** + * Suffix the options of a node's variant-selecting param that are declared + * weight variants but not installed. Other params, and values that are not + * variant ids, are returned untouched. `installed` undefined means not known yet. + */ +export function withWeightVariantAvailability( + param: ParamSchema, + variants: WeightVariantsInfo | undefined, + installed: string[] | undefined, +): ParamSchema { + if (!variants || !installed || param.id !== variants.param || !param.options) return param + const variantIds = new Set(variants.options.map((option) => option.id)) + return { + ...param, + options: param.options.map((option) => { + const value = String(option.value) + if (!variantIds.has(value) || installed.includes(value)) return option + return { ...option, label: `${option.label ?? value} (not installed)` } + }), + } +} + +/** True when `value` selects a declared weight variant that is known not to be installed. */ +export function isMissingWeightVariant( + paramId: string, + value: unknown, + variants: WeightVariantsInfo | undefined, + installed: string[] | undefined, +): boolean { + if (!variants || !installed || paramId !== variants.param) return false + const id = String(value) + return variants.options.some((option) => option.id === id) && !installed.includes(id) +} From a0ade5329732275c08adb09a0164132c1bf3ea5a Mon Sep 17 00:00:00 2001 From: Lightning Pixel Date: Sat, 19 Sep 2026 17:27:27 +0200 Subject: [PATCH 32/57] fix(slicer): slice imported meshes, and stop rotating and rescaling blindly The slicer route assumed every source was a unit-sized, Y-up glTF straight from a generator, but the Export action is reachable for any mesh in the viewer. - Imported meshes (served through /optimize/serve-file) were excluded outright, so the most direct "I have a model, slice it" path offered no action at all and gave no hint why. They are now sliceable via an exact-membership registry of the files the user picked themselves this session, which keeps the route closed to arbitrary absolute paths rather than widening its path guard. - The Y->Z rotation now applies only to glTF sources. STL/OBJ/PLY are already Z-up, and import converts them to GLB without touching the axes, so the original extension decides -- not the container's. - Normalising to 50 mm now happens only for unit-sized meshes. A mesh that already carries a real-world size is the user's own, and silently shrinking a 180 mm part would waste a print. - slicer:open no longer claims it can detect a missing OrcaSlicer: on Windows an unregistered scheme still makes ShellExecuteEx succeed, so that error branch could never run. test_normalizes_longest_edge_to_default_print_size encoded the unconditional rescale, so its fixture becomes a unit-sized mesh, matching real generator output. --- api/routers/export.py | 96 ++++++++++++++++++----- api/routers/optimize.py | 6 ++ api/services/imported_sources.py | 53 +++++++++++++ api/tests/test_export_router.py | 84 +++++++++++++++++++- electron/main/ipc-handlers.ts | 9 ++- src/areas/generate/GeneratePage.tsx | 4 +- src/areas/generate/orcaSlicerLink.test.ts | 35 ++++++++- src/areas/generate/orcaSlicerLink.ts | 51 ++++++++++-- 8 files changed, 303 insertions(+), 35 deletions(-) create mode 100644 api/services/imported_sources.py diff --git a/api/routers/export.py b/api/routers/export.py index f40a9045..6cd03d61 100644 --- a/api/routers/export.py +++ b/api/routers/export.py @@ -2,11 +2,13 @@ import binascii import io import math +from pathlib import Path import trimesh from fastapi import APIRouter, HTTPException from fastapi.responses import Response, FileResponse +from services import imported_sources from services.generator_registry import WORKSPACE_DIR router = APIRouter(tags=["export"]) @@ -26,6 +28,15 @@ # in OrcaSlicer as needed. DEFAULT_PRINT_LONGEST_MM = 50.0 +# ...but this route also serves meshes the user authored or imported, which DO +# carry a real-world size. Silently resizing a 180 mm part down to 50 mm wastes +# a print, so only rescale what is small enough to be unit-sized AI output. +UNIT_SCALE_MAX = 5.0 + +# Source formats whose up-axis is Y (the glTF convention). Everything else this +# route accepts — STL, OBJ, PLY — is conventionally Z-up already. +GLTF_SUFFIXES = {".glb", ".gltf"} + def _to_single_mesh(loaded: object) -> "trimesh.Trimesh": """Flatten a loaded GLB into one Trimesh, baking scene-graph node transforms. @@ -64,6 +75,62 @@ def _scale_to_print_size(mesh: "trimesh.Trimesh", longest_mm: float = DEFAULT_PR mesh.apply_scale(longest_mm / longest) +def _normalize_print_scale(mesh: "trimesh.Trimesh") -> bool: + """Rescale ``mesh`` only if it looks unit-sized; return whether it was rescaled. + + A mesh whose longest edge already exceeds ``UNIT_SCALE_MAX`` is assumed to + carry a real-world size the user chose, and is left untouched. + """ + extents = mesh.extents + longest = float(max(extents)) if extents is not None and len(extents) else 0.0 + if not math.isfinite(longest) or not (1e-9 < longest <= UNIT_SCALE_MAX): + return False + _scale_to_print_size(mesh) + return True + + +def _resolve_slicer_source(token: str) -> tuple[Path, str]: + """Decode ``token`` into an existing source file and its ORIGINAL suffix. + + Two kinds of source are accepted: + + * a workspace-relative path, confined to the workspace by ancestry; + * an absolute path, but ONLY when the user imported that exact file this + session (see ``services.imported_sources``). Membership there is an + equality test on the resolved path, so this grants no traversal and does + not widen the route to arbitrary disk paths. + + The returned suffix is the format the user actually supplied — for an import + that is the pre-conversion extension, which is what decides the up-axis. + """ + try: + padded = token + "=" * (-len(token) % 4) + decoded = base64.urlsafe_b64decode(padded.encode("ascii")).decode("utf-8") + except (binascii.Error, UnicodeDecodeError, ValueError): + raise HTTPException(400, "Malformed source token") + + candidate = Path(decoded) + if candidate.is_absolute(): + original_suffix = imported_sources.source_suffix(candidate) + if original_suffix is None: + raise HTTPException(400, "Invalid path") + full_path = candidate.resolve() + if not full_path.is_file(): + raise HTTPException(404, f"File not found: {decoded}") + return full_path, original_suffix + + # Containment check via ancestry, not string prefix: `startswith` would let a + # sibling like `-other/...` slip through, and `..` escapes resolve + # outside the workspace and fail this check. + workspace = WORKSPACE_DIR.resolve() + full_path = (workspace / decoded).resolve() + if full_path != workspace and workspace not in full_path.parents: + raise HTTPException(400, "Invalid path") + if not full_path.is_file(): + raise HTTPException(404, f"File not found: {decoded}") + return full_path, full_path.suffix.lower() + + @router.get("/slicer/{fmt}/{token}/{filename}") def export_for_slicer(fmt: str, token: str, filename: str): """Serve a generated GLB converted to a slicer-importable mesh, at a URL @@ -74,8 +141,10 @@ def export_for_slicer(fmt: str, token: str, filename: str): downloads the URL and derives the import filename — and therefore the mesh format — from the URL's FINAL path segment, so a query string (``?path=...``) would corrupt the parsed extension and the model would silently fail to - import. ``token`` is the url-safe-base64 of the workspace-relative source - path; ``filename`` (e.g. ``model.stl``) is what OrcaSlicer names the download. + import. ``token`` is the url-safe-base64 of the source path — workspace- + relative, or absolute for a file the user imported this session (see + ``_resolve_slicer_source``); ``filename`` (e.g. ``model.stl``) is what + OrcaSlicer names the download. """ fmt = fmt.lower() if fmt not in SLICER_FORMATS: @@ -83,28 +152,17 @@ def export_for_slicer(fmt: str, token: str, filename: str): if not filename.lower().endswith(f".{fmt}"): raise HTTPException(400, "Filename must end with the requested format extension") - try: - padded = token + "=" * (-len(token) % 4) - rel_path = base64.urlsafe_b64decode(padded.encode("ascii")).decode("utf-8") - except (binascii.Error, UnicodeDecodeError, ValueError): - raise HTTPException(400, "Malformed source token") - - # Containment check via ancestry, not string prefix: `startswith` would let a - # sibling like `-other/...` slip through, and `..` escapes resolve - # outside the workspace and fail this check. - workspace = WORKSPACE_DIR.resolve() - full_path = (workspace / rel_path).resolve() - if full_path != workspace and workspace not in full_path.parents: - raise HTTPException(400, "Invalid path") - if not full_path.is_file(): - raise HTTPException(404, f"File not found: {rel_path}") + full_path, source_suffix = _resolve_slicer_source(token) mesh = _to_single_mesh(trimesh.load(str(full_path))) # glTF/GLB is Y-up; OrcaSlicer's world is Z-up. Rotate +90° about X so the # model imports standing upright instead of on its side. (Modly's own viewer # rests generated meshes on the Y=0 plane, confirming Y is the up axis.) - mesh.apply_transform(trimesh.transformations.rotation_matrix(math.pi / 2, [1, 0, 0])) - _scale_to_print_size(mesh) + # STL/OBJ/PLY sources are already Z-up, so rotating them would do the very + # thing this corrects — lay an upright model on its side. + if source_suffix in GLTF_SUFFIXES: + mesh.apply_transform(trimesh.transformations.rotation_matrix(math.pi / 2, [1, 0, 0])) + _normalize_print_scale(mesh) data = mesh.export(file_type=fmt) if isinstance(data, str): diff --git a/api/routers/optimize.py b/api/routers/optimize.py index 6081c704..de1afa77 100644 --- a/api/routers/optimize.py +++ b/api/routers/optimize.py @@ -21,6 +21,7 @@ from urllib.parse import quote from pydantic import BaseModel +from services import imported_sources from services.generator_registry import WORKSPACE_DIR router = APIRouter(tags=["optimize"]) @@ -420,6 +421,7 @@ async def import_mesh_by_path(body: ImportByPathRequest): if ext == "glb": # Serve the original file directly — no copy + imported_sources.register(file_path, file_path) return {"url": f"/optimize/serve-file?path={quote(str(file_path))}"} # Mesh ply / obj / stl: convert to GLB in a temp directory (not the workspace) @@ -427,6 +429,10 @@ async def import_mesh_by_path(body: ImportByPathRequest): output_path = os.path.join(tmp_dir, "mesh.glb") loaded = trimesh.load(str(file_path)) loaded.export(output_path) + # Remember the ORIGINAL extension: this conversion changes the container but + # not the axes, so a .stl imported here still holds Z-up data despite the + # .glb suffix, and must not be rotated as if it were glTF. + imported_sources.register(output_path, file_path) return {"url": f"/optimize/serve-file?path={quote(output_path)}"} diff --git a/api/services/imported_sources.py b/api/services/imported_sources.py new file mode 100644 index 00000000..a21dded5 --- /dev/null +++ b/api/services/imported_sources.py @@ -0,0 +1,53 @@ +"""Registry of meshes the user explicitly imported from outside the workspace. + +`/optimize/import-by-path` serves files the user picked through the OS file +dialog, either in place or converted into a temp dir — never from the workspace. +The slicer export route confines itself to the workspace by design, so those +imports would be unsliceable without widening that guard to arbitrary absolute +paths, which would be a real regression. + +This registry is the narrow alternative: a route may accept an absolute path +only if it is an EXACT member here, i.e. a file the user chose themselves in +this session. Membership is an equality test, never a prefix test, so it grants +no traversal. + +It also remembers each file's ORIGINAL suffix. An imported `.stl` is converted +to GLB on import without any axis change, so the container extension alone would +mislead a consumer into applying the glTF Y-up -> Z-up rotation to data that is +already Z-up. + +Session-scoped and in-memory: it is emptied when the backend restarts, so a +deeplink replayed after a restart is rejected rather than silently served. Links +are consumed within a second of the click, so this is not worth persisting. +""" + +from pathlib import Path + +# resolved absolute path (str) -> original suffix, lowercase, with the dot +_SOURCES: dict[str, str] = {} + + +def register(served_path: str | Path, original_path: str | Path) -> str: + """Record a user-imported file as sliceable and return its resolved path. + + ``served_path`` is what gets served (possibly a converted temp GLB); + ``original_path`` is the file the user actually picked, whose suffix decides + the source format. + """ + resolved = str(Path(served_path).resolve()) + _SOURCES[resolved] = Path(original_path).suffix.lower() + return resolved + + +def source_suffix(path: str | Path) -> str | None: + """Original suffix of a registered import, or None if it is not registered.""" + return _SOURCES.get(str(Path(path).resolve())) + + +def is_registered(path: str | Path) -> bool: + return source_suffix(path) is not None + + +def clear() -> None: + """Drop every entry — for tests.""" + _SOURCES.clear() diff --git a/api/tests/test_export_router.py b/api/tests/test_export_router.py index bc07188f..0972e836 100644 --- a/api/tests/test_export_router.py +++ b/api/tests/test_export_router.py @@ -13,6 +13,7 @@ import trimesh import routers.export as export_router + from services import imported_sources HAVE_TRIMESH = True except Exception: # noqa: BLE001 @@ -34,9 +35,10 @@ def setUp(self) -> None: self.workspace = Path(self._tmp.name).resolve() self._orig_workspace = export_router.WORKSPACE_DIR export_router.WORKSPACE_DIR = self.workspace - # A box that is tallest along Y (glTF up-axis). Exported to GLB, it - # reloads as a Scene so the flatten path is exercised too. - box = trimesh.creation.box(extents=[10.0, 30.0, 10.0]) + # A box that is tallest along Y (glTF up-axis) and unit-sized, matching + # what image-to-3D generators emit. Exported to GLB, it reloads as a + # Scene so the flatten path is exercised too. + box = trimesh.creation.box(extents=[0.3, 1.0, 0.3]) self.rel = "Workflows/hero.glb" (self.workspace / "Workflows").mkdir(parents=True, exist_ok=True) box.export(str(self.workspace / self.rel)) @@ -64,6 +66,24 @@ def test_normalizes_longest_edge_to_default_print_size(self) -> None: longest = float(max(_load_stl(resp).extents)) self.assertAlmostEqual(longest, export_router.DEFAULT_PRINT_LONGEST_MM, places=3) + def test_leaves_a_real_world_sized_mesh_alone(self) -> None: + # A mesh that already carries a physical size is the user's own: silently + # shrinking a 180 mm part to 50 mm would waste a print. + big = trimesh.creation.box(extents=[40.0, 180.0, 40.0]) + rel = "Workflows/part.glb" + big.export(str(self.workspace / rel)) + resp = export_router.export_for_slicer("stl", _token(rel), "model.stl") + self.assertAlmostEqual(float(max(_load_stl(resp).extents)), 180.0, places=2) + + def test_does_not_rotate_a_z_up_stl_source(self) -> None: + # STL is conventionally Z-up already; the glTF Y->Z rotation would lay an + # upright model on its side. + rel = "Workflows/upright.stl" + trimesh.creation.box(extents=[10.0, 10.0, 30.0]).export(str(self.workspace / rel)) + resp = export_router.export_for_slicer("stl", _token(rel), "model.stl") + ex = _load_stl(resp).extents + self.assertEqual(int(np.argmax(ex)), 2, f"expected Z to stay the tallest axis, got extents {ex}") + def test_rejects_unsupported_format(self) -> None: with self.assertRaises(HTTPException) as ctx: export_router.export_for_slicer("glb", _token(self.rel), "model.glb") @@ -126,5 +146,63 @@ def test_scale_ignores_degenerate_mesh(self) -> None: export_router._scale_to_print_size(mesh) # must not raise +@unittest.skipUnless(HAVE_TRIMESH, "trimesh not installed") +class ImportedSourceSlicerTests(unittest.TestCase): + """Meshes imported from outside the workspace are sliceable — but only the + exact files the user picked, and never with the wrong up-axis.""" + + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self.outside = Path(self._tmp.name).resolve() + self._orig_workspace = export_router.WORKSPACE_DIR + # A workspace elsewhere, so nothing here is reachable as a relative path. + self._ws_tmp = tempfile.TemporaryDirectory() + export_router.WORKSPACE_DIR = Path(self._ws_tmp.name).resolve() + imported_sources.clear() + # Unit-sized and tallest along Y, as a glTF export would be. + self.mesh_path = self.outside / "imported.glb" + trimesh.creation.box(extents=[0.3, 1.0, 0.3]).export(str(self.mesh_path)) + + def tearDown(self) -> None: + export_router.WORKSPACE_DIR = self._orig_workspace + imported_sources.clear() + self._tmp.cleanup() + self._ws_tmp.cleanup() + + def test_rejects_an_absolute_path_that_was_never_imported(self) -> None: + # The whole point of the registry: an absolute path alone buys nothing. + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + self.assertEqual(ctx.exception.status_code, 400) + + def test_serves_a_registered_import(self) -> None: + imported_sources.register(self.mesh_path, self.mesh_path) + resp = export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + self.assertEqual(resp.media_type, "model/stl") + self.assertGreater(len(_load_stl(resp).faces), 0) + + def test_rotates_a_registered_gltf_import(self) -> None: + imported_sources.register(self.mesh_path, self.mesh_path) + resp = export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + ex = _load_stl(resp).extents + self.assertEqual(int(np.argmax(ex)), 2, f"expected Z to be the tallest axis, got extents {ex}") + + def test_does_not_rotate_an_import_that_was_stl_before_conversion(self) -> None: + # `import-by-path` converts STL/OBJ/PLY to GLB without touching the axes, + # so the .glb container here still holds Z-up data. Rotating it would be + # exactly the bug the rotation exists to prevent. + imported_sources.register(self.mesh_path, self.outside / "original.stl") + resp = export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + ex = _load_stl(resp).extents + self.assertEqual(int(np.argmax(ex)), 1, f"expected Y to stay the tallest axis, got extents {ex}") + + def test_registered_but_deleted_file_is_404(self) -> None: + imported_sources.register(self.mesh_path, self.mesh_path) + self.mesh_path.unlink() + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + self.assertEqual(ctx.exception.status_code, 404) + + if __name__ == "__main__": unittest.main() diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index f097a559..3c027340 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -596,8 +596,13 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe ipcMain.handle('shell:openExternal', (_, url: string) => shell.openExternal(url)) // Open a model in OrcaSlicer via its orcaslicer://open?file= deeplink. - // Returns success/error so the renderer can surface a fallback (e.g. when - // OrcaSlicer is not installed and no app is registered for the scheme). + // + // The returned error only covers the shell refusing the call outright. It is + // NOT an install check: on Windows an unregistered scheme still makes + // ShellExecuteEx succeed — the OS shows its own "You'll need a new app to open + // this orcaslicer link" dialog and this resolves with success. Detecting a + // missing OrcaSlicer would take a per-platform handler probe (registry on + // Windows), so the renderer must not promise the user that it knows. ipcMain.handle('slicer:open', async (_, url: string): Promise<{ success: boolean; error?: string }> => { if (typeof url !== 'string' || !url.startsWith('orcaslicer://')) { return { success: false, error: 'slicer:open requires an orcaslicer:// URL' } diff --git a/src/areas/generate/GeneratePage.tsx b/src/areas/generate/GeneratePage.tsx index f970993c..87c7a18d 100644 --- a/src/areas/generate/GeneratePage.tsx +++ b/src/areas/generate/GeneratePage.tsx @@ -689,7 +689,9 @@ export default function GeneratePage(): JSX.Element { const link = buildOrcaSlicerDeepLink(apiUrl, currentJob.outputUrl) const result = await window.electron.slicer.open(link) if (!result.success) { - showError(result.error ?? 'Could not open OrcaSlicer. Make sure it is installed.') + // Deliberately not "make sure it is installed": the main process cannot + // tell a missing OrcaSlicer from a working one (see slicer:open). + showError(result.error ?? 'Could not open OrcaSlicer.') } } catch (err) { showError(err instanceof Error ? err.message : 'Could not open OrcaSlicer.') diff --git a/src/areas/generate/orcaSlicerLink.test.ts b/src/areas/generate/orcaSlicerLink.test.ts index 2f7e47d3..514618eb 100644 --- a/src/areas/generate/orcaSlicerLink.test.ts +++ b/src/areas/generate/orcaSlicerLink.test.ts @@ -38,12 +38,43 @@ test('strips a trailing slash from the api origin', () => { assert.equal(modelUrl, `http://localhost:8765/export/slicer/stl/${encodeWorkspacePathToken('a.glb')}/model.stl`) }) -test('canOpenInOrcaSlicer accepts workspace meshes and rejects splats, imports, and empty', () => { +test('canOpenInOrcaSlicer accepts sliceable workspace meshes and rejects splats and empty', () => { assert.equal(canOpenInOrcaSlicer('/workspace/Foo/hero.glb'), true) + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/part.stl'), true) + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/part.obj'), true) + // Gaussian splats are point clouds, not printable meshes. assert.equal(canOpenInOrcaSlicer('/workspace/Foo/scan.ply'), false) assert.equal(canOpenInOrcaSlicer('/workspace/Foo/scan.splat'), false) - assert.equal(canOpenInOrcaSlicer('/optimize/serve-file?path=/tmp/x.glb'), false) + // A format the slicer route cannot convert. + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/rig.fbx'), false) assert.equal(canOpenInOrcaSlicer(undefined), false) + assert.equal(canOpenInOrcaSlicer(''), false) +}) + +test('canOpenInOrcaSlicer accepts an imported mesh served from outside the workspace', () => { + // `Import` serves user-picked files through /optimize/serve-file, never from + // the workspace — excluding those made the action invisible for the most + // direct "I have a model, slice it" path. + assert.equal(canOpenInOrcaSlicer('/optimize/serve-file?path=%2Ftmp%2Fx.glb'), true) + assert.equal(canOpenInOrcaSlicer('/optimize/serve-file?path=C%3A%5CUsers%5CMe%5Cmage.glb'), true) + assert.equal(canOpenInOrcaSlicer('/optimize/serve-file?path=%2Ftmp%2Fscan.splat'), false) +}) + +test('deeplink for an imported mesh tokenises the decoded absolute path', () => { + const absolute = String.raw`C:\Users\Me\Downloads\mage_v2.glb` + const outputUrl = `/optimize/serve-file?path=${encodeURIComponent(absolute)}` + const modelUrl = decodeURIComponent( + buildOrcaSlicerDeepLink('http://localhost:8765', outputUrl).slice('orcaslicer://open?file='.length), + ) + assert.equal( + modelUrl, + `http://localhost:8765/export/slicer/stl/${encodeWorkspacePathToken(absolute)}/model.stl`, + ) + assert.ok(!modelUrl.includes('?'), 'model URL must not contain a query string') +}) + +test('buildOrcaSlicerDeepLink refuses an output it cannot slice', () => { + assert.throws(() => buildOrcaSlicerDeepLink('http://localhost:8765', '/workspace/Foo/scan.splat')) }) test('SLICER_FORMAT is a format OrcaSlicer can import', () => { diff --git a/src/areas/generate/orcaSlicerLink.ts b/src/areas/generate/orcaSlicerLink.ts index 2bade12e..d5a76c9b 100644 --- a/src/areas/generate/orcaSlicerLink.ts +++ b/src/areas/generate/orcaSlicerLink.ts @@ -11,6 +11,15 @@ /** Format handed to OrcaSlicer. STL is universal and OrcaSlicer auto-repairs it. */ export const SLICER_FORMAT = 'stl' +/** Prefix of an imported mesh served from outside the workspace. */ +const SERVE_FILE_PREFIX = '/optimize/serve-file?path=' + +/** + * Source formats the slicer route can convert. Gaussian splats (`.splat`, and + * the `.ply` they are delivered in) are point clouds, not printable meshes. + */ +const SLICEABLE_SOURCE = /\.(glb|gltf|obj|stl)$/i + /** URL-safe base64 (no padding) of a UTF-8 string — matches the API's token decode. */ export function encodeWorkspacePathToken(workspacePath: string): string { const bytes = new TextEncoder().encode(workspacePath) @@ -20,24 +29,50 @@ export function encodeWorkspacePathToken(workspacePath: string): string { } /** - * Whether a generation output can be opened in OrcaSlicer: it must be a mesh - * served from the workspace (Gaussian splats and non-workspace imports are not - * sliceable through this route). + * The source path the slicer route should convert, or `undefined` when the + * output cannot be sliced. + * + * Two output shapes reach the viewer: a workspace URL (generated or + * workflow-produced meshes) and a `serve-file` URL (meshes the user imported + * from elsewhere on disk). Both are sliceable — the API accepts the absolute + * path of an import because it recorded that the user picked it themselves. + */ +function sliceableSourcePath(outputUrl: string | undefined): string | undefined { + if (!outputUrl) return undefined + + if (outputUrl.startsWith('/workspace/')) { + const workspacePath = outputUrl.slice('/workspace/'.length) + return SLICEABLE_SOURCE.test(workspacePath) ? workspacePath : undefined + } + + if (outputUrl.startsWith(SERVE_FILE_PREFIX)) { + const absolutePath = decodeURIComponent(outputUrl.slice(SERVE_FILE_PREFIX.length)) + return SLICEABLE_SOURCE.test(absolutePath) ? absolutePath : undefined + } + + return undefined +} + +/** + * Whether a generation output can be opened in OrcaSlicer: it must be a mesh in + * a format the slicer route converts, served either from the workspace or as a + * user-selected import. */ export function canOpenInOrcaSlicer(outputUrl: string | undefined): boolean { - if (!outputUrl) return false - return outputUrl.startsWith('/workspace/') && !/\.(ply|splat)$/i.test(outputUrl) + return sliceableSourcePath(outputUrl) !== undefined } /** * Build the `orcaslicer://open?file=...` deeplink for a generated mesh. * * @param apiUrl Modly backend origin, e.g. `http://localhost:8765` - * @param outputUrl workspace URL of the mesh, e.g. `/workspace/Foo/hero.glb` + * @param outputUrl workspace or serve-file URL of the mesh + * @throws if `outputUrl` is not sliceable — guard with {@link canOpenInOrcaSlicer} */ export function buildOrcaSlicerDeepLink(apiUrl: string, outputUrl: string): string { - const workspacePath = outputUrl.replace(/^\/workspace\//, '') - const token = encodeWorkspacePathToken(workspacePath) + const sourcePath = sliceableSourcePath(outputUrl) + if (!sourcePath) throw new Error(`Not sliceable: ${outputUrl}`) + const token = encodeWorkspacePathToken(sourcePath) const base = apiUrl.replace(/\/+$/, '') const modelUrl = `${base}/export/slicer/${SLICER_FORMAT}/${token}/model.${SLICER_FORMAT}` return `orcaslicer://open?file=${encodeURIComponent(modelUrl)}` From 3166012cf6d0cc30dfffe69106a0cb216bf09752 Mon Sep 17 00:00:00 2001 From: sckid1108 Date: Mon, 21 Sep 2026 22:24:46 -0700 Subject: [PATCH 33/57] fix(workflows): map audio slots of multi-input nodes to the file path In executeExtensionNode the multi-input branch resolved every slot's file path but only assigned mesh and image slots; an audio slot was computed and dropped, so a node declaring e.g. inputs ["audio", "text"] failed at run time with " needs an incoming audio connection" although preflight passed. Extract the slot-to-path assignment into slotInputs.ts (pure, tested) and give audio the same primary-path rule as the first image slot. Co-Authored-By: Claude Fable 5.1 --- src/areas/workflows/slotInputs.test.mjs | 54 +++++++++++++++++++++++++ src/areas/workflows/slotInputs.ts | 36 +++++++++++++++++ src/areas/workflows/workflowRunStore.ts | 15 +++---- 3 files changed, 95 insertions(+), 10 deletions(-) create mode 100644 src/areas/workflows/slotInputs.test.mjs create mode 100644 src/areas/workflows/slotInputs.ts diff --git a/src/areas/workflows/slotInputs.test.mjs b/src/areas/workflows/slotInputs.test.mjs new file mode 100644 index 00000000..16adecd2 --- /dev/null +++ b/src/areas/workflows/slotInputs.test.mjs @@ -0,0 +1,54 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +// slotInputs.ts has no runtime imports, so esbuild bundles it standalone. +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-slotinputs-test-')), 'slotInputs.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('src/areas/workflows/slotInputs.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { assignSlotFilePaths } = loadModule() + +test('a mesh slot becomes the mesh path', () => { + const r = assignSlotFilePaths(['mesh', 'image'], ['a.glb', 'b.png']) + assert.equal(r.nodeInputMeshPath, 'a.glb') + assert.equal(r.nodeInputPath, 'b.png') + assert.deepEqual(r.extraImagePaths, []) +}) + +test('the first image slot is the primary path and later images are extras', () => { + const r = assignSlotFilePaths(['image', 'image', 'image'], ['1.png', '2.png', '3.png']) + assert.equal(r.nodeInputPath, '1.png') + assert.deepEqual(r.extraImagePaths, ['2.png', '3.png']) + assert.equal(r.nodeInputMeshPath, undefined) +}) + +test('an audio slot becomes the primary path', () => { + // Regression: audio slots were resolved and then dropped, so a multi-input + // node such as [audio, text] failed with "needs an incoming audio connection". + const r = assignSlotFilePaths(['audio', 'text'], ['song.wav', undefined]) + assert.equal(r.nodeInputPath, 'song.wav') + assert.equal(r.nodeInputMeshPath, undefined) + assert.deepEqual(r.extraImagePaths, []) +}) + +test('text slots never claim a file path, and empty slots are skipped', () => { + const r = assignSlotFilePaths(['text', 'image'], ['leaked.png', undefined]) + assert.equal(r.nodeInputPath, undefined) + assert.equal(r.nodeInputMeshPath, undefined) + assert.deepEqual(r.extraImagePaths, []) +}) diff --git a/src/areas/workflows/slotInputs.ts b/src/areas/workflows/slotInputs.ts new file mode 100644 index 00000000..c07705e9 --- /dev/null +++ b/src/areas/workflows/slotInputs.ts @@ -0,0 +1,36 @@ +// Multi-input slot → file-path assignment for extension nodes. +// +// `inputTypes` is the node's declared `inputs` (one per handle); `inputPaths` is +// the file path resolved for each slot, indexed the same way (undefined where the +// slot carries text or nothing). Pure so it can be tested without the store. + +export type SlotInputType = 'image' | 'text' | 'mesh' | 'audio' + +export interface SlotFilePaths { + /** Primary file: what the extension receives as `filePath` when no mesh is present. */ + nodeInputPath?: string + /** Mesh slot, if any. Takes precedence as `filePath`; the primary file then rides in params. */ + nodeInputMeshPath?: string + /** Every image beyond the first resolved image slot. */ + extraImagePaths: string[] +} + +export function assignSlotFilePaths( + inputTypes: readonly SlotInputType[], + inputPaths: readonly (string | undefined)[], +): SlotFilePaths { + const out: SlotFilePaths = { extraImagePaths: [] } + for (let i = 0; i < inputTypes.length; i++) { + const fp = inputPaths[i] + if (!fp) continue + if (inputTypes[i] === 'mesh') { + out.nodeInputMeshPath = fp + } else if (inputTypes[i] === 'image') { + if (!out.nodeInputPath) out.nodeInputPath = fp + else out.extraImagePaths.push(fp) + } else if (inputTypes[i] === 'audio') { + if (!out.nodeInputPath) out.nodeInputPath = fp + } + } + return out +} diff --git a/src/areas/workflows/workflowRunStore.ts b/src/areas/workflows/workflowRunStore.ts index d6e6d8fc..15cd9ace 100644 --- a/src/areas/workflows/workflowRunStore.ts +++ b/src/areas/workflows/workflowRunStore.ts @@ -6,6 +6,7 @@ import { showCompletionNotification } from '@shared/utils/notification' import type { WorkflowExtension } from './mockExtensions' import type { Workflow, WFNode, WFEdge } from '@shared/types/electron.d' import { isBranchStarter, isSceneOutput, resolveDataSource, reachesSceneOutput, nearestUpstreamWaits } from './nodeBehaviors' +import { assignSlotFilePaths } from './slotInputs' // ─── Types ──────────────────────────────────────────────────────────────────── @@ -344,16 +345,10 @@ async function executeExtensionNode( } } - for (let i = 0; i < inputTypes.length; i++) { - const fp = inputPaths[i] - if (!fp) continue - if (inputTypes[i] === 'mesh') { - nodeInputMeshPath = fp - } else if (inputTypes[i] === 'image') { - if (!nodeInputPath) nodeInputPath = fp - else extraImagePaths.push(fp) - } - } + const slots = assignSlotFilePaths(inputTypes, inputPaths) + nodeInputPath = slots.nodeInputPath + nodeInputMeshPath = slots.nodeInputMeshPath + extraImagePaths.push(...slots.extraImagePaths) } else { for (const edge of incomingEdges) { const src = resolveSource(edge.source) From 3f77304f721d37af98b21272331809829af39a4c Mon Sep 17 00:00:00 2001 From: yi111 <153097222+Yi-111-a@users.noreply.github.com> Date: Fri, 25 Sep 2026 22:03:56 +0800 Subject: [PATCH 34/57] fix(generate): accept image drops from Explorer --- electron/main/model-download-preload.test.mjs | 6 ++- .../preload/artifact-registry-preload.test.ts | 26 +++++++++- electron/preload/electron-api.ts | 6 ++- electron/preload/index.ts | 4 +- .../generate/components/WorkflowPanel.tsx | 47 +++++++++++++------ src/shared/types/electron.d.ts | 1 + 6 files changed, 71 insertions(+), 19 deletions(-) diff --git a/electron/main/model-download-preload.test.mjs b/electron/main/model-download-preload.test.mjs index 8d9a16b7..d75eae0f 100644 --- a/electron/main/model-download-preload.test.mjs +++ b/electron/main/model-download-preload.test.mjs @@ -29,7 +29,11 @@ test('renderer model actions send only the model node id', async () => { on: () => {}, removeAllListeners: () => {}, } - const api = createElectronApi(ipc, { setZoomFactor: () => {} }) + const api = createElectronApi( + ipc, + { setZoomFactor: () => {} }, + { getPathForFile: () => '' }, + ) await api.model.isDownloaded('pixal3d/generate') await api.model.hasLocalData('pixal3d/generate') diff --git a/electron/preload/artifact-registry-preload.test.ts b/electron/preload/artifact-registry-preload.test.ts index 511589e4..9dec46fd 100644 --- a/electron/preload/artifact-registry-preload.test.ts +++ b/electron/preload/artifact-registry-preload.test.ts @@ -13,7 +13,7 @@ test('preload exposes scoped workspace library list/read/open methods', async () send: () => undefined, on: () => undefined, removeAllListeners: () => undefined, - }, { setZoomFactor: () => undefined }) + }, { setZoomFactor: () => undefined }, { getPathForFile: () => '' }) await api.workspace.library.list() await api.workspace.library.read({ @@ -43,3 +43,27 @@ test('preload exposes scoped workspace library list/read/open methods', async () }, ]) }) + +test('preload resolves dropped file paths through Electron webUtils', () => { + const files: unknown[] = [] + const calls: string[] = [] + const api = createElectronApi({ + invoke: async (channel: string) => { + calls.push(channel) + return null + }, + send: () => undefined, + on: () => undefined, + removeAllListeners: () => undefined, + }, { setZoomFactor: () => undefined }, { + getPathForFile: (file) => { + files.push(file) + return 'C:\\images\\input.png' + }, + }) + const file = {} as Parameters[0] + + assert.equal(api.fs.getPathForFile(file), 'C:\\images\\input.png') + assert.deepEqual(files, [file]) + assert.deepEqual(calls, []) +}) diff --git a/electron/preload/electron-api.ts b/electron/preload/electron-api.ts index af0b1265..e9ee6334 100644 --- a/electron/preload/electron-api.ts +++ b/electron/preload/electron-api.ts @@ -17,7 +17,9 @@ export interface WebFrameLike { setZoomFactor(factor: number): void } -export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFrameLike) { +export type WebUtilsLike = Pick + +export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFrameLike, webUtils: WebUtilsLike) { return { // Window controls window: { @@ -73,6 +75,8 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra // File system dialogs + local file reading fs: { + getPathForFile: (file: Parameters[0]): string => + webUtils.getPathForFile(file), selectImage: (): Promise => ipcRenderer.invoke('fs:selectImage') as Promise, selectMeshFile: (): Promise => diff --git a/electron/preload/index.ts b/electron/preload/index.ts index ec34ebf2..03b80363 100644 --- a/electron/preload/index.ts +++ b/electron/preload/index.ts @@ -1,6 +1,6 @@ -import { contextBridge, ipcRenderer, webFrame } from 'electron' +import { contextBridge, ipcRenderer, webFrame, webUtils } from 'electron' import { createElectronApi } from './electron-api' // Expose a typed API to the renderer process via window.electron -contextBridge.exposeInMainWorld('electron', createElectronApi(ipcRenderer, webFrame)) +contextBridge.exposeInMainWorld('electron', createElectronApi(ipcRenderer, webFrame, webUtils)) diff --git a/src/areas/generate/components/WorkflowPanel.tsx b/src/areas/generate/components/WorkflowPanel.tsx index 2995ddea..121f185f 100644 --- a/src/areas/generate/components/WorkflowPanel.tsx +++ b/src/areas/generate/components/WorkflowPanel.tsx @@ -13,6 +13,7 @@ import { useWorkflowRunStore } from '@areas/workflows/workflowRunStore' import { useWaitButton } from '@areas/workflows/useWaitButton' import { buildAllWorkflowExtensions, getWorkflowExtension } from '@areas/workflows/mockExtensions' import { validateWorkflowPreflight } from '@areas/workflows/preflight' +import { mimeFromPath } from '@areas/workflows/nodes/imageUtils' import type { WorkflowExtension } from '@areas/workflows/mockExtensions' import type { Workflow, WFNode, WFEdge, ParamSchema } from '@shared/types/electron.d' import { PICKER_LABELS, openParamPicker, resolvePickerIntent } from '@shared/utils/paramPicker' @@ -29,6 +30,8 @@ const TYPE_COLOR: Record = { text: '#fbbf24', } +const SUPPORTED_IMAGE_TYPES = new Set(['image/jpeg', 'image/png', 'image/webp']) + // ─── Helpers ────────────────────────────────────────────────────────────────── function topoSortNodes(nodes: Workflow['nodes'], edges: Workflow['edges']): WFNode[] { @@ -54,13 +57,6 @@ function topoSortNodes(nodes: Workflow['nodes'], edges: Workflow['edges']): WFNo return result } -function mimeFromPath(p: string): string { - const ext = p.split('.').pop()?.toLowerCase() ?? '' - if (ext === 'jpg' || ext === 'jpeg') return 'image/jpeg' - if (ext === 'webp') return 'image/webp' - return 'image/png' -} - // ─── Param field ────────────────────────────────────────────────────────────── const inputCls = 'w-full bg-zinc-800 border border-zinc-700/80 rounded-md px-2 py-1 text-[11px] text-zinc-200 focus:outline-none focus:border-accent/60' @@ -212,17 +208,40 @@ function ImageParamRow({ nodeId, nodes, onPatch }: { nodeId: string; nodes: Flow const node = nodes.find((n) => n.id === nodeId) const data = node?.data as { params: Record } | undefined const preview = data?.params.preview as string | undefined + const showToast = useAppStore((state) => state.showToast) + const loadRequest = useRef(0) + + const applyImagePath = useCallback(async (path: string | null) => { + const request = ++loadRequest.current + if (!path) return + try { + const base64 = await window.electron.fs.readFileBase64(path) + if (request !== loadRequest.current) return + const src = `data:${mimeFromPath(path)};base64,${base64}` + onPatch(nodeId, { params: { ...(data?.params ?? {}), filePath: path, preview: src } }) + } catch { + if (request === loadRequest.current) showToast('Unable to load the selected image') + } + }, [nodeId, data?.params, onPatch, showToast]) const browse = useCallback(async () => { - const p = await window.electron.fs.selectImage() - if (!p) return - const base64 = await window.electron.fs.readFileBase64(p) - const src = `data:${mimeFromPath(p)};base64,${base64}` - onPatch(nodeId, { params: { ...(data?.params ?? {}), filePath: p, preview: src } }) - }, [nodeId, data?.params, onPatch]) + await applyImagePath(await window.electron.fs.selectImage()) + }, [applyImagePath]) return ( -
+
{ + event.preventDefault() + event.dataTransfer.dropEffect = 'copy' + }} + onDrop={(event) => { + event.preventDefault() + const file = event.dataTransfer.files[0] + if (!file || !SUPPORTED_IMAGE_TYPES.has(file.type)) return + void applyImagePath(window.electron.fs.getPathForFile(file)) + }} + >
diff --git a/src/shared/types/electron.d.ts b/src/shared/types/electron.d.ts index 5119a1d4..ffe434ed 100644 --- a/src/shared/types/electron.d.ts +++ b/src/shared/types/electron.d.ts @@ -184,6 +184,7 @@ declare global { offLog: () => void } fs: { + getPathForFile: (file: File) => string selectImage: () => Promise selectMeshFile: () => Promise saveModel: (defaultName: string) => Promise From 6da39f941e825ddf2a9f3746fc0a26b2b820f6b9 Mon Sep 17 00:00:00 2001 From: kevin9327 <5299031+kevin9327@users.noreply.github.com> Date: Sun, 13 Sep 2026 06:50:52 +0900 Subject: [PATCH 35/57] fix(export,optimize): read WORKSPACE_DIR dynamically so edits work after the workspace moves export.py and optimize.py bound the workspace path with `from services.generator_registry import WORKSPACE_DIR`, capturing it at import time. POST /settings/paths rebinds that module global when the workspace is moved in Settings, but these routers kept resolving paths against the old folder. A model generated after the move could not be exported, decimated, smoothed or transformed ("File not found"), and results for inputs outside the workspace were written into the old folder's Workflows/. Read registry.WORKSPACE_DIR at call time instead, the same way generation.py and ply_to_splat already do. Co-Authored-By: Claude Opus 5 --- api/routers/export.py | 10 +- api/routers/optimize.py | 24 ++--- api/tests/test_export_router.py | 12 +-- api/tests/test_mesh_routers_workspace.py | 119 +++++++++++++++++++++++ api/tests/test_optimize_mesh_ops.py | 10 +- 5 files changed, 149 insertions(+), 26 deletions(-) create mode 100644 api/tests/test_mesh_routers_workspace.py diff --git a/api/routers/export.py b/api/routers/export.py index 6cd03d61..0d8e3429 100644 --- a/api/routers/export.py +++ b/api/routers/export.py @@ -9,7 +9,9 @@ from fastapi.responses import Response, FileResponse from services import imported_sources -from services.generator_registry import WORKSPACE_DIR +# Import the module (not the name) so WORKSPACE_DIR is read at call time: the +# settings endpoint rebinds it when the user moves the workspace. +import services.generator_registry as registry router = APIRouter(tags=["export"]) @@ -122,7 +124,7 @@ def _resolve_slicer_source(token: str) -> tuple[Path, str]: # Containment check via ancestry, not string prefix: `startswith` would let a # sibling like `-other/...` slip through, and `..` escapes resolve # outside the workspace and fail this check. - workspace = WORKSPACE_DIR.resolve() + workspace = registry.WORKSPACE_DIR.resolve() full_path = (workspace / decoded).resolve() if full_path != workspace and workspace not in full_path.parents: raise HTTPException(400, "Invalid path") @@ -181,8 +183,8 @@ def export_mesh(fmt: str, path: str): if fmt not in SUPPORTED: raise HTTPException(400, f"Unsupported format: {fmt}. Supported: {', '.join(SUPPORTED)}") - full_path = (WORKSPACE_DIR / path).resolve() - if not str(full_path).startswith(str(WORKSPACE_DIR.resolve())): + full_path = (registry.WORKSPACE_DIR / path).resolve() + if not str(full_path).startswith(str(registry.WORKSPACE_DIR.resolve())): raise HTTPException(400, "Invalid path") if not full_path.exists(): raise HTTPException(404, f"File not found: {path}") diff --git a/api/routers/optimize.py b/api/routers/optimize.py index 7ae622c5..843bb596 100644 --- a/api/routers/optimize.py +++ b/api/routers/optimize.py @@ -12,7 +12,9 @@ from pydantic import BaseModel, Field from services import imported_sources -from services.generator_registry import WORKSPACE_DIR +# Import the module (not the name) so WORKSPACE_DIR is read at call time: the +# settings endpoint rebinds it when the user moves the workspace. +import services.generator_registry as registry from services.mesh_ops import ( MeshOpContext, MeshOpExecutionError, @@ -53,8 +55,8 @@ def _resolve_input_path(raw_path: str) -> Path: raise HTTPException(404, f"File not found: {raw_path}") return resolved - resolved = (WORKSPACE_DIR / raw_path).resolve() - if not str(resolved).startswith(str(WORKSPACE_DIR.resolve())): + resolved = (registry.WORKSPACE_DIR / raw_path).resolve() + if not str(resolved).startswith(str(registry.WORKSPACE_DIR.resolve())): raise HTTPException(400, "Invalid path") if not resolved.exists(): raise HTTPException(404, f"File not found: {raw_path}") @@ -62,12 +64,12 @@ def _resolve_input_path(raw_path: str) -> Path: def _operation_output_path(input_path: Path, output_name: str) -> Path: - workspace = WORKSPACE_DIR.resolve() + workspace = registry.WORKSPACE_DIR.resolve() resolved_input = input_path.resolve() output_dir = ( input_path.parent if resolved_input == workspace or workspace in resolved_input.parents - else WORKSPACE_DIR / "Workflows" + else registry.WORKSPACE_DIR / "Workflows" ) output_dir.mkdir(parents=True, exist_ok=True) return output_dir / output_name @@ -81,7 +83,7 @@ def _run_operation( preserve_visuals: bool = False, ) -> MeshOpResult: context = MeshOpContext( - workspace_dir=WORKSPACE_DIR, + workspace_dir=registry.WORKSPACE_DIR, temp_dir=Path(tempfile.gettempdir()), output_path=output_path, preserve_visuals=preserve_visuals, @@ -101,7 +103,7 @@ def _run_operation( def _operation_response(result: MeshOpResult) -> dict[str, object]: output_path = result.file_path.resolve() try: - relative_path = output_path.relative_to(WORKSPACE_DIR.resolve()).as_posix() + relative_path = output_path.relative_to(registry.WORKSPACE_DIR.resolve()).as_posix() except ValueError: payload: dict[str, object] = {"path": str(output_path)} else: @@ -187,12 +189,12 @@ def transform_mesh(body: TransformRequest): stem = input_path.stem output_name = f"{stem}_xf_{uuid.uuid4().hex[:8]}.glb" - output_dir = input_path.parent if str(input_path).startswith(str(WORKSPACE_DIR.resolve())) else WORKSPACE_DIR / "Workflows" + output_dir = input_path.parent if str(input_path).startswith(str(registry.WORKSPACE_DIR.resolve())) else registry.WORKSPACE_DIR / "Workflows" output_dir.mkdir(parents=True, exist_ok=True) output_path = output_dir / output_name loaded.export(str(output_path)) - rel = output_path.relative_to(WORKSPACE_DIR).as_posix() + rel = output_path.relative_to(registry.WORKSPACE_DIR).as_posix() return {"url": f"/workspace/{rel}"} @@ -408,8 +410,8 @@ def export_mesh(path: str, format: str): if format not in ("obj", "stl", "ply"): raise HTTPException(400, "Supported formats: obj, stl, ply") - input_path = (WORKSPACE_DIR / path).resolve() - if not str(input_path).startswith(str(WORKSPACE_DIR.resolve())): + input_path = (registry.WORKSPACE_DIR / path).resolve() + if not str(input_path).startswith(str(registry.WORKSPACE_DIR.resolve())): raise HTTPException(400, "Invalid path") if not input_path.exists(): raise HTTPException(404, f"File not found: {path}") diff --git a/api/tests/test_export_router.py b/api/tests/test_export_router.py index 0972e836..2bed2a68 100644 --- a/api/tests/test_export_router.py +++ b/api/tests/test_export_router.py @@ -33,8 +33,8 @@ class ExportForSlicerTests(unittest.TestCase): def setUp(self) -> None: self._tmp = tempfile.TemporaryDirectory() self.workspace = Path(self._tmp.name).resolve() - self._orig_workspace = export_router.WORKSPACE_DIR - export_router.WORKSPACE_DIR = self.workspace + self._orig_workspace = export_router.registry.WORKSPACE_DIR + export_router.registry.WORKSPACE_DIR = self.workspace # A box that is tallest along Y (glTF up-axis) and unit-sized, matching # what image-to-3D generators emit. Exported to GLB, it reloads as a # Scene so the flatten path is exercised too. @@ -44,7 +44,7 @@ def setUp(self) -> None: box.export(str(self.workspace / self.rel)) def tearDown(self) -> None: - export_router.WORKSPACE_DIR = self._orig_workspace + export_router.registry.WORKSPACE_DIR = self._orig_workspace self._tmp.cleanup() def test_converts_glb_to_stl_with_download_filename(self) -> None: @@ -154,17 +154,17 @@ class ImportedSourceSlicerTests(unittest.TestCase): def setUp(self) -> None: self._tmp = tempfile.TemporaryDirectory() self.outside = Path(self._tmp.name).resolve() - self._orig_workspace = export_router.WORKSPACE_DIR + self._orig_workspace = export_router.registry.WORKSPACE_DIR # A workspace elsewhere, so nothing here is reachable as a relative path. self._ws_tmp = tempfile.TemporaryDirectory() - export_router.WORKSPACE_DIR = Path(self._ws_tmp.name).resolve() + export_router.registry.WORKSPACE_DIR = Path(self._ws_tmp.name).resolve() imported_sources.clear() # Unit-sized and tallest along Y, as a glTF export would be. self.mesh_path = self.outside / "imported.glb" trimesh.creation.box(extents=[0.3, 1.0, 0.3]).export(str(self.mesh_path)) def tearDown(self) -> None: - export_router.WORKSPACE_DIR = self._orig_workspace + export_router.registry.WORKSPACE_DIR = self._orig_workspace imported_sources.clear() self._tmp.cleanup() self._ws_tmp.cleanup() diff --git a/api/tests/test_mesh_routers_workspace.py b/api/tests/test_mesh_routers_workspace.py new file mode 100644 index 00000000..0febfb68 --- /dev/null +++ b/api/tests/test_mesh_routers_workspace.py @@ -0,0 +1,119 @@ +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +import trimesh +from fastapi import HTTPException + +import routers.export as export_router +import routers.optimize as optimize_router +import services.generator_registry as registry +from services.mesh_ops import MeshOpResult + +IDENTITY = [ + [1.0, 0.0, 0.0, 0.0], + [0.0, 1.0, 0.0, 0.0], + [0.0, 0.0, 1.0, 0.0], + [0.0, 0.0, 0.0, 1.0], +] + + +class MeshRoutersAfterWorkspaceMoveTests(unittest.TestCase): + """Export and mesh-edit endpoints must resolve paths against the workspace as + it is *now*. POST /settings/paths rebinds registry.WORKSPACE_DIR when the user + moves the workspace; a name captured at import keeps pointing at the old + folder, so a model generated after the move can't be exported or edited.""" + + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self.root = Path(self._tmp.name) + self._prev_ws = registry.WORKSPACE_DIR + # The user moved the workspace: the registry global now points here. + registry.WORKSPACE_DIR = self.root / "new_workspace" + # Keep the test hermetic against import-time bindings: if a router still + # holds its own WORKSPACE_DIR name, point it at an empty "old" folder in + # the temp tree so the assertions -- not the real workspace -- catch it. + (self.root / "old_workspace").mkdir() + self._stale = [] + for module in (export_router, optimize_router): + if hasattr(module, "WORKSPACE_DIR"): + self._stale.append((module, module.WORKSPACE_DIR)) + module.WORKSPACE_DIR = self.root / "old_workspace" + + mesh_dir = registry.WORKSPACE_DIR / "MyColl" + mesh_dir.mkdir(parents=True) + trimesh.creation.box().export(mesh_dir / "mesh.glb") + + def tearDown(self) -> None: + registry.WORKSPACE_DIR = self._prev_ws + for module, value in self._stale: + module.WORKSPACE_DIR = value + self._tmp.cleanup() + + def test_export_router_converts_a_mesh_in_the_moved_workspace(self) -> None: + response = export_router.export_mesh("stl", "MyColl/mesh.glb") + self.assertEqual(response.status_code, 200) + self.assertGreater(len(response.body), 0) + + def test_optimize_export_converts_a_mesh_in_the_moved_workspace(self) -> None: + response = optimize_router.export_mesh(path="MyColl/mesh.glb", format="obj") + self.assertEqual(response.status_code, 200) + self.assertIn(b"v ", response.body) + + def test_transform_writes_its_result_into_the_moved_workspace(self) -> None: + result = optimize_router.transform_mesh( + optimize_router.TransformRequest(path="MyColl/mesh.glb", matrix=IDENTITY) + ) + self.assertTrue(result["url"].startswith("/workspace/MyColl/mesh_xf_")) + written = registry.WORKSPACE_DIR / result["url"].removeprefix("/workspace/") + self.assertTrue(written.is_file()) + + def test_decimate_and_smooth_read_their_input_from_the_moved_workspace(self) -> None: + # /optimize/mesh and /optimize/smooth resolve their input through this helper. + resolved = optimize_router._resolve_input_path("MyColl/mesh.glb") + self.assertEqual(resolved, (registry.WORKSPACE_DIR / "MyColl" / "mesh.glb").resolve()) + + def test_decimate_and_smooth_write_their_result_into_the_moved_workspace(self) -> None: + # The backends (meshoptimizer, pymeshlab) aren't available here, so stand in + # for the mesh-ops registry and check the router's own path handling: the + # workspace it hands the operation, where the output goes, and the URL. + class _RecordingRegistry: + def __init__(self) -> None: + self.contexts = [] + + def run(self, operation_id, input_path, params, context): + self.contexts.append(context) + context.output_path.touch() + return MeshOpResult(context.output_path, {"face_count": 12}) + + ops = _RecordingRegistry() + with patch.object(optimize_router, "mesh_ops_registry", ops): + decimated = optimize_router.optimize_mesh( + optimize_router.OptimizeRequest(path="MyColl/mesh.glb", target_faces=500) + ) + smoothed = optimize_router.smooth_mesh( + optimize_router.SmoothRequest(path="MyColl/mesh.glb", iterations=2) + ) + + self.assertEqual(decimated["url"], "/workspace/MyColl/mesh_opt500.glb") + self.assertEqual(smoothed["url"], "/workspace/MyColl/mesh_smooth2.glb") + for context in ops.contexts: + self.assertEqual(context.workspace_dir, registry.WORKSPACE_DIR) + self.assertEqual(context.output_path.parent, registry.WORKSPACE_DIR / "MyColl") + + def test_a_path_leaving_the_workspace_is_still_refused(self) -> None: + # Reading the live workspace must not loosen the containment check. + trimesh.creation.box().export(self.root / "outside.glb") + calls = ( + lambda: export_router.export_mesh("stl", "../outside.glb"), + lambda: optimize_router.export_mesh(path="../outside.glb", format="obj"), + ) + for call in calls: + with self.assertRaises(HTTPException) as raised: + call() + self.assertEqual(raised.exception.status_code, 400) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_optimize_mesh_ops.py b/api/tests/test_optimize_mesh_ops.py index a1f0c5fe..6c702ddb 100644 --- a/api/tests/test_optimize_mesh_ops.py +++ b/api/tests/test_optimize_mesh_ops.py @@ -39,7 +39,7 @@ def test_generic_list_and_run_routes_use_the_shared_registry(self) -> None: registry = _FakeRegistry(output_path) with ( - patch.object(optimize, "WORKSPACE_DIR", workspace), + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), patch.object(optimize, "mesh_ops_registry", registry), ): descriptions = optimize.list_mesh_operations() @@ -75,7 +75,7 @@ def test_legacy_routes_delegate_with_their_existing_clamps_and_names(self) -> No registry = _FakeRegistry(fallback_output) with ( - patch.object(optimize, "WORKSPACE_DIR", workspace), + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), patch.object(optimize, "mesh_ops_registry", registry), ): optimize_response = optimize.optimize_mesh( @@ -115,7 +115,7 @@ def run(self, operation_id, input_path, params, context): input_path = workspace / "input.glb" input_path.touch() with ( - patch.object(optimize, "WORKSPACE_DIR", workspace), + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), patch.object(optimize, "mesh_ops_registry", MissingRegistry()), self.assertRaises(HTTPException) as raised, ): @@ -158,7 +158,7 @@ def operation(input_path, params, context): input_path = workspace / "input.glb" input_path.touch() with ( - patch.object(optimize, "WORKSPACE_DIR", workspace), + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), patch.object(optimize, "mesh_ops_registry", registry), ): optimize.run_mesh_operation( @@ -191,7 +191,7 @@ def _assert_operation_error(self, registry, expected_status: int) -> None: input_path = workspace / "input.glb" input_path.touch() with ( - patch.object(optimize, "WORKSPACE_DIR", workspace), + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), patch.object(optimize, "mesh_ops_registry", registry), self.assertRaises(HTTPException) as raised, ): From 2320b0b9ee59b9a59062cbabd9c97b8cdd42c4e4 Mon Sep 17 00:00:00 2001 From: DrHepa Date: Mon, 21 Sep 2026 19:50:54 +0200 Subject: [PATCH 36/57] feat(workflows): support scene model artifacts --- README.md | 12 + api/README.md | 1 + api/routers/generation.py | 107 +++++++-- api/runner.py | 50 ++++- api/schemas/generation.py | 13 +- api/services/artifact_input.py | 29 +++ api/services/extension_process.py | 55 ++++- api/services/generator_registry.py | 53 ++++- api/services/generators/base.py | 21 +- api/services/scene_input.py | 138 ++++++++++++ api/tests/test_extension_process.py | 28 +++ api/tests/test_generation_router.py | 3 + api/tests/test_generator_registry.py | 47 ++++ api/tests/test_runner.py | 25 +++ api/tests/test_scene_generation.py | 115 ++++++++++ api/tests/test_scene_input.py | 67 ++++++ .../main/artifact-registry-service.test.ts | 12 + electron/main/artifact-registry-service.ts | 3 +- .../main/extension-install-utils.test.mjs | 44 ++++ electron/main/extension-install-utils.ts | 34 ++- electron/main/ipc-handlers.ts | 15 +- src/areas/workflows/WorkflowsPage.tsx | 12 +- src/areas/workflows/mockExtensions.ts | 6 +- src/areas/workflows/nodes/ExtensionNode.tsx | 2 + src/areas/workflows/nodes/LoadSceneNode.tsx | 134 ++++++++++++ src/areas/workflows/preflight.test.mjs | 35 +++ src/areas/workflows/preflight.ts | 26 ++- src/areas/workflows/workflowRunStore.ts | 64 ++++-- src/areas/workflows/workflowSceneRun.test.mjs | 60 +++++ .../workflows/workflowSceneSource.test.mjs | 29 +++ src/areas/workflows/workflowSceneSource.ts | 205 ++++++++++++++++++ src/shared/stores/workflowsStore.ts | 2 +- src/shared/types/artifacts.ts | 14 ++ src/shared/types/electron.d.ts | 6 +- 34 files changed, 1390 insertions(+), 77 deletions(-) create mode 100644 api/services/artifact_input.py create mode 100644 api/services/scene_input.py create mode 100644 api/tests/test_scene_generation.py create mode 100644 api/tests/test_scene_input.py create mode 100644 src/areas/workflows/nodes/LoadSceneNode.tsx create mode 100644 src/areas/workflows/workflowSceneRun.test.mjs create mode 100644 src/areas/workflows/workflowSceneSource.test.mjs create mode 100644 src/areas/workflows/workflowSceneSource.ts diff --git a/README.md b/README.md index b162cf23..95a0b6ff 100644 --- a/README.md +++ b/README.md @@ -153,6 +153,18 @@ original behavior. ## Workflows Start with a basic workflow first. For example, on the "Workflows" tab, try: Image -> Generate Mesh -> Add to Scene. Make sure there is a connection between each of the steps. Go to the "Generate" tab, make sure the workflow is selected, then click on "Generate 3D Model". Click on "Settings/Logs/Errors" to see any issues. +Model extensions may also declare `scene` as a node input or output. A scene is +a workspace directory containing `scene-manifest.json` with schema +`modly.scene-manifest.v1`; it is not an arbitrary JSON file. Use the **Load +Scene** workflow node to select and validate an existing scene directory. +Scene-capable generators implement `generate_artifact(input_kind, +artifact_path, ...)`; legacy image generators and `POST /generate/from-image` +remain unchanged. The generic `POST /generate/from-artifact` boundary currently +accepts only `scene`, leaving future artifact kinds to separate reviewed changes. +For this first contract, `scene` is model-only and must be declared as the single +`input` value (not inside `inputs`); process and mixed-input scene nodes are rejected. +Model nodes may still accept multiple images and produce a scene. + ## Modly CLI diff --git a/api/README.md b/api/README.md index ee45bc3f..cdcd4d6b 100644 --- a/api/README.md +++ b/api/README.md @@ -29,6 +29,7 @@ uvicorn main:app --host 127.0.0.1 --port 8765 --reload | GET | `/model/status` | Model download / load status | | GET | `/model/download` | SSE stream of download progress | | POST | `/generate/from-image` | Start image-to-3D job | +| POST | `/generate/from-artifact` | Start a typed-artifact model job (`scene` only) | | GET | `/generate/status/{job_id}` | Poll job status | ## Model diff --git a/api/routers/generation.py b/api/routers/generation.py index 8481deb4..7c355014 100644 --- a/api/routers/generation.py +++ b/api/routers/generation.py @@ -4,7 +4,8 @@ import time import traceback import uuid -from typing import Dict +from pathlib import Path +from typing import Dict, Optional, Union from fastapi import APIRouter, File, Form, UploadFile, HTTPException, BackgroundTasks from services.generators.base import smooth_progress, GenerationCancelled @@ -14,7 +15,8 @@ # binding captured at import would keep writing output to the old directory. import services.generator_registry as registry from services.generator_registry import generator_registry -from schemas.generation import JobStatus +from schemas.generation import GenerateFromArtifactRequest, JobStatus +from services.artifact_input import TypedArtifactInput, validate_artifact_input router = APIRouter(tags=["generation"]) @@ -105,6 +107,7 @@ async def generate_from_image( # Verify the requested model exists in the registry try: generator_registry.get_generator(model_id) + output_kind = generator_registry.get_manifest(model_id).get("output", "mesh") except ValueError as e: raise HTTPException(400, str(e)) @@ -131,11 +134,49 @@ async def generate_from_image( _jobs[job_id] = job _cancel_events[job_id] = threading.Event() - background_tasks.add_task(_run_generation, job_id, image_bytes, full_params, collection) + background_tasks.add_task( + _run_generation, job_id, image_bytes, full_params, collection, output_kind, model_id + ) return {"job_id": job_id} +_RESERVED_ARTIFACT_PARAMS = { + "artifact_path", "input_kind", "input_path", "scene_path", "scene_manifest_path", +} + + +@router.post("/from-artifact") +async def generate_from_artifact( + request: GenerateFromArtifactRequest, + background_tasks: BackgroundTasks, +): + """Queue a validated typed artifact without serializing it as image bytes.""" + try: + manifest = generator_registry.get_manifest(request.model_id) + except (KeyError, ValueError) as exc: + raise HTTPException(400, str(exc)) from exc + declared = manifest.get("inputs") or [manifest.get("input", "image")] + if request.input_kind not in declared: + raise HTTPException(400, f"Model {request.model_id} does not accept {request.input_kind} input") + try: + artifact = validate_artifact_input(registry.WORKSPACE_DIR, request.input_kind, request.input_path) + except ValueError as exc: + raise HTTPException(400, str(exc)) from exc + + params = {k: v for k, v in request.params.items() if k not in _RESERVED_ARTIFACT_PARAMS} + params["scene_manifest_path"] = str(artifact.path) + collection = sanitize_collection(request.collection) + job_id = str(uuid.uuid4()) + _purge_old_jobs() + _jobs[job_id] = JobStatus(job_id=job_id, status="pending", progress=0) + _cancel_events[job_id] = threading.Event() + background_tasks.add_task( + _run_generation, job_id, artifact, params, collection, manifest.get("output", "mesh"), request.model_id + ) + return {"job_id": job_id} + + @router.get("/status/{job_id}") async def job_status(job_id: str): @@ -170,7 +211,14 @@ async def cancel_job(job_id: str): return {"cancelled": True} -async def _run_generation(job_id: str, image_bytes: bytes, params: dict, collection: str = "Default") -> None: +async def _run_generation( + job_id: str, + model_input: Union[bytes, TypedArtifactInput], + params: dict, + collection: str = "Default", + output_kind: str = "mesh", + model_id: Optional[str] = None, +) -> None: job = _jobs[job_id] job.status = "running" @@ -189,8 +237,14 @@ def progress_cb(pct: int, step: str = "") -> None: # Check if the model needs to be loaded BEFORE calling get_active(), # because get_active() loads the model in a blocking manner. # active_status() is an instantaneous operation (simple dict lookup). - if not generator_registry.active_status()["loaded"]: - active = generator_registry.active_status() + get_generator = (lambda: generator_registry.get_ready_generator(model_id)) \ + if model_id is not None else generator_registry.get_active + status_reader = getattr(generator_registry, "model_status", None) + status = (status_reader(model_id) if model_id is not None and status_reader + else generator_registry.active_status() if model_id is None + else {"name": model_id, "downloaded": True, "loaded": False}) + if not status["loaded"]: + active = status model_name = active['name'] init_label = f"Downloading {model_name}…" if not active['downloaded'] else f"Loading {model_name}…" progress_cb(0, init_label) @@ -202,11 +256,11 @@ def progress_cb(pct: int, step: str = "") -> None: ) load_thread.start() try: - gen = await loop.run_in_executor(None, generator_registry.get_active) + gen = await loop.run_in_executor(None, get_generator) finally: stop_load_evt.set() else: - gen = await loop.run_in_executor(None, generator_registry.get_active) + gen = await loop.run_in_executor(None, get_generator) if job_id in _cancelled: return @@ -217,18 +271,39 @@ def progress_cb(pct: int, step: str = "") -> None: gen.outputs_dir = coll_dir cancel_event = _cancel_events.get(job_id) - import inspect - supports_cancel = "cancel_event" in inspect.signature(gen.generate).parameters - output_path = await loop.run_in_executor( - None, - lambda: gen.generate(image_bytes, params, progress_cb, cancel_event) - if supports_cancel - else gen.generate(image_bytes, params, progress_cb), - ) + if isinstance(model_input, TypedArtifactInput): + # Revalidate just before crossing the inference boundary. The + # subprocess runner repeats this check inside the worker. + from services.artifact_input import revalidate_artifact_input + model_input = revalidate_artifact_input(registry.WORKSPACE_DIR, model_input) + import inspect + supports_cancel = "cancel_event" in inspect.signature(gen.generate_artifact).parameters + output_path = await loop.run_in_executor( + None, + lambda: gen.generate_artifact(model_input.kind, model_input.path, params, progress_cb, cancel_event) + if supports_cancel else gen.generate_artifact(model_input.kind, model_input.path, params, progress_cb), + ) + else: + import inspect + supports_cancel = "cancel_event" in inspect.signature(gen.generate).parameters + output_path = await loop.run_in_executor( + None, + lambda: gen.generate(model_input, params, progress_cb, cancel_event) + if supports_cancel else gen.generate(model_input, params, progress_cb), + ) if job_id in _cancelled: return + output_path = Path(output_path).resolve(strict=True) + if output_kind == "scene": + from services.scene_input import validate_scene_input + try: + output_relative = output_path.relative_to(registry.WORKSPACE_DIR.resolve()) + output_path = validate_scene_input(registry.WORKSPACE_DIR, output_relative.as_posix()) + except (OSError, ValueError) as exc: + raise ValueError("Generated scene output is not a valid workspace scene") from exc + job.status = "done" job.progress = 100 _completed_at[job_id] = time.monotonic() diff --git a/api/runner.py b/api/runner.py index 0dd21392..0dd17ddb 100644 --- a/api/runner.py +++ b/api/runner.py @@ -151,6 +151,32 @@ def _apply_manifest_metadata(gen, manifest: dict, node: dict) -> None: gen._params_schema = node.get("params_schema") or manifest.get("params_schema", []) +def decode_model_input(msg: dict): + """Decode legacy image bytes or revalidate a typed artifact in the worker.""" + if "input" not in msg: + return base64.b64decode(msg["image_b64"]) + value = msg["input"] + if not isinstance(value, dict) or set(value) != {"kind", "path"}: + raise ValueError("Typed artifact input must contain exactly kind and path") + from services.artifact_input import TypedArtifactInput, revalidate_artifact_input + typed = TypedArtifactInput(kind=value.get("kind"), path=Path(value.get("path", ""))) + return revalidate_artifact_input(WORKSPACE_DIR, typed) + + +def validate_requested_model(msg: dict, manifest: dict, node: dict) -> None: + """Reject requests routed to a worker for a different manifest node.""" + requested = msg.get("model_id") + if requested is None: # Backward compatibility with already-running legacy hosts. + return + expected = manifest["id"] + if node.get("id"): + expected = f"{expected}/{node['id']}" + if requested != expected: + raise ValueError( + f"Generation request model '{requested}' does not match worker model '{expected}'" + ) + + # ------------------------------------------------------------------ # # Main loop # ------------------------------------------------------------------ # @@ -201,10 +227,17 @@ def main() -> None: # ---- generate -------------------------------------------- elif action == "generate": + validate_requested_model(msg, manifest, node) cancel_evt = threading.Event() _cancel[rid] = cancel_evt - image_bytes = base64.b64decode(msg["image_b64"]) + model_input = decode_model_input(msg) params = msg.get("params", {}) + if hasattr(model_input, "kind"): + if not isinstance(params, dict): + raise ValueError("Model params must be an object") + reserved = {"artifact_path", "input_kind", "input_path", "scene_path", "scene_manifest_path"} + params = {key: value for key, value in params.items() if key not in reserved} + params["scene_manifest_path"] = str(model_input.path) if msg.get("outputs_dir"): gen.outputs_dir = Path(msg["outputs_dir"]) gen.outputs_dir.mkdir(parents=True, exist_ok=True) @@ -217,7 +250,20 @@ def progress_cb(pct: int, step: str = "") -> None: send({"type": "log", "level": "warning", "message": ("Model was not loaded (earlier setup failure?); " "reloaded before generating.")}) - output_path = gen.generate(image_bytes, params, progress_cb, cancel_evt) + if hasattr(model_input, "kind"): + output_path = gen.generate_artifact( + model_input.kind, model_input.path, params, progress_cb, cancel_evt + ) + else: + output_path = gen.generate(model_input, params, progress_cb, cancel_evt) + if node.get("output") == "scene": + from services.scene_input import validate_scene_input + resolved_output = Path(output_path).resolve(strict=True) + try: + relative_output = resolved_output.relative_to(WORKSPACE_DIR.resolve(strict=True)) + except (OSError, ValueError) as exc: + raise ValueError("Generated scene output is outside the workspace") from exc + output_path = validate_scene_input(WORKSPACE_DIR, relative_output.as_posix()) send({"type": "done", "id": rid, "output_path": str(output_path)}) except Exception as exc: # Detect GenerationCancelled by name to avoid import issues diff --git a/api/schemas/generation.py b/api/schemas/generation.py index 7ed6ca62..04c18c85 100644 --- a/api/schemas/generation.py +++ b/api/schemas/generation.py @@ -1,5 +1,5 @@ -from typing import Literal, Optional -from pydantic import BaseModel +from typing import Any, Literal, Optional +from pydantic import BaseModel, Field class JobStatus(BaseModel): @@ -9,3 +9,12 @@ class JobStatus(BaseModel): step: Optional[str] = None # Human-readable current step output_url: Optional[str] = None error: Optional[str] = None + + +class GenerateFromArtifactRequest(BaseModel): + """Generic typed-artifact request. Only scene is public in this release.""" + input_kind: Literal["scene"] + input_path: str + model_id: str + collection: str = "Workflows" + params: dict[str, Any] = Field(default_factory=dict) diff --git a/api/services/artifact_input.py b/api/services/artifact_input.py new file mode 100644 index 00000000..165e29ec --- /dev/null +++ b/api/services/artifact_input.py @@ -0,0 +1,29 @@ +"""Typed model artifact inputs shared by the API, worker bridge, and runner.""" +from dataclasses import dataclass +from pathlib import Path + +from services.scene_input import validate_scene_input + +SUPPORTED_ARTIFACT_INPUTS = frozenset({"scene"}) + + +@dataclass(frozen=True) +class TypedArtifactInput: + kind: str + path: Path + + +def validate_artifact_input(workspace: Path, kind: str, input_path: str) -> TypedArtifactInput: + if kind not in SUPPORTED_ARTIFACT_INPUTS: + raise ValueError(f"Unsupported artifact input kind: {kind}") + if kind == "scene": + return TypedArtifactInput(kind="scene", path=validate_scene_input(workspace, input_path)) + raise ValueError(f"Unsupported artifact input kind: {kind}") + + +def revalidate_artifact_input(workspace: Path, value: TypedArtifactInput) -> TypedArtifactInput: + try: + relative = value.path.resolve(strict=True).relative_to(workspace.resolve(strict=True)) + except (OSError, ValueError) as exc: + raise ValueError("Artifact input is outside the workspace") from exc + return validate_artifact_input(workspace, value.kind, relative.as_posix()) diff --git a/api/services/extension_process.py b/api/services/extension_process.py index 67565d36..3a39bfc5 100644 --- a/api/services/extension_process.py +++ b/api/services/extension_process.py @@ -291,16 +291,18 @@ def generate( progress_cb: Optional[Callable[[int, str], None]] = None, cancel_event: Optional[threading.Event] = None, ) -> Path: - from services.generators.base import GenerationCancelled + return self._generate_request( + {"image_b64": base64.b64encode(image_bytes).decode()}, + params, progress_cb, cancel_event, + ) - req_id = str(uuid.uuid4()) - self._send({ - "action": "generate", - "id": req_id, - "image_b64": base64.b64encode(image_bytes).decode(), - "params": params, - "outputs_dir": str(self.outputs_dir) if self.outputs_dir else None, - }) + def _receive_generation( + self, + req_id: str, + progress_cb: Optional[Callable[[int, str], None]], + cancel_event: Optional[threading.Event], + ) -> Path: + from services.generators.base import GenerationCancelled # Grace period after sending a cooperative cancel before hard-killing # the subprocess. Long enough to let generators that check cancel_event @@ -377,6 +379,41 @@ def generate( elif t == "log": print(f"[{self.MODEL_ID}] {msg.get('message', '')}", file=sys.stderr) + def generate_artifact( + self, + input_kind: str, + artifact_path: Path, + params: dict, + progress_cb: Optional[Callable[[int, str], None]] = None, + cancel_event: Optional[threading.Event] = None, + ) -> Path: + """Send a typed artifact envelope to the isolated runner.""" + from services.artifact_input import TypedArtifactInput, revalidate_artifact_input + from services.generator_registry import WORKSPACE_DIR + + validated = revalidate_artifact_input( + WORKSPACE_DIR, TypedArtifactInput(kind=input_kind, path=artifact_path) + ) + return self._generate_request( + {"input": {"kind": validated.kind, "path": str(validated.path)}}, + params, progress_cb, cancel_event, + ) + + def _generate_request( + self, + input_payload: dict, + params: dict, + progress_cb: Optional[Callable[[int, str], None]], + cancel_event: Optional[threading.Event], + ) -> Path: + req_id = str(uuid.uuid4()) + self._send({ + "action": "generate", "id": req_id, "model_id": self.MODEL_ID, + **input_payload, "params": params, + "outputs_dir": str(self.outputs_dir) if self.outputs_dir else None, + }) + return self._receive_generation(req_id, progress_cb, cancel_event) + def params_schema(self) -> list: return self._params_schema diff --git a/api/services/generator_registry.py b/api/services/generator_registry.py index 348a42cb..de1750ca 100644 --- a/api/services/generator_registry.py +++ b/api/services/generator_registry.py @@ -452,6 +452,25 @@ def _discover_extensions( node for node in raw_nodes if isinstance(node, dict) and node.get("id") ] + allowed_io = {"image", "text", "mesh", "audio", "scene"} + for node in nodes: + declared_inputs = node.get("inputs") or [node.get("input", "image")] + if (not isinstance(declared_inputs, list) + or any(value not in allowed_io for value in declared_inputs)): + raise ValueError( + f'model node "{node.get("id", "unknown")}" has an unsupported input type' + ) + if node.get("output", "mesh") not in allowed_io: + raise ValueError( + f'model node "{node.get("id", "unknown")}" has an unsupported output type' + ) + if "scene" in declared_inputs and ( + "inputs" in node or node.get("input", "image") != "scene" + ): + raise ValueError( + f'model node "{node.get("id", "unknown")}" must declare scene ' + 'as its single input field' + ) # Markers left while setup or runtime registration is unfinished: # the folder is not ready to be loaded. The readable manifest lets @@ -540,6 +559,7 @@ def _discover_extensions( "hf_include_prefixes": node.get("hf_include_prefixes", []), "params_schema": node.get("params_schema", manifest.get("params_schema", [])), "input": node.get("input", "image"), + "inputs": node.get("inputs"), "output": node.get("output", "mesh"), } if model_sources is not None: @@ -612,6 +632,9 @@ def initialize( ) # Subprocess mode: wrap in ExtensionProcess gen = ExtensionProcess(ext_dir, manifest) + # Pin the subprocess envelope to the exact registry key; + # multi-node workers must never fall back to an extension ID. + gen.MODEL_ID = model_id gen.model_dir = MODELS_DIR / model_id gen.outputs_dir = WORKSPACE_DIR else: @@ -704,10 +727,13 @@ def _assert_not_quarantined(model_id: str) -> None: def get_active(self) -> BaseGenerator: """Returns the active generator. Downloads and loads if necessary.""" - self._assert_not_quarantined(self._active_id) - gen = self._generators[self._active_id] - downloaded = self._is_downloaded(self._active_id, gen) - if "model_sources" in self._manifests[self._active_id] and not downloaded: + return self.get_ready_generator(self._active_id) + + def get_ready_generator(self, model_id: str) -> BaseGenerator: + """Load and return exactly ``model_id`` without consulting active state.""" + gen = self.get_generator(model_id) + downloaded = self._is_downloaded(model_id, gen) + if "model_sources" in self._manifests[model_id] and not downloaded: raise RuntimeError( "Model sources are incomplete. Download this node's weights " "from the Modly Models page before generation." @@ -724,6 +750,16 @@ def get_active(self) -> BaseGenerator: gen.load() return gen + def model_status(self, model_id: str) -> dict: + gen = self.get_generator(model_id) + manifest = self._manifests[model_id] + return { + "id": model_id, + "name": manifest.get("name", gen.DISPLAY_NAME), + "downloaded": self._is_downloaded(model_id, gen), + "loaded": gen.is_loaded(), + } + def get_generator(self, model_id: str) -> BaseGenerator: self._assert_not_quarantined(model_id) if model_id not in self._generators: @@ -765,14 +801,7 @@ def _is_downloaded(self, model_id: str, gen: BaseGenerator) -> bool: return gen.is_downloaded() def active_status(self) -> dict: - gen = self._generators[self._active_id] - manifest = self._manifests[self._active_id] - return { - "id": self._active_id, - "name": manifest.get("name", gen.DISPLAY_NAME), - "downloaded": self._is_downloaded(self._active_id, gen), - "loaded": gen.is_loaded(), - } + return self.model_status(self._active_id) def all_status(self) -> list: result = [] diff --git a/api/services/generators/base.py b/api/services/generators/base.py index fd62ceef..0b5169a2 100644 --- a/api/services/generators/base.py +++ b/api/services/generators/base.py @@ -143,7 +143,6 @@ def is_loaded(self) -> bool: # Inference # ------------------------------------------------------------------ # - @abstractmethod def generate( self, image_bytes: bytes, @@ -157,7 +156,25 @@ def generate( progress_cb(percent: int, step_label: str) cancel_event: set this to interrupt generation between steps. """ - ... + raise NotImplementedError( + f"{type(self).__name__} does not implement legacy image generation" + ) + + def generate_artifact( + self, + input_kind: str, + artifact_path: Path, + params: dict, + progress_cb: Optional[Callable[[int, str], None]] = None, + cancel_event: Optional[threading.Event] = None, + ) -> Path: + """Generate from a validated typed artifact. + + New extensions should override this method. The default delegates to + ``generate`` with the canonical path so scene-capable extensions built + against the pre-release contract remain compatible. + """ + return self.generate(artifact_path, params, progress_cb, cancel_event) # type: ignore[arg-type] def _check_cancelled(self, cancel_event: Optional[threading.Event]) -> None: """Raises GenerationCancelled if cancel_event is set.""" diff --git a/api/services/scene_input.py b/api/services/scene_input.py new file mode 100644 index 00000000..c5395050 --- /dev/null +++ b/api/services/scene_input.py @@ -0,0 +1,138 @@ +"""Validation shared by the API and isolated model runner for scene inputs.""" +import json +import math +import re +import stat +from pathlib import Path, PurePosixPath, PureWindowsPath + +SCHEMA = "modly.scene-manifest.v1" +MANIFEST = "scene-manifest.json" +MAX_MANIFEST_BYTES = 1024 * 1024 +MAX_REFERENCES = 4096 +MAX_REFERENCED_FILE_BYTES = 16 * 1024**3 +MAX_REFERENCED_TOTAL_BYTES = 64 * 1024**3 + + +def _relative(value: str, *, allow_dot: bool = False) -> Path: + if not isinstance(value, str) or not value or value != value.strip() or "\x00" in value: + raise ValueError("Scene path must be a nonempty workspace-relative path") + value = value.replace("\\", "/") + if value == "." and allow_dot: + return Path(".") + if (PurePosixPath(value).is_absolute() or PureWindowsPath(value).is_absolute() + or re.match(r"^[A-Za-z][A-Za-z0-9+.-]*:", value) + or re.search(r"%(?:25|2e|2f|5c|00)", value, re.I) + or re.search(r"%(?![0-9a-f]{2})", value, re.I) + or any(part in ("", ".", "..") for part in value.split("/"))): + raise ValueError("Scene path must be a safe workspace-relative path") + return Path(*value.split("/")) + + +def _inside(path: Path, root: Path) -> Path: + try: + relative = path.relative_to(root) + except ValueError as exc: + raise ValueError("Scene path escapes its allowed root") from exc + current = root + for part in relative.parts: + current = current / part + try: + info = current.lstat() + except OSError as exc: + raise ValueError("Scene referenced path is missing or unreadable") from exc + is_reparse = bool(getattr(info, "st_file_attributes", 0) & getattr(stat, "FILE_ATTRIBUTE_REPARSE_POINT", 0)) + if current.is_symlink() or is_reparse: + raise ValueError("Scene referenced path must not use symlinks or reparse points") + try: + resolved = path.resolve(strict=True) + except OSError as exc: + raise ValueError("Scene referenced path is missing or unreadable") from exc + if not resolved.is_relative_to(root): + raise ValueError("Scene path escapes its allowed root") + return resolved + + +def validate_scene_input(workspace: Path, scene_path: str) -> Path: + """Return a canonical manifest Path, rejecting traversal and symlink escapes. + + V1 sceneRoot is workspace-relative, except '.' denotes the manifest directory. + An asset's path and preview references are sceneRoot-relative; workspacePath + is always workspace-relative. Existence never determines which base applies. + """ + root = workspace.resolve(strict=True) + relative = _relative(scene_path) + requested = root / relative + if requested.name != MANIFEST: + if requested.suffix.lower() == ".json": + raise ValueError(f"Scene input accepts a directory or {MANIFEST}") + requested = requested / MANIFEST + try: + manifest_path = _inside(requested, root) + if not manifest_path.is_file(): + raise ValueError("Scene manifest is not a file") + if manifest_path.stat().st_size > MAX_MANIFEST_BYTES: + raise ValueError("Scene manifest exceeds the 1 MiB limit") + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + raise ValueError("Scene manifest is missing or invalid JSON") from exc + if not isinstance(manifest, dict) or manifest.get("schema") != SCHEMA: + raise ValueError(f"Scene manifest schema must be {SCHEMA}") + scene_root = _relative(manifest.get("sceneRoot"), allow_dot=True) + # v1 sceneRoot is workspace-relative; only '.' means the manifest's directory. + # Never select a base according to which unrelated path happens to exist. + scene_root_candidate = manifest_path.parent if scene_root == Path(".") else root / scene_root + root_path = _inside(scene_root_candidate, root) + if not root_path.is_dir(): + raise ValueError("Scene root is not a directory") + if not isinstance(manifest.get("assets"), list): + raise ValueError("Scene manifest assets must be an array") + if len(manifest["assets"]) > MAX_REFERENCES: + raise ValueError("Scene manifest contains too many asset references") + total_bytes = 0 + for asset in manifest["assets"]: + # Assets may be opaque extension metadata. Validate only references the + # host understands; never silently resolve a supplied path outside workspace. + if isinstance(asset, dict): + for field, base in (("workspacePath", root), ("path", root_path)): + if field not in asset: + continue + target = _inside(base / _relative(asset[field]), root) + if not target.is_file(): + raise ValueError(f"Scene asset {field} is not a file") + size = target.stat().st_size + if size > MAX_REFERENCED_FILE_BYTES: + raise ValueError("Scene referenced file exceeds the size limit") + total_bytes += size + if total_bytes > MAX_REFERENCED_TOTAL_BYTES: + raise ValueError("Scene referenced files exceed the total size limit") + preview = manifest.get("preview", {}) + if not isinstance(preview, dict): + raise ValueError("Scene preview must be an object") + for name in ("image", "video"): + if name in preview: + relative_preview = _relative(preview[name]) + target = _inside(root_path / relative_preview, root) + if not target.is_file(): + raise ValueError(f"Scene preview {name} is not a file") + view = manifest.get("initialView") + if view is not None: + if not isinstance(view, dict): + raise ValueError("Scene initialView must be an object") + for field in ("position", "target", "up"): + triple = view.get(field) + if triple is None and field == "up": + continue + if (not isinstance(triple, list) or len(triple) != 3 + or any(not isinstance(n, (int, float)) or isinstance(n, bool) + or not math.isfinite(n) for n in triple)): + raise ValueError(f"Scene initialView {field} must be a finite numeric triple") + if view["position"] == view["target"] or view.get("up") == [0, 0, 0]: + raise ValueError("Scene initialView has degenerate camera vectors") + return manifest_path + + +def revalidate_scene_manifest(workspace: Path, manifest_path: Path) -> Path: + """Recheck the file immediately before model invocation in the worker.""" + root = workspace.resolve(strict=True) + candidate = _inside(manifest_path, root) + return validate_scene_input(root, candidate.relative_to(root).as_posix()) diff --git a/api/tests/test_extension_process.py b/api/tests/test_extension_process.py index 348e293f..7db13e4a 100644 --- a/api/tests/test_extension_process.py +++ b/api/tests/test_extension_process.py @@ -2,6 +2,9 @@ import platform import queue import unittest +import json +import tempfile +from unittest.mock import patch from pathlib import Path from services.extension_process import ExtensionProcess, _venv_python @@ -12,6 +15,31 @@ def _make_proc() -> ExtensionProcess: class ExtensionProcessTests(unittest.TestCase): + def test_generation_envelope_pins_worker_model_id(self) -> None: + proc = _make_proc() + sent = [] + proc._send = sent.append + proc._receive_generation = lambda *args: Path("result.glb") + proc._generate_request({"image_b64": ""}, {}, None, None) + self.assertEqual(sent[0]["model_id"], "demo") + + def test_generate_artifact_sends_typed_scene_without_image_bytes(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + workspace = Path(tmp) / "workspace" + scene = workspace / "Workflows" / "room" + scene.mkdir(parents=True) + manifest = scene / "scene-manifest.json" + manifest.write_text(json.dumps({ + "schema": "modly.scene-manifest.v1", "sceneRoot": ".", "assets": [], + })) + proc = _make_proc() + calls = [] + proc._generate_request = lambda payload, params, progress, cancel: calls.append((payload, params)) or manifest + with patch("services.generator_registry.WORKSPACE_DIR", workspace): + result = proc.generate_artifact("scene", manifest, {"quality": "high"}) + self.assertEqual(result, manifest) + self.assertEqual(calls, [({"input": {"kind": "scene", "path": str(manifest.resolve())}}, {"quality": "high"})]) + def test_read_loop_writes_sentinel_to_own_queue_only(self) -> None: proc = _make_proc() diff --git a/api/tests/test_generation_router.py b/api/tests/test_generation_router.py index 20fdda94..5e03c241 100644 --- a/api/tests/test_generation_router.py +++ b/api/tests/test_generation_router.py @@ -49,6 +49,9 @@ def get_active(self) -> _FakeGenerator: def get_generator(self, model_id: str) -> _FakeGenerator: return self._gen + def get_manifest(self, model_id: str) -> dict: + return {"output": "mesh"} + def switch_model(self, model_id: str) -> None: pass diff --git a/api/tests/test_generator_registry.py b/api/tests/test_generator_registry.py index ff9d090c..a4814b9e 100644 --- a/api/tests/test_generator_registry.py +++ b/api/tests/test_generator_registry.py @@ -155,6 +155,53 @@ def test_legacy_generator_supports_eager_and_lazy_sibling_imports(self) -> None: self.registry.reload() self.assertNotIn(str(extension.resolve()), sys.path) + def test_scene_io_is_registered_but_capture_and_video_are_rejected(self) -> None: + for extension_id, input_kind in (("scene-io", "scene"), ("capture-io", "capture"), ("video-io", "video")): + extension = self._make_extension(extension_id) + manifest = { + "id": extension_id, "name": extension_id, "type": "model", + "generator_class": "TestGenerator", + "nodes": [{"id": "generate", "input": input_kind, "output": "scene"}], + } + (extension / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + (extension / "generator.py").write_text( + "from services.generators.base import BaseGenerator\n" + "class TestGenerator(BaseGenerator):\n" + " def load(self): self._model = object()\n" + " def generate(self, value, params, progress_cb=None, cancel_event=None): return self.outputs_dir\n", + encoding="utf-8", + ) + + self.registry.initialize() + self.assertEqual(self.registry.get_manifest("scene-io/generate")["input"], "scene") + self.assertIn("capture-io/generate", self.registry.load_errors()) + self.assertIn("video-io/generate", self.registry.load_errors()) + + def test_scene_input_rejects_multi_input_shapes_but_image_multi_can_output_scene(self) -> None: + cases = { + "scene-mixed": {"input": "scene", "inputs": ["scene", "text"], "output": "mesh"}, + "scene-array": {"input": "scene", "inputs": ["scene"], "output": "mesh"}, + "images-scene": {"input": "image", "inputs": ["image", "image"], "output": "scene"}, + } + for extension_id, node in cases.items(): + extension = self._make_extension(extension_id) + (extension / "manifest.json").write_text(json.dumps({ + "id": extension_id, "name": extension_id, "type": "model", + "generator_class": "TestGenerator", + "nodes": [{"id": "generate", **node}], + }), encoding="utf-8") + (extension / "generator.py").write_text( + "from services.generators.base import BaseGenerator\n" + "class TestGenerator(BaseGenerator):\n" + " def load(self): self._model = object()\n" + " def generate(self, value, params, progress_cb=None, cancel_event=None): return self.outputs_dir\n", + encoding="utf-8", + ) + self.registry.initialize() + self.assertIn("scene-mixed/generate", self.registry.load_errors()) + self.assertIn("scene-array/generate", self.registry.load_errors()) + self.assertIn("images-scene/generate", self.registry._generators) + def test_declared_sources_block_generation_even_when_generator_overrides_readiness(self) -> None: extension = self._make_extension("multi-source") manifest = { diff --git a/api/tests/test_runner.py b/api/tests/test_runner.py index 8fce3d31..a666f87a 100644 --- a/api/tests/test_runner.py +++ b/api/tests/test_runner.py @@ -5,6 +5,7 @@ import json import tempfile import importlib +from unittest.mock import patch from contextlib import redirect_stdout from pathlib import Path @@ -20,6 +21,30 @@ class RunnerTests(unittest.TestCase): + def test_decode_typed_scene_revalidates_worker_workspace_and_keeps_legacy_image(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + workspace = Path(tmp) / "workspace" + scene = workspace / "Workflows" / "room" + scene.mkdir(parents=True) + manifest = scene / "scene-manifest.json" + manifest.write_text(json.dumps({ + "schema": "modly.scene-manifest.v1", "sceneRoot": ".", "assets": [], + })) + with patch.object(runner, "WORKSPACE_DIR", workspace): + typed = runner.decode_model_input({"input": {"kind": "scene", "path": str(manifest)}}) + self.assertEqual(typed.kind, "scene") + self.assertEqual(typed.path, manifest.resolve()) + self.assertEqual(runner.decode_model_input({"image_b64": "aW1hZ2U="}), b"image") + with self.assertRaises(ValueError): + runner.decode_model_input({"input": {"kind": "video", "path": str(manifest)}}) + + def test_runner_model_envelope_rejects_cross_node_dispatch(self) -> None: + manifest = {"id": "pixal3d"} + node = {"id": "worldsculpt"} + runner.validate_requested_model({"model_id": "pixal3d/worldsculpt"}, manifest, node) + with self.assertRaisesRegex(ValueError, "does not match"): + runner.validate_requested_model({"model_id": "pixal3d/generate"}, manifest, node) + def test_select_node_uses_model_dir_override(self) -> None: manifest = { "nodes": [ diff --git a/api/tests/test_scene_generation.py b/api/tests/test_scene_generation.py new file mode 100644 index 00000000..a804171a --- /dev/null +++ b/api/tests/test_scene_generation.py @@ -0,0 +1,115 @@ +import asyncio +import json +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +from fastapi import BackgroundTasks, HTTPException +from pydantic import ValidationError + +import routers.generation as generation +import services.generator_registry as registry +from schemas.generation import GenerateFromArtifactRequest + + +class _Registry: + def __init__(self): + self.switched = False + def get_generator(self, model_id): return object() + def get_manifest(self, model_id): return {"input": "scene"} + def switch_model(self, model_id): self.switched = True + + +class SceneGenerationTests(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + self.workspace = Path(self.tmp.name) / "workspace" + self.scene = self.workspace / "Workflows" / "room" + self.scene.mkdir(parents=True) + self.manifest = self.scene / "scene-manifest.json" + self.manifest.write_text(json.dumps({"schema": "modly.scene-manifest.v1", "sceneRoot": ".", "assets": []})) + self.registry = _Registry() + self.patches = [patch.object(generation, "generator_registry", self.registry), patch.object(registry, "WORKSPACE_DIR", self.workspace)] + for item in self.patches: item.start() + + def tearDown(self): + for item in reversed(self.patches): item.stop() + generation._jobs.clear(); generation._cancel_events.clear(); generation._cancelled.clear(); generation._completed_at.clear() + self.tmp.cleanup() + + def test_generic_route_queues_typed_scene_and_strips_reserved_params(self): + tasks = BackgroundTasks() + result = asyncio.run(generation.generate_from_artifact(GenerateFromArtifactRequest( + input_kind="scene", input_path="Workflows/room", model_id="demo/scene", + params={"artifact_path": "/etc/passwd", "input_kind": "image", "quality": "high"}, + ), tasks)) + queued = tasks.tasks[0] + self.assertEqual(queued.args[1].kind, "scene") + self.assertEqual(queued.args[1].path, self.manifest.resolve()) + self.assertEqual(queued.args[2]["scene_manifest_path"], str(self.manifest.resolve())) + self.assertNotIn("artifact_path", queued.args[2]) + self.assertNotIn("input_kind", queued.args[2]) + self.assertEqual(queued.args[5], "demo/scene") + self.assertEqual(result["job_id"], queued.args[0]) + + def test_generic_route_rejects_unsupported_kind_and_model_mismatch(self): + for kind in ("video", "capture", "image"): + with self.subTest(kind=kind), self.assertRaises(ValidationError): + GenerateFromArtifactRequest( + input_kind=kind, input_path="Workflows/room", model_id="demo/scene") + self.registry.get_manifest = lambda _model_id: {"input": "image"} + with self.assertRaises(HTTPException) as caught: + asyncio.run(generation.generate_from_artifact(GenerateFromArtifactRequest( + input_kind="scene", input_path="Workflows/room", model_id="demo/image"), BackgroundTasks())) + self.assertEqual(caught.exception.status_code, 400) + + def test_rejects_traversal_before_switch_or_queue(self): + with self.assertRaises(HTTPException): + asyncio.run(generation.generate_from_artifact(GenerateFromArtifactRequest( + input_kind="scene", input_path="../outside", model_id="demo/scene"), BackgroundTasks())) + self.assertFalse(self.registry.switched) + + def test_queued_scene_job_is_pinned_to_requested_model(self): + calls = [] + + class Generator: + outputs_dir = None + def is_loaded(self): return True + def generate_artifact(self, kind, path, params, progress_cb, cancel_event=None): + calls.append(("model-a", kind, path)) + output = Path(self.outputs_dir) / "result.glb" + output.write_bytes(b"glb") + return output + + generator = Generator() + registry_stub = type("Registry", (), { + "get_ready_generator": lambda self, model_id: generator if model_id == "demo/a" else (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), + "get_active": lambda self: (_ for _ in ()).throw(AssertionError("mutable active model must not be used")), + })() + job_id = "pinned-scene" + generation._jobs[job_id] = generation.JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = __import__("threading").Event() + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(generation._run_generation( + job_id, generation.TypedArtifactInput("scene", self.manifest.resolve()), {}, + "Workflows", "mesh", "demo/a", + )) + self.assertEqual(calls[0][0], "model-a") + self.assertEqual(generation._jobs[job_id].status, "done") + + def test_missing_pinned_model_fails_actionably(self): + registry_stub = type("Registry", (), { + "get_ready_generator": lambda self, model_id: (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), + "get_active": lambda self: (_ for _ in ()).throw(AssertionError("must not use active model")), + })() + job_id = "missing-scene" + generation._jobs[job_id] = generation.JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = __import__("threading").Event() + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(generation._run_generation( + job_id, generation.TypedArtifactInput("scene", self.manifest.resolve()), {}, + "Workflows", "mesh", "demo/missing", + )) + self.assertEqual(generation._jobs[job_id].status, "error") + self.assertIn("Unknown model ID: demo/missing", generation._jobs[job_id].error) diff --git a/api/tests/test_scene_input.py b/api/tests/test_scene_input.py new file mode 100644 index 00000000..43bd9946 --- /dev/null +++ b/api/tests/test_scene_input.py @@ -0,0 +1,67 @@ +import json +import tempfile +import unittest +from pathlib import Path + +from services.scene_input import validate_scene_input, revalidate_scene_manifest + + +class SceneInputTests(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + self.workspace = Path(self.tmp.name) / "workspace" + self.scene = self.workspace / "Workflows" / "room" + self.scene.mkdir(parents=True) + (self.scene / "model.glb").write_bytes(b"mesh") + self.manifest = self.scene / "scene-manifest.json" + self.manifest.write_text(json.dumps({ + "schema": "modly.scene-manifest.v1", + "sceneRoot": ".", + "assets": [{"path": "model.glb"}], + })) + + def tearDown(self): + self.tmp.cleanup() + + def test_directory_and_manifest_are_canonical_paths(self): + expected = self.manifest.resolve() + self.assertEqual(validate_scene_input(self.workspace, "Workflows/room"), expected) + self.assertEqual(validate_scene_input(self.workspace, "Workflows/room/scene-manifest.json"), expected) + self.assertEqual(revalidate_scene_manifest(self.workspace, expected), expected) + + def test_rejects_traversal_absolute_encoded_and_non_manifest_json(self): + for path in ("../outside", "/etc/passwd", "C:/outside", "Workflows/room/../room", "Workflows/room/other.json", "Workflows/%2e%2e"): + with self.subTest(path=path), self.assertRaises(ValueError): + validate_scene_input(self.workspace, path) + + def test_rejects_symlinks_missing_assets_and_oversized_manifest(self): + outside = Path(self.tmp.name) / "outside" + outside.mkdir() + (outside / "scene-manifest.json").write_text(self.manifest.read_text()) + (self.workspace / "Workflows" / "link").symlink_to(outside, target_is_directory=True) + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/link") + + data = json.loads(self.manifest.read_text()) + data["assets"] = [{"path": "missing.glb"}] + self.manifest.write_text(json.dumps(data)) + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/room") + + self.manifest.write_text(" " * (1024 * 1024 + 1)) + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/room") + + def test_rejects_malformed_manifest_and_asset_paths(self): + original = json.loads(self.manifest.read_text()) + for patch in ( + {"schema": "wrong"}, + {"assets": "bad"}, + {"assets": [{"path": "../escape.glb"}]}, + {"preview": {"image": None}}, + {"initialView": {"position": [0, 0, 0], "target": [0, 0, 0]}}, + ): + with self.subTest(patch=patch): + self.manifest.write_text(json.dumps({**original, **patch})) + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/room") diff --git a/electron/main/artifact-registry-service.test.ts b/electron/main/artifact-registry-service.test.ts index 882120d4..276e90a1 100644 --- a/electron/main/artifact-registry-service.test.ts +++ b/electron/main/artifact-registry-service.test.ts @@ -64,6 +64,18 @@ test('lists Workflows and Exports assets while skipping hidden, cache, and inter assert.equal(result.success && result.entries.find((entry) => entry.workspacePath.endsWith('exported.ply'))?.openable, false) })) +test('registers a generated scene directory through its canonical manifest artifact', () => withWorkspace(async (workspaceDir) => { + await mkdir(path.join(workspaceDir, 'Workflows/world'), { recursive: true }) + await writeFile(path.join(workspaceDir, 'Workflows/world/scene-manifest.json'), JSON.stringify({ + schema: 'modly.scene-manifest.v1', sceneRoot: '.', assets: [], + })) + const result = await listWorkspaceAssetLibrary({ workspaceDir }) + assert.equal(result.success, true) + const scene = result.success && result.entries.find((entry) => entry.workspacePath === 'Workflows/world/scene-manifest.json') + assert.equal(scene && scene.capability, 'scene-manifest') + assert.equal(scene && scene.state, 'ready') +})) + test('reads and opens only safe GLB/GLTF workspace assets', () => withWorkspace(async (workspaceDir) => { await mkdir(path.join(workspaceDir, 'Workflows/checkpoints'), { recursive: true }) await mkdir(path.join(workspaceDir, 'Exports'), { recursive: true }) diff --git a/electron/main/artifact-registry-service.ts b/electron/main/artifact-registry-service.ts index 3c26de73..2acb657f 100644 --- a/electron/main/artifact-registry-service.ts +++ b/electron/main/artifact-registry-service.ts @@ -124,7 +124,7 @@ export function classifyAssetLibraryCandidate(candidate: AssetLibraryClassificat if (candidate.workspacePath.endsWith('.world.json')) { return { capability: 'generated-world', state: 'ready', previewKind: 'text', openable: false, nonOpenableReason: 'Generated worlds are list-only in this release.' } } - if (candidate.workspacePath.endsWith('.scene.json')) { + if (candidate.workspacePath.endsWith('.scene.json') || candidate.workspacePath.endsWith('/scene-manifest.json')) { return { capability: 'scene-manifest', state: 'ready', previewKind: 'text', openable: false, nonOpenableReason: 'Scene manifests are list-only in this release.' } } if (INTRINSIC_MOTION_EXTENSIONS.has(extension)) { @@ -153,6 +153,7 @@ function objectField(value: unknown): Record | undefined { function manifestCapabilityFor(workspacePath: string): 'generated-world' | 'scene-manifest' | undefined { if (workspacePath.endsWith('.world.json')) return 'generated-world' if (workspacePath.endsWith('.scene.json')) return 'scene-manifest' + if (workspacePath.endsWith('/scene-manifest.json')) return 'scene-manifest' return undefined } diff --git a/electron/main/extension-install-utils.test.mjs b/electron/main/extension-install-utils.test.mjs index 84139f9a..4dcf42da 100644 --- a/electron/main/extension-install-utils.test.mjs +++ b/electron/main/extension-install-utils.test.mjs @@ -85,6 +85,50 @@ test('validateInstallManifest accepts multi-source nodes and preserves legacy sh }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) }) +test('validateInstallManifest accepts scene IO and rejects undeclared future artifact kinds', () => { + const mod = loadModule() + const files = { hasEntryFile: () => false, hasGeneratorFile: () => true } + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'scene-model', generator_class: 'Generator', + nodes: [{ id: 'normalize', input: 'scene', output: 'scene' }], + }, files, 'repository')) + for (const input of ['capture', 'video']) { + assert.throws(() => mod.validateInstallManifest({ + id: 'future-model', generator_class: 'Generator', + nodes: [{ id: 'future', input, output: 'scene' }], + }, files, 'repository'), /supported artifact type/) + } +}) + +test('scene is model-only, single-input, while image-multi to scene stays valid', () => { + const mod = loadModule() + const modelFiles = { hasEntryFile: () => false, hasGeneratorFile: () => true } + const processFiles = { hasEntryFile: () => true, hasGeneratorFile: () => false } + for (const node of [ + { id: 'mixed', input: 'scene', inputs: ['scene', 'text'], output: 'mesh' }, + { id: 'duplicate', input: 'scene', inputs: ['scene', 'scene'], output: 'mesh' }, + { id: 'hidden', input: 'image', inputs: ['scene'], output: 'mesh' }, + ]) { + assert.throws(() => mod.validateInstallManifest({ id: 'bad', generator_class: 'Generator', nodes: [node] }, modelFiles, 'repository'), /scene.*single|single.*scene/i) + } + for (const node of [ + { id: 'input', input: 'scene', output: 'mesh' }, + { id: 'output', input: 'image', output: 'scene' }, + ]) { + assert.throws(() => mod.validateInstallManifest({ id: 'proc', type: 'process', entry: 'processor.js', nodes: [node] }, processFiles, 'repository'), /scene.*model|model.*scene/i) + } + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'images-to-scene', generator_class: 'Generator', + nodes: [{ id: 'prepare', input: 'image', inputs: ['image', 'image'], output: 'scene' }], + }, modelFiles, 'repository')) + for (const output of ['scene', 'mesh']) { + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: `scene-to-${output}`, generator_class: 'Generator', + nodes: [{ id: 'generate', input: 'scene', output }], + }, modelFiles, 'repository')) + } +}) + test('validateInstallManifest rejects malformed or process model_sources', () => { const mod = loadModule() const source = { diff --git a/electron/main/extension-install-utils.ts b/electron/main/extension-install-utils.ts index 05b965b0..41e9798c 100644 --- a/electron/main/extension-install-utils.ts +++ b/electron/main/extension-install-utils.ts @@ -10,7 +10,13 @@ export interface InstallManifest { entry?: string generator_class?: string model_sources?: unknown - nodes?: Array<{ id?: string; model_sources?: unknown } & ModelSourceNode> + nodes?: Array<{ + id?: string + input?: unknown + inputs?: unknown + output?: unknown + model_sources?: unknown + } & ModelSourceNode> } export interface ValidatedInstallManifest { @@ -27,6 +33,21 @@ export interface ExtensionReloadPayload { errors: Record } +export function assertSupportedSceneNodeShape( + kind: 'model' | 'process', + node: { id?: string; input?: unknown; inputs?: unknown; output?: unknown }, + declaredInputs: unknown[], + output: unknown, +): void { + const usesSceneInput = declaredInputs.includes('scene') + if (kind === 'process' && (usesSceneInput || output === 'scene')) { + throw new Error('manifest.json: scene input and output are supported only for model nodes') + } + if (kind === 'model' && usesSceneInput && (node.inputs !== undefined || node.input !== 'scene')) { + throw new Error(`manifest.json: ${node.id ?? 'node'} must declare scene as its single input field`) + } +} + export type IncompleteInstallRecoveryAction = | 'none' | 'remove-incomplete' @@ -45,11 +66,22 @@ export function validateInstallManifest( const isProcess = manifest.type === 'process' const entryFile = manifest.entry ?? 'processor.js' const nodes = Array.isArray(manifest.nodes) ? manifest.nodes.filter((node) => node?.id) : [] + const allowedIo = new Set(['image', 'text', 'mesh', 'audio', 'scene']) if (manifest.model_sources !== undefined) { throw new Error('manifest.json: model_sources must be declared on a model node') } for (const node of Array.isArray(manifest.nodes) ? manifest.nodes : []) { + const declaredInputs = node.inputs === undefined ? [node.input ?? 'image'] : node.inputs + if (!Array.isArray(declaredInputs) || declaredInputs.length === 0 + || declaredInputs.some((value) => typeof value !== 'string' || !allowedIo.has(value))) { + throw new Error(`manifest.json: ${node.id ?? 'node'}.input must use a supported artifact type`) + } + const output = node.output ?? 'mesh' + if (typeof output !== 'string' || !allowedIo.has(output)) { + throw new Error(`manifest.json: ${node.id ?? 'node'}.output must use a supported artifact type`) + } + assertSupportedSceneNodeShape(isProcess ? 'process' : 'model', node, declaredInputs, output) if (node.model_sources === undefined) continue if (isProcess) { throw new Error('manifest.json: model_sources is supported only for model nodes') diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 4fa24788..9e5131de 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -49,6 +49,7 @@ import { validateExtensionReloadPayload, validateExistingExtensionReplacement, validateInstallManifest, + assertSupportedSceneNodeShape, } from './extension-install-utils' import { beginExtensionRegistrationTransaction, @@ -844,10 +845,10 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe nodes?: { id: string name?: string - input?: 'mesh' | 'image' | 'text' | 'audio' - inputs?: ('mesh' | 'image' | 'text' | 'audio')[] + input?: 'mesh' | 'image' | 'text' | 'audio' | 'scene' + inputs?: ('mesh' | 'image' | 'text' | 'audio' | 'scene')[] input_labels?: string[] - output?: 'mesh' | 'image' | 'text' | 'audio' + output?: 'mesh' | 'image' | 'text' | 'audio' | 'scene' params_schema?: unknown[] param_defaults?: Record hf_repo?: string @@ -873,7 +874,15 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe if (parsed.model_sources !== undefined) { throw new Error('manifest.json: model_sources must be declared on a model node') } + const allowedIo = new Set(['image', 'text', 'mesh', 'audio', 'scene']) const nodes = (parsed.nodes ?? []).map(n => { + const declaredInputs = n.inputs ?? [n.input ?? 'image'] + for (const input of declaredInputs) { + if (!allowedIo.has(input)) throw new Error(`manifest.json: unsupported node input type "${input}"`) + } + const output = n.output ?? 'mesh' + if (!allowedIo.has(output)) throw new Error(`manifest.json: unsupported node output type "${output}"`) + assertSupportedSceneNodeShape(parsed.type === 'process' ? 'process' : 'model', n, declaredInputs, output) if (parsed.type === 'process' && n.model_sources !== undefined) { throw new Error('manifest.json: model_sources is supported only for model nodes') } diff --git a/src/areas/workflows/WorkflowsPage.tsx b/src/areas/workflows/WorkflowsPage.tsx index d15c06ff..32452774 100644 --- a/src/areas/workflows/WorkflowsPage.tsx +++ b/src/areas/workflows/WorkflowsPage.tsx @@ -27,6 +27,7 @@ import ImageNode from './nodes/ImageNode' import TextNode from './nodes/TextNode' import AddToSceneNode from './nodes/AddToSceneNode' import Load3DMeshNode from './nodes/Load3DMeshNode' +import LoadSceneNode from './nodes/LoadSceneNode' import PreviewImageNode from './nodes/PreviewImageNode' import ImagePreviewNode from './nodes/ImagePreviewNode' import WaitNode from './nodes/WaitNode' @@ -38,7 +39,7 @@ import WorkflowEdge from './nodes/WorkflowEdge' const DRAG_KEY = 'modly/extension-id' const DRAG_NODE_KEY = 'modly/node-type' -const NODE_TYPES = { extensionNode: ExtensionNode, imageNode: ImageNode, textNode: TextNode, outputNode: AddToSceneNode, meshNode: Load3DMeshNode, previewNode: PreviewImageNode, imagePreviewNode: ImagePreviewNode, waitNode: WaitNode, whileNode: WhileNode, forEachNode: ForEachNode } +const NODE_TYPES = { extensionNode: ExtensionNode, imageNode: ImageNode, textNode: TextNode, outputNode: AddToSceneNode, meshNode: Load3DMeshNode, sceneNode: LoadSceneNode, previewNode: PreviewImageNode, imagePreviewNode: ImagePreviewNode, waitNode: WaitNode, whileNode: WhileNode, forEachNode: ForEachNode } // Loop-container node types: resizable frames whose children form a loop body. // (For Each iterators are plain source nodes, not containers.) @@ -62,14 +63,15 @@ function findWhileContainerAt(nodes: Node[], pos: { x: number; y: number }): Nod // ─── IO badge ───────────────────────────────────────────────────────────────── -const IO_STYLES: Record<'image' | 'text' | 'mesh' | 'audio', string> = { +const IO_STYLES: Record<'image' | 'text' | 'mesh' | 'audio' | 'scene', string> = { audio: 'bg-emerald-500/15 text-emerald-400 border-emerald-500/25', image: 'bg-sky-500/15 text-sky-400 border-sky-500/25', mesh: 'bg-violet-500/15 text-violet-400 border-violet-500/25', text: 'bg-amber-500/15 text-amber-400 border-amber-500/25', + scene: 'bg-emerald-500/15 text-emerald-400 border-emerald-500/25', } -function IoBadge({ type }: { type: 'image' | 'text' | 'mesh' | 'audio' }) { +function IoBadge({ type }: { type: 'image' | 'text' | 'mesh' | 'audio' | 'scene' }) { return ( {type} @@ -99,6 +101,7 @@ const PANEL_BUILTIN_NODES = [ { type: 'imageNode', label: 'Image', color: '#38bdf8', icon: <> }, { type: 'textNode', label: 'Text', color: '#fbbf24', icon: <> }, { type: 'meshNode', label: 'Load 3D Mesh', color: '#a78bfa', icon: <> }, + { type: 'sceneNode', label: 'Load Scene', color: '#34d399', icon: <> }, { type: 'outputNode', label: 'Add to Scene', color: '#a78bfa', icon: <> }, { type: 'previewNode', label: 'Preview Views', color: '#38bdf8', icon: <> }, { type: 'imagePreviewNode', label: 'Preview Image', color: '#38bdf8', icon: <> }, @@ -346,6 +349,7 @@ const BUILTIN_NODES = [ { type: 'imageNode', label: 'Image', color: '#38bdf8', description: 'Image input' }, { type: 'textNode', label: 'Text', color: '#fbbf24', description: 'Text input' }, { type: 'meshNode', label: 'Load 3D Mesh', color: '#a78bfa', description: 'Load a 3D mesh file or use current model' }, + { type: 'sceneNode', label: 'Load Scene', color: '#34d399', description: 'Load and validate a workspace scene directory' }, { type: 'outputNode', label: 'Add to Scene', color: '#a78bfa', description: 'Output node — adds the mesh to the 3D scene' }, { type: 'previewNode', label: 'Preview Views', color: '#38bdf8', description: 'Displays multi-view image outputs in a 2×3 grid' }, { type: 'imagePreviewNode', label: 'Preview Image', color: '#38bdf8', description: 'Displays a single image output in the workflow' }, @@ -703,6 +707,7 @@ function getNodeOutputType(node: Node | undefined, allExts: WorkflowExtension[]) if (!node) return undefined if (node.type === 'imageNode') return 'image' if (node.type === 'meshNode') return 'mesh' + if (node.type === 'sceneNode') return 'scene' if (node.type === 'textNode') return 'text' if (node.type === 'imagePreviewNode') return 'image' return allExts.find((e) => e.id === (node.data as WFNodeData)?.extensionId)?.output @@ -1374,6 +1379,7 @@ const MINI_NODE_TINTS: Record = { imageNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, textNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, meshNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, + sceneNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, extensionNode: { fill: 'rgba(167,139,250,0.24)', stroke: '#a78bfa' }, outputNode: { fill: 'rgba(56,189,248,0.22)', stroke: '#38bdf8' }, previewNode: { fill: 'rgba(56,189,248,0.22)', stroke: '#38bdf8' }, diff --git a/src/areas/workflows/mockExtensions.ts b/src/areas/workflows/mockExtensions.ts index 2bbc8cc6..be727fea 100644 --- a/src/areas/workflows/mockExtensions.ts +++ b/src/areas/workflows/mockExtensions.ts @@ -10,10 +10,10 @@ export interface WorkflowExtension { nodeId: string // "node_id" name: string description: string - input: 'image' | 'text' | 'mesh' | 'audio' - inputs?: ('image' | 'text' | 'mesh' | 'audio')[] // multi-input; overrides input when set + input: 'image' | 'text' | 'mesh' | 'audio' | 'scene' + inputs?: ('image' | 'text' | 'mesh' | 'audio' | 'scene')[] // multi-input; overrides input when set inputLabels?: string[] // display labels per input slot - output: 'image' | 'text' | 'mesh' | 'audio' + output: 'image' | 'text' | 'mesh' | 'audio' | 'scene' params: ParamSchema[] builtin: boolean type: 'model' | 'process' diff --git a/src/areas/workflows/nodes/ExtensionNode.tsx b/src/areas/workflows/nodes/ExtensionNode.tsx index abe2f9f6..4be11a14 100644 --- a/src/areas/workflows/nodes/ExtensionNode.tsx +++ b/src/areas/workflows/nodes/ExtensionNode.tsx @@ -16,6 +16,7 @@ const HANDLE_COLOR: Record = { image: '#38bdf8', mesh: '#a78bfa', text: '#fbbf24', + scene: '#34d399', } const TAG_CLS: Record = { @@ -23,6 +24,7 @@ const TAG_CLS: Record = { image: 'border-sky-500/30 bg-sky-500/10 text-sky-400', mesh: 'border-violet-500/30 bg-violet-500/10 text-violet-400', text: 'border-amber-500/30 bg-amber-500/10 text-amber-400', + scene: 'border-emerald-500/30 bg-emerald-500/10 text-emerald-400', } // ─── Param control ──────────────────────────────────────────────────────────── diff --git a/src/areas/workflows/nodes/LoadSceneNode.tsx b/src/areas/workflows/nodes/LoadSceneNode.tsx new file mode 100644 index 00000000..2ac7dded --- /dev/null +++ b/src/areas/workflows/nodes/LoadSceneNode.tsx @@ -0,0 +1,134 @@ +import { useCallback, useLayoutEffect, useRef, useState } from 'react' +import { Handle, Position, useReactFlow } from '@xyflow/react' +import type { WFNodeData } from '@shared/types/electron.d' + +import BaseNode from './BaseNode' +import { resolveSceneSourceManifest } from '../workflowSceneSource' + +const OUTPUT_COLOR = '#34d399' + +async function validateAndPersistScenePath(args: { + id: string + data: WFNodeData + nextPath: string + updateNodeData: ReturnType['updateNodeData'] +}): Promise { + const settings = await window.electron.settings.get() + const resolution = await resolveSceneSourceManifest({ + scenePath: args.nextPath, + workspaceDir: settings.workspaceDir, + readFileBase64: window.electron.fs.readFileBase64, + }) + + if (!resolution.ok) { + args.updateNodeData(args.id, { + params: { + ...args.data.params, + path: args.nextPath, + manifestPath: undefined, + sceneRoot: undefined, + error: resolution.error, + }, + }) + return + } + + args.updateNodeData(args.id, { + params: { + ...args.data.params, + path: resolution.inputWorkspacePath, + manifestPath: resolution.manifestWorkspacePath, + sceneRoot: resolution.sceneRoot, + sourceKind: resolution.sourceKind, + error: undefined, + }, + }) +} + +export default function LoadSceneNode({ id, data, selected }: { id: string; data: WFNodeData; selected?: boolean }) { + const { updateNodeData } = useReactFlow() + const ioRowRef = useRef(null) + const [handleTop, setHandleTop] = useState('50%') + + useLayoutEffect(() => { + if (ioRowRef.current) { + const center = ioRowRef.current.offsetTop + ioRowRef.current.offsetHeight / 2 + setHandleTop(`${center}px`) + } + }, []) + + const scenePath = typeof data.params.path === 'string' ? data.params.path : '' + const manifestPath = typeof data.params.manifestPath === 'string' ? data.params.manifestPath : undefined + const sceneRoot = typeof data.params.sceneRoot === 'string' ? data.params.sceneRoot : undefined + const error = typeof data.params.error === 'string' ? data.params.error : undefined + + const browseDirectory = useCallback(async () => { + const path = await window.electron.fs.selectDirectory() + if (!path) return + await validateAndPersistScenePath({ id, data, nextPath: path, updateNodeData }) + }, [id, data, updateNodeData]) + + const validatePath = useCallback(async () => { + if (!scenePath.trim()) return + await validateAndPersistScenePath({ id, data, nextPath: scenePath, updateNodeData }) + }, [id, data, scenePath, updateNodeData]) + + return ( + + + + + + + } + subheader={ +
+ scene +
+ } + handles={ + + } + > +
+ updateNodeData(id, { params: { ...data.params, path: event.target.value } })} + className="nodrag w-full rounded-lg border border-zinc-700 bg-zinc-800 px-2.5 py-2 text-[10px] text-zinc-200 placeholder-zinc-600 focus:outline-none focus:border-emerald-500/40" + /> +
+ + +
+ {manifestPath ? ( +
+
Manifest: {manifestPath}
+ {sceneRoot &&
sceneRoot: {sceneRoot}
} +
+ ) : ( +
+ Loads an existing workspace scene manifest for downstream scene nodes. +
+ )} + {error &&
{error}
} +
+
+ ) +} diff --git a/src/areas/workflows/preflight.test.mjs b/src/areas/workflows/preflight.test.mjs index 75b09d71..3b4df8aa 100644 --- a/src/areas/workflows/preflight.test.mjs +++ b/src/areas/workflows/preflight.test.mjs @@ -131,3 +131,38 @@ test('multi-input extension requires every declared input type', () => { assert.ok(!issues.some((i) => i.key === 'proc:missing:image')) assert.ok(issues.some((i) => i.key === 'proc:missing:text')) }) + +test('scene input requires a validated Load Scene source and rejects image wiring', () => { + const { validateWorkflowPreflight } = loadModule() + const model = { id: 'model', type: 'extensionNode', position: { x: 0, y: 0 }, data: { extensionId: 'pack/process-node' } } + const scene = { id: 'scene', type: 'sceneNode', position: { x: 0, y: 0 }, data: { params: { manifestPath: 'Workflows/room/scene-manifest.json' } } } + for (const output of ['scene', 'mesh']) { + const extension = ext({ input: 'scene', output, type: 'model' }) + assert.deepEqual(validateWorkflowPreflight(wf([scene, model], [{ id: 'scene-edge', source: 'scene', target: 'model' }]), [extension]), []) + } + + const extension = ext({ input: 'scene', output: 'scene', type: 'model' }) + const issues = validateWorkflowPreflight(wf([imageNode(), model], [{ id: 'image-edge', source: 'img', target: 'model' }]), [extension]) + assert.ok(issues.some((issue) => issue.key === 'model:missing:scene')) + assert.ok(issues.some((issue) => issue.key === 'model:type:image-edge')) +}) + +test('Load Scene must be validated before a workflow can run', () => { + const { validateWorkflowPreflight } = loadModule() + const scene = { id: 'scene', type: 'sceneNode', position: { x: 0, y: 0 }, data: { params: { path: 'Workflows/room' } } } + const issues = validateWorkflowPreflight(wf([scene], []), []) + assert.equal(issues[0].key, 'scene:scene-invalid') +}) + +test('renderer fails closed for unsupported process and mixed scene node shapes', () => { + const { validateWorkflowPreflight } = loadModule() + const scene = { id: 'scene', type: 'sceneNode', position: { x: 0, y: 0 }, data: { params: { manifestPath: 'Workflows/room/scene-manifest.json' } } } + const target = { id: 'target', type: 'extensionNode', position: { x: 0, y: 0 }, data: { extensionId: 'pack/process-node' } } + for (const extension of [ + ext({ input: 'scene', output: 'mesh', type: 'process' }), + ext({ input: 'scene', inputs: ['scene', 'text'], output: 'mesh', type: 'model' }), + ]) { + const issues = validateWorkflowPreflight(wf([scene, target], [{ id: 'e', source: 'scene', target: 'target' }]), [extension]) + assert.ok(issues.some((issue) => issue.key === 'target:unsupported-scene-shape')) + } +}) diff --git a/src/areas/workflows/preflight.ts b/src/areas/workflows/preflight.ts index 3985855c..df7d35ca 100644 --- a/src/areas/workflows/preflight.ts +++ b/src/areas/workflows/preflight.ts @@ -2,7 +2,7 @@ import type { Workflow, WFNode } from '@shared/types/electron.d' import { getWorkflowExtension, type WorkflowExtension } from './mockExtensions' import { isPassthrough, isBranchConsumer, resolveDataSource, nearestUpstreamWaits } from './nodeBehaviors' -type DataType = 'image' | 'text' | 'mesh' | 'audio' +type DataType = 'image' | 'text' | 'mesh' | 'audio' | 'scene' export interface WorkflowPreflightIssue { key: string @@ -14,6 +14,7 @@ function nodeLabel(node: WFNode, allExtensions: WorkflowExtension[]): string { if (node.type === 'imageNode') return 'Image' if (node.type === 'textNode') return 'Text' if (node.type === 'meshNode') return 'Load 3D Mesh' + if (node.type === 'sceneNode') return 'Load Scene' if (node.type === 'outputNode') return 'Add to Scene' if (node.type === 'previewNode') return 'Preview Views' if (node.type === 'imagePreviewNode') return 'Preview Image' @@ -28,6 +29,7 @@ function nodeLabel(node: WFNode, allExtensions: WorkflowExtension[]): string { } function formatType(type: DataType): string { + if (type === 'scene') return 'scene' if (type === 'mesh') return 'mesh' if (type === 'image') return 'image' if (type === 'audio') return 'audio' @@ -44,6 +46,7 @@ function getNodeOutputType(node: WFNode, allExtensions: WorkflowExtension[]): Da if (node.type === 'imageNode') return 'image' if (node.type === 'textNode') return 'text' if (node.type === 'meshNode' || node.type === 'outputNode') return 'mesh' + if (node.type === 'sceneNode') return 'scene' if (node.type === 'previewNode') return 'image' if (node.type === 'imagePreviewNode') return 'image' if (node.type === 'forEachNode') { @@ -96,6 +99,13 @@ export function validateWorkflowPreflight( }) } + if (node.type === 'sceneNode' && !((node.data.params?.manifestPath as string | undefined)?.trim())) { + pushIssue(issues, { + key: `${node.id}:scene-invalid`, nodeId: node.id, + message: 'Load Scene needs a validated scene directory.', + }) + } + // A node fed by two different Wait branches can't be scheduled into a single // branch — it would run before either branch produces its mesh. if ( @@ -121,6 +131,20 @@ export function validateWorkflowPreflight( continue } + const usesSceneInput = ext.input === 'scene' || ext.inputs?.includes('scene') === true + const unsupportedSceneShape = + (ext.type === 'process' && (usesSceneInput || ext.output === 'scene')) + || (ext.type === 'model' && usesSceneInput + && (ext.inputs !== undefined || ext.input !== 'scene')) + if (unsupportedSceneShape) { + pushIssue(issues, { + key: `${node.id}:unsupported-scene-shape`, + nodeId: node.id, + message: `${ext.name} uses an unsupported scene input or output declaration.`, + }) + continue + } + const incomingEdges = workflow.edges.filter((edge) => edge.target === node.id) const requiredTypes = [...new Set((ext.inputs ?? [ext.input]) as DataType[])] diff --git a/src/areas/workflows/workflowRunStore.ts b/src/areas/workflows/workflowRunStore.ts index 15cd9ace..82166635 100644 --- a/src/areas/workflows/workflowRunStore.ts +++ b/src/areas/workflows/workflowRunStore.ts @@ -306,6 +306,14 @@ async function executeExtensionNode( selectedImagePath, selectedImageData } = ctx const ext = getWorkflowExtension(node.data.extensionId ?? '', allExtensions) + if (ext) { + const usesSceneInput = ext.input === 'scene' || ext.inputs?.includes('scene') === true + if ((ext.type === 'process' && (usesSceneInput || ext.output === 'scene')) + || (ext.type === 'model' && usesSceneInput + && (ext.inputs !== undefined || ext.input !== 'scene'))) { + throw new Error(`${ext.name} uses an unsupported scene input or output declaration`) + } + } // Freshest params at the moment the node starts (so loop iterations / Retry pick // up edits made while paused, not the values captured at run start). const liveParams = _liveParams.current.get(node.id) ?? node.data.params ?? {} @@ -318,6 +326,7 @@ async function executeExtensionNode( let nodeInputPath: string | undefined let nodeInputText: string | undefined let nodeInputMeshPath: string | undefined + let nodeInputScenePath: string | undefined // Per-slot texts for multi-text-input nodes (e.g. positive/negative prompts). // Indexed by target handle: input-0 → texts[0], input-1 → texts[1]. const nodeInputTexts: (string | undefined)[] = [] @@ -354,21 +363,24 @@ async function executeExtensionNode( const src = resolveSource(edge.source) if (src?.filePath !== undefined) nodeInputPath = src.filePath if (src?.text !== undefined && src.text.trim().length > 0) nodeInputText = src.text + if (src?.outputType === 'scene') nodeInputScenePath = src.filePath } } const isModelNode = ext?.type === 'model' if (isModelNode) { + const isSceneInput = ext?.inputs ? ext.inputs.includes('scene') : ext?.input === 'scene' const isTextInput = ext?.inputs ? ext.inputs.every((i) => i === 'text') : ext?.input === 'text' - const activeImagePath = isTextInput ? undefined : (nodeInputPath ?? selectedImagePath) - if (!isTextInput && !selectedImageData && (!activeImagePath || activeImagePath.trim().length === 0)) { + if (isSceneInput && !nodeInputScenePath) throw new Error(`${ext?.name ?? 'Model'} needs an incoming scene connection`) + const activeImagePath = (isTextInput || isSceneInput) ? undefined : (nodeInputPath ?? selectedImagePath) + if (!isTextInput && !isSceneInput && !selectedImageData && (!activeImagePath || activeImagePath.trim().length === 0)) { throw new Error('No input image selected for model node') } let blob: Blob let fname: string - if (isTextInput || (selectedImageData && nodeInputPath === undefined)) { + if (isTextInput || isSceneInput || (selectedImageData && nodeInputPath === undefined)) { const base64 = selectedImageData && nodeInputPath === undefined ? selectedImageData : 'iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==' // 1x1 transparent PNG @@ -401,21 +413,30 @@ async function executeExtensionNode( ) const effectiveParams = { ...schemaDefaults, ...liveParams } - const fd = new FormData() - fd.append('image', blob, fname) - fd.append('model_id', node.data.extensionId ?? '') - fd.append('collection', 'Workflows') - fd.append('remesh', 'none') - fd.append('enable_texture', 'false') - fd.append('texture_resolution', '1024') - fd.append('params', JSON.stringify({ ...effectiveParams, ...extraParams })) - setRunState((s) => ({ ...s, blockProgress: 5, blockStep: 'Submitting to model…' })) - - const { data } = await client.post<{ job_id: string }>( - '/generate/from-image', fd, - { headers: { 'Content-Type': 'multipart/form-data' } }, - ) + let submission: { data: { job_id: string } } + if (isSceneInput) { + const normalized = nodeInputScenePath!.replace(/\\/g, '/') + const inputPath = normalized.startsWith(`${workspaceDir}/`) + ? normalized.slice(workspaceDir.length + 1) + : normalized.replace(/^\/workspace\//, '') + submission = await client.post('/generate/from-artifact', { + input_kind: 'scene', input_path: inputPath, + model_id: node.data.extensionId ?? '', collection: 'Workflows', + params: { ...effectiveParams, ...extraParams }, + }) + } else { + const fd = new FormData() + fd.append('image', blob, fname) + fd.append('model_id', node.data.extensionId ?? '') + fd.append('collection', 'Workflows') + fd.append('remesh', 'none') + fd.append('enable_texture', 'false') + fd.append('texture_resolution', '1024') + fd.append('params', JSON.stringify({ ...effectiveParams, ...extraParams })) + submission = await client.post('/generate/from-image', fd, { headers: { 'Content-Type': 'multipart/form-data' } }) + } + const { data } = submission _activeJobId.current = data.job_id while (true) { @@ -594,7 +615,7 @@ export const useWorkflowRunStore = create((set, get) => { if (!outputUrl) { for (const [, o] of ctx.nodeOutputs) { if (o.filePath) { - if (o.outputType === 'audio') { + if (o.outputType === 'audio' || o.outputType === 'scene') { outputPath = o.filePath continue } @@ -802,6 +823,13 @@ export const useWorkflowRunStore = create((set, get) => { if (fp) nodeOutputs.set(node.id, { filePath: fp, outputType: 'mesh' }) } } + if (node.type === 'sceneNode') { + const manifestPath = node.data.params?.manifestPath as string | undefined + if (manifestPath) nodeOutputs.set(node.id, { + filePath: `${workspaceDir}/${manifestPath.replace(/^\/+/, '')}`, + outputType: 'scene', + }) + } } const ctx: RunContext = { diff --git a/src/areas/workflows/workflowSceneRun.test.mjs b/src/areas/workflows/workflowSceneRun.test.mjs new file mode 100644 index 00000000..a2c2b2d0 --- /dev/null +++ b/src/areas/workflows/workflowSceneRun.test.mjs @@ -0,0 +1,60 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { build } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +const dir = mkdtempSync(join(tmpdir(), 'modly-scene-run-')) +const stub = (name, source) => { const path = join(dir, name); writeFileSync(path, source); return path } +const appStoreStub = stub('app.ts', ` +export const appState: any = { apiUrl: 'http://modly.test', currentJob: null, + setCurrentJob(value: any) { this.currentJob = value }, + updateCurrentJob(value: any) { this.currentJob = { ...(this.currentJob ?? {}), ...value } } } +export const useAppStore: any = (selector: any) => selector(appState) +useAppStore.getState = () => appState +`) +const axiosStub = stub('axios.ts', `const axios: any = { create: () => (globalThis as any).__sceneClient }; export default axios; export type AxiosInstance = any`) +const extStub = stub('ext.ts', `export const getWorkflowExtension = (id: string, all: any[]) => all.find((value) => value.id === id); export type WorkflowExtension = any`) +const notifyStub = stub('notify.ts', `export const showCompletionNotification = async () => {}`) +const aliases = new Map([ + ['axios', axiosStub], ['@shared/stores/appStore', appStoreStub], + ['./mockExtensions', extStub], ['@shared/utils/notification', notifyStub], +]) +const outfile = join(dir, 'store.cjs') +writeFileSync(outfile, (await build({ + entryPoints: [resolve('src/areas/workflows/workflowRunStore.ts')], bundle: true, + platform: 'node', format: 'cjs', write: false, + plugins: [{ name: 'aliases', setup(build) { build.onResolve({ filter: /.*/ }, (args) => aliases.has(args.path) ? { path: aliases.get(args.path) } : null) } }], +})).outputFiles[0].text) +const { useWorkflowRunStore } = createRequire(import.meta.url)(outfile) + +test('scene model uses typed artifact route and preserves scene output', async () => { + const posts = [] + globalThis.window = { electron: { + settings: { get: async () => ({ workspaceDir: '/workspace' }) }, + fs: { deleteDirectory: async () => ({ success: true }), listFiles: async () => [], readFileBase64: async () => { throw new Error('scene must not be read as image bytes') } }, + } } + globalThis.__sceneClient = { + post: async (url, body) => { posts.push({ url, body }); return { data: { job_id: 'scene-job' } } }, + get: async () => ({ data: { status: 'done', progress: 100, output_url: '/workspace/Workflows/result/scene-manifest.json' } }), + } + const workflow = { + id: 'wf', name: 'Scene', description: '', createdAt: '', updatedAt: '', + nodes: [ + { id: 'source', type: 'sceneNode', position: { x: 0, y: 0 }, data: { enabled: true, params: { manifestPath: 'Workflows/input/scene-manifest.json' } } }, + { id: 'model', type: 'extensionNode', position: { x: 1, y: 0 }, data: { enabled: true, extensionId: 'pixal/world', params: {} } }, + ], + edges: [{ id: 'e', source: 'source', target: 'model' }], + } + const extension = { id: 'pixal/world', name: 'World', type: 'model', input: 'scene', output: 'scene', params: [] } + await useWorkflowRunStore.getState().run(workflow, [extension]) + assert.equal(posts[0].url, '/generate/from-artifact') + assert.deepEqual(posts[0].body, { + input_kind: 'scene', input_path: 'Workflows/input/scene-manifest.json', + model_id: 'pixal/world', collection: 'Workflows', params: {}, + }) + assert.equal(useWorkflowRunStore.getState().runState.outputPath, '/workspace/Workflows/result/scene-manifest.json') + assert.equal(useWorkflowRunStore.getState().runState.outputUrl, undefined) +}) diff --git a/src/areas/workflows/workflowSceneSource.test.mjs b/src/areas/workflows/workflowSceneSource.test.mjs new file mode 100644 index 00000000..ad2a0280 --- /dev/null +++ b/src/areas/workflows/workflowSceneSource.test.mjs @@ -0,0 +1,29 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +const outfile = join(mkdtempSync(join(tmpdir(), 'modly-scene-source-')), 'scene.cjs') +writeFileSync(outfile, buildSync({ entryPoints: [resolve('src/areas/workflows/workflowSceneSource.ts')], bundle: true, platform: 'node', format: 'cjs', write: false }).outputFiles[0].text) +const { resolveSceneSourceManifest } = createRequire(import.meta.url)(outfile) +const encoded = Buffer.from(JSON.stringify({ schema: 'modly.scene-manifest.v1', sceneRoot: '.', assets: [] })).toString('base64') + +test('Load Scene resolves directory and manifest without image bytes', async () => { + for (const scenePath of ['Workflows/room', 'Workflows/room/scene-manifest.json']) { + const result = await resolveSceneSourceManifest({ scenePath, workspaceDir: '/workspace', readFileBase64: async () => encoded }) + assert.equal(result.ok, true) + assert.equal(result.manifestWorkspacePath, 'Workflows/room/scene-manifest.json') + } +}) + +test('Load Scene refuses unsafe paths before reading', async () => { + let reads = 0 + for (const scenePath of ['../outside', '/etc/passwd', 'C:/outside', 'Workflows/%2e%2e', 'Workflows/room/']) { + const result = await resolveSceneSourceManifest({ scenePath, workspaceDir: '/workspace', readFileBase64: async () => { reads++; return encoded } }) + assert.equal(result.ok, false, scenePath) + } + assert.equal(reads, 0) +}) diff --git a/src/areas/workflows/workflowSceneSource.ts b/src/areas/workflows/workflowSceneSource.ts new file mode 100644 index 00000000..5873b692 --- /dev/null +++ b/src/areas/workflows/workflowSceneSource.ts @@ -0,0 +1,205 @@ +import type { SceneArtifactManifestInitialView, SceneArtifactManifestPreview, SceneArtifactManifestV1 } from '../../shared/types/artifacts' + +export const SCENE_MANIFEST_FILE_NAME = 'scene-manifest.json' + +export type SceneSourceKind = 'manifest' | 'directory' + +export type ResolveSceneSourceSuccess = { + ok: true + sourceKind: SceneSourceKind + inputWorkspacePath: string + manifestWorkspacePath: string + manifestAbsolutePath: string + sceneRoot: string + manifest: SceneArtifactManifestV1 +} + +export type ResolveSceneSourceFailure = { + ok: false + error: string +} + +export type ResolveSceneSourceResult = ResolveSceneSourceSuccess | ResolveSceneSourceFailure + +type ResolveSceneSourceArgs = { + scenePath: string + workspaceDir: string + readFileBase64: (filePath: string) => Promise +} + +function isAbsolutePath(value: string): boolean { + return value.startsWith('/') || /^[A-Za-z]:\//.test(value) +} + +function trimTrailingSlashes(value: string): string { + return value.replace(/\/+$/, '') +} + +function isSafeRelativePath(value: unknown, allowDot = false): value is string { + if (typeof value !== 'string' || !value || value !== value.trim() || value.includes('\u0000')) return false + const normalized = value.replace(/\\/g, '/') + if (allowDot && normalized === '.') return true + if (isAbsolutePath(normalized) || /^[A-Za-z][A-Za-z0-9+.-]*:/.test(normalized) + || /%(?:25|2e|2f|5c|00)/i.test(normalized) || /%(?![0-9a-f]{2})/i.test(normalized)) return false + return normalized.split('/').every((segment) => segment.length > 0 && segment !== '.' && segment !== '..') +} + +function normalizeWorkspaceRelativePath(value: string | undefined, workspaceDir: string): string | undefined { + const normalizedValue = value?.replace(/\\/g, '/') + if (!normalizedValue) return undefined + + const normalizedWorkspace = trimTrailingSlashes(workspaceDir.replace(/\\/g, '/')) + let relativePath: string | undefined + + if (normalizedValue.startsWith('/workspace/')) { + relativePath = normalizedValue.slice('/workspace/'.length) + } else if (normalizedValue === normalizedWorkspace) { + return undefined + } else if (normalizedValue.startsWith(`${normalizedWorkspace}/`)) { + relativePath = normalizedValue.slice(normalizedWorkspace.length + 1) + } else if (!isAbsolutePath(normalizedValue)) { + relativePath = normalizedValue + } + + if (!relativePath) return undefined + return isSafeRelativePath(relativePath) ? relativePath : undefined +} + +function resolveSceneSourceKind(inputWorkspacePath: string): SceneSourceKind | undefined { + if (inputWorkspacePath === SCENE_MANIFEST_FILE_NAME || inputWorkspacePath.endsWith(`/${SCENE_MANIFEST_FILE_NAME}`)) { + return 'manifest' + } + if (inputWorkspacePath.toLowerCase().endsWith('.json')) return undefined + return 'directory' +} + +function decodeBase64Utf8(base64: string): string { + const bytes = Uint8Array.from(atob(base64), (char) => char.charCodeAt(0)) + return new TextDecoder().decode(bytes) +} + +function isPlainObject(value: unknown): value is Record { + return typeof value === 'object' && value !== null && !Array.isArray(value) +} + +function isSafePreviewReference(value: unknown): value is string { + return isSafeRelativePath(value) +} + +function isSceneManifestPreview(value: unknown): value is SceneArtifactManifestPreview { + return isPlainObject(value) + && (!('image' in value) || isSafePreviewReference(value.image)) + && (!('video' in value) || isSafePreviewReference(value.video)) +} + +function isFiniteTriple(value: unknown): value is [number, number, number] { + return Array.isArray(value) && value.length === 3 + && value.every((component) => typeof component === 'number' && Number.isFinite(component)) +} + +function isSceneManifestInitialView(value: unknown): value is SceneArtifactManifestInitialView { + if (!isPlainObject(value)) return false + const { position, target, up } = value + if (!isFiniteTriple(position) || !isFiniteTriple(target)) return false + if (position.every((component, index) => component === target[index])) return false + return up === undefined || (isFiniteTriple(up) && up.some((component) => component !== 0)) +} + +function normalizeSceneRoot(sceneRoot: unknown): string | undefined { + return isSafeRelativePath(sceneRoot, true) ? sceneRoot.replace(/\\/g, '/') : undefined +} + +function validateSceneManifest(manifest: unknown): { ok: true; manifest: SceneArtifactManifestV1; sceneRoot: string } | { ok: false; error: string } { + if (!isPlainObject(manifest)) { + return { ok: false, error: 'Scene manifest must be a JSON object.' } + } + if (manifest.schema !== 'modly.scene-manifest.v1') { + return { ok: false, error: 'Scene manifest schema must be modly.scene-manifest.v1.' } + } + + const rawSceneRoot = manifest.sceneRoot + const sceneRoot = normalizeSceneRoot(rawSceneRoot) + if (typeof rawSceneRoot !== 'string' || !sceneRoot) { + return { ok: false, error: 'Scene manifest sceneRoot must be a safe relative path.' } + } + if (!Array.isArray(manifest.assets)) { + return { ok: false, error: 'Scene manifest assets must be an array.' } + } + if (manifest.assets.some((asset) => isPlainObject(asset) + && (('workspacePath' in asset && !isSafeRelativePath(asset.workspacePath)) + || ('path' in asset && !isSafeRelativePath(asset.path))))) { + return { ok: false, error: 'Scene manifest asset paths must be safe relative file references.' } + } + + const { preview, initialView, ...metadata } = manifest + if (preview !== undefined && !isPlainObject(preview)) { + return { ok: false, error: 'Scene manifest preview must be a JSON object.' } + } + if (preview !== undefined && !isSceneManifestPreview(preview)) { + return { ok: false, error: 'Scene manifest preview image/video must be safe relative file references.' } + } + if (initialView !== undefined && !isPlainObject(initialView)) { + return { ok: false, error: 'Scene manifest initialView must be a JSON object.' } + } + if (initialView !== undefined && !isSceneManifestInitialView(initialView)) { + return { ok: false, error: 'Scene manifest initialView requires distinct finite numeric position/target triples and optional non-zero finite up.' } + } + + return { + ok: true, + sceneRoot, + manifest: { + ...metadata, + schema: 'modly.scene-manifest.v1', + sceneRoot: rawSceneRoot, + assets: manifest.assets, + ...(preview !== undefined ? { preview } : {}), + ...(initialView !== undefined ? { initialView } : {}), + }, + } +} + +export async function resolveSceneSourceManifest(args: ResolveSceneSourceArgs): Promise { + const inputWorkspacePath = normalizeWorkspaceRelativePath(args.scenePath, args.workspaceDir) + if (!inputWorkspacePath) { + return { ok: false, error: 'Load Scene requires a safe workspace-relative scene path.' } + } + + const sourceKind = resolveSceneSourceKind(inputWorkspacePath) + if (!sourceKind) { + return { ok: false, error: `Load Scene accepts ${SCENE_MANIFEST_FILE_NAME} or a scene directory.` } + } + + const manifestWorkspacePath = sourceKind === 'manifest' + ? inputWorkspacePath + : `${inputWorkspacePath}/${SCENE_MANIFEST_FILE_NAME}` + const normalizedWorkspace = trimTrailingSlashes(args.workspaceDir.replace(/\\/g, '/')) + const manifestAbsolutePath = `${normalizedWorkspace}/${manifestWorkspacePath}` + + let manifestRaw: string + try { + manifestRaw = decodeBase64Utf8(await args.readFileBase64(manifestAbsolutePath)) + } catch (error) { + return { ok: false, error: `Unable to read scene manifest: ${String(error)}` } + } + + let parsed: unknown + try { + parsed = JSON.parse(manifestRaw) + } catch (error) { + return { ok: false, error: `Scene manifest is not valid JSON: ${String(error)}` } + } + + const validation = validateSceneManifest(parsed) + if (!validation.ok) return validation + + return { + ok: true, + sourceKind, + inputWorkspacePath, + manifestWorkspacePath, + manifestAbsolutePath, + sceneRoot: validation.sceneRoot, + manifest: validation.manifest, + } +} diff --git a/src/shared/stores/workflowsStore.ts b/src/shared/stores/workflowsStore.ts index b8753375..f6fb829c 100644 --- a/src/shared/stores/workflowsStore.ts +++ b/src/shared/stores/workflowsStore.ts @@ -98,7 +98,7 @@ interface LegacyWorkflow { // Source-only nodes have no target handle; sink-only nodes have no source handle. // An edge into/out of the wrong side can't resolve a handle and makes React Flow // warn ("Couldn't create edge for target handle id: null") on every render. -export const NODE_TYPES_WITHOUT_TARGET = new Set(['imageNode', 'textNode', 'meshNode', 'inputNode', 'forEachNode']) +export const NODE_TYPES_WITHOUT_TARGET = new Set(['imageNode', 'textNode', 'meshNode', 'sceneNode', 'inputNode', 'forEachNode']) export const NODE_TYPES_WITHOUT_SOURCE = new Set(['outputNode', 'previewNode']) function sanitizeEdges(nodes: WFNode[], edges: WFEdge[]): WFEdge[] { diff --git a/src/shared/types/artifacts.ts b/src/shared/types/artifacts.ts index a8dfc6eb..57c6b8d4 100644 --- a/src/shared/types/artifacts.ts +++ b/src/shared/types/artifacts.ts @@ -4,3 +4,17 @@ export interface ArtifactProvenance { source?: string [key: string]: unknown } + +export interface SceneArtifactManifestPreview { image?: string; video?: string } +export interface SceneArtifactManifestInitialView { + position: [number, number, number] + target: [number, number, number] + up?: [number, number, number] +} +export interface SceneArtifactManifestV1 { + schema: 'modly.scene-manifest.v1' + sceneRoot: string + assets: unknown[] + preview?: SceneArtifactManifestPreview + initialView?: SceneArtifactManifestInitialView +} diff --git a/src/shared/types/electron.d.ts b/src/shared/types/electron.d.ts index ffe434ed..c6cca213 100644 --- a/src/shared/types/electron.d.ts +++ b/src/shared/types/electron.d.ts @@ -14,10 +14,10 @@ import type { export interface ExtensionNode { id: string name: string - input: 'image' | 'text' | 'mesh' | 'audio' - inputs?: ('image' | 'text' | 'mesh' | 'audio')[] // multi-input nodes; overrides input when set + input: 'image' | 'text' | 'mesh' | 'audio' | 'scene' + inputs?: ('image' | 'text' | 'mesh' | 'audio' | 'scene')[] // multi-input nodes; overrides input when set inputLabels?: string[] // display labels per input slot (e.g. positive/negative) - output: 'image' | 'text' | 'mesh' | 'audio' + output: 'image' | 'text' | 'mesh' | 'audio' | 'scene' paramsSchema: ParamSchema[] paramDefaults?: Record hfRepo?: string From 10cdc15ce1a6c717f1e6e99d5427c44a7c30dcd6 Mon Sep 17 00:00:00 2001 From: DrHepa Date: Thu, 1 Oct 2026 22:52:43 +0200 Subject: [PATCH 37/57] fix(workflows): address scene artifact review --- api/routers/generation.py | 104 ++++++-- api/routers/model.py | 8 +- api/routers/workflow_runs.py | 38 +-- api/services/generator_registry.py | 148 ++++++------ api/services/generators/base.py | 10 +- api/tests/test_base_generator.py | 8 +- api/tests/test_generator_registry.py | 37 ++- api/tests/test_model_router.py | 45 ++++ api/tests/test_scene_generation.py | 224 +++++++++++++++++- api/tests/test_scene_input.py | 8 +- api/tests/test_workflow_runs_lifecycle.py | 107 +++++++++ api/tests/test_workflow_runs_router.py | 5 +- electron/main/bounded-file-reader.test.mjs | 48 ++++ electron/main/bounded-file-reader.ts | 40 ++++ .../main/extension-install-utils.test.mjs | 6 +- electron/main/extension-install-utils.ts | 11 +- electron/main/ipc-handlers.ts | 22 +- src/areas/workflows/WorkflowsPage.tsx | 8 +- src/areas/workflows/nodes/ExtensionNode.tsx | 4 +- src/areas/workflows/nodes/LoadSceneNode.tsx | 67 +++--- src/areas/workflows/workflowRunStore.ts | 6 +- .../workflows/workflowSceneSource.test.mjs | 43 +++- src/areas/workflows/workflowSceneSource.ts | 41 ++++ 23 files changed, 817 insertions(+), 221 deletions(-) create mode 100644 electron/main/bounded-file-reader.test.mjs create mode 100644 electron/main/bounded-file-reader.ts diff --git a/api/routers/generation.py b/api/routers/generation.py index 7c355014..02879218 100644 --- a/api/routers/generation.py +++ b/api/routers/generation.py @@ -4,6 +4,7 @@ import time import traceback import uuid +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from typing import Dict, Optional, Union from fastapi import APIRouter, File, Form, UploadFile, HTTPException, BackgroundTasks @@ -28,6 +29,14 @@ _cancelled: set = set() _cancel_events: Dict[str, threading.Event] = {} _completed_at: Dict[str, float] = {} +_job_generators: Dict[str, object] = {} +# A pinned generation owns the complete switch/load/generate lifecycle. Keeping +# that lifecycle on one dedicated worker provides process-wide serialization +# without parking default-executor workers on a blocking lock acquisition. +_pinned_generation_executor = ThreadPoolExecutor( + max_workers=1, + thread_name_prefix="modly-pinned-generation", +) _JOB_TTL = 1800 # purge terminal jobs after 30 minutes @@ -39,6 +48,7 @@ def _purge_old_jobs() -> None: _jobs.pop(jid, None) _cancelled.discard(jid) _cancel_events.pop(jid, None) + _job_generators.pop(jid, None) _completed_at.pop(jid, None) @@ -111,8 +121,6 @@ async def generate_from_image( except ValueError as e: raise HTTPException(400, str(e)) - generator_registry.switch_model(model_id) - # Parse model-specific params from JSON and merge with common fields try: model_params = json.loads(params) @@ -197,11 +205,10 @@ async def cancel_job(job_id: str): if job.status in ("pending", "running"): job.status = "cancelled" _completed_at[job_id] = time.monotonic() - # Kill the active generator subprocess immediately so inference stops now. - # _run_generation will catch the resulting exception, see job_id in _cancelled, - # and return cleanly without setting an error status. + # Kill only the subprocess bound to this job. A queued cancellation must not + # terminate whichever earlier job currently owns the active generator. try: - gen = generator_registry._generators.get(generator_registry._active_id) + gen = _job_generators.get(job_id) if gen is not None and hasattr(gen, "_proc") and gen._proc and gen._proc.poll() is None: gen._proc.kill() gen._loaded = False @@ -219,6 +226,50 @@ async def _run_generation( output_kind: str = "mesh", model_id: Optional[str] = None, ) -> None: + # Pinned jobs share one model lifecycle. Switching is deliberately deferred + # until this job runs on the dedicated worker: request-time switches can + # otherwise unload a running job or leave A loading beside B. Crucially, + # queued jobs are executor work items rather than default-executor threads + # blocked on a lock, so cancelling a waiter cannot orphan queue ownership or + # starve the worker that performs generation. + loop = asyncio.get_running_loop() + executor = _pinned_generation_executor if model_id is not None else None + future = loop.run_in_executor( + executor, + _run_generation_impl, + job_id, + model_input, + params, + collection, + output_kind, + model_id, + ) + try: + await future + except asyncio.CancelledError: + # asyncio cancellation attempts to cancel a queued concurrent future. + # If it has already begun, the event lets the generator stop safely. + _cancelled.add(job_id) + cancel_event = _cancel_events.get(job_id) + if cancel_event is not None: + cancel_event.set() + job = _jobs.get(job_id) + if job is not None and job.status in ("pending", "running"): + job.status = "cancelled" + _completed_at[job_id] = time.monotonic() + raise + + +def _run_generation_impl( + job_id: str, + model_input: Union[bytes, TypedArtifactInput], + params: dict, + collection: str, + output_kind: str, + model_id: Optional[str], +) -> None: + if job_id in _cancelled: + return job = _jobs[job_id] job.status = "running" @@ -232,17 +283,15 @@ def progress_cb(pct: int, step: str = "") -> None: job.step = step try: - loop = asyncio.get_running_loop() - - # Check if the model needs to be loaded BEFORE calling get_active(), - # because get_active() loads the model in a blocking manner. - # active_status() is an instantaneous operation (simple dict lookup). - get_generator = (lambda: generator_registry.get_ready_generator(model_id)) \ + # Check if the model needs to be loaded BEFORE calling the generator + # getter, because that call can load the model in a blocking manner. + get_generator = (lambda: generator_registry.activate_ready_generator(model_id)) \ if model_id is not None else generator_registry.get_active - status_reader = getattr(generator_registry, "model_status", None) - status = (status_reader(model_id) if model_id is not None and status_reader - else generator_registry.active_status() if model_id is None - else {"name": model_id, "downloaded": True, "loaded": False}) + if model_id is not None: + _job_generators[job_id] = generator_registry.get_generator(model_id) + status_reader = (lambda: generator_registry.model_status(model_id)) \ + if model_id is not None else generator_registry.active_status + status = status_reader() if not status["loaded"]: active = status model_name = active['name'] @@ -256,12 +305,13 @@ def progress_cb(pct: int, step: str = "") -> None: ) load_thread.start() try: - gen = await loop.run_in_executor(None, get_generator) + gen = get_generator() finally: stop_load_evt.set() else: - gen = await loop.run_in_executor(None, get_generator) + gen = get_generator() + _job_generators[job_id] = gen if job_id in _cancelled: return @@ -278,18 +328,18 @@ def progress_cb(pct: int, step: str = "") -> None: model_input = revalidate_artifact_input(registry.WORKSPACE_DIR, model_input) import inspect supports_cancel = "cancel_event" in inspect.signature(gen.generate_artifact).parameters - output_path = await loop.run_in_executor( - None, - lambda: gen.generate_artifact(model_input.kind, model_input.path, params, progress_cb, cancel_event) - if supports_cancel else gen.generate_artifact(model_input.kind, model_input.path, params, progress_cb), + output_path = ( + gen.generate_artifact(model_input.kind, model_input.path, params, progress_cb, cancel_event) + if supports_cancel + else gen.generate_artifact(model_input.kind, model_input.path, params, progress_cb) ) else: import inspect supports_cancel = "cancel_event" in inspect.signature(gen.generate).parameters - output_path = await loop.run_in_executor( - None, - lambda: gen.generate(model_input, params, progress_cb, cancel_event) - if supports_cancel else gen.generate(model_input, params, progress_cb), + output_path = ( + gen.generate(model_input, params, progress_cb, cancel_event) + if supports_cancel + else gen.generate(model_input, params, progress_cb) ) if job_id in _cancelled: @@ -328,3 +378,5 @@ def progress_cb(pct: int, step: str = "") -> None: job.status = "error" job.error = tb.strip() _completed_at[job_id] = time.monotonic() + finally: + _job_generators.pop(job_id, None) diff --git a/api/routers/model.py b/api/routers/model.py index 0b40d155..7d34f6cd 100644 --- a/api/routers/model.py +++ b/api/routers/model.py @@ -59,20 +59,20 @@ def _check_download_control(control: dict[str, threading.Event]) -> None: @router.get("/status") async def model_status(): """Status of the active model.""" - return generator_registry.active_status() + return await asyncio.to_thread(generator_registry.active_status) @router.get("/all") async def all_models_status(): """Status of all known models (downloaded, loaded, required VRAM).""" - return generator_registry.all_status() + return await asyncio.to_thread(generator_registry.all_status) @router.get("/params") async def model_params(model_id: Optional[str] = None): """Parameter schema of the active model (or a specified model).""" try: - return generator_registry.params_schema(model_id) + return await asyncio.to_thread(generator_registry.params_schema, model_id) except KeyError: raise HTTPException(404, f"Unknown model ID: {model_id}") @@ -81,7 +81,7 @@ async def model_params(model_id: Optional[str] = None): async def switch_model(model_id: str): """Switch the active model.""" try: - generator_registry.switch_model(model_id) + await asyncio.to_thread(generator_registry.switch_model, model_id) return {"active": model_id} except ValueError as e: raise HTTPException(400, str(e)) diff --git a/api/routers/workflow_runs.py b/api/routers/workflow_runs.py index 9b78a205..ba654a53 100644 --- a/api/routers/workflow_runs.py +++ b/api/routers/workflow_runs.py @@ -1,6 +1,5 @@ import json import threading -import time import uuid from typing import Optional from fastapi import APIRouter, BackgroundTasks, File, Form, HTTPException, UploadFile @@ -9,11 +8,10 @@ from routers.generation import ( VALID_REMESH_MODES, _cancel_events, - _cancelled, - _completed_at, _jobs, _purge_old_jobs, _run_generation, + cancel_job, sanitize_collection, ) from schemas.generation import JobStatus @@ -59,10 +57,7 @@ async def create_run_from_image( **model_params, } - # Same constraint /generate/from-image enforces on this field, checked before touching - # the registry below for the same reason that endpoint checks it first: switch_model() - # unloads whatever generator is currently active, and a request rejected for a bad - # remesh value should not pay for -- or force a reload after -- evicting it. + # Keep the same request validation as /generate/from-image before filing a job. if full_params["remesh"] not in VALID_REMESH_MODES: raise HTTPException(400, "remesh must be 'quad', 'triangle', or 'none'") @@ -70,11 +65,10 @@ async def create_run_from_image( try: generator_registry.get_generator(model_id) + output_kind = generator_registry.get_manifest(model_id).get("output", "mesh") except ValueError as e: raise HTTPException(400, str(e)) - generator_registry.switch_model(model_id) - job_id = str(uuid.uuid4()) image_bytes = await image.read() @@ -83,7 +77,9 @@ async def create_run_from_image( _jobs[job_id] = JobStatus(job_id=job_id, status="pending", progress=0) _cancel_events[job_id] = threading.Event() - background_tasks.add_task(_run_generation, job_id, image_bytes, full_params, collection) + background_tasks.add_task( + _run_generation, job_id, image_bytes, full_params, collection, output_kind, model_id + ) return {"run_id": job_id, "status": "pending"} @@ -111,24 +107,4 @@ async def get_run(run_id: str): @router.post("/{run_id}/cancel") async def cancel_run(run_id: str): - job = _jobs.get(run_id) - if not job: - raise HTTPException(404, f"Run {run_id} not found") - - _cancelled.add(run_id) - if run_id in _cancel_events: - _cancel_events[run_id].set() - if job.status in ("pending", "running"): - job.status = "cancelled" - _completed_at[run_id] = time.monotonic() - - try: - gen = generator_registry._generators.get(generator_registry._active_id) - if gen is not None and hasattr(gen, "_proc") and gen._proc and gen._proc.poll() is None: - gen._proc.kill() - gen._loaded = False - gen._proc = None - except Exception: - pass - - return {"cancelled": True} + return await cancel_job(run_id) diff --git a/api/services/generator_registry.py b/api/services/generator_registry.py index de1750ca..40c08db7 100644 --- a/api/services/generator_registry.py +++ b/api/services/generator_registry.py @@ -452,19 +452,13 @@ def _discover_extensions( node for node in raw_nodes if isinstance(node, dict) and node.get("id") ] - allowed_io = {"image", "text", "mesh", "audio", "scene"} for node in nodes: - declared_inputs = node.get("inputs") or [node.get("input", "image")] - if (not isinstance(declared_inputs, list) - or any(value not in allowed_io for value in declared_inputs)): - raise ValueError( - f'model node "{node.get("id", "unknown")}" has an unsupported input type' - ) - if node.get("output", "mesh") not in allowed_io: - raise ValueError( - f'model node "{node.get("id", "unknown")}" has an unsupported output type' - ) - if "scene" in declared_inputs and ( + declared_inputs = node.get("inputs") + uses_scene_input = ( + node.get("input", "image") == "scene" + or (isinstance(declared_inputs, list) and "scene" in declared_inputs) + ) + if uses_scene_input and ( "inputs" in node or node.get("input", "image") != "scene" ): raise ValueError( @@ -606,6 +600,7 @@ def __init__(self) -> None: self._generators: Dict[str, BaseGenerator] = {} self._manifests: Dict[str, dict] = {} self._errors: Dict[str, str] = {} + self._lifecycle_lock = threading.RLock() self._legacy_imports = _LegacyImportManager() self._active_id: str = os.environ.get("SELECTED_MODEL_ID", "sf3d") @@ -727,38 +722,47 @@ def _assert_not_quarantined(model_id: str) -> None: def get_active(self) -> BaseGenerator: """Returns the active generator. Downloads and loads if necessary.""" - return self.get_ready_generator(self._active_id) + with self._lifecycle_lock: + return self.get_ready_generator(self._active_id) def get_ready_generator(self, model_id: str) -> BaseGenerator: """Load and return exactly ``model_id`` without consulting active state.""" - gen = self.get_generator(model_id) - downloaded = self._is_downloaded(model_id, gen) - if "model_sources" in self._manifests[model_id] and not downloaded: - raise RuntimeError( - "Model sources are incomplete. Download this node's weights " - "from the Modly Models page before generation." - ) - if not gen.is_loaded(): - if not downloaded: - if isinstance(gen, ExtensionProcess): - # Let the subprocess handle its own download logic during - # load() — some extensions (e.g. mv-adapter) need custom - # multi-repo downloads that the standard HF endpoint can't do. - pass - else: - gen._auto_download() - gen.load() - return gen + with self._lifecycle_lock: + gen = self.get_generator(model_id) + downloaded = self._is_downloaded(model_id, gen) + if "model_sources" in self._manifests[model_id] and not downloaded: + raise RuntimeError( + "Model sources are incomplete. Download this node's weights " + "from the Modly Models page before generation." + ) + if not gen.is_loaded(): + if not downloaded: + if isinstance(gen, ExtensionProcess): + # Let the subprocess handle its own download logic during + # load() — some extensions (e.g. mv-adapter) need custom + # multi-repo downloads that the standard HF endpoint can't do. + pass + else: + gen._auto_download() + gen.load() + return gen + + def activate_ready_generator(self, model_id: str) -> BaseGenerator: + """Atomically make ``model_id`` active and return it ready for inference.""" + with self._lifecycle_lock: + self.switch_model(model_id) + return self.get_ready_generator(model_id) def model_status(self, model_id: str) -> dict: - gen = self.get_generator(model_id) - manifest = self._manifests[model_id] - return { - "id": model_id, - "name": manifest.get("name", gen.DISPLAY_NAME), - "downloaded": self._is_downloaded(model_id, gen), - "loaded": gen.is_loaded(), - } + with self._lifecycle_lock: + gen = self.get_generator(model_id) + manifest = self._manifests[model_id] + return { + "id": model_id, + "name": manifest.get("name", gen.DISPLAY_NAME), + "downloaded": self._is_downloaded(model_id, gen), + "loaded": gen.is_loaded(), + } def get_generator(self, model_id: str) -> BaseGenerator: self._assert_not_quarantined(model_id) @@ -777,16 +781,17 @@ def get_manifest(self, model_id: str) -> dict: def switch_model(self, model_id: str) -> None: """Switches the active model. Unloads the previous one if different.""" - self._assert_not_quarantined(model_id) - if model_id not in self._generators: - raise ValueError( - f"Unknown model ID: '{model_id}'. " - f"Available: {list(self._generators.keys())}" - ) - if model_id != self._active_id: - if self._active_id in self._generators: - self._generators[self._active_id].unload() - self._active_id = model_id + with self._lifecycle_lock: + self._assert_not_quarantined(model_id) + if model_id not in self._generators: + raise ValueError( + f"Unknown model ID: '{model_id}'. " + f"Available: {list(self._generators.keys())}" + ) + if model_id != self._active_id: + if self._active_id in self._generators: + self._generators[self._active_id].unload() + self._active_id = model_id # ------------------------------------------------------------------ # # Status @@ -801,31 +806,34 @@ def _is_downloaded(self, model_id: str, gen: BaseGenerator) -> bool: return gen.is_downloaded() def active_status(self) -> dict: - return self.model_status(self._active_id) + with self._lifecycle_lock: + return self.model_status(self._active_id) def all_status(self) -> list: - result = [] - for model_id, gen in self._generators.items(): - manifest = self._manifests[model_id] - result.append({ - "id": model_id, - "name": manifest.get("name", gen.DISPLAY_NAME), - "description": manifest.get("description", ""), - "version": manifest.get("version", ""), - "vram_gb": manifest.get("vram_gb", gen.VRAM_GB), - "hf_repo": manifest.get("hf_repo", ""), - "tags": manifest.get("tags", []), - "downloaded": self._is_downloaded(model_id, gen), - "loaded": gen.is_loaded(), - "active": model_id == self._active_id, - }) - return result + with self._lifecycle_lock: + result = [] + for model_id, gen in self._generators.items(): + manifest = self._manifests[model_id] + result.append({ + "id": model_id, + "name": manifest.get("name", gen.DISPLAY_NAME), + "description": manifest.get("description", ""), + "version": manifest.get("version", ""), + "vram_gb": manifest.get("vram_gb", gen.VRAM_GB), + "hf_repo": manifest.get("hf_repo", ""), + "tags": manifest.get("tags", []), + "downloaded": self._is_downloaded(model_id, gen), + "loaded": gen.is_loaded(), + "active": model_id == self._active_id, + }) + return result def params_schema(self, model_id: Optional[str] = None) -> list: - target_id = model_id or self._active_id - if target_id not in self._generators: - raise KeyError(target_id) - return self._generators[target_id].params_schema() + with self._lifecycle_lock: + target_id = model_id or self._active_id + if target_id not in self._generators: + raise KeyError(target_id) + return self._generators[target_id].params_schema() # ------------------------------------------------------------------ # # Paths update & shutdown diff --git a/api/services/generators/base.py b/api/services/generators/base.py index 0b5169a2..adeb9802 100644 --- a/api/services/generators/base.py +++ b/api/services/generators/base.py @@ -170,11 +170,13 @@ def generate_artifact( ) -> Path: """Generate from a validated typed artifact. - New extensions should override this method. The default delegates to - ``generate`` with the canonical path so scene-capable extensions built - against the pre-release contract remain compatible. + Typed-artifact extensions must implement this explicitly. Falling back + to ``generate`` would pass a filesystem ``Path`` to the legacy + image-bytes ABI and fail far from the actual contract violation. """ - return self.generate(artifact_path, params, progress_cb, cancel_event) # type: ignore[arg-type] + raise NotImplementedError( + f"{type(self).__name__} does not implement {input_kind} artifact generation" + ) def _check_cancelled(self, cancel_event: Optional[threading.Event]) -> None: """Raises GenerationCancelled if cancel_event is set.""" diff --git a/api/tests/test_base_generator.py b/api/tests/test_base_generator.py index 580b3d25..9adfe00c 100644 --- a/api/tests/test_base_generator.py +++ b/api/tests/test_base_generator.py @@ -43,6 +43,12 @@ def test_download_progress_print_is_ascii_safe(self) -> None: self.assertTrue(first_line.endswith("...")) self.assertTrue(all(ord(ch) < 128 for ch in printed)) + def test_typed_artifact_generation_requires_an_explicit_override(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + gen = _DownloadGen(Path(tmp) / "model", Path(tmp) / "out") + with self.assertRaisesRegex(NotImplementedError, "scene artifact generation"): + gen.generate_artifact("scene", Path(tmp) / "scene-manifest.json", {}) + if __name__ == "__main__": - unittest.main() \ No newline at end of file + unittest.main() diff --git a/api/tests/test_generator_registry.py b/api/tests/test_generator_registry.py index a4814b9e..76691e1b 100644 --- a/api/tests/test_generator_registry.py +++ b/api/tests/test_generator_registry.py @@ -155,7 +155,7 @@ def test_legacy_generator_supports_eager_and_lazy_sibling_imports(self) -> None: self.registry.reload() self.assertNotIn(str(extension.resolve()), sys.path) - def test_scene_io_is_registered_but_capture_and_video_are_rejected(self) -> None: + def test_scene_and_existing_custom_io_types_are_registered(self) -> None: for extension_id, input_kind in (("scene-io", "scene"), ("capture-io", "capture"), ("video-io", "video")): extension = self._make_extension(extension_id) manifest = { @@ -174,8 +174,8 @@ def test_scene_io_is_registered_but_capture_and_video_are_rejected(self) -> None self.registry.initialize() self.assertEqual(self.registry.get_manifest("scene-io/generate")["input"], "scene") - self.assertIn("capture-io/generate", self.registry.load_errors()) - self.assertIn("video-io/generate", self.registry.load_errors()) + self.assertEqual(self.registry.get_manifest("capture-io/generate")["input"], "capture") + self.assertEqual(self.registry.get_manifest("video-io/generate")["input"], "video") def test_scene_input_rejects_multi_input_shapes_but_image_multi_can_output_scene(self) -> None: cases = { @@ -258,6 +258,37 @@ def test_declared_sources_block_generation_even_when_generator_overrides_readine with self.assertRaisesRegex(RuntimeError, "Model sources are incomplete"): self.registry.get_active() + def test_activate_ready_generator_switches_before_loading_exact_model(self) -> None: + class Generator: + DISPLAY_NAME = "test" + def __init__(self): + self.loaded = False + self.unloads = 0 + def is_downloaded(self): return True + def is_loaded(self): return self.loaded + def load(self): self.loaded = True + def unload(self): + self.loaded = False + self.unloads += 1 + + first = Generator() + second = Generator() + first.loaded = True + self.registry._generators = {"demo/a": first, "demo/b": second} + self.registry._manifests = { + "demo/a": {"name": "A"}, + "demo/b": {"name": "B"}, + } + self.registry._active_id = "demo/a" + + selected = self.registry.activate_ready_generator("demo/b") + + self.assertIs(selected, second) + self.assertEqual(self.registry._active_id, "demo/b") + self.assertFalse(first.loaded) + self.assertEqual(first.unloads, 1) + self.assertTrue(second.loaded) + def test_reload_preserves_legacy_path_owned_by_the_host(self) -> None: extension = self._make_extension("host-owned-path") self._write_manifest(extension, extension_id="host-owned-path") diff --git a/api/tests/test_model_router.py b/api/tests/test_model_router.py index 3abaef0b..6196f3cb 100644 --- a/api/tests/test_model_router.py +++ b/api/tests/test_model_router.py @@ -2,6 +2,7 @@ import json import sys import tempfile +import threading import types import unittest from pathlib import Path @@ -197,5 +198,49 @@ def test_composite_model_unload_route_uses_path_converter(self) -> None: self.assertEqual(model_router.Request.__module__, "urllib.request") +class ModelLifecycleEndpointTests(unittest.TestCase): + def test_status_and_switch_do_not_block_the_event_loop_while_lifecycle_is_contended(self) -> None: + class ContendedRegistry: + def __init__(self): + self.lock = threading.Lock() + def active_status(self): + with self.lock: + return {"id": "demo/a", "loaded": True} + def switch_model(self, model_id): + with self.lock: + return None + + registry = ContendedRegistry() + previous = model_router.generator_registry + model_router.generator_registry = registry + + async def assert_responsive(call): + registry.lock.acquire() + release = threading.Timer(0.25, registry.lock.release) + release.start() + try: + start = asyncio.get_running_loop().time() + task = asyncio.create_task(call()) + await asyncio.sleep(0.02) + elapsed = asyncio.get_running_loop().time() - start + self.assertLess(elapsed, 0.15) + return await task + finally: + release.join() + + async def run(): + status = await assert_responsive(model_router.model_status) + switched = await assert_responsive(lambda: model_router.switch_model("demo/b")) + return status, switched + + try: + status, switched = asyncio.run(run()) + finally: + model_router.generator_registry = previous + + self.assertEqual(status["id"], "demo/a") + self.assertEqual(switched, {"active": "demo/b"}) + + if __name__ == "__main__": unittest.main() diff --git a/api/tests/test_scene_generation.py b/api/tests/test_scene_generation.py index a804171a..14288e6e 100644 --- a/api/tests/test_scene_generation.py +++ b/api/tests/test_scene_generation.py @@ -1,7 +1,9 @@ import asyncio import json import tempfile +import threading import unittest +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from unittest.mock import patch @@ -35,7 +37,7 @@ def setUp(self): def tearDown(self): for item in reversed(self.patches): item.stop() - generation._jobs.clear(); generation._cancel_events.clear(); generation._cancelled.clear(); generation._completed_at.clear() + generation._jobs.clear(); generation._cancel_events.clear(); generation._cancelled.clear(); generation._completed_at.clear(); generation._job_generators.clear() self.tmp.cleanup() def test_generic_route_queues_typed_scene_and_strips_reserved_params(self): @@ -52,6 +54,7 @@ def test_generic_route_queues_typed_scene_and_strips_reserved_params(self): self.assertNotIn("input_kind", queued.args[2]) self.assertEqual(queued.args[5], "demo/scene") self.assertEqual(result["job_id"], queued.args[0]) + self.assertFalse(self.registry.switched) def test_generic_route_rejects_unsupported_kind_and_model_mismatch(self): for kind in ("video", "capture", "image"): @@ -84,7 +87,9 @@ def generate_artifact(self, kind, path, params, progress_cb, cancel_event=None): generator = Generator() registry_stub = type("Registry", (), { - "get_ready_generator": lambda self, model_id: generator if model_id == "demo/a" else (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), + "model_status": lambda self, model_id: {"name": model_id, "downloaded": True, "loaded": True}, + "get_generator": lambda self, model_id: generator if model_id == "demo/a" else (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), + "activate_ready_generator": lambda self, model_id: generator if model_id == "demo/a" else (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), "get_active": lambda self: (_ for _ in ()).throw(AssertionError("mutable active model must not be used")), })() job_id = "pinned-scene" @@ -100,7 +105,9 @@ def generate_artifact(self, kind, path, params, progress_cb, cancel_event=None): def test_missing_pinned_model_fails_actionably(self): registry_stub = type("Registry", (), { - "get_ready_generator": lambda self, model_id: (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), + "model_status": lambda self, model_id: {"name": model_id, "downloaded": True, "loaded": False}, + "get_generator": lambda self, model_id: object(), + "activate_ready_generator": lambda self, model_id: (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), "get_active": lambda self: (_ for _ in ()).throw(AssertionError("must not use active model")), })() job_id = "missing-scene" @@ -113,3 +120,214 @@ def test_missing_pinned_model_fails_actionably(self): )) self.assertEqual(generation._jobs[job_id].status, "error") self.assertIn("Unknown model ID: demo/missing", generation._jobs[job_id].error) + + def test_interleaved_pinned_jobs_serialize_model_lifecycle_and_cancel_exact_job(self): + class Proc: + def __init__(self, owner): + self.owner = owner + self.alive = True + def poll(self): + return None if self.alive else 0 + def kill(self): + self.alive = False + self.owner.killed = True + + class Generator: + def __init__(self, model_id): + self.model_id = model_id + self.outputs_dir = None + self.loaded = False + self.started = threading.Event() + self.killed = False + self._loaded = True + self._proc = Proc(self) + def is_loaded(self): return self.loaded + def is_downloaded(self): return True + def load(self): self.loaded = True + def unload(self): self.loaded = False + def generate_artifact(self, kind, path, params, progress_cb, cancel_event=None): + self.started.set() + if self.model_id == "demo/a": + if cancel_event is None or not cancel_event.wait(2): + raise AssertionError("first job was not cancelled") + raise generation.GenerationCancelled() + output = Path(self.outputs_dir) / "result.glb" + output.write_bytes(b"glb") + return output + + class Registry: + def __init__(self): + self.generators = {model_id: Generator(model_id) for model_id in ("demo/a", "demo/b")} + self.active_id = "demo/a" + def get_generator(self, model_id): return self.generators[model_id] + def model_status(self, model_id): + gen = self.generators[model_id] + return {"name": model_id, "downloaded": True, "loaded": gen.loaded} + def switch_model(self, model_id): + if model_id != self.active_id: + self.generators[self.active_id].unload() + self.active_id = model_id + def activate_ready_generator(self, model_id): + self.switch_model(model_id) + gen = self.generators[model_id] + gen.load() + if sum(candidate.loaded for candidate in self.generators.values()) != 1: + raise AssertionError("more than one generator is resident") + return gen + def get_active(self): raise AssertionError("pinned jobs must not consult mutable active state") + + registry_stub = Registry() + # Reproduce the request interleaving: A switches first, then B switches + # before either background task has started. + registry_stub.switch_model("demo/a") + registry_stub.switch_model("demo/b") + for job_id in ("job-a", "job-b"): + generation._jobs[job_id] = generation.JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = threading.Event() + + async def run_both(): + first = asyncio.create_task(generation._run_generation( + "job-a", generation.TypedArtifactInput("scene", self.manifest.resolve()), {}, + "Workflows", "mesh", "demo/a", + )) + second = asyncio.create_task(generation._run_generation( + "job-b", generation.TypedArtifactInput("scene", self.manifest.resolve()), {}, + "Workflows", "mesh", "demo/b", + )) + started = await asyncio.to_thread(registry_stub.generators["demo/a"].started.wait, 1) + self.assertTrue(started) + await asyncio.sleep(0.05) + self.assertFalse(registry_stub.generators["demo/b"].started.is_set()) + await generation.cancel_job("job-a") + await asyncio.gather(first, second) + + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(run_both()) + + self.assertTrue(registry_stub.generators["demo/a"].killed) + self.assertFalse(registry_stub.generators["demo/b"].killed) + self.assertEqual(generation._jobs["job-a"].status, "cancelled") + self.assertEqual(generation._jobs["job-b"].status, "done") + self.assertEqual(registry_stub.active_id, "demo/b") + self.assertFalse(registry_stub.generators["demo/a"].loaded) + self.assertTrue(registry_stub.generators["demo/b"].loaded) + + def test_cancelling_queued_generation_does_not_wedge_later_jobs(self): + release_first = threading.Event() + first_started = threading.Event() + + class Generator: + outputs_dir = None + + def __init__(self): + self.calls = 0 + + def generate(self, image_bytes, params, progress_cb, cancel_event=None): + self.calls += 1 + call_number = self.calls + if call_number == 1: + first_started.set() + if not release_first.wait(2): + raise AssertionError("first generation was not released") + output = Path(self.outputs_dir) / f"result-{call_number}.glb" + output.write_bytes(b"glb") + return output + + generator = Generator() + registry_stub = type("Registry", (), { + "get_generator": lambda self, model_id: generator, + "model_status": lambda self, model_id: { + "name": model_id, "downloaded": True, "loaded": True, + }, + "activate_ready_generator": lambda self, model_id: generator, + "get_active": lambda self: (_ for _ in ()).throw( + AssertionError("pinned jobs must not use mutable active state") + ), + })() + for job_id in ("blocking", "cancelled-waiter", "after-cancel"): + generation._jobs[job_id] = generation.JobStatus( + job_id=job_id, status="pending", progress=0, + ) + generation._cancel_events[job_id] = threading.Event() + + async def run_scenario(): + first = asyncio.create_task(generation._run_generation( + "blocking", b"image", {}, "Workflows", "mesh", "demo/a", + )) + started = await asyncio.to_thread(first_started.wait, 1) + self.assertTrue(started) + + queued = asyncio.create_task(generation._run_generation( + "cancelled-waiter", b"image", {}, "Workflows", "mesh", "demo/b", + )) + await asyncio.sleep(0) + queued.cancel() + with self.assertRaises(asyncio.CancelledError): + await queued + + release_first.set() + await asyncio.wait_for(first, 1) + await asyncio.wait_for(generation._run_generation( + "after-cancel", b"image", {}, "Workflows", "mesh", "demo/c", + ), 1) + + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(run_scenario()) + + self.assertEqual(generation._jobs["cancelled-waiter"].status, "cancelled") + self.assertEqual(generation._jobs["after-cancel"].status, "done") + self.assertEqual(generator.calls, 2) + + def test_pinned_jobs_do_not_starve_small_default_executor(self): + class Generator: + outputs_dir = None + + def __init__(self): + self.calls = 0 + + def generate(self, image_bytes, params, progress_cb, cancel_event=None): + self.calls += 1 + output = Path(self.outputs_dir) / f"result-{self.calls}.glb" + output.write_bytes(b"glb") + return output + + generator = Generator() + registry_stub = type("Registry", (), { + "get_generator": lambda self, model_id: generator, + "model_status": lambda self, model_id: { + "name": model_id, "downloaded": True, "loaded": True, + }, + "activate_ready_generator": lambda self, model_id: generator, + "get_active": lambda self: (_ for _ in ()).throw( + AssertionError("pinned jobs must not use mutable active state") + ), + })() + job_ids = ("small-pool-a", "small-pool-b", "small-pool-c") + for job_id in job_ids: + generation._jobs[job_id] = generation.JobStatus( + job_id=job_id, status="pending", progress=0, + ) + generation._cancel_events[job_id] = threading.Event() + + default_executor = ThreadPoolExecutor(max_workers=2) + + async def run_scenario(): + asyncio.get_running_loop().set_default_executor(default_executor) + await asyncio.wait_for(asyncio.gather(*( + generation._run_generation( + job_id, b"image", {}, "Workflows", "mesh", f"demo/{job_id}", + ) + for job_id in job_ids + )), 2) + + try: + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(run_scenario()) + finally: + default_executor.shutdown(wait=True) + + self.assertEqual(generator.calls, 3) + self.assertEqual( + [generation._jobs[job_id].status for job_id in job_ids], + ["done", "done", "done"], + ) diff --git a/api/tests/test_scene_input.py b/api/tests/test_scene_input.py index 43bd9946..7d570d10 100644 --- a/api/tests/test_scene_input.py +++ b/api/tests/test_scene_input.py @@ -34,14 +34,18 @@ def test_rejects_traversal_absolute_encoded_and_non_manifest_json(self): with self.subTest(path=path), self.assertRaises(ValueError): validate_scene_input(self.workspace, path) - def test_rejects_symlinks_missing_assets_and_oversized_manifest(self): + def test_rejects_symlink_escape_when_supported(self): outside = Path(self.tmp.name) / "outside" outside.mkdir() (outside / "scene-manifest.json").write_text(self.manifest.read_text()) - (self.workspace / "Workflows" / "link").symlink_to(outside, target_is_directory=True) + try: + (self.workspace / "Workflows" / "link").symlink_to(outside, target_is_directory=True) + except OSError as exc: + self.skipTest(f"directory symlinks are unavailable on this host: {exc}") with self.assertRaises(ValueError): validate_scene_input(self.workspace, "Workflows/link") + def test_rejects_missing_assets_and_oversized_manifest(self): data = json.loads(self.manifest.read_text()) data["assets"] = [{"path": "missing.glb"}] self.manifest.write_text(json.dumps(data)) diff --git a/api/tests/test_workflow_runs_lifecycle.py b/api/tests/test_workflow_runs_lifecycle.py index f8fe7371..7360cc03 100644 --- a/api/tests/test_workflow_runs_lifecycle.py +++ b/api/tests/test_workflow_runs_lifecycle.py @@ -1,12 +1,15 @@ import asyncio +import tempfile import threading import time import unittest +from pathlib import Path from fastapi import BackgroundTasks import routers.generation as generation import routers.workflow_runs as workflow_runs +import services.generator_registry as registry_module from schemas.generation import JobStatus @@ -30,6 +33,9 @@ class _FakeRegistry: def get_generator(self, model_id: str) -> object: return object() + def get_manifest(self, model_id: str) -> dict: + return {"output": "mesh"} + def switch_model(self, model_id: str) -> None: pass @@ -40,6 +46,7 @@ def _clear_job_stores() -> None: generation._cancel_events, generation._cancelled, generation._completed_at, + generation._job_generators, ): store.clear() @@ -91,6 +98,106 @@ def test_cancel_run_records_completion_so_it_can_be_purged(self) -> None: # Without a _completed_at stamp the purge sweep can never evict a cancelled run. self.assertIn(run_id, generation._completed_at) + def test_workflow_and_generate_routes_share_exact_model_queue_and_cancellation(self) -> None: + class Proc: + def __init__(self, owner): + self.owner = owner + self.alive = True + def poll(self): return None if self.alive else 0 + def kill(self): + self.alive = False + self.owner.killed = True + + class Generator: + def __init__(self, model_id): + self.model_id = model_id + self.outputs_dir = None + self.loaded = False + self.started = threading.Event() + self.killed = False + self._loaded = True + self._proc = Proc(self) + def is_loaded(self): return self.loaded + def load(self): self.loaded = True + def unload(self): self.loaded = False + def generate(self, image_bytes, params, progress_cb, cancel_event=None): + self.started.set() + if self.model_id == "demo/workflow": + if cancel_event is None or not cancel_event.wait(2): + raise AssertionError("workflow job was not cancelled") + raise generation.GenerationCancelled() + output = Path(self.outputs_dir) / "result.glb" + output.write_bytes(b"glb") + return output + + class Registry: + def __init__(self): + ids = ("demo/workflow", "demo/generate") + self.generators = {model_id: Generator(model_id) for model_id in ids} + self.active_id = "demo/generate" + def get_generator(self, model_id): return self.generators[model_id] + def get_manifest(self, model_id): return {"output": "mesh"} + def model_status(self, model_id): + gen = self.generators[model_id] + return {"name": model_id, "downloaded": True, "loaded": gen.loaded} + def activate_ready_generator(self, model_id): + if model_id != self.active_id: + self.generators[self.active_id].unload() + self.active_id = model_id + gen = self.generators[model_id] + gen.load() + if sum(candidate.loaded for candidate in self.generators.values()) != 1: + raise AssertionError("more than one generator is resident") + return gen + def get_active(self): raise AssertionError("queued routes must use their pinned model id") + + fake_registry = Registry() + previous_generation_registry = generation.generator_registry + previous_workflow_registry = workflow_runs.generator_registry + previous_workspace = registry_module.WORKSPACE_DIR + with tempfile.TemporaryDirectory() as tmp: + registry_module.WORKSPACE_DIR = Path(tmp) + generation.generator_registry = fake_registry + workflow_runs.generator_registry = fake_registry + + async def run_cross_route_jobs(): + workflow_background = BackgroundTasks() + workflow_response = await workflow_runs.create_run_from_image( + workflow_background, _FakeUpload(), "demo/workflow", "Workflows", "{}", + ) + generate_background = BackgroundTasks() + generate_response = await generation.generate_from_image( + generate_background, _FakeUpload(), "demo/generate", "Workflows", + "quad", False, 1024, "{}", + ) + self.assertEqual(workflow_background.tasks[0].args[5], "demo/workflow") + self.assertEqual(generate_background.tasks[0].args[5], "demo/generate") + + first = asyncio.create_task(workflow_background.tasks[0]()) + second = asyncio.create_task(generate_background.tasks[0]()) + started = await asyncio.to_thread( + fake_registry.generators["demo/workflow"].started.wait, 1, + ) + self.assertTrue(started) + await asyncio.sleep(0.05) + self.assertFalse(fake_registry.generators["demo/generate"].started.is_set()) + await workflow_runs.cancel_run(workflow_response["run_id"]) + await asyncio.gather(first, second) + return workflow_response["run_id"], generate_response["job_id"] + + try: + workflow_job_id, generate_job_id = asyncio.run(run_cross_route_jobs()) + finally: + generation.generator_registry = previous_generation_registry + workflow_runs.generator_registry = previous_workflow_registry + registry_module.WORKSPACE_DIR = previous_workspace + + self.assertTrue(fake_registry.generators["demo/workflow"].killed) + self.assertFalse(fake_registry.generators["demo/generate"].killed) + self.assertEqual(generation._jobs[workflow_job_id].status, "cancelled") + self.assertEqual(generation._jobs[generate_job_id].status, "done") + self.assertEqual(fake_registry.active_id, "demo/generate") + if __name__ == "__main__": unittest.main() diff --git a/api/tests/test_workflow_runs_router.py b/api/tests/test_workflow_runs_router.py index f045bef8..7537a484 100644 --- a/api/tests/test_workflow_runs_router.py +++ b/api/tests/test_workflow_runs_router.py @@ -24,6 +24,9 @@ class _FakeRegistry: def get_generator(self, model_id: str) -> object: return object() + def get_manifest(self, model_id: str) -> dict: + return {"output": "mesh"} + def switch_model(self, model_id: str) -> None: pass @@ -79,7 +82,7 @@ def _collection_forwarded(self, collection: str) -> str: params="{}", ) ) - # add_task(_run_generation, job_id, image_bytes, full_params, collection) + # add_task(_run_generation, ..., collection, output_kind, model_id) return background.tasks[0].args[3] def test_collection_is_forwarded_to_the_run(self) -> None: diff --git a/electron/main/bounded-file-reader.test.mjs b/electron/main/bounded-file-reader.test.mjs new file mode 100644 index 00000000..03071180 --- /dev/null +++ b/electron/main/bounded-file-reader.test.mjs @@ -0,0 +1,48 @@ +import assert from 'node:assert/strict' +import { readFileSync } from 'node:fs' +import { resolve } from 'node:path' +import test from 'node:test' + +const { + MAX_SCENE_MANIFEST_BYTES, + readLocalFileBase64, +} = await import(new URL('./bounded-file-reader.ts', import.meta.url).href) + +test('the authoritative readFileBase64 IPC delegates to the bounded reader', () => { + const handlers = readFileSync(resolve('electron/main/ipc-handlers.ts'), 'utf8') + assert.match( + handlers, + /ipcMain\.handle\('fs:readFileBase64',[\s\S]*?readLocalFileBase64\(filePath\)/, + ) +}) + +test('oversized scene manifests are rejected before their bytes are read', async () => { + let reads = 0 + await assert.rejects( + readLocalFileBase64('/workspace/room/scene-manifest.json', { + statFile: async () => ({ size: MAX_SCENE_MANIFEST_BYTES + 1, isFile: () => true }), + readBytes: async () => { reads++; return new Uint8Array() }, + }), + /exceeds the 1 MiB limit/, + ) + assert.equal(reads, 0) +}) + +test('scene manifests at the bound are read and encoded', async () => { + const bytes = new TextEncoder().encode('{"schema":"modly.scene-manifest.v1"}') + const encoded = await readLocalFileBase64('/workspace/room/scene-manifest.json', { + statFile: async () => ({ size: MAX_SCENE_MANIFEST_BYTES, isFile: () => true }), + readBytes: async () => bytes, + }) + assert.equal(encoded, Buffer.from(bytes).toString('base64')) +}) + +test('a manifest that grows after stat is rejected before base64 transfer', async () => { + await assert.rejects( + readLocalFileBase64('/workspace/room/scene-manifest.json', { + statFile: async () => ({ size: 1, isFile: () => true }), + readBytes: async () => new Uint8Array(MAX_SCENE_MANIFEST_BYTES + 1), + }), + /exceeds the 1 MiB limit/, + ) +}) diff --git a/electron/main/bounded-file-reader.ts b/electron/main/bounded-file-reader.ts new file mode 100644 index 00000000..f7b18ba4 --- /dev/null +++ b/electron/main/bounded-file-reader.ts @@ -0,0 +1,40 @@ +import { basename } from 'node:path' +import { readFile, stat } from 'node:fs/promises' + +export const SCENE_MANIFEST_FILE_NAME = 'scene-manifest.json' +export const MAX_SCENE_MANIFEST_BYTES = 1024 * 1024 + +type FileReadDependencies = { + statFile: (filePath: string) => Promise<{ size: number; isFile: () => boolean }> + readBytes: (filePath: string) => Promise +} + +const DEFAULT_DEPENDENCIES: FileReadDependencies = { + statFile: stat, + readBytes: readFile, +} + +export async function readLocalFileBase64( + filePath: string, + dependencies: FileReadDependencies = DEFAULT_DEPENDENCIES, +): Promise { + if (typeof filePath !== 'string' || filePath.trim().length === 0) { + throw new Error('fs:readFileBase64 requires a non-empty file path') + } + + const isSceneManifest = basename(filePath) === SCENE_MANIFEST_FILE_NAME + if (isSceneManifest) { + const fileInfo = await dependencies.statFile(filePath) + if (!fileInfo.isFile()) throw new Error('Scene manifest is not a file') + if (fileInfo.size > MAX_SCENE_MANIFEST_BYTES) { + throw new Error('Scene manifest exceeds the 1 MiB limit') + } + } + + const bytes = await dependencies.readBytes(filePath) + // Recheck the transferred bytes in case the file grew between stat and read. + if (isSceneManifest && bytes.byteLength > MAX_SCENE_MANIFEST_BYTES) { + throw new Error('Scene manifest exceeds the 1 MiB limit') + } + return Buffer.from(bytes).toString('base64') +} diff --git a/electron/main/extension-install-utils.test.mjs b/electron/main/extension-install-utils.test.mjs index 4dcf42da..3c5ac9bb 100644 --- a/electron/main/extension-install-utils.test.mjs +++ b/electron/main/extension-install-utils.test.mjs @@ -85,7 +85,7 @@ test('validateInstallManifest accepts multi-source nodes and preserves legacy sh }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) }) -test('validateInstallManifest accepts scene IO and rejects undeclared future artifact kinds', () => { +test('validateInstallManifest accepts scene IO without rejecting third-party artifact kinds', () => { const mod = loadModule() const files = { hasEntryFile: () => false, hasGeneratorFile: () => true } assert.doesNotThrow(() => mod.validateInstallManifest({ @@ -93,10 +93,10 @@ test('validateInstallManifest accepts scene IO and rejects undeclared future art nodes: [{ id: 'normalize', input: 'scene', output: 'scene' }], }, files, 'repository')) for (const input of ['capture', 'video']) { - assert.throws(() => mod.validateInstallManifest({ + assert.doesNotThrow(() => mod.validateInstallManifest({ id: 'future-model', generator_class: 'Generator', nodes: [{ id: 'future', input, output: 'scene' }], - }, files, 'repository'), /supported artifact type/) + }, files, 'repository')) } }) diff --git a/electron/main/extension-install-utils.ts b/electron/main/extension-install-utils.ts index 41e9798c..c50e85e0 100644 --- a/electron/main/extension-install-utils.ts +++ b/electron/main/extension-install-utils.ts @@ -66,21 +66,12 @@ export function validateInstallManifest( const isProcess = manifest.type === 'process' const entryFile = manifest.entry ?? 'processor.js' const nodes = Array.isArray(manifest.nodes) ? manifest.nodes.filter((node) => node?.id) : [] - const allowedIo = new Set(['image', 'text', 'mesh', 'audio', 'scene']) - if (manifest.model_sources !== undefined) { throw new Error('manifest.json: model_sources must be declared on a model node') } for (const node of Array.isArray(manifest.nodes) ? manifest.nodes : []) { - const declaredInputs = node.inputs === undefined ? [node.input ?? 'image'] : node.inputs - if (!Array.isArray(declaredInputs) || declaredInputs.length === 0 - || declaredInputs.some((value) => typeof value !== 'string' || !allowedIo.has(value))) { - throw new Error(`manifest.json: ${node.id ?? 'node'}.input must use a supported artifact type`) - } + const declaredInputs = Array.isArray(node.inputs) ? node.inputs : [node.input ?? 'image'] const output = node.output ?? 'mesh' - if (typeof output !== 'string' || !allowedIo.has(output)) { - throw new Error(`manifest.json: ${node.id ?? 'node'}.output must use a supported artifact type`) - } assertSupportedSceneNodeShape(isProcess ? 'process' : 'model', node, declaredInputs, output) if (node.model_sources === undefined) continue if (isProcess) { diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 9e5131de..09f1b649 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -68,6 +68,7 @@ import { } from './extension-install-recovery' import { registerWorkspaceAssetLibraryIpcHandlers } from './artifact-registry-service' import { updatesSupported } from './updater' +import { readLocalFileBase64 } from './bounded-file-reader' type WindowGetter = () => BrowserWindow | null const pExecFile = promisify(execFile) @@ -379,13 +380,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }) // Read local file → base64 (bypasses file:// restrictions in the renderer) - ipcMain.handle('fs:readFileBase64', async (_, filePath: string) => { - if (typeof filePath !== 'string' || filePath.trim().length === 0) { - throw new Error('fs:readFileBase64 requires a non-empty file path') - } - const buffer = await readFile(filePath) - return buffer.toString('base64') - }) + ipcMain.handle('fs:readFileBase64', (_, filePath: string) => readLocalFileBase64(filePath)) ipcMain.handle('fs:readScreenshotDataUrl', async (_, filename: string) => { const filePath = app.isPackaged @@ -845,10 +840,10 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe nodes?: { id: string name?: string - input?: 'mesh' | 'image' | 'text' | 'audio' | 'scene' - inputs?: ('mesh' | 'image' | 'text' | 'audio' | 'scene')[] + input?: string + inputs?: string[] input_labels?: string[] - output?: 'mesh' | 'image' | 'text' | 'audio' | 'scene' + output?: string params_schema?: unknown[] param_defaults?: Record hf_repo?: string @@ -874,14 +869,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe if (parsed.model_sources !== undefined) { throw new Error('manifest.json: model_sources must be declared on a model node') } - const allowedIo = new Set(['image', 'text', 'mesh', 'audio', 'scene']) const nodes = (parsed.nodes ?? []).map(n => { - const declaredInputs = n.inputs ?? [n.input ?? 'image'] - for (const input of declaredInputs) { - if (!allowedIo.has(input)) throw new Error(`manifest.json: unsupported node input type "${input}"`) - } + const declaredInputs = Array.isArray(n.inputs) ? n.inputs : [n.input ?? 'image'] const output = n.output ?? 'mesh' - if (!allowedIo.has(output)) throw new Error(`manifest.json: unsupported node output type "${output}"`) assertSupportedSceneNodeShape(parsed.type === 'process' ? 'process' : 'model', n, declaredInputs, output) if (parsed.type === 'process' && n.model_sources !== undefined) { throw new Error('manifest.json: model_sources is supported only for model nodes') diff --git a/src/areas/workflows/WorkflowsPage.tsx b/src/areas/workflows/WorkflowsPage.tsx index 32452774..8fc4cea6 100644 --- a/src/areas/workflows/WorkflowsPage.tsx +++ b/src/areas/workflows/WorkflowsPage.tsx @@ -68,7 +68,7 @@ const IO_STYLES: Record<'image' | 'text' | 'mesh' | 'audio' | 'scene', string> = image: 'bg-sky-500/15 text-sky-400 border-sky-500/25', mesh: 'bg-violet-500/15 text-violet-400 border-violet-500/25', text: 'bg-amber-500/15 text-amber-400 border-amber-500/25', - scene: 'bg-emerald-500/15 text-emerald-400 border-emerald-500/25', + scene: 'bg-pink-500/15 text-pink-400 border-pink-500/25', } function IoBadge({ type }: { type: 'image' | 'text' | 'mesh' | 'audio' | 'scene' }) { @@ -101,7 +101,7 @@ const PANEL_BUILTIN_NODES = [ { type: 'imageNode', label: 'Image', color: '#38bdf8', icon: <> }, { type: 'textNode', label: 'Text', color: '#fbbf24', icon: <> }, { type: 'meshNode', label: 'Load 3D Mesh', color: '#a78bfa', icon: <> }, - { type: 'sceneNode', label: 'Load Scene', color: '#34d399', icon: <> }, + { type: 'sceneNode', label: 'Load Scene', color: '#f472b6', icon: <> }, { type: 'outputNode', label: 'Add to Scene', color: '#a78bfa', icon: <> }, { type: 'previewNode', label: 'Preview Views', color: '#38bdf8', icon: <> }, { type: 'imagePreviewNode', label: 'Preview Image', color: '#38bdf8', icon: <> }, @@ -349,7 +349,7 @@ const BUILTIN_NODES = [ { type: 'imageNode', label: 'Image', color: '#38bdf8', description: 'Image input' }, { type: 'textNode', label: 'Text', color: '#fbbf24', description: 'Text input' }, { type: 'meshNode', label: 'Load 3D Mesh', color: '#a78bfa', description: 'Load a 3D mesh file or use current model' }, - { type: 'sceneNode', label: 'Load Scene', color: '#34d399', description: 'Load and validate a workspace scene directory' }, + { type: 'sceneNode', label: 'Load Scene', color: '#f472b6', description: 'Load and validate a workspace scene directory' }, { type: 'outputNode', label: 'Add to Scene', color: '#a78bfa', description: 'Output node — adds the mesh to the 3D scene' }, { type: 'previewNode', label: 'Preview Views', color: '#38bdf8', description: 'Displays multi-view image outputs in a 2×3 grid' }, { type: 'imagePreviewNode', label: 'Preview Image', color: '#38bdf8', description: 'Displays a single image output in the workflow' }, @@ -1379,7 +1379,7 @@ const MINI_NODE_TINTS: Record = { imageNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, textNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, meshNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, - sceneNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, + sceneNode: { fill: 'rgba(244,114,182,0.22)', stroke: '#f472b6' }, extensionNode: { fill: 'rgba(167,139,250,0.24)', stroke: '#a78bfa' }, outputNode: { fill: 'rgba(56,189,248,0.22)', stroke: '#38bdf8' }, previewNode: { fill: 'rgba(56,189,248,0.22)', stroke: '#38bdf8' }, diff --git a/src/areas/workflows/nodes/ExtensionNode.tsx b/src/areas/workflows/nodes/ExtensionNode.tsx index 4be11a14..e5dc3143 100644 --- a/src/areas/workflows/nodes/ExtensionNode.tsx +++ b/src/areas/workflows/nodes/ExtensionNode.tsx @@ -16,7 +16,7 @@ const HANDLE_COLOR: Record = { image: '#38bdf8', mesh: '#a78bfa', text: '#fbbf24', - scene: '#34d399', + scene: '#f472b6', } const TAG_CLS: Record = { @@ -24,7 +24,7 @@ const TAG_CLS: Record = { image: 'border-sky-500/30 bg-sky-500/10 text-sky-400', mesh: 'border-violet-500/30 bg-violet-500/10 text-violet-400', text: 'border-amber-500/30 bg-amber-500/10 text-amber-400', - scene: 'border-emerald-500/30 bg-emerald-500/10 text-emerald-400', + scene: 'border-pink-500/30 bg-pink-500/10 text-pink-400', } // ─── Param control ──────────────────────────────────────────────────────────── diff --git a/src/areas/workflows/nodes/LoadSceneNode.tsx b/src/areas/workflows/nodes/LoadSceneNode.tsx index 2ac7dded..40f3af1c 100644 --- a/src/areas/workflows/nodes/LoadSceneNode.tsx +++ b/src/areas/workflows/nodes/LoadSceneNode.tsx @@ -1,18 +1,26 @@ import { useCallback, useLayoutEffect, useRef, useState } from 'react' import { Handle, Position, useReactFlow } from '@xyflow/react' -import type { WFNodeData } from '@shared/types/electron.d' +import type { ReactFlowInstance } from '@xyflow/react' +import type { WFNode, WFNodeData } from '@shared/types/electron.d' import BaseNode from './BaseNode' -import { resolveSceneSourceManifest } from '../workflowSceneSource' +import { + applySceneValidationResult, + invalidateValidatedScenePath, + resolveSceneSourceManifest, +} from '../workflowSceneSource' -const OUTPUT_COLOR = '#34d399' +const OUTPUT_COLOR = '#f472b6' async function validateAndPersistScenePath(args: { id: string - data: WFNodeData nextPath: string - updateNodeData: ReturnType['updateNodeData'] + updateNodeData: ReactFlowInstance['updateNodeData'] }): Promise { + args.updateNodeData(args.id, (node) => ({ + params: invalidateValidatedScenePath(node.data.params, args.nextPath), + })) + const settings = await window.electron.settings.get() const resolution = await resolveSceneSourceManifest({ scenePath: args.nextPath, @@ -20,33 +28,14 @@ async function validateAndPersistScenePath(args: { readFileBase64: window.electron.fs.readFileBase64, }) - if (!resolution.ok) { - args.updateNodeData(args.id, { - params: { - ...args.data.params, - path: args.nextPath, - manifestPath: undefined, - sceneRoot: undefined, - error: resolution.error, - }, - }) - return - } - - args.updateNodeData(args.id, { - params: { - ...args.data.params, - path: resolution.inputWorkspacePath, - manifestPath: resolution.manifestWorkspacePath, - sceneRoot: resolution.sceneRoot, - sourceKind: resolution.sourceKind, - error: undefined, - }, + args.updateNodeData(args.id, (node) => { + const params = applySceneValidationResult(node.data.params, args.nextPath, resolution) + return params ? { params } : {} }) } export default function LoadSceneNode({ id, data, selected }: { id: string; data: WFNodeData; selected?: boolean }) { - const { updateNodeData } = useReactFlow() + const { updateNodeData } = useReactFlow() const ioRowRef = useRef(null) const [handleTop, setHandleTop] = useState('50%') @@ -65,13 +54,13 @@ export default function LoadSceneNode({ id, data, selected }: { id: string; data const browseDirectory = useCallback(async () => { const path = await window.electron.fs.selectDirectory() if (!path) return - await validateAndPersistScenePath({ id, data, nextPath: path, updateNodeData }) - }, [id, data, updateNodeData]) + await validateAndPersistScenePath({ id, nextPath: path, updateNodeData }) + }, [id, updateNodeData]) const validatePath = useCallback(async () => { if (!scenePath.trim()) return - await validateAndPersistScenePath({ id, data, nextPath: scenePath, updateNodeData }) - }, [id, data, scenePath, updateNodeData]) + await validateAndPersistScenePath({ id, nextPath: scenePath, updateNodeData }) + }, [id, scenePath, updateNodeData]) return ( - scene + scene
} handles={ @@ -106,20 +95,20 @@ export default function LoadSceneNode({ id, data, selected }: { id: string; data type="text" value={scenePath} placeholder="Scenes/castle or Scenes/castle/scene-manifest.json" - onChange={(event) => updateNodeData(id, { params: { ...data.params, path: event.target.value } })} - className="nodrag w-full rounded-lg border border-zinc-700 bg-zinc-800 px-2.5 py-2 text-[10px] text-zinc-200 placeholder-zinc-600 focus:outline-none focus:border-emerald-500/40" + onChange={(event) => updateNodeData(id, { params: invalidateValidatedScenePath(data.params, event.target.value) })} + className="nodrag w-full rounded-lg border border-zinc-700 bg-zinc-800 px-2.5 py-2 text-[10px] text-zinc-200 placeholder-zinc-600 focus:outline-none focus:border-pink-500/40" />
- -
{manifestPath ? ( -
-
Manifest: {manifestPath}
+
+
Manifest: {manifestPath}
{sceneRoot &&
sceneRoot: {sceneRoot}
}
) : ( diff --git a/src/areas/workflows/workflowRunStore.ts b/src/areas/workflows/workflowRunStore.ts index 82166635..ba6faf70 100644 --- a/src/areas/workflows/workflowRunStore.ts +++ b/src/areas/workflows/workflowRunStore.ts @@ -7,6 +7,7 @@ import type { WorkflowExtension } from './mockExtensions' import type { Workflow, WFNode, WFEdge } from '@shared/types/electron.d' import { isBranchStarter, isSceneOutput, resolveDataSource, reachesSceneOutput, nearestUpstreamWaits } from './nodeBehaviors' import { assignSlotFilePaths } from './slotInputs' +import type { SlotInputType } from './slotInputs' // ─── Types ──────────────────────────────────────────────────────────────────── @@ -337,7 +338,10 @@ async function executeExtensionNode( const incomingEdges = workflow.edges.filter((e) => e.target === node.id) if (ext?.inputs && ext.inputs.length > 1) { - const inputTypes = ext.inputs + // Scene is intentionally a single-input-only model contract. The guard + // above rejects it before this multi-slot path, and filtering it here also + // narrows the remaining declarations to assignSlotFilePaths' exact ABI. + const inputTypes: SlotInputType[] = ext.inputs.filter((input) => input !== 'scene') // Resolved by target handle first, then typed by that slot's declared input -- // not by the arrival order of `incomingEdges`, which does not match slot order. const inputPaths = new Array(inputTypes.length).fill(undefined) diff --git a/src/areas/workflows/workflowSceneSource.test.mjs b/src/areas/workflows/workflowSceneSource.test.mjs index ad2a0280..778e13d3 100644 --- a/src/areas/workflows/workflowSceneSource.test.mjs +++ b/src/areas/workflows/workflowSceneSource.test.mjs @@ -8,7 +8,11 @@ import { join, resolve } from 'node:path' const outfile = join(mkdtempSync(join(tmpdir(), 'modly-scene-source-')), 'scene.cjs') writeFileSync(outfile, buildSync({ entryPoints: [resolve('src/areas/workflows/workflowSceneSource.ts')], bundle: true, platform: 'node', format: 'cjs', write: false }).outputFiles[0].text) -const { resolveSceneSourceManifest } = createRequire(import.meta.url)(outfile) +const { + applySceneValidationResult, + invalidateValidatedScenePath, + resolveSceneSourceManifest, +} = createRequire(import.meta.url)(outfile) const encoded = Buffer.from(JSON.stringify({ schema: 'modly.scene-manifest.v1', sceneRoot: '.', assets: [] })).toString('base64') test('Load Scene resolves directory and manifest without image bytes', async () => { @@ -27,3 +31,40 @@ test('Load Scene refuses unsafe paths before reading', async () => { } assert.equal(reads, 0) }) + +test('editing a validated Load Scene path clears every derived scene reference', () => { + const params = invalidateValidatedScenePath({ + path: 'Workflows/old', + manifestPath: 'Workflows/old/scene-manifest.json', + sceneRoot: '.', + sourceKind: 'directory', + }, 'Workflows/new') + + assert.equal(params.path, 'Workflows/new') + assert.equal(params.manifestPath, undefined) + assert.equal(params.sceneRoot, undefined) + assert.equal(params.sourceKind, undefined) +}) + +test('an async validation result cannot overwrite a subsequently edited path', async () => { + const oldResolution = await resolveSceneSourceManifest({ + scenePath: 'Workflows/old', + workspaceDir: '/workspace', + readFileBase64: async () => encoded, + }) + assert.equal(oldResolution.ok, true) + + const currentParams = invalidateValidatedScenePath({ + path: 'Workflows/old', + manifestPath: 'Workflows/old/scene-manifest.json', + sceneRoot: '.', + sourceKind: 'directory', + }, 'Workflows/new') + + assert.equal( + applySceneValidationResult(currentParams, 'Workflows/old', oldResolution), + undefined, + ) + assert.equal(currentParams.path, 'Workflows/new') + assert.equal(currentParams.manifestPath, undefined) +}) diff --git a/src/areas/workflows/workflowSceneSource.ts b/src/areas/workflows/workflowSceneSource.ts index 5873b692..efe2b3ed 100644 --- a/src/areas/workflows/workflowSceneSource.ts +++ b/src/areas/workflows/workflowSceneSource.ts @@ -21,6 +21,47 @@ export type ResolveSceneSourceFailure = { export type ResolveSceneSourceResult = ResolveSceneSourceSuccess | ResolveSceneSourceFailure +export function applySceneValidationResult( + currentParams: Record, + validatedPath: string, + resolution: ResolveSceneSourceResult, +): Record | undefined { + if (currentParams.path !== validatedPath) return undefined + + if (!resolution.ok) { + return { + ...currentParams, + manifestPath: undefined, + sceneRoot: undefined, + sourceKind: undefined, + error: resolution.error, + } + } + + return { + ...currentParams, + path: resolution.inputWorkspacePath, + manifestPath: resolution.manifestWorkspacePath, + sceneRoot: resolution.sceneRoot, + sourceKind: resolution.sourceKind, + error: undefined, + } +} + +export function invalidateValidatedScenePath( + params: Record, + nextPath: string, +): Record { + return { + ...params, + path: nextPath, + manifestPath: undefined, + sceneRoot: undefined, + sourceKind: undefined, + error: undefined, + } +} + type ResolveSceneSourceArgs = { scenePath: string workspaceDir: string From e75f19bfeb2412b265cd45ce8247976d1ee6e090 Mon Sep 17 00:00:00 2001 From: mojo Date: Thu, 1 Oct 2026 21:51:06 -0400 Subject: [PATCH 38/57] Add files via upload --- WorkflowPanel.tsx | 798 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 798 insertions(+) create mode 100644 WorkflowPanel.tsx diff --git a/WorkflowPanel.tsx b/WorkflowPanel.tsx new file mode 100644 index 00000000..e5906ae5 --- /dev/null +++ b/WorkflowPanel.tsx @@ -0,0 +1,798 @@ +import { useCallback, useEffect, useMemo, useRef, useState } from 'react' +import { + ReactFlowProvider, + useNodesState, useEdgesState, useReactFlow, + type Node as FlowNode, type Edge as FlowEdge, +} from '@xyflow/react' +// ReactFlowProvider wraps EmbeddedCanvas so useReactFlow() works in param rows +import { useWorkflowsStore } from '@shared/stores/workflowsStore' +import { useAppStore } from '@shared/stores/appStore' +import { useExtensionsStore } from '@shared/stores/extensionsStore' +import { useNavStore } from '@shared/stores/navStore' +import { useWorkflowRunStore } from '@areas/workflows/workflowRunStore' +import { useWaitButton } from '@areas/workflows/useWaitButton' +import { buildAllWorkflowExtensions, getWorkflowExtension } from '@areas/workflows/mockExtensions' +import { validateWorkflowPreflight } from '@areas/workflows/preflight' +import { mimeFromPath } from '@areas/workflows/nodes/imageUtils' +import type { WorkflowExtension } from '@areas/workflows/mockExtensions' +import type { Workflow, WFNode, WFEdge, ParamSchema } from '@shared/types/electron.d' +import { PICKER_LABELS, openParamPicker, resolvePickerIntent } from '@shared/utils/paramPicker' +import { PickerIcon } from '@shared/components/ui' +import ChatPanel from './ChatPanel' + +type PanelMode = 'basic' | 'chat' + +// ─── Constants ──────────────────────────────────────────────────────────────── + +const TYPE_COLOR: Record = { + image: '#38bdf8', + mesh: '#a78bfa', + text: '#fbbf24', +} + +const SUPPORTED_IMAGE_TYPES = new Set(['image/jpeg', 'image/png', 'image/webp']) + +// ─── Helpers ────────────────────────────────────────────────────────────────── + +function topoSortNodes(nodes: Workflow['nodes'], edges: Workflow['edges']): WFNode[] { + const nodeMap = new Map(nodes.map((n) => [n.id, n])) + const inDegree = new Map(nodes.map((n) => [n.id, 0])) + const adj = new Map(nodes.map((n) => [n.id, [] as string[]])) + for (const e of edges) { + if (!nodeMap.has(e.source) || !nodeMap.has(e.target)) continue + adj.get(e.source)!.push(e.target) + inDegree.set(e.target, (inDegree.get(e.target) ?? 0) + 1) + } + const queue = nodes.filter((n) => (inDegree.get(n.id) ?? 0) === 0) + const result: WFNode[] = [] + while (queue.length > 0) { + const node = queue.shift()! + result.push(node) + for (const neighbor of adj.get(node.id) ?? []) { + const deg = (inDegree.get(neighbor) ?? 0) - 1 + inDegree.set(neighbor, deg) + if (deg === 0) queue.push(nodeMap.get(neighbor)!) + } + } + return result +} + +// ─── Param field ────────────────────────────────────────────────────────────── + +const inputCls = 'w-full bg-zinc-800 border border-zinc-700/80 rounded-md px-2 py-1 text-[11px] text-zinc-200 focus:outline-none focus:border-accent/60' + +function IntInput({ value, onChange, className }: { value: number; onChange: (v: number) => void; className: string }) { + const [text, setText] = useState(String(value)) + const prevValue = useRef(value) + if (prevValue.current !== value && parseInt(text, 10) !== value) { + prevValue.current = value + setText(String(value)) + } + return ( + { + const raw = e.target.value + if (raw !== '' && raw !== '-' && !/^-?\d+$/.test(raw)) return + setText(raw) + const n = parseInt(raw, 10) + if (!isNaN(n)) { prevValue.current = n; onChange(n) } + }} + className={className} + /> + ) +} + +function FloatInput({ value, onChange, className, min, max, step, label, defaultValue }: { + value: number + onChange: (v: number) => void + className: string + min?: number + max?: number + step?: number + label: string + defaultValue: number +}) { + const [text, setText] = useState(String(value)) + const prevValue = useRef(value) + if (prevValue.current !== value && parseFloat(text.replace(',', '.')) !== value) { + prevValue.current = value + setText(String(value)) + } + const sliderMin = typeof min === 'number' ? min : 0 + const sliderMax = typeof max === 'number' ? max : 0 + const hasSlider = typeof min === 'number' && typeof max === 'number' && sliderMax > sliderMin + const sliderStep = typeof step === 'number' && step > 0 + ? step + : hasSlider ? (sliderMax - sliderMin) / 100 : undefined + const parsedValue = typeof value === 'number' ? value : Number.parseFloat(String(value)) + const sliderValue = Number.isFinite(parsedValue) ? parsedValue : defaultValue + const numberInput = ( + { + const raw = e.target.value.replace(',', '.') + if (raw !== '' && raw !== '-' && raw !== '.' && !/^-?\d*\.?\d*$/.test(raw)) return + setText(e.target.value) + const num = parseFloat(raw) + if (!isNaN(num)) { prevValue.current = num; onChange(num) } + }} + className={hasSlider ? `${className.replace('w-full', 'w-16 shrink-0 text-center')} nodrag` : className} + /> + ) + if (!hasSlider) return numberInput + + return ( +
+ { + const num = e.currentTarget.valueAsNumber + if (Number.isFinite(num)) { setText(String(num)); prevValue.current = num; onChange(num) } + }} + aria-label={`${label} slider`} + style={{ accentColor: '#38bdf8', cursor: 'pointer' }} + className="nodrag min-w-0 flex-1" + /> + {numberInput} +
+ ) +} + +function ParamField({ param, value, onChange }: { + param: ParamSchema + value: number | string + onChange: (v: number | string) => void +}) { + if (param.type === 'select') { + return ( + + ) + } + if (param.type === 'string') { + const intent = resolvePickerIntent(param) + return ( +
+ onChange(e.target.value)} className={`${inputCls} flex-1`} /> + +
+ ) + } + if (param.type === 'float') { + return onChange(v)} className={inputCls} + min={param.min} max={param.max} step={param.step} label={param.label} + defaultValue={typeof param.default === 'number' ? param.default : 0} /> + } + // int + return onChange(v)} className={inputCls} /> +} + +// ─── Workflow dropdown ──────────────────────────────────────────────────────── + +function WorkflowDropdown({ workflows, value, onChange }: { + workflows: Workflow[] + value: string | null + onChange: (id: string) => void +}) { + const [open, setOpen] = useState(false) + const selected = workflows.find((w) => w.id === value) + + if (workflows.length === 0) { + return ( +
+ No workflows yet +
+ ) + } + + return ( +
+ + + {open && ( +
+ {workflows.map((wf, i) => ( + + ))} +
+ )} +
+ ) +} + +// ─── Node param rows ────────────────────────────────────────────────────────── +// These components receive nodes + onPatch directly from EmbeddedCanvas +// to avoid relying on the React Flow store (which requires a mounted ). + +type PatchFn = (nodeId: string, patch: Record) => void + +function ImageParamRow({ nodeId, nodes, onPatch }: { nodeId: string; nodes: FlowNode[]; onPatch: PatchFn }) { + const node = nodes.find((n) => n.id === nodeId) + const data = node?.data as { params: Record } | undefined + const preview = data?.params.preview as string | undefined + const showToast = useAppStore((state) => state.showToast) + const loadRequest = useRef(0) + + const applyImagePath = useCallback(async (path: string | null) => { + const request = ++loadRequest.current + if (!path) return + try { + const base64 = await window.electron.fs.readFileBase64(path) + if (request !== loadRequest.current) return + const src = `data:${mimeFromPath(path)};base64,${base64}` + onPatch(nodeId, { params: { ...(data?.params ?? {}), filePath: path, preview: src } }) + } catch { + if (request === loadRequest.current) showToast('Unable to load the selected image') + } + }, [nodeId, data?.params, onPatch, showToast]) + + const browse = useCallback(async () => { + await applyImagePath(await window.electron.fs.selectImage()) + }, [applyImagePath]) + + return ( +
{ + event.preventDefault() + event.dataTransfer.dropEffect = 'copy' + }} + onDrop={(event) => { + event.preventDefault() + const file = event.dataTransfer.files[0] + if (!file || !SUPPORTED_IMAGE_TYPES.has(file.type)) return + void applyImagePath(window.electron.fs.getPathForFile(file)) + }} + > +
+ + + + + Image +
+ {preview ? ( + + ) : ( + + )} +
+ ) +} + +function MeshParamRow({ nodeId, nodes, onPatch }: { nodeId: string; nodes: FlowNode[]; onPatch: PatchFn }) { + const node = nodes.find((n) => n.id === nodeId) + const data = node?.data as { params: Record } | undefined + const source = (data?.params.source as 'file' | 'current' | undefined) ?? 'file' + const fileName = data?.params.fileName as string | undefined + + const browse = useCallback(async () => { + const p = await window.electron.fs.selectMeshFile() + if (!p) return + const name = p.split(/[\\/]/).pop() ?? p + onPatch(nodeId, { params: { ...(data?.params ?? {}), filePath: p, fileName: name, source: 'file' } }) + }, [nodeId, data?.params, onPatch]) + + const toggleSource = useCallback(() => { + const next = source === 'file' ? 'current' : 'file' + onPatch(nodeId, { params: { ...(data?.params ?? {}), source: next } }) + }, [nodeId, data?.params, source, onPatch]) + + return ( +
+
+ + + + Load 3D Mesh +
+ + {/* Toggle: use current model */} + + + {source === 'file' ? ( + fileName ? ( + + ) : ( + + ) + ) : ( +
+ Uses the model currently loaded in the 3D viewer +
+ )} +
+ ) +} + +function TextParamRow({ nodeId, nodes, onPatch }: { nodeId: string; nodes: FlowNode[]; onPatch: PatchFn }) { + const node = nodes.find((n) => n.id === nodeId) + const data = node?.data as { params: Record } | undefined + const text = (data?.params.text as string | undefined) ?? '' + + return ( +
+
+ + + + Text +
+