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 1/2] 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 2/2] 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<{