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 11184ddc..04eb14ea 100644 --- a/README.md +++ b/README.md @@ -69,6 +69,9 @@ npm run build ## Platform notes +- AMD GPUs are supported through ROCm on Linux and Windows: a Radeon card is detected + automatically and extensions are steered to ROCm PyTorch wheels, with no ROCm install + required. See [docs/running-on-amd-rocm.md](docs/running-on-amd-rocm.md). - macOS support targets Apple Silicon only. - macOS uses native window controls. Windows and Linux keep the existing custom controls. - The top bar includes a live RAM indicator sourced from the main process. @@ -106,11 +109,188 @@ Modly supports external model and process extensions. Each extension is a GitHub ![Install models](docs/install-models.png) +### Multiple Hugging Face repositories per model node + +A model node whose weights are split across repositories can declare +`model_sources`. Modly validates every source, downloads them sequentially in +one Models-page action, and considers the node installed only when every +declared check exists. + +```json +{ + "id": "generate", + "model_sources": [ + { + "id": "primary", + "provider": "huggingface", + "repo_id": "org/main-model", + "destination": ".", + "checks": ["model.safetensors"] + }, + { + "id": "encoder", + "provider": "huggingface", + "repo_id": "org/encoder", + "revision": "v1.0", + "destination": "auxiliary/encoder", + "include_prefixes": ["config.json", "model.safetensors"], + "checks": ["config.json", "model.safetensors"] + } + ] +} +``` + +`destination`, filters, and checks use safe POSIX paths relative to the node's +model directory. Every check must name a regular, non-empty file included by +that source's filters; invalid plans fail before any file is downloaded. Pin a +tag or commit in `revision` when reproducible weights are required. The only +supported provider is `huggingface`. Existing nodes that use `hf_repo`, +`download_check`, `hf_include_prefixes`, and `hf_skip_prefixes` keep their +original behavior. + +### 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 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. + +### Separately installable weight variants + +A model node that publishes the same weights in several variants (quantizations, +precisions…) can declare `weight_variants` next to `hf_repo`. The Extensions page +lists every variant under the node, and each one is downloaded or deleted on its own. + +```json +{ + "id": "generate", + "hf_repo": "org/model-gguf", + "download_check": "pipeline.json", + "hf_include_prefixes": ["pipeline.json", "encoder/", "dit/"], + "params_schema": [ + { + "id": "quant", + "label": "Quantization", + "type": "select", + "default": "Q5_K_M", + "options": [ + { "value": "Q4_K_M", "label": "Q4_K_M" }, + { "value": "Q5_K_M", "label": "Q5_K_M" } + ] + } + ], + "weight_variants": { + "param": "quant", + "default": "Q5_K_M", + "options": [ + { + "id": "Q4_K_M", + "label": "Q4_K_M", + "size_gb": 2.4, + "vram_gb": 6, + "include_prefixes": ["dit/model_Q4_K_M.gguf"], + "checks": ["dit/model_Q4_K_M.gguf"] + }, + { + "id": "Q5_K_M", + "include_prefixes": ["dit/model_Q5_K_M.gguf"], + "checks": ["dit/model_Q5_K_M.gguf"] + } + ] + } +} +``` + +- `param` names the `params_schema` select whose values are the variant ids. That + param must exist on the node (or on the extension, as its fallback), and when it + declares `options` they must cover every variant id. +- `size_gb` (download size) and `vram_gb` (approximate VRAM the variant needs) are + optional positive numbers, shown next to the variant when present. +- Every install downloads the shared files (`hf_include_prefixes`, with every + variant's files excluded automatically) plus one variant: the one asked for, or the + `default` one — the first option when `default` is omitted. Files already complete + on disk are skipped, so adding a second variant only fetches that variant. + Only declared variants are excluded from the shared pass: keep + `hf_include_prefixes` narrow enough that a variant the repository publishes but + the manifest does not declare (e.g. an extra `dit/model_Q8_0.gguf`) is not + downloaded with every install. +- A variant is installed when all of its `checks` exist; the node is installed once + its `download_check` and at least one variant are present. +- Generation fails with an explicit message when the selected variant is not + installed, and the node's selector labels those options `(not installed)`. +- `include_prefixes` and `checks` are safe POSIX paths relative to the node's model + directory. Prefixes of two variants cannot overlap, `download_check` stays outside + every variant, and `weight_variants` cannot be combined with `model_sources`. + --- ## Workflows Start with a basic workflow first. For example, on the "Workflows" tab, try: Image -> Generate Mesh -> Add to Scene. Make sure there is a connection between each of the steps. Go to the "Generate" tab, make sure the workflow is selected, then click on "Generate 3D Model". Click on "Settings/Logs/Errors" to see any issues. +Model extensions may also declare `scene` as a node input or output. A scene is +a workspace directory containing `scene-manifest.json` with schema +`modly.scene-manifest.v1`; it is not an arbitrary JSON file. Use the **Load +Scene** workflow node to select and validate an existing scene directory. +Scene-capable generators implement `generate_artifact(input_kind, +artifact_path, ...)`; legacy image generators and `POST /generate/from-image` +remain unchanged. The generic `POST /generate/from-artifact` boundary currently +accepts only `scene`, leaving future artifact kinds to separate reviewed changes. +For this first contract, `scene` is model-only and must be declared as the single +`input` value (not inside `inputs`); process and mixed-input scene nodes are rejected. +Model nodes may still accept multiple images and produce a scene. + ## Modly CLI diff --git a/api/README.md b/api/README.md index ee45bc3f..cdcd4d6b 100644 --- a/api/README.md +++ b/api/README.md @@ -29,6 +29,7 @@ uvicorn main:app --host 127.0.0.1 --port 8765 --reload | GET | `/model/status` | Model download / load status | | GET | `/model/download` | SSE stream of download progress | | POST | `/generate/from-image` | Start image-to-3D job | +| POST | `/generate/from-artifact` | Start a typed-artifact model job (`scene` only) | | GET | `/generate/status/{job_id}` | Poll job status | ## Model diff --git a/api/main.py b/api/main.py index 382ed67e..958ba2a0 100644 --- a/api/main.py +++ b/api/main.py @@ -12,7 +12,7 @@ from services.stdio_utf8 import ensure_utf8_stdio ensure_utf8_stdio() # must run before any print/logging hits the pipe -from routers import generation, model, optimize, status, settings, extensions, export, workflow_runs, agent +from routers import generation, model, optimize, status, settings, extensions, export, workflow_runs, agent, llm @asynccontextmanager @@ -21,8 +21,10 @@ async def lifespan(app: FastAPI): from services.generator_registry import generator_registry generator_registry.initialize() yield - # Shutdown: unload all models + # Shutdown: unload all models and stop the local LLM server generator_registry.unload_all() + from services.llm_server import llama_pool + llama_pool.unload_all(force=True) class _StatusFilter(logging.Filter): @@ -34,7 +36,7 @@ def filter(self, record): app = FastAPI( title="Modly API", - version="0.4.2", + version="0.4.3", lifespan=lifespan, ) @@ -57,6 +59,7 @@ def filter(self, record): app.include_router(export.router, prefix="/export") app.include_router(workflow_runs.router, prefix="/workflow-runs") app.include_router(agent.router) +app.include_router(llm.router, prefix="/llm") # Serve generated files from workspace — dynamic so path changes take effect immediately @app.get("/workspace/{full_path:path}") diff --git a/api/resources/llm_catalog.json b/api/resources/llm_catalog.json new file mode 100644 index 00000000..928f4a2a --- /dev/null +++ b/api/resources/llm_catalog.json @@ -0,0 +1,169 @@ +[ + { + "id": "qwen3.5-4b", + "name": "Qwen 3.5 4B", + "description": "Newest small Qwen. Tops independent tool-calling tests for its size and stays fast on modest GPUs.", + "hf_repo": "unsloth/Qwen3.5-4B-GGUF", + "hf_filename": "Qwen3.5-4B-Q4_K_M.gguf", + "size_bytes": 2740937888, + "quant": "Q4_K_M", + "ctx": 16384, + "vram_estimate_mb": 4600, + "ngl_suggestion": 99, + "tags": [ + "fast", + "4b" + ], + "sampling": { + "temperature": 0.7, + "top_p": 0.8, + "top_k": 20, + "presence_penalty": 1.5 + } + }, + { + "id": "qwen3.5-9b", + "name": "Qwen 3.5 9B", + "description": "Best balance of reliability and speed on an 8 GB+ GPU. Noticeably steadier than 4B over long conversations.", + "hf_repo": "unsloth/Qwen3.5-9B-GGUF", + "hf_filename": "Qwen3.5-9B-Q4_K_M.gguf", + "size_bytes": 5680522464, + "quant": "Q4_K_M", + "ctx": 16384, + "vram_estimate_mb": 7600, + "ngl_suggestion": 99, + "tags": [ + "balanced", + "9b" + ], + "sampling": { + "temperature": 0.7, + "top_p": 0.8, + "top_k": 20, + "presence_penalty": 1.5 + } + }, + { + "id": "qwen3-4b", + "name": "Qwen 3 4B Instruct", + "description": "Latest Qwen generation. Fast, excellent tool-calling for its size — the recommended starting point.", + "hf_repo": "unsloth/Qwen3-4B-Instruct-2507-GGUF", + "hf_filename": "Qwen3-4B-Instruct-2507-Q4_K_M.gguf", + "size_bytes": 2497281120, + "quant": "Q4_K_M", + "ctx": 16384, + "vram_estimate_mb": 4200, + "ngl_suggestion": 99, + "tags": [ + "fast", + "4b", + "default" + ] + }, + { + "id": "qwen3-8b", + "name": "Qwen 3 8B", + "description": "Great quality/speed balance with hybrid thinking. Fits in 8 GB VRAM.", + "hf_repo": "unsloth/Qwen3-8B-GGUF", + "hf_filename": "Qwen3-8B-Q4_K_M.gguf", + "size_bytes": 5027784512, + "quant": "Q4_K_M", + "ctx": 16384, + "vram_estimate_mb": 7000, + "ngl_suggestion": 99, + "tags": [ + "balanced", + "8b", + "thinking" + ] + }, + { + "id": "qwen3-14b", + "name": "Qwen 3 14B", + "description": "High quality with hybrid thinking (reasons before answering when useful). Needs ~11 GB VRAM.", + "hf_repo": "unsloth/Qwen3-14B-GGUF", + "hf_filename": "Qwen3-14B-Q4_K_M.gguf", + "size_bytes": 9001753984, + "quant": "Q4_K_M", + "ctx": 16384, + "vram_estimate_mb": 11000, + "ngl_suggestion": 99, + "tags": [ + "quality", + "14b", + "thinking" + ] + }, + { + "id": "gpt-oss-20b", + "name": "GPT-OSS 20B (OpenAI)", + "description": "OpenAI's open-weight MoE (3.6B active params — fast for its size). Top-tier tool calling and reasoning. Best with 16 GB VRAM.", + "hf_repo": "ggml-org/gpt-oss-20b-GGUF", + "hf_filename": "gpt-oss-20b-MXFP4.gguf", + "size_bytes": 12109566624, + "quant": "MXFP4", + "ctx": 16384, + "vram_estimate_mb": 13500, + "ngl_suggestion": 99, + "tags": [ + "quality", + "20b", + "thinking", + "moe" + ] + }, + { + "id": "qwen3-vl-4b", + "name": "Qwen 3 VL 4B (vision, light)", + "description": "Small vision model — sees attached images while fitting in ~5 GB VRAM. Pick this over the 8B on smaller GPUs.", + "hf_repo": "unsloth/Qwen3-VL-4B-Instruct-GGUF", + "hf_filename": "Qwen3-VL-4B-Instruct-Q4_K_M.gguf", + "hf_mmproj_filename": "mmproj-F16.gguf", + "mmproj_size_bytes": 836180640, + "size_bytes": 2497282336, + "quant": "Q4_K_M", + "ctx": 16384, + "vram_estimate_mb": 4600, + "ngl_suggestion": 99, + "tags": [ + "vision", + "4b", + "fast" + ] + }, + { + "id": "qwen3-vl-8b", + "name": "Qwen 3 VL 8B (vision)", + "description": "Sees images — the agent can look at attached pictures before picking a workflow. Needs ~8 GB VRAM.", + "hf_repo": "unsloth/Qwen3-VL-8B-Instruct-GGUF", + "hf_filename": "Qwen3-VL-8B-Instruct-Q4_K_M.gguf", + "hf_mmproj_filename": "mmproj-F16.gguf", + "mmproj_size_bytes": 1159030336, + "size_bytes": 5027785568, + "quant": "Q4_K_M", + "ctx": 16384, + "vram_estimate_mb": 7800, + "ngl_suggestion": 99, + "tags": [ + "vision", + "8b" + ] + }, + { + "id": "cadquery-coder-7b", + "name": "CadQuery Coder 7B (CAD-specialized)", + "description": "Qwen2.5-Coder-7B fine-tuned on CadQuery generation — best pick for the Text to CAD node. Community model, CC-BY-NC-SA license.", + "hf_repo": "yuvit-batra/qwen2.5-coder-7b-cadquery-gguf", + "hf_filename": "qwen2.5-coder-7b-cadquery-Q4_K_M.gguf", + "size_bytes": 4680000000, + "quant": "Q4_K_M", + "ctx": 16384, + "vram_estimate_mb": 6200, + "ngl_suggestion": 99, + "tags": [ + "code", + "cad", + "7b" + ] + } +] \ No newline at end of file diff --git a/api/routers/agent.py b/api/routers/agent.py index 3eeb2a73..1fb389eb 100644 --- a/api/routers/agent.py +++ b/api/routers/agent.py @@ -1,11 +1,19 @@ """ -Agent chat endpoint — runs an Ollama-powered tool-use loop against Modly's API. +Agent chat endpoint — runs a tool-use loop against Modly's API, on the managed +local llama.cpp server or any OpenAI-compatible provider. """ +import asyncio +import json import re import uuid +from typing import Optional + import httpx from fastapi import APIRouter -from pydantic import BaseModel +from pydantic import BaseModel, Field + +from services import llm_server +from services.llm_server import llama_pool router = APIRouter(prefix="/agent", tags=["agent"]) @@ -267,7 +275,8 @@ async def execute_tool(name: str, arguments: dict, context: dict) -> tuple[str, return f"Available models:\n{lines}", None elif name == "unload_models": - await client.post(f"{MODLY_API}/model/unload-all") + r = await client.post(f"{MODLY_API}/model/unload-all") + r.raise_for_status() return "All 3D generation models have been unloaded from VRAM.", None elif name == "get_mesh_info": @@ -373,13 +382,19 @@ async def execute_tool(name: str, arguments: dict, context: dict) -> tuple[str, class ChatMessage(BaseModel): role: str content: str - images: list[str] = [] + images: list[str] = [] # data URLs + + +class ProviderConfig(BaseModel): + type: str = "local" # "local" | "external" + base_url: Optional[str] = None # external only, e.g. https://api.openai.com/v1 + api_key: Optional[str] = None class AgentChatRequest(BaseModel): messages: list[ChatMessage] - ollama_url: str = "http://localhost:11434" - model: str = "qwen2.5:3b" + model: str = Field(default_factory=llm_server.default_model_id) # local: catalog/custom id — external: provider model name + provider: ProviderConfig = ProviderConfig() context: dict = {} thinking: str = "auto" # "auto" | "on" | "off" @@ -397,9 +412,10 @@ class AgentChatResponse(BaseModel): def _extract_thinking(msg: dict) -> tuple[str, str | None]: - """Return (clean_content, thinking_text). Handles both Ollama native field and tags.""" - content = msg.get("content", "") - thinking = msg.get("thinking") or None + """Return (clean_content, thinking_text). Handles llama-server's + reasoning_content field and inline tags.""" + content = msg.get("content") or "" + thinking = msg.get("reasoning_content") or None if not thinking: match = re.search(r"(.*?)", content, re.DOTALL) if match: @@ -408,18 +424,100 @@ def _extract_thinking(msg: dict) -> tuple[str, str | None]: return content, thinking -@router.get("/models") -async def list_ollama_models(ollama_url: str = "http://localhost:11434"): - async with httpx.AsyncClient(timeout=5.0) as client: +def _auth_headers(base_url: str, api_key: Optional[str]) -> dict: + headers: dict[str, str] = {} + if api_key: + headers["Authorization"] = f"Bearer {api_key}" + if "anthropic" in base_url: + # Anthropic's OpenAI-compat layer also accepts the native headers. + headers["x-api-key"] = api_key + headers["anthropic-version"] = "2023-06-01" + return headers + + +class ExternalModelsRequest(BaseModel): + base_url: str + api_key: str = "" + + +@router.post("/external/models") +async def list_external_models(req: ExternalModelsRequest): + """Proxy the provider's /models listing (avoids CORS issues from the renderer). + + POST with the key in the body, never a GET query string: uvicorn's access log + records the full request line, and that log ends up in runtime.log. + """ + headers = _auth_headers(req.base_url, req.api_key) + async with httpx.AsyncClient(timeout=10.0) as client: try: - r = await client.get(f"{ollama_url}/api/tags") + r = await client.get(f"{req.base_url.rstrip('/')}/models", headers=headers) r.raise_for_status() - models = [m["name"] for m in r.json().get("models", [])] - return {"models": models} + data = r.json().get("data", []) + return {"models": sorted(m["id"] for m in data if isinstance(m, dict) and m.get("id"))} except Exception: return {"models": []} +async def _unload_llm_after_workflow(request: AgentChatRequest, actions_done: list[ActionDone]) -> None: + """Free the local LLM's VRAM once a workflow has been dispatched, so the + workflow gets the full GPU. Best-effort — never fail the chat over it.""" + if request.provider.type != "local": + return + if not any(a.tool == "run_workflow" for a in actions_done): + return + try: + await asyncio.to_thread(llama_pool.unload_all) + except Exception: + pass + + +def _error_detail(r: httpx.Response) -> str: + """The provider's own message (OpenAI-style `{"error": {"message": …}}`), else the raw body.""" + try: + err = r.json().get("error") + if isinstance(err, dict) and err.get("message"): + return str(err["message"])[:300] + if isinstance(err, str): + return err[:300] + except (ValueError, AttributeError): + pass + return r.text[:300] + + +def _tool_arguments(raw) -> dict: + """OpenAI-compatible servers send tool arguments as a JSON string.""" + if isinstance(raw, dict): + return raw + try: + parsed = json.loads(raw or "{}") + except (TypeError, ValueError): + return {} + return parsed if isinstance(parsed, dict) else {} + + +def _user_entry(m: ChatMessage, send_images: bool) -> dict: + if not (m.images and send_images): + return {"role": m.role, "content": m.content} + parts: list[dict] = [{"type": "text", "text": m.content}] + for data_url in m.images: + parts.append({"type": "image_url", "image_url": {"url": data_url}}) + return {"role": m.role, "content": parts} + + +def _drop_images(messages: list[dict]) -> bool: + """Replace image parts with a text note, in place. True if any were dropped.""" + dropped = False + for m in messages: + if not isinstance(m.get("content"), list): + continue + texts = [p["text"] for p in m["content"] if p.get("type") == "text"] + if any(p.get("type") == "image_url" for p in m["content"]): + dropped = True + texts.append("[An image was attached, but this model reads text only.]") + m["content"] = "\n".join(texts) + return dropped + + @router.post("/chat", response_model=AgentChatResponse) async def agent_chat(request: AgentChatRequest): messages: list[dict] = [{"role": "system", "content": SYSTEM_PROMPT}] @@ -451,67 +549,93 @@ async def agent_chat(request: AgentChatRequest): ), }) + # ONE system message, first. Qwen3.5's chat template raises "System message + # must be at the beginning", which llama-server returns as a flat HTTP 400; + # other templates silently drop the extra ones. + messages = [{"role": "system", "content": "\n\n".join(m["content"] for m in messages)}] + + slot = None + send_images = True + extra: dict = {} + if request.provider.type == "local": + try: + spec = llm_server.resolve_model(request.model) + # hold=True: the slot stays claimed for the whole loop, so neither the + # idle reaper nor a model loading elsewhere can evict it mid-answer. + slot = await asyncio.to_thread(llama_pool.ensure, request.model, spec, True) + except Exception as e: + return AgentChatResponse(message=f"Could not start the local LLM: {e}") + base_url = slot.base_url + headers: dict = {} + send_images = spec["vision"] + extra.update(llm_server.sampling_for(request.model)) + if request.thinking == "off": + extra["chat_template_kwargs"] = {"enable_thinking": False} + else: + if not request.provider.base_url: + return AgentChatResponse(message="No provider URL configured. Set one in Settings → Agent.") + base_url = request.provider.base_url.rstrip("/") + headers = _auth_headers(base_url, request.provider.api_key) + for m in request.messages: - entry: dict = {"role": m.role, "content": m.content} - if m.images: - entry["images"] = m.images - messages.append(entry) + messages.append(_user_entry(m, send_images)) actions_done: list[ActionDone] = [] all_thinking: list[str] = [] + final: AgentChatResponse | None = None - # Build Ollama think param - ollama_extra: dict = {} - if request.thinking == "on": - ollama_extra["think"] = True - elif request.thinking == "off": - ollama_extra["think"] = False - - async with httpx.AsyncClient(timeout=120.0) as client: - for _ in range(10): # max tool-call rounds - r = await client.post( - f"{request.ollama_url}/api/chat", - json={"model": request.model, "messages": messages, "tools": TOOLS, "stream": False, **ollama_extra}, - ) - - if r.status_code != 200: - return AgentChatResponse( - message=f"Ollama error ({r.status_code}). Is Ollama running at {request.ollama_url}?", - ) - - msg = r.json()["message"] - messages.append(msg) - - clean_content, thinking_text = _extract_thinking(msg) - if thinking_text: - all_thinking.append(thinking_text) - - tool_calls = msg.get("tool_calls") or [] - if not tool_calls: - combined_thinking = "\n\n---\n\n".join(all_thinking) if all_thinking else None - return AgentChatResponse( - message=clean_content, - actions=actions_done, - thinking=combined_thinking, - ) - - for tc in tool_calls: - fn = tc["function"] - result_text, payload = await execute_tool(fn["name"], fn.get("arguments") or {}, request.context) - actions_done.append(ActionDone(tool=fn["name"], result=result_text, payload=payload)) - messages.append({"role": "tool", "content": result_text}) - - has_workflow = any(a.tool == "run_workflow" for a in actions_done) - if has_workflow: - # Unload LLM from VRAM immediately so the workflow has full GPU memory - try: - await client.post( - f"{request.ollama_url}/api/generate", - json={"model": request.model, "keep_alive": 0}, - timeout=5.0, + try: + async with httpx.AsyncClient(timeout=120.0) as client: + for _ in range(10): # max tool-call rounds + r = await client.post( + f"{base_url}/chat/completions", + headers=headers, + json={"model": request.model, "messages": messages, "tools": TOOLS, "stream": False, **extra}, ) - except Exception: - pass + # An external model's vision support is unknown up front, and a + # text-only one rejects image parts with a 400 that failed the + # whole turn. Retry once with the images described as omitted. + if r.status_code == 400 and request.provider.type != "local" and _drop_images(messages): + r = await client.post( + f"{base_url}/chat/completions", + headers=headers, + json={"model": request.model, "messages": messages, "tools": TOOLS, "stream": False, **extra}, + ) - combined_thinking = "\n\n---\n\n".join(all_thinking) if all_thinking else None - return AgentChatResponse(message="Reached maximum tool iterations.", actions=actions_done, thinking=combined_thinking) + if r.status_code != 200: + return AgentChatResponse(message=f"LLM error ({r.status_code}): {_error_detail(r)}") + + msg = r.json()["choices"][0]["message"] + tool_calls = msg.get("tool_calls") or [] + assistant: dict = {"role": "assistant", "content": msg.get("content") or ""} + if tool_calls: + assistant["tool_calls"] = tool_calls + messages.append(assistant) + + clean_content, thinking_text = _extract_thinking(msg) + if thinking_text: + all_thinking.append(thinking_text) + + if not tool_calls: + final = AgentChatResponse(message=clean_content, actions=actions_done) + break + + for tc in tool_calls: + fn = tc["function"] + result_text, payload = await execute_tool(fn["name"], _tool_arguments(fn.get("arguments")), request.context) + actions_done.append(ActionDone(tool=fn["name"], result=result_text, payload=payload)) + messages.append({"role": "tool", "tool_call_id": tc.get("id", ""), "content": result_text}) + except httpx.HTTPError as e: + return AgentChatResponse(message=f"Could not reach the LLM at {base_url}: {e}", actions=actions_done) + finally: + if slot is not None: + slot.release() + + # After the release: unload_all() spares a held slot, so unloading while + # still holding ours would keep the agent's own model on the GPU. + await _unload_llm_after_workflow(request, actions_done) + + if final is None: + final = AgentChatResponse(message="Reached maximum tool iterations.", actions=actions_done) + final.thinking = "\n\n---\n\n".join(all_thinking) if all_thinking else None + return final diff --git a/api/routers/export.py b/api/routers/export.py index 2a2f2bf3..854d8bf2 100644 --- a/api/routers/export.py +++ b/api/routers/export.py @@ -1,23 +1,190 @@ +import base64 +import binascii import io +import math +from pathlib import Path import trimesh from fastapi import APIRouter, HTTPException from fastapi.responses import Response, FileResponse -from services.generator_registry import WORKSPACE_DIR +from services import imported_sources +# Import the module (not the name) so WORKSPACE_DIR is read at call time: the +# settings endpoint rebinds it when the user moves the workspace. +import services.generator_registry as registry router = APIRouter(tags=["export"]) SUPPORTED = {"glb", "stl", "obj", "ply"} +# Formats OrcaSlicer's importer accepts (see the orcaslicer://open contract). +# GLB is deliberately excluded — OrcaSlicer cannot import glTF/GLB, so a .glb +# deeplink downloads but silently fails to slice. +SLICER_FORMATS = {"stl", "obj"} +SLICER_MEDIA_TYPES = {"stl": "model/stl", "obj": "text/plain"} + +# Image-to-3D output has no inherent physical scale (a single photo carries no +# real-world size), and AI generators emit roughly unit-sized meshes — which +# import into a slicer as an invisible ~1 mm speck. Normalise the longest +# bounding-box edge to a sane, obviously-printable default; the user rescales +# in OrcaSlicer as needed. +DEFAULT_PRINT_LONGEST_MM = 50.0 + +# ...but this route also serves meshes the user authored or imported, which DO +# carry a real-world size. Silently resizing a 180 mm part down to 50 mm wastes +# a print, so only rescale what is small enough to be unit-sized AI output. +UNIT_SCALE_MAX = 5.0 + +# Source formats whose up-axis is Y (the glTF convention). Everything else this +# route accepts — STL, OBJ, PLY — is conventionally Z-up already. +GLTF_SUFFIXES = {".glb", ".gltf"} + + +def _to_single_mesh(loaded: object) -> "trimesh.Trimesh": + """Flatten a loaded GLB into one Trimesh, baking scene-graph node transforms. + + ``trimesh.util.concatenate(scene.geometry.values())`` would DROP the node + transforms and misassemble a multi-node scene, so flatten at the scene level + where the graph transforms are applied. + """ + if isinstance(loaded, trimesh.Trimesh): + return loaded + if isinstance(loaded, trimesh.Scene): + if len(loaded.geometry) == 0: + raise HTTPException(422, "Mesh contains no geometry") + # Bake the scene-graph node transforms into a single mesh. The spelling + # varies across trimesh versions — to_mesh()/to_geometry() are the modern + # APIs (4.6+); dump(concatenate=True) is the pre-removal fallback for 4.5. + for flatten in (lambda s: s.to_mesh(), lambda s: s.to_geometry(), lambda s: s.dump(concatenate=True)): + try: + result = flatten(loaded) + except (AttributeError, TypeError): + continue + if isinstance(result, trimesh.Trimesh): + return result + if isinstance(result, (list, tuple)) and result: + return trimesh.util.concatenate(result) + # Fallback: concatenate the geometry as-is (may ignore node transforms). + return trimesh.util.concatenate(list(loaded.geometry.values())) + raise HTTPException(422, "Unsupported mesh contents") + + +def _scale_to_print_size(mesh: "trimesh.Trimesh", longest_mm: float = DEFAULT_PRINT_LONGEST_MM) -> None: + """Uniformly scale ``mesh`` in place so its longest bbox edge is ``longest_mm``.""" + extents = mesh.extents + longest = float(max(extents)) if extents is not None and len(extents) else 0.0 + if longest > 1e-9 and math.isfinite(longest): + mesh.apply_scale(longest_mm / longest) + + +def _normalize_print_scale(mesh: "trimesh.Trimesh") -> bool: + """Rescale ``mesh`` only if it looks unit-sized; return whether it was rescaled. + + A mesh whose longest edge already exceeds ``UNIT_SCALE_MAX`` is assumed to + carry a real-world size the user chose, and is left untouched. + """ + extents = mesh.extents + longest = float(max(extents)) if extents is not None and len(extents) else 0.0 + if not math.isfinite(longest) or not (1e-9 < longest <= UNIT_SCALE_MAX): + return False + _scale_to_print_size(mesh) + return True + + +def _resolve_slicer_source(token: str) -> tuple[Path, str]: + """Decode ``token`` into an existing source file and its ORIGINAL suffix. + + Two kinds of source are accepted: + + * a workspace-relative path, confined to the workspace by ancestry; + * an absolute path, but ONLY when the user imported that exact file this + session (see ``services.imported_sources``). Membership there is an + equality test on the resolved path, so this grants no traversal and does + not widen the route to arbitrary disk paths. + + The returned suffix is the format the user actually supplied — for an import + that is the pre-conversion extension, which is what decides the up-axis. + """ + try: + padded = token + "=" * (-len(token) % 4) + decoded = base64.urlsafe_b64decode(padded.encode("ascii")).decode("utf-8") + except (binascii.Error, UnicodeDecodeError, ValueError): + raise HTTPException(400, "Malformed source token") + + candidate = Path(decoded) + if candidate.is_absolute(): + original_suffix = imported_sources.source_suffix(candidate) + if original_suffix is None: + raise HTTPException(400, "Invalid path") + full_path = candidate.resolve() + if not full_path.is_file(): + raise HTTPException(404, f"File not found: {decoded}") + return full_path, original_suffix + + # Containment check via ancestry, not string prefix: `startswith` would let a + # sibling like `-other/...` slip through, and `..` escapes resolve + # outside the workspace and fail this check. + workspace = registry.WORKSPACE_DIR.resolve() + full_path = (workspace / decoded).resolve() + if full_path != workspace and workspace not in full_path.parents: + raise HTTPException(400, "Invalid path") + if not full_path.is_file(): + raise HTTPException(404, f"File not found: {decoded}") + return full_path, full_path.suffix.lower() + + +@router.get("/slicer/{fmt}/{token}/{filename}") +def export_for_slicer(fmt: str, token: str, filename: str): + """Serve a generated GLB converted to a slicer-importable mesh, at a URL + shaped for OrcaSlicer's ``orcaslicer://open?file=`` deeplink. + + The URL is intentionally path-only and ends in the real filename+extension + (e.g. ``/export/slicer/stl//model.stl``). OrcaSlicer + downloads the URL and derives the import filename — and therefore the mesh + format — from the URL's FINAL path segment, so a query string (``?path=...``) + would corrupt the parsed extension and the model would silently fail to + import. ``token`` is the url-safe-base64 of the source path — workspace- + relative, or absolute for a file the user imported this session (see + ``_resolve_slicer_source``); ``filename`` (e.g. ``model.stl``) is what + OrcaSlicer names the download. + """ + fmt = fmt.lower() + if fmt not in SLICER_FORMATS: + raise HTTPException(400, f"Unsupported slicer format: {fmt}. Supported: {', '.join(sorted(SLICER_FORMATS))}") + if not filename.lower().endswith(f".{fmt}"): + raise HTTPException(400, "Filename must end with the requested format extension") + + full_path, source_suffix = _resolve_slicer_source(token) + + mesh = _to_single_mesh(trimesh.load(str(full_path))) + # glTF/GLB is Y-up; OrcaSlicer's world is Z-up. Rotate +90° about X so the + # model imports standing upright instead of on its side. (Modly's own viewer + # rests generated meshes on the Y=0 plane, confirming Y is the up axis.) + # STL/OBJ/PLY sources are already Z-up, so rotating them would do the very + # thing this corrects — lay an upright model on its side. + if source_suffix in GLTF_SUFFIXES: + mesh.apply_transform(trimesh.transformations.rotation_matrix(math.pi / 2, [1, 0, 0])) + _normalize_print_scale(mesh) + + data = mesh.export(file_type=fmt) + if isinstance(data, str): + data = data.encode("utf-8") + return Response( + content=data, + media_type=SLICER_MEDIA_TYPES.get(fmt, "application/octet-stream"), + # Fixed name (not the client-supplied segment) — keeps arbitrary input out + # of the response header. OrcaSlicer names the file from the URL anyway. + headers={"Content-Disposition": f'attachment; filename="model.{fmt}"'}, + ) + @router.get("/{fmt}") def export_mesh(fmt: str, path: str): if fmt not in SUPPORTED: raise HTTPException(400, f"Unsupported format: {fmt}. Supported: {', '.join(SUPPORTED)}") - full_path = (WORKSPACE_DIR / path).resolve() - if not str(full_path).startswith(str(WORKSPACE_DIR.resolve())): + full_path = (registry.WORKSPACE_DIR / path).resolve() + if not registry.is_within_workspace(full_path): raise HTTPException(400, "Invalid path") if not full_path.exists(): raise HTTPException(404, f"File not found: {path}") diff --git a/api/routers/extensions.py b/api/routers/extensions.py index d826a923..fdfbf0e0 100644 --- a/api/routers/extensions.py +++ b/api/routers/extensions.py @@ -1,6 +1,11 @@ import asyncio +import json +import os +import platform as platform_module +import re import subprocess import sys +from pathlib import Path from fastapi import APIRouter, Body, HTTPException router = APIRouter(tags=["extensions"]) @@ -18,7 +23,8 @@ async def reload_extensions(payload: dict | None = Body(default=None)): candidate = payload.get("validationCapability") if isinstance(candidate, dict): validation_capability = candidate - generator_registry.reload(validation_capability) + # Off the event loop: reload waits for any in-progress model load. + await asyncio.to_thread(generator_registry.reload, validation_capability) return { "reloaded": True, "models": list(generator_registry._generators.keys()), @@ -47,17 +53,45 @@ async def setup_extension(ext_id: str): # No setup.py → legacy extension, nothing to do return {"status": "skipped", "reason": "no setup.py"} - # Detect GPU compute capability - gpu_sm = _detect_gpu_sm() + # Detect GPU compute capability. NVIDIA keeps detection priority, exactly + # like electron/main/gpu-detect.ts: a Ryzen APU beside an NVIDIA dGPU must + # resolve to CUDA on both code paths. + gpu_sm, cuda_version = _detect_nvidia_gpu() + gfx_target = "" if gpu_sm else _detect_gfx_target() + flavor = "cuda" if gpu_sm else ("rocm" if gfx_target else "cpu") + + # Pass arguments as JSON so setup.py sees torch_flavor. The keys mirror + # runExtensionSetup in electron/main/ipc-handlers.ts exactly — setup.py + # scripts read the same contract whichever side launched them. Note this + # endpoint is a fallback: Electron normally runs setup.py itself, and only + # that path gets the ROCm index rewriting for extensions that ignore + # torch_flavor. + args = json.dumps({ + "python_exe": sys.executable, + "ext_dir": str(ext_dir), + "gpu_sm": gpu_sm, + "cuda_version": cuda_version, + "accelerator": flavor, + "torch_flavor": flavor, + "gfx_target": gfx_target, + "torch_index_url": _rocm_index_url() if flavor == "rocm" else "", + "platform": sys.platform, + "arch": _node_arch(), + }) # Run setup.py using Modly's embedded Python (sys.executable) loop = asyncio.get_running_loop() result = await loop.run_in_executor( None, lambda: subprocess.run( - [sys.executable, str(setup_py), sys.executable, str(ext_dir), str(gpu_sm)], + [sys.executable, str(setup_py), args], capture_output=True, text=True, + # pip's own output carries box-drawing characters; the locale codec + # (cp1252 on Windows) raises on some of them and would turn a + # successful setup into a 500 with no usable message. + encoding="utf-8", + errors="replace", ) ) @@ -65,9 +99,10 @@ async def setup_extension(ext_id: str): raise HTTPException(500, f"setup.py failed:\n{result.stderr}") return { - "status": "ok", - "gpu_sm": gpu_sm, - "output": result.stdout, + "status": "ok", + "gpu_sm": gpu_sm, + "gfx_target": gfx_target, + "output": result.stdout, } @@ -78,13 +113,120 @@ async def extension_errors(): return generator_registry.load_errors() -def _detect_gpu_sm() -> int: - """Returns GPU compute capability as integer (e.g. 86 for SM 8.6), or 0 if no GPU.""" +def _detect_nvidia_gpu() -> tuple[int, int]: + """ + Returns (compute capability, max CUDA version) — e.g. (86, 124) — or (0, 0) + when there is no NVIDIA GPU. + + Asks nvidia-smi rather than torch: this process runs in Modly's main venv, + which has no torch at all (see api/requirements.txt), so a torch import + always failed here and silently reported every machine as CPU-only. Asking + the driver also sidesteps the ROCm ambiguity — PyTorch's HIP build answers + the whole torch.cuda API, reporting (12, 0) for a gfx1200 Radeon exactly + like an sm_120 Blackwell. Mirrors parseNvidiaSmi in + electron/main/gpu-detect.ts, including the driver → CUDA version table. + """ + try: + result = subprocess.run( + ["nvidia-smi", "--query-gpu=compute_cap,driver_version", "--format=csv,noheader"], + capture_output=True, + text=True, + timeout=15, + ) + except (OSError, subprocess.SubprocessError): + return 0, 0 + if result.returncode != 0: + return 0, 0 + + line = result.stdout.strip().splitlines()[0].strip() if result.stdout.strip() else "" + if not line: + return 0, 0 + parts = [part.strip() for part in line.split(",")] + + try: + sm = round(float(parts[0]) * 10) + except (ValueError, IndexError): + sm = 86 + try: + driver_major = int((parts[1] if len(parts) > 1 else "0").split(".")[0]) + except ValueError: + driver_major = 0 + + cuda_version = 118 # safe minimum + for threshold, version in ( + (570, 128), (560, 126), (555, 125), (550, 124), + (545, 123), (535, 122), (530, 121), (525, 120), (520, 118), + ): + if driver_major >= threshold: + cuda_version = version + break + return sm, cuda_version + + +def _rocm_index_url() -> str: + """The pip index a ROCm torch install must come from. Mirrors + resolveRocmTorchSpec in electron/main/gpu-detect.ts.""" + override = os.environ.get("MODLY_ROCM_INDEX", "").strip() + if override: + return override + if sys.platform == "win32": + return "https://repo.amd.com/rocm/whl-multi-arch/" + return "https://download.pytorch.org/whl/rocm7.2" + + +def _node_arch() -> str: + """platform.machine() mapped onto Node's process.arch vocabulary, so + setup.py sees the same values whichever side launched it.""" + machine = platform_module.machine().lower() + if machine in ("x86_64", "amd64"): + return "x64" + if machine in ("aarch64", "arm64"): + return "arm64" + return machine + + +def _detect_gfx_target() -> str: + """ + Returns the ROCm compute target (e.g. "gfx1200"), or "" when there is no AMD GPU. + + Reads the kernel's KFD topology rather than asking torch: this process runs + in Modly's main venv, which has no torch at all (see api/requirements.txt). + The amdgpu driver publishes the target on its own, so no ROCm install is + needed either. Mirrors electron/main/gpu-detect.ts. + """ + kfd_nodes = Path("/sys/class/kfd/kfd/topology/nodes") + if not Path("/dev/kfd").exists() or not kfd_nodes.is_dir(): + return "" + + def _prop(text: str, key: str) -> int: + match = re.search(rf"^{key}\s+(\d+)\s*$", text, re.M) + return int(match.group(1)) if match else 0 + try: - import torch - if torch.cuda.is_available(): - major, minor = torch.cuda.get_device_capability(0) - return major * 10 + minor - except Exception: - pass - return 0 + nodes = sorted(kfd_nodes.iterdir(), key=lambda p: int(p.name) if p.name.isdigit() else 0) + except OSError: + return "" + + # Among GPU nodes the largest simd_count wins, mirroring parseKfdGfxTarget + # in electron/main/gpu-detect.ts: on an APU + dGPU machine the APU commonly + # gets the lower node number, and the discrete card has more SIMDs. + best_target, best_simd = "", 0 + for node in nodes: + try: + text = (node / "properties").read_text() + except OSError: + continue + # Node 0 is the CPU node (simd_count 0) and carries no compute target. + simd_count = _prop(text, "simd_count") + if simd_count <= 0: + continue + # major*10000 + minor*100 + step, minor and step read as hex digits. + version = _prop(text, "gfx_target_version") + if version <= 0: + continue + major, minor, step = version // 10000, (version % 10000) // 100, version % 100 + if major <= 0 or minor > 15 or step > 15: + continue + if simd_count > best_simd: + best_target, best_simd = f"gfx{major}{minor:x}{step:x}", simd_count + return best_target diff --git a/api/routers/generation.py b/api/routers/generation.py index ea8b25e0..73d0baae 100644 --- a/api/routers/generation.py +++ b/api/routers/generation.py @@ -4,13 +4,24 @@ import time import traceback import uuid -from typing import Dict +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Dict, Optional, Union from fastapi import APIRouter, File, Form, UploadFile, HTTPException, BackgroundTasks from services.generators.base import smooth_progress, GenerationCancelled import re as _re -from services.generator_registry import generator_registry, WORKSPACE_DIR -from schemas.generation import JobStatus +# Import the module (not the name) so WORKSPACE_DIR is read at call time: the +# settings endpoint rebinds it when the user relocates the workspace, and a +# binding captured at import would keep writing output to the old directory. +import services.generator_registry as registry +from services.generator_registry import generator_registry +from schemas.generation import GenerateFromArtifactRequest, JobStatus +from services.artifact_input import ( + RESERVED_ARTIFACT_PARAMS, + TypedArtifactInput, + validate_artifact_input, +) router = APIRouter(tags=["generation"]) @@ -22,6 +33,14 @@ _cancelled: set = set() _cancel_events: Dict[str, threading.Event] = {} _completed_at: Dict[str, float] = {} +_job_generators: Dict[str, object] = {} +# A pinned generation owns the complete switch/load/generate lifecycle. Keeping +# that lifecycle on one dedicated worker provides process-wide serialization +# without parking default-executor workers on a blocking lock acquisition. +_pinned_generation_executor = ThreadPoolExecutor( + max_workers=1, + thread_name_prefix="modly-pinned-generation", +) _JOB_TTL = 1800 # purge terminal jobs after 30 minutes @@ -33,6 +52,7 @@ def _purge_old_jobs() -> None: _jobs.pop(jid, None) _cancelled.discard(jid) _cancel_events.pop(jid, None) + _job_generators.pop(jid, None) _completed_at.pop(jid, None) @@ -70,7 +90,9 @@ def sanitize_collection(collection: str) -> str: return "Default" try: - (WORKSPACE_DIR / collection).resolve().relative_to(WORKSPACE_DIR.resolve()) + (registry.WORKSPACE_DIR / collection).resolve().relative_to( + registry.WORKSPACE_DIR.resolve() + ) except (OSError, ValueError): return "Default" @@ -99,11 +121,10 @@ async def generate_from_image( # Verify the requested model exists in the registry try: generator_registry.get_generator(model_id) + output_kind = generator_registry.get_manifest(model_id).get("output", "mesh") except ValueError as e: raise HTTPException(400, str(e)) - generator_registry.switch_model(model_id) - # Parse model-specific params from JSON and merge with common fields try: model_params = json.loads(params) @@ -125,8 +146,41 @@ async def generate_from_image( _jobs[job_id] = job _cancel_events[job_id] = threading.Event() - background_tasks.add_task(_run_generation, job_id, image_bytes, full_params, collection) + background_tasks.add_task( + _run_generation, job_id, image_bytes, full_params, collection, output_kind, model_id + ) + + return {"job_id": job_id} + +@router.post("/from-artifact") +async def generate_from_artifact( + request: GenerateFromArtifactRequest, + background_tasks: BackgroundTasks, +): + """Queue a validated typed artifact without serializing it as image bytes.""" + try: + manifest = generator_registry.get_manifest(request.model_id) + except (KeyError, ValueError) as exc: + raise HTTPException(400, str(exc)) from exc + declared = manifest.get("inputs") or [manifest.get("input", "image")] + if request.input_kind not in declared: + raise HTTPException(400, f"Model {request.model_id} does not accept {request.input_kind} input") + try: + artifact = validate_artifact_input(registry.WORKSPACE_DIR, request.input_kind, request.input_path) + except ValueError as exc: + raise HTTPException(400, str(exc)) from exc + + params = {k: v for k, v in request.params.items() if k not in RESERVED_ARTIFACT_PARAMS} + params["scene_manifest_path"] = str(artifact.path) + collection = sanitize_collection(request.collection) + job_id = str(uuid.uuid4()) + _purge_old_jobs() + _jobs[job_id] = JobStatus(job_id=job_id, status="pending", progress=0) + _cancel_events[job_id] = threading.Event() + background_tasks.add_task( + _run_generation, job_id, artifact, params, collection, manifest.get("output", "mesh"), request.model_id + ) return {"job_id": job_id} @@ -150,11 +204,10 @@ async def cancel_job(job_id: str): if job.status in ("pending", "running"): job.status = "cancelled" _completed_at[job_id] = time.monotonic() - # Kill the active generator subprocess immediately so inference stops now. - # _run_generation will catch the resulting exception, see job_id in _cancelled, - # and return cleanly without setting an error status. + # Kill only the subprocess bound to this job. A queued cancellation must not + # terminate whichever earlier job currently owns the active generator. try: - gen = generator_registry._generators.get(generator_registry._active_id) + gen = _job_generators.get(job_id) if gen is not None and hasattr(gen, "_proc") and gen._proc and gen._proc.poll() is None: gen._proc.kill() gen._loaded = False @@ -164,9 +217,66 @@ async def cancel_job(job_id: str): return {"cancelled": True} -async def _run_generation(job_id: str, image_bytes: bytes, params: dict, collection: str = "Default") -> None: +async def _run_generation( + job_id: str, + model_input: Union[bytes, TypedArtifactInput], + params: dict, + collection: str = "Default", + output_kind: str = "mesh", + model_id: Optional[str] = None, +) -> None: + # Pinned jobs share one model lifecycle. Switching is deliberately deferred + # until this job runs on the dedicated worker: request-time switches can + # otherwise unload a running job or leave A loading beside B. Crucially, + # queued jobs are executor work items rather than default-executor threads + # blocked on a lock, so cancelling a waiter cannot orphan queue ownership or + # starve the worker that performs generation. + loop = asyncio.get_running_loop() + executor = _pinned_generation_executor if model_id is not None else None + # Shown while this job waits behind another one on the single worker; + # _run_generation_impl clears it as soon as the job actually starts. + queued_job = _jobs.get(job_id) + if queued_job is not None and executor is not None: + queued_job.step = "Waiting for the previous generation…" + future = loop.run_in_executor( + executor, + _run_generation_impl, + job_id, + model_input, + params, + collection, + output_kind, + model_id, + ) + try: + await future + except asyncio.CancelledError: + # asyncio cancellation attempts to cancel a queued concurrent future. + # If it has already begun, the event lets the generator stop safely. + _cancelled.add(job_id) + cancel_event = _cancel_events.get(job_id) + if cancel_event is not None: + cancel_event.set() + job = _jobs.get(job_id) + if job is not None and job.status in ("pending", "running"): + job.status = "cancelled" + _completed_at[job_id] = time.monotonic() + raise + + +def _run_generation_impl( + job_id: str, + model_input: Union[bytes, TypedArtifactInput], + params: dict, + collection: str, + output_kind: str, + model_id: Optional[str], +) -> None: + if job_id in _cancelled: + return job = _jobs[job_id] job.status = "running" + job.step = None def progress_cb(pct: int, step: str = "") -> None: # Monotonic: the loading phase walks the bar up on a background thread and @@ -178,13 +288,21 @@ def progress_cb(pct: int, step: str = "") -> None: job.step = step try: - loop = asyncio.get_running_loop() - - # Check if the model needs to be loaded BEFORE calling get_active(), - # because get_active() loads the model in a blocking manner. - # active_status() is an instantaneous operation (simple dict lookup). - if not generator_registry.active_status()["loaded"]: - active = generator_registry.active_status() + # Refuse a selected weight variant that is not installed before loading + # anything. Uses the job's model: the switch to it happens below. + generator_registry.assert_weight_variant_installed(params, model_id) + + # Check if the model needs to be loaded BEFORE calling the generator + # getter, because that call can load the model in a blocking manner. + get_generator = (lambda: generator_registry.activate_ready_generator(model_id)) \ + if model_id is not None else generator_registry.get_active + if model_id is not None: + _job_generators[job_id] = generator_registry.get_generator(model_id) + status_reader = (lambda: generator_registry.model_status(model_id)) \ + if model_id is not None else generator_registry.active_status + status = status_reader() + if not status["loaded"]: + active = status model_name = active['name'] init_label = f"Downloading {model_name}…" if not active['downloaded'] else f"Loading {model_name}…" progress_cb(0, init_label) @@ -196,38 +314,60 @@ def progress_cb(pct: int, step: str = "") -> None: ) load_thread.start() try: - gen = await loop.run_in_executor(None, generator_registry.get_active) + gen = get_generator() finally: stop_load_evt.set() else: - gen = await loop.run_in_executor(None, generator_registry.get_active) + gen = get_generator() + _job_generators[job_id] = gen if job_id in _cancelled: return # Direct output to the collection subfolder - coll_dir = WORKSPACE_DIR / collection + coll_dir = registry.WORKSPACE_DIR / collection coll_dir.mkdir(parents=True, exist_ok=True) gen.outputs_dir = coll_dir cancel_event = _cancel_events.get(job_id) - import inspect - supports_cancel = "cancel_event" in inspect.signature(gen.generate).parameters - output_path = await loop.run_in_executor( - None, - lambda: gen.generate(image_bytes, params, progress_cb, cancel_event) - if supports_cancel - else gen.generate(image_bytes, params, progress_cb), - ) + if isinstance(model_input, TypedArtifactInput): + # Revalidate just before crossing the inference boundary. The + # subprocess runner repeats this check inside the worker. + from services.artifact_input import revalidate_artifact_input + model_input = revalidate_artifact_input(registry.WORKSPACE_DIR, model_input) + import inspect + supports_cancel = "cancel_event" in inspect.signature(gen.generate_artifact).parameters + output_path = ( + gen.generate_artifact(model_input.kind, model_input.path, params, progress_cb, cancel_event) + if supports_cancel + else gen.generate_artifact(model_input.kind, model_input.path, params, progress_cb) + ) + else: + import inspect + supports_cancel = "cancel_event" in inspect.signature(gen.generate).parameters + output_path = ( + gen.generate(model_input, params, progress_cb, cancel_event) + if supports_cancel + else gen.generate(model_input, params, progress_cb) + ) if job_id in _cancelled: return + output_path = Path(output_path).resolve(strict=True) + if output_kind == "scene": + from services.scene_input import validate_scene_input + try: + output_relative = output_path.relative_to(registry.WORKSPACE_DIR.resolve()) + output_path = validate_scene_input(registry.WORKSPACE_DIR, output_relative.as_posix()) + except (OSError, ValueError) as exc: + raise ValueError("Generated scene output is not a valid workspace scene") from exc + job.status = "done" job.progress = 100 _completed_at[job_id] = time.monotonic() try: - rel = output_path.relative_to(WORKSPACE_DIR) + rel = output_path.relative_to(registry.WORKSPACE_DIR) job.output_url = f"/workspace/{rel.as_posix()}" except ValueError: job.output_url = f"/workspace/{collection}/{output_path.name}" @@ -247,3 +387,5 @@ def progress_cb(pct: int, step: str = "") -> None: job.status = "error" job.error = tb.strip() _completed_at[job_id] = time.monotonic() + finally: + _job_generators.pop(job_id, None) diff --git a/api/routers/llm.py b/api/routers/llm.py new file mode 100644 index 00000000..33e1534c --- /dev/null +++ b/api/routers/llm.py @@ -0,0 +1,418 @@ +""" +Local LLM engine endpoints — model catalog, GGUF downloads, llama-server lifecycle. +Reuses the streamed/resumable downloader from routers.model. +""" +import asyncio +import json +import threading +from pathlib import Path +from typing import Optional +import httpx +from fastapi import APIRouter, HTTPException +from fastapi.responses import StreamingResponse +from pydantic import BaseModel + +from routers.model import DownloadCancelled, DownloadPaused, _download_file_streamed +from services import llm_server +from services.llm_server import llama_pool + +router = APIRouter(tags=["llm"]) + +_controls: dict[str, dict[str, threading.Event]] = {} + + +class _DownloadState: + """Tracks one model's in-flight download independent of any single SSE + connection, so closing/reopening the Model Library modal reattaches to the + same download instead of racing a second one against the same .part file.""" + + def __init__(self) -> None: + self.last_msg: dict = {} + self.subscribers: list[asyncio.Queue] = [] + self.task: asyncio.Task | None = None + + +_downloads: dict[str, _DownloadState] = {} + + +def _new_control(key: str) -> dict[str, threading.Event]: + control: dict[str, threading.Event] = {"pause": threading.Event(), "cancel": threading.Event()} + _controls[key] = control + return control + + +def _check_control(control: dict[str, threading.Event]) -> None: + if control["cancel"].is_set(): + raise DownloadCancelled() + if control["pause"].is_set(): + raise DownloadPaused() + + +def _fmt(data: dict) -> str: + return f"data: {json.dumps(data)}\n\n" + + +# ─── Catalog / status ───────────────────────────────────────────────────────── + +@router.get("/models") +async def list_models(tag: Optional[str] = None, downloaded: bool = False): + """Model library. `tag` filters catalog entries by category (e.g. tag=code + for coder models); custom user GGUFs are always included since their + capabilities are unknown. `downloaded=true` keeps only ready-to-use models.""" + models = llm_server.list_models() + if tag: + models = [m for m in models if m.get("source") == "custom" or tag in (m.get("tags") or [])] + if downloaded: + models = [m for m in models if m.get("downloaded")] + return {"models": models} + + +@router.get("/status") +async def status(): + return { + "binary_installed": llm_server.binary_installed(), + "has_nvidia_gpu": llm_server.has_nvidia_gpu(), + "vram_gb": llm_server.detect_vram_gb() or None, + "models_dir": str(llm_server.LLM_MODELS_DIR), + "server": await asyncio.to_thread(llama_pool.snapshot), + } + + +class LlmConfigRequest(BaseModel): + max_models: str | int # "auto" or 1..MAX_SLOT_PORTS + + +@router.get("/config") +async def get_config(): + return { + "max_models": llm_server.load_config().get("max_models", "auto"), + "resolved_max_models": llm_server.resolve_max_models(), + "vram_gb": llm_server.detect_vram_gb() or None, + "max_slot_ports": llm_server.MAX_SLOT_PORTS, + } + + +@router.post("/config") +async def set_config(request: LlmConfigRequest): + value: str | int = request.max_models + if isinstance(value, str) and value.lower() != "auto": + try: + value = int(value) + except ValueError: + raise HTTPException(status_code=422, detail="max_models must be 'auto' or an integer") + if isinstance(value, int) and not (1 <= value <= llm_server.MAX_SLOT_PORTS): + raise HTTPException(status_code=422, detail=f"max_models must be between 1 and {llm_server.MAX_SLOT_PORTS}") + cfg = llm_server.load_config() + cfg["max_models"] = value + llm_server.save_config(cfg) + # A lowered limit applies immediately — small-VRAM users count on it. + await asyncio.to_thread(llama_pool.enforce_limit) + return {"max_models": value, "resolved_max_models": llm_server.resolve_max_models()} + + +@router.post("/unload") +async def unload(): + await asyncio.to_thread(llama_pool.unload_all) + return {"unloaded": True} + + +@router.delete("/models/{model_id}") +async def delete_model(model_id: str): + try: + spec = llm_server.resolve_model(model_id) + except KeyError as e: + raise HTTPException(status_code=404, detail=str(e)) + if await asyncio.to_thread(llama_pool.is_loaded, model_id): + await asyncio.to_thread(llama_pool.unload, model_id) + spec["gguf_path"].unlink(missing_ok=True) + if spec.get("mmproj_path"): + spec["mmproj_path"].unlink(missing_ok=True) + return {"deleted": True} + + +# ─── Chat completion through the managed server ─────────────────────────────── + +class LlmChatRequest(BaseModel): + model: str + messages: list[dict] + temperature: Optional[float] = None + max_tokens: Optional[int] = None + stream: Optional[bool] = False + + +@router.post("/chat") +async def llm_chat(request: LlmChatRequest): + """Load `model` (hot-swapping if needed) and run one OpenAI-format chat completion. + + Used by workflow nodes (LLM, Text to CAD, …) so they share the managed + llama-server instead of loading their own copy of the model. When + `stream` is set, the llama-server SSE chunks are proxied straight through so + the caller can render tokens as they arrive. + """ + try: + spec = llm_server.resolve_model(request.model) + except KeyError as e: + raise HTTPException(status_code=404, detail=str(e)) + try: + # hold=True: the slot comes back claimed, so a model loading in another + # thread cannot evict it between here and the request below. + slot = await asyncio.to_thread(llama_pool.ensure, request.model, spec, True) + except Exception as e: + raise HTTPException(status_code=503, detail=f"Could not start the local LLM: {e}") + + payload: dict = {"model": request.model, "messages": request.messages, "stream": bool(request.stream)} + if request.temperature is not None: + payload["temperature"] = request.temperature + if request.max_tokens is not None: + payload["max_tokens"] = request.max_tokens + + # The claim spans the whole completion, not just its end: a generation longer + # than the idle TTL (a Text-to-CAD node can sit here for minutes) used to be + # reaped mid-answer, since only a *finished* call marked the slot as used. + if request.stream: + async def proxy(): + try: + async with httpx.AsyncClient(timeout=600.0) as client: + async with client.stream("POST", f"{slot.base_url}/chat/completions", json=payload) as r: + if r.status_code != 200: + body = (await r.aread()).decode("utf-8", "replace")[:300] + yield f"data: {json.dumps({'error': f'llama-server error ({r.status_code}): {body}'})}\n\n" + return + async for line in r.aiter_lines(): + if line.startswith("data: "): + yield f"{line}\n\n" + finally: + slot.release() + return StreamingResponse(proxy(), media_type="text/event-stream") + + try: + async with httpx.AsyncClient(timeout=600.0) as client: + r = await client.post(f"{slot.base_url}/chat/completions", json=payload) + if r.status_code != 200: + raise HTTPException(status_code=502, detail=f"llama-server error ({r.status_code}): {r.text[:300]}") + return r.json() + finally: + slot.release() + + +# ─── GGUF download (SSE) ────────────────────────────────────────────────────── + +@router.post("/download/pause") +async def pause_download(model_id: str): + control = _controls.get(model_id) + if control: + control["pause"].set() + return {"paused": True} + + +@router.post("/download/cancel") +async def cancel_download(model_id: str): + """Cancel an in-flight download, or clean up a paused one. + + A paused download has already left `_run_download` — its control is gone + from `_controls`, so setting the cancel event would be a no-op. Reporting + success there left the user with a multi-GB .part file they believed they + had deleted, so the partial files are removed here instead.""" + control = _controls.get(model_id) + if control: + control["cancel"].set() + return {"cancelled": True, "removed_partials": []} + + entry = next((e for e in llm_server.load_catalog() if e["id"] == model_id), None) + removed = _discard_incomplete(entry) if entry else [] + _downloads.pop(model_id, None) + return {"cancelled": bool(removed), "removed": removed} + + +def _model_files(entry: dict) -> list[Path]: + """Every file a model needs on disk: the weights, plus the vision projector.""" + paths = [llm_server.LLM_MODELS_DIR / entry["hf_filename"]] + if entry.get("hf_mmproj_filename"): + paths.append(llm_server.LLM_MODELS_DIR / llm_server.mmproj_local_name(entry)) + return paths + + +def _discard_incomplete(entry: dict) -> list[str]: + """Drop what a cancelled download left behind, and report what went. + + Not just the .part: a vision model downloads weights then projector, so + cancelling during the second one left 2.5 GB of finished weights on disk + under a model still reported as `downloaded: false` — no trash button is + offered for those, so the space could not be reclaimed from the UI at all. + A model whose files are all present is complete, not in flight, and is + never touched here (deleting it is what DELETE /llm/models/{id} is for).""" + paths = _model_files(entry) + if all(p.exists() for p in paths): + return [] + removed = [] + for path in paths: + for target in (path.with_suffix(path.suffix + ".part"), path): + if target.exists(): + target.unlink(missing_ok=True) + removed.append(target.name) + return removed + + +def _is_terminal(msg: dict) -> bool: + return msg.get("status") == "done" or "error" in msg or msg.get("cancelled") or msg.get("paused") + + +async def _run_download(model_id: str, entry: dict, control: dict[str, threading.Event], state: _DownloadState) -> None: + """Owns one model's download end to end, independent of any SSE connection. + Broadcasts progress to every currently-attached watcher (see `_broadcast`).""" + loop = asyncio.get_running_loop() + + def _broadcast(msg: dict) -> None: + # _progress (below) runs in a worker thread — hop back onto the loop. + def _do() -> None: + state.last_msg = msg + for q in list(state.subscribers): + q.put_nowait(msg) + loop.call_soon_threadsafe(_do) + + from huggingface_hub import hf_hub_url + + # Vision models ship a companion mmproj file — download it alongside the + # weights, stored under a per-model local name (HF names collide). + files: list[tuple[str, str, int]] = [(entry["hf_filename"], entry["hf_filename"], entry.get("size_bytes") or 0)] + if entry.get("hf_mmproj_filename"): + files.append(( + entry["hf_mmproj_filename"], + llm_server.mmproj_local_name(entry), + entry.get("mmproj_size_bytes") or 0, + )) + grand_total = sum(size for _, _, size in files) + + try: + _broadcast({"percent": 0, "status": "Starting download…"}) + completed_bytes = 0 + + for index, (hf_filename, local_filename, _size) in enumerate(files): + url = hf_hub_url(repo_id=entry["hf_repo"], filename=hf_filename) + + def _progress(msg: dict, _base: int = completed_bytes) -> None: + done = _base + msg.get("bytesDownloaded", 0) + if grand_total: + pct = min(99, round(done / grand_total * 100)) + # The streaming helper reports per-file figures. A vision + # model downloads two files, so the bar (combined percent) + # and the label baked into `status` (per-file percent) drew + # two different numbers side by side. Restate all of it + # against the combined total; `file`/`fileIndex` stay + # per-file, which is what they're for. + msg["percent"] = pct + msg["bytesDownloaded"] = done + msg["totalBytes"] = grand_total + status = msg.get("status") or "" + if status.endswith("%"): + msg["status"] = f"{status.rsplit(' ', 1)[0]} {pct}%" + _broadcast(msg) + + await loop.run_in_executor( + None, + lambda url=url, filename=local_filename, cb=_progress: _download_file_streamed( + url=url, + filename=filename, + dest_dir=str(llm_server.LLM_MODELS_DIR), + file_index=index + 1, + total_files=len(files), + base_percent=0, + progress_cb=cb, + control=control, + ), + ) + completed_bytes += _size + + _broadcast({"percent": 100, "status": "done"}) + + except DownloadPaused: + # Carry the last percent so the modal's bar stays where it stopped + # instead of snapping back to 0 behind the Resume button. + _broadcast({"paused": True, "status": "paused", "percent": state.last_msg.get("percent", 0)}) + except DownloadCancelled: + _discard_incomplete(entry) + _broadcast({"cancelled": True, "status": "cancelled", "percent": 0}) + except Exception as exc: + _broadcast({"error": str(exc)}) + finally: + if _controls.get(model_id) is control: + _controls.pop(model_id, None) + + +@router.get("/download") +async def download_model(model_id: str): + entry = next((e for e in llm_server.load_catalog() if e["id"] == model_id), None) + if entry is None: + raise HTTPException(status_code=404, detail=f"Unknown catalog model: {model_id}") + + # Reconnect-safe: if this model is already downloading, attach a new watcher + # instead of starting a second download racing the same .part file — this is + # what lets the Model Library modal be closed and reopened mid-download. + state = _downloads.get(model_id) + if state is None or state.task is None or state.task.done(): + control = _new_control(model_id) + state = _DownloadState() + _downloads[model_id] = state + state.task = asyncio.create_task(_run_download(model_id, entry, control, state)) + + queue: asyncio.Queue[dict] = asyncio.Queue() + state.subscribers.append(queue) + if state.last_msg: + queue.put_nowait(state.last_msg) # replay current progress immediately on (re)connect + + async def stream(): + try: + while True: + msg = await queue.get() + yield _fmt(msg) + if _is_terminal(msg): + break + finally: + if queue in state.subscribers: + state.subscribers.remove(queue) + + return StreamingResponse(stream(), media_type="text/event-stream") + + +# ─── Engine (llama-server binary) install (SSE) ─────────────────────────────── + +@router.get("/binary/install") +async def install_binary(): + if llm_server.binary_installed(): + async def already(): + yield _fmt({"percent": 100, "status": "done"}) + return StreamingResponse(already(), media_type="text/event-stream") + + control = _new_control("__binary__") + + async def stream(): + loop = asyncio.get_running_loop() + queue: asyncio.Queue[dict] = asyncio.Queue() + + def _progress(msg: dict) -> None: + loop.call_soon_threadsafe(queue.put_nowait, msg) + + try: + future = loop.run_in_executor( + None, + lambda: llm_server.install_binary(_progress, lambda: _check_control(control)), + ) + while not future.done(): + try: + msg = await asyncio.wait_for(queue.get(), timeout=2.0) + except asyncio.TimeoutError: + continue + else: + yield _fmt(msg) + await future + while not queue.empty(): + yield _fmt(queue.get_nowait()) + except (DownloadPaused, DownloadCancelled): + yield _fmt({"cancelled": True, "status": "cancelled"}) + except Exception as exc: + yield _fmt({"error": str(exc)}) + finally: + if _controls.get("__binary__") is control: + _controls.pop("__binary__", None) + + return StreamingResponse(stream(), media_type="text/event-stream") diff --git a/api/routers/model.py b/api/routers/model.py index 4f04718b..aab4c3f4 100644 --- a/api/routers/model.py +++ b/api/routers/model.py @@ -8,9 +8,18 @@ from typing import Optional from urllib.error import HTTPError, URLError from urllib.request import Request, urlopen -from fastapi import APIRouter, HTTPException +from fastapi import APIRouter, Header, 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, + resolve_source_destination_at_root, + resolve_weight_storage_root, + validate_source_file_plan, +) router = APIRouter(tags=["model"]) @@ -52,20 +61,20 @@ def _check_download_control(control: dict[str, threading.Event]) -> None: @router.get("/status") async def model_status(): """Status of the active model.""" - return generator_registry.active_status() + return await asyncio.to_thread(generator_registry.active_status) @router.get("/all") async def all_models_status(): """Status of all known models (downloaded, loaded, required VRAM).""" - return generator_registry.all_status() + return await asyncio.to_thread(generator_registry.all_status) @router.get("/params") async def model_params(model_id: Optional[str] = None): """Parameter schema of the active model (or a specified model).""" try: - return generator_registry.params_schema(model_id) + return await asyncio.to_thread(generator_registry.params_schema, model_id) except KeyError: raise HTTPException(404, f"Unknown model ID: {model_id}") @@ -74,7 +83,7 @@ async def model_params(model_id: Optional[str] = None): async def switch_model(model_id: str): """Switch the active model.""" try: - generator_registry.switch_model(model_id) + await asyncio.to_thread(generator_registry.switch_model, model_id) return {"active": model_id} except ValueError as e: raise HTTPException(400, str(e)) @@ -83,7 +92,11 @@ async def switch_model(model_id: str): @router.post("/unload-all") async def unload_all_models(): """Unloads all models from memory to free VRAM/RAM.""" - generator_registry.unload_all() + # Off the event loop: unloading waits for any in-progress model load. + await asyncio.to_thread(generator_registry.unload_all) + # The local LLMs hold VRAM too; "Free memory" left them resident. + from services.llm_server import llama_pool + await asyncio.to_thread(llama_pool.unload_all) # Force Python to release memory back to the OS import gc gc.collect() @@ -97,15 +110,24 @@ async def unload_all_models(): return {"unloaded": True} -@router.post("/unload/{model_id}") +@router.post("/unload/{model_id:path}") async def unload_model(model_id: str): """Unloads a model from memory so its files can be safely deleted.""" try: 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") @@ -122,13 +144,167 @@ async def cancel_hf_download(model_id: str): return {"cancelled": True} +@router.post("/hf-download-sources") +async def hf_download_sources(request: FastAPIRequest, model_id: str): + """Download sources into one validated node or extension-shared target.""" + try: + body = await request.json() + if not isinstance(body, dict): + raise ValueError("Request body must be an object") + raw_sources = body.get("sources") + if raw_sources is None: + raise ValueError("sources are required") + sources = normalize_model_sources({"model_sources": raw_sources}) + model_root = resolve_weight_storage_root(registry_module.MODELS_DIR, model_id) + destinations = { + source["id"]: resolve_source_destination_at_root( + model_root, source["destination"] + ) + for source in sources + } + except (TypeError, ValueError) as exc: + raise HTTPException(400, str(exc)) from exc + + authorization = request.headers.get("authorization", "") + hf_token = ( + authorization[7:].strip() + if authorization.lower().startswith("bearer ") + else os.environ.get("HUGGING_FACE_HUB_TOKEN") + or os.environ.get("HF_TOKEN") + or None + ) + control = _new_download_control(model_id) + + async def stream(): + loop = asyncio.get_running_loop() + + def _fmt(data: dict) -> str: + return f"data: {json.dumps(data)}\n\n" + + try: + yield _fmt({"percent": 0, "status": "Listing repository files..."}) + files_by_source: dict[str, list[str]] = {} + + for source in sources: + _check_download_control(control) + + def _list_files(current=source): + from huggingface_hub import list_repo_files + listed = list_repo_files( + current["repo_id"], + revision=current.get("revision"), + token=hf_token, + ) + include = current.get("include_prefixes", []) + skip = current.get("skip_prefixes", []) + return [ + filename for filename in listed + if (not include or any(filename.startswith(prefix) for prefix in include)) + if not any(filename.startswith(prefix) for prefix in skip) + ] + + files = await loop.run_in_executor(None, _list_files) + if not files: + raise RuntimeError( + f'No files found in Hugging Face repo: {source["repo_id"]}' + ) + destination = destinations[source["id"]] + for filename in files: + resolve_download_path(destination, filename) + files_by_source[source["id"]] = files + + validate_source_file_plan(sources, files_by_source) + planned_files = [ + (source, filename) + for source in sources + for filename in files_by_source[source["id"]] + ] + total = len(planned_files) + yield _fmt({"percent": 1, "status": f"Downloading {total} files..."}) + + from huggingface_hub import hf_hub_url + + for index, (source, filename) in enumerate(planned_files): + _check_download_control(control) + display_file = f'{source["id"]}/{filename}' + base_percent = 1 + round(index / total * 94) + yield _fmt({ + "percent": base_percent, + "file": display_file, + "fileIndex": index + 1, + "totalFiles": total, + "status": f"Starting {display_file}", + "bytesDownloaded": 0, + "stalledSeconds": 0, + }) + + queue: asyncio.Queue[dict] = asyncio.Queue() + + def _progress(message: dict) -> None: + message["file"] = display_file + loop.call_soon_threadsafe(queue.put_nowait, message) + + url = hf_hub_url( + repo_id=source["repo_id"], + filename=filename, + revision=source.get("revision"), + ) + future = loop.run_in_executor( + None, + lambda: _download_file_streamed( + url=url, + filename=filename, + dest_dir=str(destinations[source["id"]]), + file_index=index + 1, + total_files=total, + base_percent=base_percent, + progress_cb=_progress, + control=control, + token=hf_token, + ), + ) + while not future.done(): + try: + message = await asyncio.wait_for(queue.get(), timeout=2.0) + except asyncio.TimeoutError: + continue + yield _fmt(message) + + final_size = await future + _check_download_control(control) + yield _fmt({ + "percent": 1 + round((index + 1) / total * 94), + "file": display_file, + "fileIndex": index + 1, + "totalFiles": total, + "status": "Downloaded", + "bytesDownloaded": final_size, + "stalledSeconds": 0, + }) + + yield _fmt({"percent": 100, "status": "done"}) + except DownloadPaused: + yield _fmt({"paused": True, "status": "paused"}) + except DownloadCancelled: + for part in model_root.rglob("*.part"): + part.unlink(missing_ok=True) + yield _fmt({"cancelled": True, "status": "cancelled"}) + except Exception as exc: + yield _fmt({"error": str(exc)}) + finally: + if _download_controls.get(model_id) is control: + _download_controls.pop(model_id, None) + + return StreamingResponse(stream(), media_type="text/event-stream") + + @router.get("/hf-download") async def hf_download( repo_id: str, model_id: str, skip_prefixes: Optional[str] = None, include_prefixes: Optional[str] = None, - token: Optional[str] = None, + x_hf_token: Optional[str] = Header(default=None), ): """ Streams a HuggingFace Hub model download via SSE. @@ -137,14 +313,19 @@ async def hf_download( skip_prefixes: JSON-encoded list of path prefixes to exclude. include_prefixes: JSON-encoded list of path prefixes to include (whitelist). - token: HuggingFace access token for gated repos (from Electron settings). - All three fall back to the extension's manifest / environment when not supplied. + X-HF-Token: HuggingFace access token for gated repos (from Electron settings). + A HEADER, not a query param: uvicorn logs the full request + line to stdout, python-bridge.ts pipes that into runtime.log, + and `log:readAll` hands that file to the user for bug + reports — the token used to travel all the way there. + All fall back to the extension's manifest / environment when not supplied. SSE format: data: {"percent": 0-100, "file": "...", "status": "..."} """ + token = x_hf_token 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: @@ -295,7 +476,12 @@ def _download_file_streamed( final_path.parent.mkdir(parents=True, exist_ok=True) if final_path.exists(): - return final_path.stat().st_size + if not final_path.is_file(): + raise RuntimeError(f"Download target is not a regular file: {filename}") + existing_size = final_path.stat().st_size + if existing_size > 0: + return existing_size + final_path.unlink() # Explicit token (from caller) > env vars > none hf_token = ( diff --git a/api/routers/optimize.py b/api/routers/optimize.py index 6081c704..c325ff20 100644 --- a/api/routers/optimize.py +++ b/api/routers/optimize.py @@ -1,27 +1,28 @@ import hashlib import os -import re -import shutil import tempfile import uuid -try: - import pymeshlab as _pymeshlab - _PYMESHLAB_AVAILABLE = True -except ImportError: - _pymeshlab = None - _PYMESHLAB_AVAILABLE = False - import numpy as np import trimesh -import trimesh.visual from fastapi import APIRouter, HTTPException, UploadFile, File from fastapi.responses import FileResponse, Response from pathlib import Path from urllib.parse import quote -from pydantic import BaseModel - -from services.generator_registry import WORKSPACE_DIR +from pydantic import BaseModel, Field + +from services import imported_sources +# Import the module (not the name) so WORKSPACE_DIR is read at call time: the +# settings endpoint rebinds it when the user moves the workspace. +import services.generator_registry as registry +from services.mesh_ops import ( + MeshOpContext, + MeshOpExecutionError, + MeshOpNotFoundError, + MeshOpResult, + MeshOpUnavailableError, + mesh_ops_registry, +) router = APIRouter(tags=["optimize"]) @@ -36,16 +37,16 @@ class SmoothRequest(BaseModel): iterations: int +class MeshOpRequest(BaseModel): + path: str + params: dict[str, object] = Field(default_factory=dict) + + class TransformRequest(BaseModel): path: str # format: "{collection}/{filename}" matrix: list[list[float]] # row-major 4x4 world transform -def _require_pymeshlab(): - if not _PYMESHLAB_AVAILABLE: - raise HTTPException(503, "pymeshlab is unavailable on this system (DLL blocked by Windows Application Control policy)") - - def _resolve_input_path(raw_path: str) -> Path: candidate = Path(raw_path) if candidate.is_absolute(): @@ -54,150 +55,119 @@ def _resolve_input_path(raw_path: str) -> Path: raise HTTPException(404, f"File not found: {raw_path}") return resolved - resolved = (WORKSPACE_DIR / raw_path).resolve() - if not str(resolved).startswith(str(WORKSPACE_DIR.resolve())): + resolved = (registry.WORKSPACE_DIR / raw_path).resolve() + if not registry.is_within_workspace(resolved): raise HTTPException(400, "Invalid path") if not resolved.exists(): raise HTTPException(404, f"File not found: {raw_path}") return resolved -@router.post("/mesh") -def optimize_mesh(body: OptimizeRequest): - _require_pymeshlab() - target_faces = max(100, min(500_000, body.target_faces)) - - input_path = _resolve_input_path(body.path) - - tmp_dir = tempfile.mkdtemp() - try: - result = _decimate(str(input_path), target_faces, tmp_dir) - finally: - shutil.rmtree(tmp_dir, ignore_errors=True) - - stem = input_path.stem - output_name = f"{stem}_opt{target_faces}.glb" - output_dir = input_path.parent if str(input_path).startswith(str(WORKSPACE_DIR.resolve())) else WORKSPACE_DIR / "Workflows" +def _operation_output_path(input_path: Path, output_name: str) -> Path: + workspace = registry.WORKSPACE_DIR.resolve() + resolved_input = input_path.resolve() + output_dir = ( + input_path.parent + if resolved_input == workspace or workspace in resolved_input.parents + else registry.WORKSPACE_DIR / "Workflows" + ) output_dir.mkdir(parents=True, exist_ok=True) - output_path = output_dir / output_name - result.export(str(output_path)) + return output_dir / output_name + + +def _run_operation( + operation_id: str, + input_path: Path, + params: dict[str, object], + output_path: Path | None = None, + preserve_visuals: bool = False, +) -> MeshOpResult: + context = MeshOpContext( + workspace_dir=registry.WORKSPACE_DIR, + temp_dir=Path(tempfile.gettempdir()), + output_path=output_path, + preserve_visuals=preserve_visuals, + ) + try: + return mesh_ops_registry.run(operation_id, input_path, params, context) + except MeshOpNotFoundError as exc: + raise HTTPException(404, f"Unknown mesh operation: {operation_id}") from exc + except FileNotFoundError as exc: + raise HTTPException(404, str(exc)) from exc + except (MeshOpUnavailableError, MeshOpExecutionError) as exc: + raise HTTPException(503, str(exc)) from exc + except (TypeError, ValueError) as exc: + raise HTTPException(400, str(exc)) from exc + + +def _operation_response(result: MeshOpResult) -> dict[str, object]: + output_path = result.file_path.resolve() + try: + relative_path = output_path.relative_to(registry.WORKSPACE_DIR.resolve()).as_posix() + except ValueError: + payload: dict[str, object] = {"path": str(output_path)} + else: + payload = { + "path": relative_path, + "url": f"/workspace/{relative_path}", + } + payload.update(result.details) + return payload - face_count = len(result.faces) - rel = output_path.relative_to(WORKSPACE_DIR).as_posix() - return {"url": f"/workspace/{rel}", "face_count": face_count} +@router.get("/ops") +def list_mesh_operations(): + return mesh_ops_registry.describe() -def _has_texture(geom: trimesh.Trimesh) -> bool: - if not isinstance(geom.visual, trimesh.visual.TextureVisuals): - return False - mat = geom.visual.material - if mat is None: - return False - # Simple material (SimpleMaterial / Material) - if getattr(mat, "image", None) is not None: - return True - # PBR material (from Trellis2 SLaT texturing and GLB imports) - if getattr(mat, "baseColorTexture", None) is not None: - return True - return False - - -def _get_texture_image(geom: trimesh.Trimesh): - """Return the base color texture image regardless of material type.""" - mat = geom.visual.material - img = getattr(mat, "image", None) - if img is not None: - return img - return getattr(mat, "baseColorTexture", None) - - -def _decimate(input_path: str, target_faces: int, tmp_dir: str) -> trimesh.Trimesh: - loaded = trimesh.load(input_path) - if isinstance(loaded, trimesh.Scene): - geoms = list(loaded.geometry.values()) - geom = trimesh.util.concatenate(geoms) if len(geoms) > 1 else geoms[0] - else: - geom = loaded - - ms = _pymeshlab.MeshSet() - - if _has_texture(geom): - # ── Textured path: OBJ intermediate to preserve UV coordinates ────── - obj_in = os.path.join(tmp_dir, "input.obj") - mtl_in = os.path.join(tmp_dir, "input.mtl") - tex_in = os.path.join(tmp_dir, "texture.png") - obj_out = os.path.join(tmp_dir, "output.obj") - - # Save texture image under a known filename (handles PBR and simple materials) - _get_texture_image(geom).save(tex_in) - - # Export OBJ (trimesh writes UV coords + MTL) - geom.export(obj_in) - - # Patch MTL so any map_Kd points to our known texture filename - if os.path.exists(mtl_in): - mtl = open(mtl_in).read() - mtl = re.sub(r"map_Kd\s+\S+", "map_Kd texture.png", mtl) - open(mtl_in, "w").write(mtl) - - ms.load_new_mesh(obj_in) - ms.meshing_decimation_quadric_edge_collapse( - targetfacenum=target_faces, - preservetexcoord=True, # ← keeps UV coordinates intact - preservenormal=True, - preservetopology=True, - autoclean=True, - ) - ms.save_current_mesh(obj_out) - # Patch output MTL too, so trimesh can find the texture on load - mtl_out = obj_out.replace(".obj", ".mtl") - if os.path.exists(mtl_out): - mtl = open(mtl_out).read() - mtl = re.sub(r"map_Kd\s+\S+", "map_Kd texture.png", mtl) - open(mtl_out, "w").write(mtl) +@router.post("/op/{op_name}") +def run_mesh_operation(op_name: str, body: MeshOpRequest): + input_path = _resolve_input_path(body.path) + return _operation_response( + _run_operation(op_name, input_path, body.params) + ) - return trimesh.load(obj_out) - else: - # ── Geometry-only path: PLY (fast, no texture to worry about) ──────── - ply_in = os.path.join(tmp_dir, "input.ply") - ply_out = os.path.join(tmp_dir, "output.ply") - - geom.export(ply_in) - ms.load_new_mesh(ply_in) - ms.meshing_decimation_quadric_edge_collapse( - targetfacenum=target_faces, - preservenormal=True, - preservetopology=True, - autoclean=True, +@router.post("/mesh") +def optimize_mesh(body: OptimizeRequest): + target_faces = max(100, min(500_000, body.target_faces)) + input_path = _resolve_input_path(body.path) + output_path = _operation_output_path( + input_path, + f"{input_path.stem}_opt{target_faces}.glb", + ) + response = _operation_response( + _run_operation( + "decimate", + input_path, + {"target_faces": target_faces}, + output_path, ) - ms.save_current_mesh(ply_out) - return trimesh.load(ply_out, force="mesh") + ) + return { + "url": response.get("url"), + "face_count": response.get("face_count", 0), + } @router.post("/smooth") def smooth_mesh(body: SmoothRequest): - _require_pymeshlab() iterations = max(1, min(20, body.iterations)) - input_path = _resolve_input_path(body.path) - - tmp_dir = tempfile.mkdtemp() - try: - result = _smooth(str(input_path), iterations, tmp_dir) - finally: - shutil.rmtree(tmp_dir, ignore_errors=True) - - stem = input_path.stem - output_name = f"{stem}_smooth{iterations}.glb" - output_dir = input_path.parent if str(input_path).startswith(str(WORKSPACE_DIR.resolve())) else WORKSPACE_DIR / "Workflows" - output_dir.mkdir(parents=True, exist_ok=True) - output_path = output_dir / output_name - result.export(str(output_path)) - - rel = output_path.relative_to(WORKSPACE_DIR).as_posix() - return {"url": f"/workspace/{rel}"} + output_path = _operation_output_path( + input_path, + f"{input_path.stem}_smooth{iterations}.glb", + ) + response = _operation_response( + _run_operation( + "smooth", + input_path, + {"iterations": iterations, "lambda_": 0.5, "mode": "laplacian"}, + output_path, + preserve_visuals=True, + ) + ) + return {"url": response.get("url")} @router.post("/transform") @@ -219,62 +189,16 @@ def transform_mesh(body: TransformRequest): stem = input_path.stem output_name = f"{stem}_xf_{uuid.uuid4().hex[:8]}.glb" - output_dir = input_path.parent if str(input_path).startswith(str(WORKSPACE_DIR.resolve())) else WORKSPACE_DIR / "Workflows" + workspace = registry.WORKSPACE_DIR.resolve() + output_dir = input_path.parent if registry.is_within_workspace(input_path.resolve()) else workspace / "Workflows" output_dir.mkdir(parents=True, exist_ok=True) output_path = output_dir / output_name loaded.export(str(output_path)) - rel = output_path.relative_to(WORKSPACE_DIR).as_posix() + rel = output_path.resolve().relative_to(workspace).as_posix() return {"url": f"/workspace/{rel}"} -def _smooth(input_path: str, iterations: int, tmp_dir: str) -> trimesh.Trimesh: - loaded = trimesh.load(input_path) - if isinstance(loaded, trimesh.Scene): - geoms = list(loaded.geometry.values()) - geom = trimesh.util.concatenate(geoms) if len(geoms) > 1 else geoms[0] - else: - geom = loaded - - ms = _pymeshlab.MeshSet() - - if _has_texture(geom): - obj_in = os.path.join(tmp_dir, "input.obj") - mtl_in = os.path.join(tmp_dir, "input.mtl") - tex_in = os.path.join(tmp_dir, "texture.png") - obj_out = os.path.join(tmp_dir, "output.obj") - - _get_texture_image(geom).save(tex_in) - geom.export(obj_in) - - if os.path.exists(mtl_in): - mtl = open(mtl_in).read() - mtl = re.sub(r"map_Kd\s+\S+", "map_Kd texture.png", mtl) - open(mtl_in, "w").write(mtl) - - ms.load_new_mesh(obj_in) - ms.apply_coord_laplacian_smoothing(stepsmoothnum=iterations) - ms.save_current_mesh(obj_out) - - mtl_out = obj_out.replace(".obj", ".mtl") - if os.path.exists(mtl_out): - mtl = open(mtl_out).read() - mtl = re.sub(r"map_Kd\s+\S+", "map_Kd texture.png", mtl) - open(mtl_out, "w").write(mtl) - - return trimesh.load(obj_out) - - else: - ply_in = os.path.join(tmp_dir, "input.ply") - ply_out = os.path.join(tmp_dir, "output.ply") - - geom.export(ply_in) - ms.load_new_mesh(ply_in) - ms.apply_coord_laplacian_smoothing(stepsmoothnum=iterations) - ms.save_current_mesh(ply_out) - return trimesh.load(ply_out, force="mesh") - - class ImportByPathRequest(BaseModel): path: str # absolute path on disk @@ -420,6 +344,7 @@ async def import_mesh_by_path(body: ImportByPathRequest): if ext == "glb": # Serve the original file directly — no copy + imported_sources.register(file_path, file_path) return {"url": f"/optimize/serve-file?path={quote(str(file_path))}"} # Mesh ply / obj / stl: convert to GLB in a temp directory (not the workspace) @@ -427,6 +352,10 @@ async def import_mesh_by_path(body: ImportByPathRequest): output_path = os.path.join(tmp_dir, "mesh.glb") loaded = trimesh.load(str(file_path)) loaded.export(output_path) + # Remember the ORIGINAL extension: this conversion changes the container but + # not the axes, so a .stl imported here still holds Z-up data despite the + # .glb suffix, and must not be rotated as if it were glTF. + imported_sources.register(output_path, file_path) return {"url": f"/optimize/serve-file?path={quote(output_path)}"} @@ -454,10 +383,8 @@ def ply_to_splat(path: str): `path` is workspace-relative (e.g. "Workflows/foo.ply"). A .splat is served as-is; a GS .ply is normalised + converted (cached by mtime + conv version). """ - import services.generator_registry as reg # dynamic: workspace dir may change at runtime - workspace = reg.WORKSPACE_DIR.resolve() - src = (workspace / path).resolve() - if not str(src).startswith(str(workspace)): + src = (registry.WORKSPACE_DIR.resolve() / path).resolve() + if not registry.is_within_workspace(src): raise HTTPException(400, "Invalid path") if not src.is_file(): raise HTTPException(404, "File not found") @@ -482,8 +409,8 @@ def export_mesh(path: str, format: str): if format not in ("obj", "stl", "ply"): raise HTTPException(400, "Supported formats: obj, stl, ply") - input_path = (WORKSPACE_DIR / path).resolve() - if not str(input_path).startswith(str(WORKSPACE_DIR.resolve())): + input_path = (registry.WORKSPACE_DIR / path).resolve() + if not registry.is_within_workspace(input_path): raise HTTPException(400, "Invalid path") if not input_path.exists(): raise HTTPException(404, f"File not found: {path}") @@ -503,4 +430,4 @@ def export_mesh(path: str, format: str): content=data, media_type=mime, headers={"Content-Disposition": f'attachment; filename="{stem}.{format}"'}, - ) \ No newline at end of file + ) diff --git a/api/routers/settings.py b/api/routers/settings.py index 4ad895f2..9106ad2d 100644 --- a/api/routers/settings.py +++ b/api/routers/settings.py @@ -1,3 +1,4 @@ +import asyncio import os from fastapi import APIRouter from pydantic import BaseModel @@ -28,9 +29,11 @@ async def get_paths(): @router.post("/paths") async def update_paths(body: PathsUpdate): - reg_module.generator_registry.update_paths( - models_dir = Path(body.models_dir) if body.models_dir else None, - workspace_dir = Path(body.workspace_dir) if body.workspace_dir else None, + # Off the event loop: changing paths waits for any in-progress model load. + await asyncio.to_thread( + reg_module.generator_registry.update_paths, + Path(body.models_dir) if body.models_dir else None, + Path(body.workspace_dir) if body.workspace_dir else None, ) return { "models_dir": str(reg_module.MODELS_DIR), diff --git a/api/routers/workflow_runs.py b/api/routers/workflow_runs.py index cb47bc92..ba654a53 100644 --- a/api/routers/workflow_runs.py +++ b/api/routers/workflow_runs.py @@ -8,9 +8,10 @@ from routers.generation import ( VALID_REMESH_MODES, _cancel_events, - _cancelled, _jobs, + _purge_old_jobs, _run_generation, + cancel_job, sanitize_collection, ) from schemas.generation import JobStatus @@ -56,10 +57,7 @@ async def create_run_from_image( **model_params, } - # Same constraint /generate/from-image enforces on this field, checked before touching - # the registry below for the same reason that endpoint checks it first: switch_model() - # unloads whatever generator is currently active, and a request rejected for a bad - # remesh value should not pay for -- or force a reload after -- evicting it. + # Keep the same request validation as /generate/from-image before filing a job. if full_params["remesh"] not in VALID_REMESH_MODES: raise HTTPException(400, "remesh must be 'quad', 'triangle', or 'none'") @@ -67,18 +65,21 @@ async def create_run_from_image( try: generator_registry.get_generator(model_id) + output_kind = generator_registry.get_manifest(model_id).get("output", "mesh") except ValueError as e: raise HTTPException(400, str(e)) - generator_registry.switch_model(model_id) - job_id = str(uuid.uuid4()) image_bytes = await image.read() + _purge_old_jobs() + _jobs[job_id] = JobStatus(job_id=job_id, status="pending", progress=0) _cancel_events[job_id] = threading.Event() - background_tasks.add_task(_run_generation, job_id, image_bytes, full_params, collection) + background_tasks.add_task( + _run_generation, job_id, image_bytes, full_params, collection, output_kind, model_id + ) return {"run_id": job_id, "status": "pending"} @@ -106,23 +107,4 @@ async def get_run(run_id: str): @router.post("/{run_id}/cancel") async def cancel_run(run_id: str): - job = _jobs.get(run_id) - if not job: - raise HTTPException(404, f"Run {run_id} not found") - - _cancelled.add(run_id) - if run_id in _cancel_events: - _cancel_events[run_id].set() - if job.status in ("pending", "running"): - job.status = "cancelled" - - try: - gen = generator_registry._generators.get(generator_registry._active_id) - if gen is not None and hasattr(gen, "_proc") and gen._proc and gen._proc.poll() is None: - gen._proc.kill() - gen._loaded = False - gen._proc = None - except Exception: - pass - - return {"cancelled": True} + return await cancel_job(run_id) diff --git a/api/runner.py b/api/runner.py index 19d0df71..93feef5d 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]) @@ -104,6 +117,46 @@ def _resolve_ready_schema(GenClass, node: dict, manifest: dict) -> list: return node.get("params_schema") or manifest.get("params_schema", []) +def _generator_is_loaded(gen) -> bool: + """ + Best-effort read of the generator's loaded state. + + Never raises: this is used on the error path, where a generator left in a + half-initialised state can make is_loaded() itself blow up. An unreadable + state is reported as "not loaded" so the caller reloads rather than reusing + a worker that may be broken. + """ + try: + return bool(gen.is_loaded()) + except Exception: + return False + + +def _ensure_model_loaded(gen) -> bool: + """ + Guarantees the model is in memory before inference. Returns True if a + reload was needed. + + Texture pipelines are built lazily on first use, and extensions typically + free the shape pipeline first to make room for them. When that setup fails + (a missing `xatlas` being the common case) the generator is left with + _model = None, yet the worker process stays alive and both ExtensionProcess + and GeneratorRegistry still consider it loaded — so load() is never called + again and every later generate() dies with + + TypeError: 'NoneType' object is not callable + + until the process is killed by hand. Re-loading here mirrors what + GeneratorRegistry.get_active() already does host-side, which makes a failed + setup recoverable on the next attempt instead of permanently poisoning the + worker. + """ + if _generator_is_loaded(gen): + return False + gen.load() + return True + + def _apply_manifest_metadata(gen, manifest: dict, node: dict) -> None: gen.hf_repo = node.get("hf_repo") or manifest.get("hf_repo", "") gen.hf_skip_prefixes = node.get("hf_skip_prefixes") or manifest.get("hf_skip_prefixes", []) @@ -111,13 +164,39 @@ def _apply_manifest_metadata(gen, manifest: dict, node: dict) -> None: gen._params_schema = node.get("params_schema") or manifest.get("params_schema", []) +def decode_model_input(msg: dict): + """Decode legacy image bytes or revalidate a typed artifact in the worker.""" + if "input" not in msg: + return base64.b64decode(msg["image_b64"]) + value = msg["input"] + if not isinstance(value, dict) or set(value) != {"kind", "path"}: + raise ValueError("Typed artifact input must contain exactly kind and path") + from services.artifact_input import TypedArtifactInput, revalidate_artifact_input + typed = TypedArtifactInput(kind=value.get("kind"), path=Path(value.get("path", ""))) + return revalidate_artifact_input(WORKSPACE_DIR, typed) + + +def validate_requested_model(msg: dict, manifest: dict, node: dict) -> None: + """Reject requests routed to a worker for a different manifest node.""" + requested = msg.get("model_id") + if requested is None: # Backward compatibility with already-running legacy hosts. + return + expected = manifest["id"] + if node.get("id"): + expected = f"{expected}/{node['id']}" + if requested != expected: + raise ValueError( + f"Generation request model '{requested}' does not match worker model '{expected}'" + ) + + # ------------------------------------------------------------------ # # Main loop # ------------------------------------------------------------------ # 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) @@ -127,11 +206,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. @@ -144,6 +221,9 @@ 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) # Active cancel events keyed by request id @@ -161,10 +241,18 @@ def main() -> None: # ---- generate -------------------------------------------- elif action == "generate": + validate_requested_model(msg, manifest, node) cancel_evt = threading.Event() _cancel[rid] = cancel_evt - image_bytes = base64.b64decode(msg["image_b64"]) + model_input = decode_model_input(msg) params = msg.get("params", {}) + if hasattr(model_input, "kind"): + if not isinstance(params, dict): + raise ValueError("Model params must be an object") + from services.artifact_input import RESERVED_ARTIFACT_PARAMS + params = {key: value for key, value in params.items() + if key not in RESERVED_ARTIFACT_PARAMS} + params["scene_manifest_path"] = str(model_input.path) if msg.get("outputs_dir"): gen.outputs_dir = Path(msg["outputs_dir"]) gen.outputs_dir.mkdir(parents=True, exist_ok=True) @@ -173,16 +261,37 @@ def progress_cb(pct: int, step: str = "") -> None: send({"type": "progress", "id": rid, "pct": pct, "step": step}) try: - output_path = gen.generate(image_bytes, params, progress_cb, cancel_evt) + if _ensure_model_loaded(gen): + send({"type": "log", "level": "warning", + "message": ("Model was not loaded (earlier setup failure?); " + "reloaded before generating.")}) + if hasattr(model_input, "kind"): + output_path = gen.generate_artifact( + model_input.kind, model_input.path, params, progress_cb, cancel_evt + ) + else: + output_path = gen.generate(model_input, params, progress_cb, cancel_evt) + if node.get("output") == "scene": + from services.scene_input import validate_scene_input + resolved_output = Path(output_path).resolve(strict=True) + try: + relative_output = resolved_output.relative_to(WORKSPACE_DIR.resolve(strict=True)) + except (OSError, ValueError) as exc: + raise ValueError("Generated scene output is outside the workspace") from exc + output_path = validate_scene_input(WORKSPACE_DIR, relative_output.as_posix()) send({"type": "done", "id": rid, "output_path": str(output_path)}) except Exception as exc: # Detect GenerationCancelled by name to avoid import issues if type(exc).__name__ == "GenerationCancelled": send({"type": "cancelled", "id": rid}) else: + # Report whether the model survived the failure so the + # host can drop its cached "loaded" flag and reload + # instead of reusing a worker whose _model is None. send({"type": "error", "id": rid, "message": str(exc), - "traceback": traceback.format_exc()}) + "traceback": traceback.format_exc(), + "loaded": _generator_is_loaded(gen)}) finally: _cancel.pop(rid, None) diff --git a/api/schemas/generation.py b/api/schemas/generation.py index 7ed6ca62..04c18c85 100644 --- a/api/schemas/generation.py +++ b/api/schemas/generation.py @@ -1,5 +1,5 @@ -from typing import Literal, Optional -from pydantic import BaseModel +from typing import Any, Literal, Optional +from pydantic import BaseModel, Field class JobStatus(BaseModel): @@ -9,3 +9,12 @@ class JobStatus(BaseModel): step: Optional[str] = None # Human-readable current step output_url: Optional[str] = None error: Optional[str] = None + + +class GenerateFromArtifactRequest(BaseModel): + """Generic typed-artifact request. Only scene is public in this release.""" + input_kind: Literal["scene"] + input_path: str + model_id: str + collection: str = "Workflows" + params: dict[str, Any] = Field(default_factory=dict) diff --git a/api/services/artifact_input.py b/api/services/artifact_input.py new file mode 100644 index 00000000..05f4661d --- /dev/null +++ b/api/services/artifact_input.py @@ -0,0 +1,34 @@ +"""Typed model artifact inputs shared by the API, worker bridge, and runner.""" +from dataclasses import dataclass +from pathlib import Path + +from services.scene_input import validate_scene_input + +SUPPORTED_ARTIFACT_INPUTS = frozenset({"scene"}) + +# Transport parameters set by the host; callers must not be able to forge them. +RESERVED_ARTIFACT_PARAMS = frozenset({ + "artifact_path", "input_kind", "input_path", "scene_path", "scene_manifest_path", +}) + + +@dataclass(frozen=True) +class TypedArtifactInput: + kind: str + path: Path + + +def validate_artifact_input(workspace: Path, kind: str, input_path: str) -> TypedArtifactInput: + if kind not in SUPPORTED_ARTIFACT_INPUTS: + raise ValueError(f"Unsupported artifact input kind: {kind}") + if kind == "scene": + return TypedArtifactInput(kind="scene", path=validate_scene_input(workspace, input_path)) + raise ValueError(f"Unsupported artifact input kind: {kind}") + + +def revalidate_artifact_input(workspace: Path, value: TypedArtifactInput) -> TypedArtifactInput: + try: + relative = value.path.resolve(strict=True).relative_to(workspace.resolve(strict=True)) + except (OSError, ValueError) as exc: + raise ValueError("Artifact input is outside the workspace") from exc + return validate_artifact_input(workspace, value.kind, relative.as_posix()) diff --git a/api/services/extension_process.py b/api/services/extension_process.py index 518f94f2..5eaed4b8 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() @@ -73,9 +74,12 @@ def _build_env(self) -> dict: env["MODELS_DIR"] = str(MODELS_DIR) env["WORKSPACE_DIR"] = str(WORKSPACE_DIR) env["MODLY_API_DIR"] = str(Path(__file__).parent.parent) + # Lets extensions call back into the Modly API (e.g. /llm/chat for the shared LLM). + env.setdefault("MODLY_API_URL", "http://127.0.0.1:8765") # Force the worker's Python stdio to UTF-8 so it matches the UTF-8 # pipe readers below regardless of the OS locale (cp1252/cp932). env["PYTHONUTF8"] = "1" + env["PYTHONIOENCODING"] = "utf-8" if sys.platform == "darwin": env.setdefault("NUMBA_DISABLE_JIT", "1") # Must be set before the subprocess's first `import torch` — PyTorch @@ -84,11 +88,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: @@ -119,6 +129,12 @@ def _start(self) -> None: stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, + # Without these, text=True decodes with the locale codec — cp1252 + # on Windows. tqdm draws its partial blocks with U+258D/U+258F, + # whose UTF-8 bytes (0x8d/0x8f) are undefined there: the decode + # raises, _stderr_loop dies, nobody drains the pipe, and the + # child blocks forever on write once it fills. errors="replace" + # keeps a stray non-UTF-8 byte from resurrecting that failure. encoding="utf-8", errors="replace", bufsize=1, @@ -184,6 +200,8 @@ def _install_missing_package(self, python: Path, module_name: str, package_name: stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, + encoding="utf-8", + errors="replace", ) except subprocess.CalledProcessError as exc: details = (exc.stderr or exc.stdout or "").strip() @@ -291,16 +309,18 @@ def generate( progress_cb: Optional[Callable[[int, str], None]] = None, cancel_event: Optional[threading.Event] = None, ) -> Path: - from services.generators.base import GenerationCancelled + return self._generate_request( + {"image_b64": base64.b64encode(image_bytes).decode()}, + params, progress_cb, cancel_event, + ) - req_id = str(uuid.uuid4()) - self._send({ - "action": "generate", - "id": req_id, - "image_b64": base64.b64encode(image_bytes).decode(), - "params": params, - "outputs_dir": str(self.outputs_dir) if self.outputs_dir else None, - }) + def _receive_generation( + self, + req_id: str, + progress_cb: Optional[Callable[[int, str], None]], + cancel_event: Optional[threading.Event], + ) -> Path: + from services.generators.base import GenerationCancelled # Grace period after sending a cooperative cancel before hard-killing # the subprocess. Long enough to let generators that check cancel_event @@ -362,6 +382,13 @@ def generate( return Path(msg["output_path"]) elif t == "error": + # A failed generation can leave the worker without a model: a + # lazy texture-setup failure frees the shape pipeline and then + # raises. The worker reports its post-failure state, so drop + # our cached flag and let GeneratorRegistry.get_active() reload + # before the next run rather than reusing a broken worker. + if msg.get("loaded") is False: + self._loaded = False raise RuntimeError(msg.get("traceback") or msg.get("message", "Unknown error")) elif t == "cancelled": @@ -370,6 +397,41 @@ def generate( elif t == "log": print(f"[{self.MODEL_ID}] {msg.get('message', '')}", file=sys.stderr) + def generate_artifact( + self, + input_kind: str, + artifact_path: Path, + params: dict, + progress_cb: Optional[Callable[[int, str], None]] = None, + cancel_event: Optional[threading.Event] = None, + ) -> Path: + """Send a typed artifact envelope to the isolated runner.""" + from services.artifact_input import TypedArtifactInput, revalidate_artifact_input + from services.generator_registry import WORKSPACE_DIR + + validated = revalidate_artifact_input( + WORKSPACE_DIR, TypedArtifactInput(kind=input_kind, path=artifact_path) + ) + return self._generate_request( + {"input": {"kind": validated.kind, "path": str(validated.path)}}, + params, progress_cb, cancel_event, + ) + + def _generate_request( + self, + input_payload: dict, + params: dict, + progress_cb: Optional[Callable[[int, str], None]], + cancel_event: Optional[threading.Event], + ) -> Path: + req_id = str(uuid.uuid4()) + self._send({ + "action": "generate", "id": req_id, "model_id": self.MODEL_ID, + **input_payload, "params": params, + "outputs_dir": str(self.outputs_dir) if self.outputs_dir else None, + }) + return self._receive_generation(req_id, progress_cb, cancel_event) + def params_schema(self) -> list: return self._params_schema diff --git a/api/services/generator_registry.py b/api/services/generator_registry.py index ecb6bb6c..2792030e 100644 --- a/api/services/generator_registry.py +++ b/api/services/generator_registry.py @@ -25,6 +25,18 @@ from services.generators.base import BaseGenerator from services.extension_process import ExtensionProcess, _venv_python +from services.model_sources import ( + missing_weight_variant, + model_sources_are_downloaded, + normalize_model_sources, + normalize_weight_group_references, + normalize_weight_groups, + normalize_weight_variants, + validate_model_node_ids, + resolve_weight_group_root, + safe_source_id, + weight_group_sources_are_downloaded, +) # ------------------------------------------------------------------ # # Global paths @@ -54,6 +66,17 @@ print(f"[Registry] EXTENSIONS_DIR = {EXTENSIONS_DIR or '(not set)'}") +def is_within_workspace(resolved_path: Path) -> bool: + """True when an already-resolved path is the current workspace or inside it. + + Compares ancestry, not string prefixes: ``startswith`` would also accept a + sibling folder such as ``-other``. Reads WORKSPACE_DIR at call + time, since moving the workspace rebinds it. + """ + workspace = WORKSPACE_DIR.resolve() + return resolved_path == workspace or workspace in resolved_path.parents + + # ------------------------------------------------------------------ # # Extension loader # ------------------------------------------------------------------ # @@ -428,6 +451,11 @@ def _discover_extensions( ext_id = manifest["id"] class_name = manifest["generator_class"] + if "model_sources" in manifest: + raise ValueError("model_sources must be declared on a model node") + + weight_groups = normalize_weight_groups(manifest) + if ext_id != ext_dir.name: message = ( f"Extension folder '{ext_dir.name}' declares mismatched " @@ -448,6 +476,50 @@ 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 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") + 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 "weight_variants" in node: + raise ValueError( + f'model node "{node_id}": weight_variants cannot be combined ' + "with 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" + ) + for node in nodes: + declared_inputs = node.get("inputs") + uses_scene_input = ( + node.get("input", "image") == "scene" + or (isinstance(declared_inputs, list) and "scene" in declared_inputs) + ) + if uses_scene_input and ( + "inputs" in node or node.get("input", "image") != "scene" + ): + raise ValueError( + f'model node "{node.get("id", "unknown")}" must declare scene ' + 'as its single input field' + ) # Markers left while setup or runtime registration is unfinished: # the folder is not ready to be loaded. The readable manifest lets @@ -462,8 +534,11 @@ def _discover_extensions( registration_authorization is not None and registration_authorization[0] == ext_id and registration_authorization[1].exists() + # Normalize like the capability's destination (resolved + # root + name): EXTENSIONS_DIR may be a Windows 8.3 short + # path. The extension folder itself is not resolved. and registration_authorization[2] - == Path(os.path.abspath(ext_dir)) + == Path(os.path.abspath(ext_dir.parent.resolve() / ext_dir.name)) ) ) ) @@ -523,6 +598,15 @@ 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 [] + weight_variants = normalize_weight_variants( + node, node.get("params_schema", manifest.get("params_schema", [])) + ) node_manifest = { **manifest, "id": f"{ext_id}/{node['id']}", @@ -535,8 +619,14 @@ def _discover_extensions( "hf_include_prefixes": node.get("hf_include_prefixes", []), "params_schema": node.get("params_schema", manifest.get("params_schema", [])), "input": node.get("input", "image"), + "inputs": node.get("inputs"), "output": node.get("output", "mesh"), + "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 + if weight_variants is not None: + node_manifest["weight_variants"] = weight_variants full_id = f"{ext_id}/{node['id']}" result[full_id] = (cls_or_None, node_manifest, ext_dir, legacy_context) if subprocess_mode: @@ -579,6 +669,11 @@ def __init__(self) -> None: self._generators: Dict[str, BaseGenerator] = {} self._manifests: Dict[str, dict] = {} self._errors: Dict[str, str] = {} + # Serializes everything that changes which generators exist or are + # loaded (switch/load, unload, reload, path changes). Read-only status + # calls deliberately skip it: a load can hold it for minutes (first-run + # downloads happen inside load()), and status must stay responsive. + self._lifecycle_lock = threading.RLock() self._legacy_imports = _LegacyImportManager() self._active_id: str = os.environ.get("SELECTED_MODEL_ID", "sf3d") @@ -605,6 +700,9 @@ def initialize( ) # Subprocess mode: wrap in ExtensionProcess gen = ExtensionProcess(ext_dir, manifest) + # Pin the subprocess envelope to the exact registry key; + # multi-node workers must never fall back to an extension ID. + gen.MODEL_ID = model_id gen.model_dir = MODELS_DIR / model_id gen.outputs_dir = WORKSPACE_DIR else: @@ -623,6 +721,15 @@ 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"] + ) + for group in manifest.get("weight_groups", []) + } + self._generators[model_id] = gen self._manifests[model_id] = manifest self._errors.pop(model_id, None) @@ -654,25 +761,26 @@ def reload(self, validation_capability: object = None) -> None: registration_authorization = _consume_registration_validation_capability( validation_capability, ) - print("[Registry] Reloading extensions...") - for gen in self._generators.values(): - if isinstance(gen, ExtensionProcess): - gen.stop() - if gen._proc is not None: - raise RuntimeError( - "Extension subprocess remained attached after stop()" - ) - else: - try: - gen.unload() - except Exception: - pass - self._generators.clear() - self._manifests.clear() - self._errors.clear() - self._remove_legacy_paths() - self.initialize(registration_authorization) - print("[Registry] Reload complete.") + with self._lifecycle_lock: + print("[Registry] Reloading extensions...") + for gen in self._generators.values(): + if isinstance(gen, ExtensionProcess): + gen.stop() + if gen._proc is not None: + raise RuntimeError( + "Extension subprocess remained attached after stop()" + ) + else: + try: + gen.unload() + except Exception: + pass + self._generators.clear() + self._manifests.clear() + self._errors.clear() + self._remove_legacy_paths() + self.initialize(registration_authorization) + print("[Registry] Reload complete.") def load_errors(self) -> Dict[str, str]: """Returns extension loading errors.""" @@ -697,19 +805,68 @@ def _assert_not_quarantined(model_id: str) -> None: def get_active(self) -> BaseGenerator: """Returns the active generator. Downloads and loads if necessary.""" - self._assert_not_quarantined(self._active_id) - gen = self._generators[self._active_id] - if not gen.is_loaded(): - if not gen.is_downloaded(): - if isinstance(gen, ExtensionProcess): - # Let the subprocess handle its own download logic during - # load() — some extensions (e.g. mv-adapter) need custom - # multi-repo downloads that the standard HF endpoint can't do. - pass - else: - gen._auto_download() - gen.load() - return gen + with self._lifecycle_lock: + return self.get_ready_generator(self._active_id) + + def get_ready_generator(self, model_id: str) -> BaseGenerator: + """Load and return exactly ``model_id`` without consulting active state.""" + with self._lifecycle_lock: + gen = self.get_generator(model_id) + downloaded = self._is_downloaded(model_id, gen) + manifest = self._manifests[model_id] + if ( + "model_sources" in manifest or manifest.get("weight_groups") + ) and not downloaded: + raise RuntimeError( + "Model sources are incomplete. Download this node's shared and private weights " + "from the Modly Models page before generation." + ) + if not gen.is_loaded(): + if not downloaded: + if isinstance(gen, ExtensionProcess): + # Let the subprocess handle its own download logic during + # load() — some extensions (e.g. mv-adapter) need custom + # multi-repo downloads that the standard HF endpoint can't do. + pass + else: + gen._auto_download() + gen.load() + return gen + + def activate_ready_generator(self, model_id: str) -> BaseGenerator: + """Atomically make ``model_id`` active and return it ready for inference.""" + with self._lifecycle_lock: + self.switch_model(model_id) + return self.get_ready_generator(model_id) + + def model_status(self, model_id: str) -> dict: + gen = self.get_generator(model_id) + manifest = self._manifests[model_id] + return { + "id": model_id, + "name": manifest.get("name", gen.DISPLAY_NAME), + "downloaded": self._is_downloaded(model_id, gen), + "loaded": gen.is_loaded(), + } + + def assert_weight_variant_installed( + self, params: dict, model_id: Optional[str] = None + ) -> None: + """Refuse generation when the weight variant selected by params is not installed. + + ``model_id`` is the job's model. Pinned jobs only switch to it once they + run, so the active model is not a reliable stand-in before that. + """ + target_id = model_id or self._active_id + manifest = self._manifests.get(target_id, {}) + option = missing_weight_variant( + MODELS_DIR, target_id, manifest.get("weight_variants"), params + ) + if option is not None: + raise RuntimeError( + f'{option["label"]} weights for {target_id} are not installed. ' + "Install them from the Extensions page, or select an installed variant." + ) def get_generator(self, model_id: str) -> BaseGenerator: self._assert_not_quarantined(model_id) @@ -728,35 +885,63 @@ def get_manifest(self, model_id: str) -> dict: def switch_model(self, model_id: str) -> None: """Switches the active model. Unloads the previous one if different.""" - self._assert_not_quarantined(model_id) - if model_id not in self._generators: - raise ValueError( - f"Unknown model ID: '{model_id}'. " - f"Available: {list(self._generators.keys())}" - ) - if model_id != self._active_id: - if self._active_id in self._generators: - self._generators[self._active_id].unload() - self._active_id = model_id + with self._lifecycle_lock: + self._assert_not_quarantined(model_id) + if model_id not in self._generators: + raise ValueError( + f"Unknown model ID: '{model_id}'. " + f"Available: {list(self._generators.keys())}" + ) + # 3D generation owns the GPU: evict the chat LLMs before anything is + # about to allocate on it. The trigger is "the target model is not + # resident", not "the target model changed" — the common case is the + # default generator, which is already `_active_id` at boot and still + # has to load its weights. Gating on the id alone let a full LLM pool + # (2 slots, ~11.6 GB of 12) sit through an entire generation. + target = self._generators[model_id] + if model_id != self._active_id or not target.is_loaded(): + from services.llm_server import llama_pool + llama_pool.unload_all() + + if model_id != self._active_id: + if self._active_id in self._generators: + self._generators[self._active_id].unload() + self._active_id = model_id # ------------------------------------------------------------------ # # Status # ------------------------------------------------------------------ # + def _is_downloaded(self, model_id: str, gen: BaseGenerator) -> bool: + manifest = self._manifests[model_id] + private_ready = True + if "model_sources" in manifest: + 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: - gen = self._generators[self._active_id] - manifest = self._manifests[self._active_id] - return { - "id": self._active_id, - "name": manifest.get("name", gen.DISPLAY_NAME), - "downloaded": gen.is_downloaded(), - "loaded": gen.is_loaded(), - } + return self.model_status(self._active_id) def all_status(self) -> list: result = [] - for model_id, gen in self._generators.items(): - manifest = self._manifests[model_id] + active_id = self._active_id + # Snapshot: a concurrent reload() may clear the dicts mid-iteration. + for model_id, gen in list(self._generators.items()): + manifest = self._manifests.get(model_id) + if manifest is None: + continue result.append({ "id": model_id, "name": manifest.get("name", gen.DISPLAY_NAME), @@ -765,17 +950,18 @@ def all_status(self) -> list: "vram_gb": manifest.get("vram_gb", gen.VRAM_GB), "hf_repo": manifest.get("hf_repo", ""), "tags": manifest.get("tags", []), - "downloaded": gen.is_downloaded(), + "downloaded": self._is_downloaded(model_id, gen), "loaded": gen.is_loaded(), - "active": model_id == self._active_id, + "active": model_id == active_id, }) return result def params_schema(self, model_id: Optional[str] = None) -> list: target_id = model_id or self._active_id - if target_id not in self._generators: + gen = self._generators.get(target_id) + if gen is None: raise KeyError(target_id) - return self._generators[target_id].params_schema() + return gen.params_schema() # ------------------------------------------------------------------ # # Paths update & shutdown @@ -785,25 +971,44 @@ def update_paths(self, models_dir: Optional[Path], workspace_dir: Optional[Path] global MODELS_DIR, WORKSPACE_DIR import services.generator_registry as _self_module - if models_dir is not None: - self.unload_all() - models_dir.mkdir(parents=True, exist_ok=True) - _self_module.MODELS_DIR = models_dir - for model_id, gen in self._generators.items(): - gen.model_dir = models_dir / model_id + with self._lifecycle_lock: + if models_dir is not None: + self.unload_all() + models_dir.mkdir(parents=True, exist_ok=True) + _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) - _self_module.WORKSPACE_DIR = workspace_dir - for gen in self._generators.values(): - gen.outputs_dir = workspace_dir + if workspace_dir is not None: + workspace_dir.mkdir(parents=True, exist_ok=True) + _self_module.WORKSPACE_DIR = workspace_dir + for gen in self._generators.values(): + gen.outputs_dir = workspace_dir def unload_all(self) -> None: - for gen in self._generators.values(): - if isinstance(gen, ExtensionProcess): - gen.stop() - else: - gen.unload() + with self._lifecycle_lock: + still_loaded = [] + for model_id, gen in self._generators.items(): + if isinstance(gen, ExtensionProcess): + gen.stop() + else: + gen.unload() + if gen.is_loaded(): + still_loaded.append(model_id) + if still_loaded: + raise RuntimeError( + "Models are still loaded; weights were preserved: " + + ", ".join(still_loaded) + ) # Singleton diff --git a/api/services/generators/base.py b/api/services/generators/base.py index fd62ceef..63ded9ad 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) @@ -90,6 +91,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 @@ -143,7 +147,6 @@ def is_loaded(self) -> bool: # Inference # ------------------------------------------------------------------ # - @abstractmethod def generate( self, image_bytes: bytes, @@ -157,7 +160,27 @@ def generate( progress_cb(percent: int, step_label: str) cancel_event: set this to interrupt generation between steps. """ - ... + raise NotImplementedError( + f"{type(self).__name__} does not implement legacy image generation" + ) + + def generate_artifact( + self, + input_kind: str, + artifact_path: Path, + params: dict, + progress_cb: Optional[Callable[[int, str], None]] = None, + cancel_event: Optional[threading.Event] = None, + ) -> Path: + """Generate from a validated typed artifact. + + Typed-artifact extensions must implement this explicitly. Falling back + to ``generate`` would pass a filesystem ``Path`` to the legacy + image-bytes ABI and fail far from the actual contract violation. + """ + raise NotImplementedError( + f"{type(self).__name__} does not implement {input_kind} artifact generation" + ) def _check_cancelled(self, cancel_event: Optional[threading.Event]) -> None: """Raises GenerationCancelled if cancel_event is set.""" diff --git a/api/services/imported_sources.py b/api/services/imported_sources.py new file mode 100644 index 00000000..a21dded5 --- /dev/null +++ b/api/services/imported_sources.py @@ -0,0 +1,53 @@ +"""Registry of meshes the user explicitly imported from outside the workspace. + +`/optimize/import-by-path` serves files the user picked through the OS file +dialog, either in place or converted into a temp dir — never from the workspace. +The slicer export route confines itself to the workspace by design, so those +imports would be unsliceable without widening that guard to arbitrary absolute +paths, which would be a real regression. + +This registry is the narrow alternative: a route may accept an absolute path +only if it is an EXACT member here, i.e. a file the user chose themselves in +this session. Membership is an equality test, never a prefix test, so it grants +no traversal. + +It also remembers each file's ORIGINAL suffix. An imported `.stl` is converted +to GLB on import without any axis change, so the container extension alone would +mislead a consumer into applying the glTF Y-up -> Z-up rotation to data that is +already Z-up. + +Session-scoped and in-memory: it is emptied when the backend restarts, so a +deeplink replayed after a restart is rejected rather than silently served. Links +are consumed within a second of the click, so this is not worth persisting. +""" + +from pathlib import Path + +# resolved absolute path (str) -> original suffix, lowercase, with the dot +_SOURCES: dict[str, str] = {} + + +def register(served_path: str | Path, original_path: str | Path) -> str: + """Record a user-imported file as sliceable and return its resolved path. + + ``served_path`` is what gets served (possibly a converted temp GLB); + ``original_path`` is the file the user actually picked, whose suffix decides + the source format. + """ + resolved = str(Path(served_path).resolve()) + _SOURCES[resolved] = Path(original_path).suffix.lower() + return resolved + + +def source_suffix(path: str | Path) -> str | None: + """Original suffix of a registered import, or None if it is not registered.""" + return _SOURCES.get(str(Path(path).resolve())) + + +def is_registered(path: str | Path) -> bool: + return source_suffix(path) is not None + + +def clear() -> None: + """Drop every entry — for tests.""" + _SOURCES.clear() diff --git a/api/services/llm_server.py b/api/services/llm_server.py new file mode 100644 index 00000000..c41fb196 --- /dev/null +++ b/api/services/llm_server.py @@ -0,0 +1,1045 @@ +""" +Local LLM engine — manages a llama.cpp `llama-server` subprocess. + +Everything lives under the agent directory (MODLY_LLM_DIR, set by Electron to +the `agent` folder beside models/, extensions/, … — ~/.modly/llm/ when the API +runs standalone): + bin/ llama-server binary + DLLs (auto-downloaded from GitHub releases) + models/ GGUF files (catalog downloads + any custom .gguf the user drops in) + logs/ one log per llama-server slot + config.json pool settings (max_models) + +Nothing is hardcoded to a machine: the binary variant is picked per-platform +(CUDA if an NVIDIA driver is present, otherwise Vulkan, otherwise CPU) and +models are chosen by the user from the catalog. +""" +import contextlib +import hashlib +import os +import re +import shutil +import subprocess +import sys +import tempfile +import threading +import time +import zipfile +from pathlib import Path +from typing import Callable, Optional +from urllib.request import Request, urlopen +import json as _json + +LLM_DIR = Path(os.environ.get("MODLY_LLM_DIR") or Path.home() / ".modly" / "llm") +BIN_DIR = LLM_DIR / "bin" +LLM_MODELS_DIR = LLM_DIR / "models" +LOGS_DIR = LLM_DIR / "logs" +SERVER_PORT = int(os.environ.get("MODLY_LLM_PORT", "8791")) + +# Modly is first and foremost a 3D-generation app: the LLM must never sit on +# VRAM it isn't using. Idle models are evicted after this many seconds. +IDLE_TTL_SECONDS = int(os.environ.get("MODLY_LLM_IDLE_TTL", "300")) + +# Multi-model pool: each loaded model gets its own llama-server process on its +# own port (SERVER_PORT .. SERVER_PORT+MAX_SLOT_PORTS-1). How many may run at +# once is user-configurable ("auto" sizes it from the GPU's VRAM — small cards +# stay at 1, exactly like the old single-slot behavior). +MAX_SLOT_PORTS = 4 +CONFIG_PATH = LLM_DIR / "config.json" + +LLM_MODELS_DIR.mkdir(parents=True, exist_ok=True) + + +def load_config() -> dict: + try: + return _json.loads(CONFIG_PATH.read_text(encoding="utf-8")) + except Exception: + return {} + + +def save_config(cfg: dict) -> None: + CONFIG_PATH.write_text(_json.dumps(cfg, indent=2), encoding="utf-8") + + +_vram_gb_cache: Optional[float] = None + + +def detect_vram_gb() -> float: + """Total VRAM of the first NVIDIA GPU in GiB (0.0 if none/unknown).""" + global _vram_gb_cache + if _vram_gb_cache is not None: + return _vram_gb_cache + try: + out = subprocess.run( + ["nvidia-smi", "--query-gpu=memory.total", "--format=csv,noheader,nounits"], + capture_output=True, text=True, timeout=10, + ).stdout.strip().splitlines() + _vram_gb_cache = round(float(out[0]) / 1024.0, 1) if out else 0.0 + except Exception: + _vram_gb_cache = 0.0 + return _vram_gb_cache + + +def resolve_max_models() -> int: + """How many llama-server processes may run at once. + + Priority: MODLY_LLM_MAX_MODELS env → user config → auto from VRAM. + Auto is deliberately conservative: 8 GB cards keep the single-slot + behavior, VRAM stays available for the 3D pipeline.""" + raw = os.environ.get("MODLY_LLM_MAX_MODELS") or load_config().get("max_models") or "auto" + if isinstance(raw, str) and raw.lower() == "auto": + # Thresholds sit just under the marketing size on purpose: a card sold + # as "12 GB" reports 12227 MiB = 11.9 GiB, so a `>= 12` test excluded + # every 12 GB card there is (5070, 4070, 3060 12G…) from the 2-slot + # tier it was written for. Same story at 20 for a 20/24 GB card. + vram = detect_vram_gb() + n = 3 if vram >= 19.5 else 2 if vram >= 11.5 else 1 + else: + try: + n = int(raw) + except (TypeError, ValueError): + n = 1 + return max(1, min(n, MAX_SLOT_PORTS)) + +_BINARY_NAME = "llama-server.exe" if sys.platform == "win32" else "llama-server" +_UA = {"User-Agent": "modly-llm"} + +# Pinned to a known-good tag by default so installs are reproducible and never +# silently pick up a broken/incompatible "latest" release. Set MODLY_LLM_RELEASE +# to another llama.cpp tag (e.g. "b10075") to override, or to "latest" to track +# the newest release. +_RELEASE_TAG = os.environ.get("MODLY_LLM_RELEASE", "b10075") +_RELEASES_API = ( + "https://api.github.com/repos/ggml-org/llama.cpp/releases/latest" + if _RELEASE_TAG.lower() == "latest" + else f"https://api.github.com/repos/ggml-org/llama.cpp/releases/tags/{_RELEASE_TAG}" +) + +CATALOG_PATH = Path(__file__).resolve().parent.parent / "resources" / "llm_catalog.json" + +DEFAULT_CTX = 16384 +DEFAULT_NGL = 99 + + +def load_catalog() -> list[dict]: + return _json.loads(CATALOG_PATH.read_text(encoding="utf-8")) + + +_FALLBACK_DEFAULT_MODEL = "qwen3-4b" + +# What the agent sends when a catalog entry declares nothing of its own. Low +# temperature is what keeps small models from emitting malformed tool-call JSON. +_DEFAULT_SAMPLING = {"temperature": 0.2} + + +def sampling_for(model_id: str) -> dict: + """Sampling parameters for one local model. + + Model families publish the settings they were tuned for, and ignoring them + is not neutral: served at a flat temperature with no presence penalty, a + Qwen 3.5 repeated the same lookup ten times in one turn and never acted. + A catalog entry can therefore carry its own `sampling` block.""" + entry = next((e for e in load_catalog() if e["id"] == model_id), None) + declared = (entry or {}).get("sampling") or {} + return {**_DEFAULT_SAMPLING, **{k: v for k, v in declared.items() if v is not None}} + + +def default_model_id() -> str: + """Catalog entry tagged "default", falling back to a known-good id if the + catalog has no such tag (keeps old behavior if llm_catalog.json changes).""" + for entry in load_catalog(): + if "default" in (entry.get("tags") or []): + return entry["id"] + return _FALLBACK_DEFAULT_MODEL + + +def mmproj_local_name(entry: dict) -> str: + """Local filename for a vision projector. HF repos all name it mmproj-F16.gguf, + which would collide across models in the shared models dir.""" + return f"mmproj-{entry['id']}.gguf" + + +def list_models() -> list[dict]: + """Catalog entries with a `downloaded` flag + custom GGUFs found on disk.""" + catalog = load_catalog() + known_files = set() + for entry in catalog: + path = LLM_MODELS_DIR / entry["hf_filename"] + downloaded = path.exists() + if entry.get("hf_mmproj_filename"): + downloaded = downloaded and (LLM_MODELS_DIR / mmproj_local_name(entry)).exists() + known_files.add(mmproj_local_name(entry)) + # Show the real download size (weights + vision projector) + entry["size_bytes"] = (entry.get("size_bytes") or 0) + (entry.get("mmproj_size_bytes") or 0) + entry["downloaded"] = downloaded + entry["source"] = "catalog" + known_files.add(entry["hf_filename"]) + + custom = [ + { + "id": f"custom:{f.name}", + "name": f.stem, + "description": "Custom model found in the models folder.", + "hf_filename": f.name, + "size_bytes": f.stat().st_size, + "downloaded": True, + "source": "custom", + "ctx": DEFAULT_CTX, + "ngl_suggestion": DEFAULT_NGL, + "tags": ["custom"], + } + for f in sorted(LLM_MODELS_DIR.glob("*.gguf")) + if f.name not in known_files and not f.name.lower().startswith("mmproj") + ] + return catalog + custom + + +def estimate_vram_mb(entry: dict) -> int: + """How much VRAM this model is expected to hold, in MiB. + + Catalog entries declare it. A user's own GGUF doesn't, so it is derived from + the file: weights land in VRAM about 1:1, plus the KV cache and compute + buffers for a 16k context. Measured against the catalog's own numbers, whose + declared/file ratio runs 1.22–1.32, so 1.25 + 500 MiB sits in the middle.""" + declared = entry.get("vram_estimate_mb") + if isinstance(declared, (int, float)) and declared > 0: + return int(declared) + size_mb = (entry.get("size_bytes") or 0) / (1024 * 1024) + return int(size_mb * 1.25 + 500) if size_mb else 0 + + +def vram_budget_mb() -> int: + """VRAM the LLM pool may fill, in MiB — total minus what the desktop needs. + 0 when no NVIDIA GPU is detected, which disables budgeting entirely.""" + total = detect_vram_gb() * 1024 + if total <= 0: + return 0 + reserve = _env_int("MODLY_LLM_VRAM_RESERVE_MB", 768) + return max(0, int(total) - reserve) + + +def _env_int(name: str, default: int) -> int: + try: + return int(os.environ.get(name) or default) + except ValueError: + return default + + +def resolve_model(model_id: str) -> dict: + """Return {gguf_path, mmproj_path, ngl, ctx, vision, vram_mb} for a catalog id + or a custom: id.""" + if model_id.startswith("custom:"): + # The name is user/agent-supplied and this path is both loaded and + # DELETEd (DELETE /llm/models/{model_id}). Confine it to the models dir: + # on Windows a backslash is a separator too, so "custom:..\..\x.gguf" + # would otherwise resolve — and unlink — outside it. + name = model_id[len("custom:"):] + path = (LLM_MODELS_DIR / name).resolve() + root = LLM_MODELS_DIR.resolve() + if path.parent != root or not path.exists() or path.suffix != ".gguf": + raise KeyError(f"Custom model not found: {model_id}") + return { + "gguf_path": path, "mmproj_path": None, "ngl": DEFAULT_NGL, + "ctx": DEFAULT_CTX, "vision": False, + "vram_mb": estimate_vram_mb({"size_bytes": path.stat().st_size}), + } + for entry in load_catalog(): + if entry["id"] == model_id: + has_mmproj = bool(entry.get("hf_mmproj_filename")) + return { + "gguf_path": LLM_MODELS_DIR / entry["hf_filename"], + "mmproj_path": LLM_MODELS_DIR / mmproj_local_name(entry) if has_mmproj else None, + "ngl": entry.get("ngl_suggestion", DEFAULT_NGL), + "ctx": entry.get("ctx", DEFAULT_CTX), + "vision": "vision" in (entry.get("tags") or []), + "vram_mb": estimate_vram_mb(entry), + } + raise KeyError(f"Unknown model id: {model_id}") + + +def binary_path() -> Path: + return BIN_DIR / _BINARY_NAME + + +def binary_installed() -> bool: + return binary_path().exists() + + +def has_nvidia_gpu() -> bool: + return shutil.which("nvidia-smi") is not None + + +# ─── Binary bootstrap ───────────────────────────────────────────────────────── + +def _fetch_release() -> dict: + with urlopen(Request(_RELEASES_API, headers=_UA), timeout=30) as r: + return _json.loads(r.read()) + + +def _cuda_ver(name: str) -> tuple[int, int]: + m = re.search(r"cuda-(?:cu)?(\d+)[._](\d+)", name.lower()) + return (int(m.group(1)), int(m.group(2))) if m else (0, 0) + + +def _driver_cuda_version() -> tuple[int, int]: + """Highest CUDA version the installed NVIDIA driver supports (0,0 if unknown).""" + try: + out = subprocess.run(["nvidia-smi"], capture_output=True, text=True, timeout=10).stdout + m = re.search(r"CUDA Version:\s*(\d+)\.(\d+)", out) + return (int(m.group(1)), int(m.group(2))) if m else (0, 0) + except Exception: + return (0, 0) + + +def _pick_assets(assets: list[dict]) -> list[dict]: + """Ordered list of assets to install for this machine. + + Windows + NVIDIA: the CUDA build (self-contained, includes CPU fallback) + plus its matching `cudart-…` runtime package. Otherwise Vulkan, then CPU. + macOS / Linux: the platform tarball (Vulkan variant preferred on Linux). + """ + import platform as _platform + + def matching(ext: tuple[str, ...], *needles: str) -> list[dict]: + return [ + a for a in assets + if a["name"].endswith(ext) and all(n in a["name"].lower() for n in needles) + ] + + if sys.platform == "win32": + if has_nvidia_gpu(): + builds = matching((".zip",), "win", "x64", "cuda") + builds = [b for b in builds if not b["name"].startswith("cudart")] + if builds: + driver = _driver_cuda_version() + # Newest build the driver can run; if the driver version is + # unknown, the oldest build has the widest compatibility. + supported = [b for b in builds if _cuda_ver(b["name"]) <= driver] + pool = supported or builds + pick = max(pool, key=lambda b: _cuda_ver(b["name"])) if supported else \ + min(pool, key=lambda b: _cuda_ver(b["name"])) + ver = _cuda_ver(pick["name"]) + cudart = [ + a for a in assets + if a["name"].startswith("cudart") and _cuda_ver(a["name"]) == ver + ] + return [pick, *cudart[:1]] + return (matching((".zip",), "win", "x64", "vulkan") + or matching((".zip",), "win", "x64", "cpu") + or matching((".zip",), "win", "x64", "avx2"))[:1] + + if sys.platform == "darwin": + arch = "arm64" if _platform.machine().lower() in ("arm64", "aarch64") else "x64" + exts = (".zip", ".tar.gz") + return (matching(exts, "macos", arch) or matching(exts, "macos"))[:1] + + exts = (".zip", ".tar.gz") + return (matching(exts, "ubuntu", "vulkan", "x64") + or matching(exts, "ubuntu", "x64") + or matching(exts, "linux", "x64"))[:1] + + +def _download_asset(asset: dict, progress_cb: Callable[[dict], None], control_check: Callable[[], None], label: str) -> Path: + suffix = ".tar.gz" if asset["name"].endswith(".tar.gz") else ".zip" + fd, tmp_name = tempfile.mkstemp(suffix=suffix) + os.close(fd) + tmp = Path(tmp_name) + sha256 = hashlib.sha256() + try: + with urlopen(Request(asset["browser_download_url"], headers=_UA), timeout=30) as resp: + total = int(resp.headers.get("Content-Length", 0)) or asset.get("size", 0) + done = 0 + last_emit = 0.0 + with open(tmp, "wb") as fh: + while chunk := resp.read(1 << 20): + control_check() + fh.write(chunk) + sha256.update(chunk) + done += len(chunk) + now = time.monotonic() + if now - last_emit >= 0.5: + progress_cb({ + "status": f"Downloading {label} ({asset['name']})", + "bytesDownloaded": done, + "totalBytes": total, + "percent": round(done / total * 100) if total else 0, + }) + last_emit = now + _verify_digest(asset, sha256.hexdigest()) + except BaseException: + tmp.unlink(missing_ok=True) # never leave a partial or rejected archive behind + raise + return tmp + + +def _verify_digest(asset: dict, actual_sha256: str) -> None: + """Refuse an archive whose bytes differ from what GitHub published for it. + + The files extracted from it are executed (llama-server and its DLLs), so a + corrupted or tampered download must not get that far. GitHub reports a + `digest` ("sha256:") for every release asset; one without it (older + uploads) cannot be checked and is accepted as before.""" + expected = asset.get("digest") or "" + if not expected.startswith("sha256:"): + return + if actual_sha256.lower() != expected[len("sha256:"):].lower(): + raise RuntimeError( + f"Checksum mismatch for {asset['name']}: the download does not match the " + "published llama.cpp release. Try installing the engine again." + ) + + +_LIB_SUFFIXES = (".so", ".dylib", ".metal") + + +def _extract_archive(tmp: Path, asset_name: str) -> int: + """Extract runtime files flat into BIN_DIR (exe/dll on Windows, bin/* + libs elsewhere).""" + BIN_DIR.mkdir(parents=True, exist_ok=True) + count = 0 + + def wanted(member_path: str, fname_low: str) -> bool: + if sys.platform == "win32": + return fname_low.endswith((".exe", ".dll")) + return "/bin/" in member_path.replace("\\", "/") or fname_low.endswith(_LIB_SUFFIXES) + + if asset_name.endswith(".tar.gz"): + import tarfile + with tarfile.open(tmp, "r:gz") as tf: + for m in tf.getmembers(): + if not m.isfile(): + continue + fname = Path(m.name).name + if not fname or not wanted(m.name, fname.lower()): + continue + src = tf.extractfile(m) + if src is None: + continue + dest = BIN_DIR / fname + with open(dest, "wb") as dst: + shutil.copyfileobj(src, dst) + if not fname.lower().endswith(_LIB_SUFFIXES): + dest.chmod(0o755) + count += 1 + return count + + with zipfile.ZipFile(tmp) as zf: + for item in zf.infolist(): + if item.is_dir(): + continue + fname = Path(item.filename).name + if not fname or not wanted(item.filename, fname.lower()): + continue + dest = BIN_DIR / fname + with zf.open(item) as src, open(dest, "wb") as dst: + shutil.copyfileobj(src, dst) + if sys.platform != "win32" and not fname.lower().endswith(_LIB_SUFFIXES): + dest.chmod(0o755) + count += 1 + return count + + +def install_binary(progress_cb: Callable[[dict], None], control_check: Callable[[], None]) -> None: + """Download and install the best llama-server build for this machine.""" + progress_cb({"status": "Fetching latest llama.cpp release…", "percent": 0}) + release = _fetch_release() + progress_cb({"status": f"Release {release['tag_name']} — selecting build for this machine…", "percent": 1}) + + assets = _pick_assets(release["assets"]) + if not assets: + raise RuntimeError("No compatible llama.cpp build found for this platform.") + + for i, asset in enumerate(assets): + label = "engine" if i == 0 else "CUDA runtime" + tmp = _download_asset(asset, progress_cb, control_check, label) + try: + count = _extract_archive(tmp, asset["name"]) + progress_cb({"status": f"Extracted {count} files from {asset['name']}"}) + finally: + tmp.unlink(missing_ok=True) + + if not binary_installed(): + raise RuntimeError("Install finished but llama-server binary is missing.") + progress_cb({"status": "done", "percent": 100}) + + +# ─── Orphan prevention ──────────────────────────────────────────────────────── +# If Modly is force-killed, a plain child process would keep its model in +# memory forever. On Windows we put llama-server in a Job object with +# KILL_ON_JOB_CLOSE (the OS kills it when Modly dies); on Linux we ask the +# kernel to deliver SIGKILL on parent death. + +_win_job_handle = None + + +def _windows_job() -> Optional[int]: + global _win_job_handle + if _win_job_handle is not None: + return _win_job_handle + try: + import ctypes + from ctypes import wintypes + + class JOBOBJECT_BASIC_LIMIT_INFORMATION(ctypes.Structure): + _fields_ = [ + ("PerProcessUserTimeLimit", ctypes.c_int64), + ("PerJobUserTimeLimit", ctypes.c_int64), + ("LimitFlags", wintypes.DWORD), + ("MinimumWorkingSetSize", ctypes.c_size_t), + ("MaximumWorkingSetSize", ctypes.c_size_t), + ("ActiveProcessLimit", wintypes.DWORD), + ("Affinity", ctypes.c_size_t), + ("PriorityClass", wintypes.DWORD), + ("SchedulingClass", wintypes.DWORD), + ] + + class IO_COUNTERS(ctypes.Structure): + _fields_ = [(n, ctypes.c_uint64) for n in ( + "ReadOperationCount", "WriteOperationCount", "OtherOperationCount", + "ReadTransferCount", "WriteTransferCount", "OtherTransferCount")] + + class JOBOBJECT_EXTENDED_LIMIT_INFORMATION(ctypes.Structure): + _fields_ = [ + ("BasicLimitInformation", JOBOBJECT_BASIC_LIMIT_INFORMATION), + ("IoInfo", IO_COUNTERS), + ("ProcessMemoryLimit", ctypes.c_size_t), + ("JobMemoryLimit", ctypes.c_size_t), + ("PeakProcessMemoryUsed", ctypes.c_size_t), + ("PeakJobMemoryUsed", ctypes.c_size_t), + ] + + kernel32 = ctypes.windll.kernel32 + job = kernel32.CreateJobObjectW(None, None) + info = JOBOBJECT_EXTENDED_LIMIT_INFORMATION() + info.BasicLimitInformation.LimitFlags = 0x2000 # JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE + kernel32.SetInformationJobObject(job, 9, ctypes.byref(info), ctypes.sizeof(info)) # JobObjectExtendedLimitInformation + _win_job_handle = job + return job + except Exception: + return None + + +def _tie_to_parent(process: subprocess.Popen) -> None: + if sys.platform == "win32": + job = _windows_job() + if job is not None: + try: + import ctypes + ctypes.windll.kernel32.AssignProcessToJobObject(job, int(process._handle)) # type: ignore[attr-defined] + except Exception: + pass + + +def _linux_preexec(): + try: + import ctypes + import signal as _signal + libc = ctypes.CDLL("libc.so.6", use_errno=True) + libc.prctl(1, _signal.SIGKILL) # PR_SET_PDEATHSIG + except Exception: + pass + + +def _kill_stale_server(port: int = SERVER_PORT) -> None: + """Kill a leftover llama-server holding `port` (e.g. after a force-kill of Modly).""" + try: + with urlopen(Request(f"http://127.0.0.1:{port}/health", headers=_UA), timeout=1) as r: + if r.status != 200: + return + except Exception: + return # nothing listening — the normal case + try: + if sys.platform == "win32": + out = subprocess.run( + ["netstat", "-ano", "-p", "tcp"], capture_output=True, text=True, timeout=10, + ).stdout + for line in out.splitlines(): + if f":{port}" in line and "LISTENING" in line.upper(): + pid = line.split()[-1] + name = subprocess.run( + ["tasklist", "/FI", f"PID eq {pid}", "/FO", "CSV", "/NH"], + capture_output=True, text=True, timeout=10, + ).stdout.lower() + if "llama-server" in name: + subprocess.run(["taskkill", "/PID", pid, "/F"], capture_output=True, timeout=10) + break + else: + out = subprocess.run( + ["lsof", "-ti", f"tcp:{port}"], capture_output=True, text=True, timeout=10, + ).stdout + for pid in out.split(): + name = subprocess.run( + ["ps", "-p", pid, "-o", "comm="], capture_output=True, text=True, timeout=10, + ).stdout.lower() + if "llama-server" in name: + subprocess.run(["kill", "-9", pid], capture_output=True, timeout=10) + except Exception: + pass + + +# ─── Managed server ─────────────────────────────────────────────────────────── + +class LlamaServerManager: + """One llama-server slot — a single loaded GGUF on a dedicated port. + + Slots are owned by LlamaPool, which decides how many run at once, evicts + idle ones (IDLE_TTL_SECONDS), and unloads everything when the 3D pipeline + needs the VRAM. + """ + + def __init__(self, port: int = SERVER_PORT) -> None: + self.port = port + self._lock = threading.Lock() + self._process: Optional[subprocess.Popen] = None + self._current_id: Optional[str] = None + self._started_at: Optional[float] = None + self._last_used: float = 0.0 + self._log_file = None + # Requests currently being answered by this slot. Guarded by its own + # lock: _lock is held for the whole of a cold start (up to 180 s), and a + # request must never wait on that to declare itself in flight. + self._inflight: int = 0 + self._inflight_lock = threading.Lock() + # Expected VRAM of whatever this slot is serving, for the pool's budget. + self.vram_mb: int = 0 + + @property + def base_url(self) -> str: + return f"http://127.0.0.1:{self.port}/v1" + + def _alive(self) -> bool: + return self._process is not None and self._process.poll() is None + + def ensure(self, model_id: str, spec: dict) -> None: + """Blocking: make sure `model_id` is loaded, swapping the server if needed. + `spec` comes from resolve_model().""" + with self._lock: + self._last_used = time.monotonic() + if self._current_id == model_id and self._alive(): + return + if not binary_installed(): + raise RuntimeError("llama-server is not installed. Install the engine in Settings → Agent.") + if not spec["gguf_path"].exists(): + raise RuntimeError(f"Model file not found: {spec['gguf_path'].name}. Download it in Settings → Agent.") + + self._terminate_locked() + _kill_stale_server(self.port) + self._spawn(spec) + self._current_id = model_id + self._started_at = time.monotonic() + self._last_used = time.monotonic() + + def touch(self) -> None: + """Mark the server as recently used (postpones idle eviction).""" + self._last_used = time.monotonic() + + @property + def busy_count(self) -> int: + with self._inflight_lock: + return self._inflight + + def hold(self) -> None: + """Claim the slot for a request that is about to start.""" + with self._inflight_lock: + self._inflight += 1 + self._last_used = time.monotonic() + + def release(self) -> None: + with self._inflight_lock: + self._inflight -= 1 + self._last_used = time.monotonic() + + @contextlib.contextmanager + def busy(self): + """Hold the slot for the duration of one request. + + _last_used only moved when a completion *finished*, so a generation + longer than IDLE_TTL_SECONDS — a Text-to-CAD node parked on /llm/chat, + or an agent round on a 14B running on CPU — looked idle to the reaper, + which terminated the server mid-answer: truncated stream, node failed + with no message. Marking the request in flight covers the whole call, + including the prompt-eval stall before the first token.""" + self.hold() + try: + yield self + finally: + self.release() + + def _spawn(self, spec: dict) -> None: + n_threads = max(1, int((os.cpu_count() or 4) * 0.8)) + cmd = [ + str(binary_path()), + "-m", str(spec["gguf_path"]), + "--port", str(self.port), + "--host", "127.0.0.1", + "-c", str(spec["ctx"]), + "-ngl", str(spec["ngl"]), + "--jinja", + "-np", "1", + "--threads", str(n_threads), + "--threads-batch", str(n_threads), + ] + if spec.get("mmproj_path"): + cmd += ["--mmproj", str(spec["mmproj_path"])] + if has_nvidia_gpu(): + # Halve KV-cache VRAM so the larger context still fits 12 GB cards. + # V-cache quantization requires flash attention (fine on CUDA). + cmd += ["-fa", "on", "-ctk", "q8_0", "-ctv", "q8_0"] + # llama-server resolves its plugin DLLs (ggml-cuda, ggml-cpu-*) relative + # to cwd, not the exe path — running from bin/ is required. + env = {**os.environ, "PATH": str(BIN_DIR) + os.pathsep + os.environ.get("PATH", "")} + kwargs: dict = {} + if sys.platform.startswith("linux"): + kwargs["preexec_fn"] = _linux_preexec + + # Overwritten on every spawn (one log per slot/port) so a crash at + # startup (OOM, corrupt GGUF) can be diagnosed instead of vanishing + # into DEVNULL. + LOGS_DIR.mkdir(parents=True, exist_ok=True) + self._log_file = open(LOGS_DIR / f"slot-{self.port}.log", "w", encoding="utf-8") + + self._process = subprocess.Popen( + cmd, + stdout=self._log_file, + stderr=subprocess.STDOUT, + cwd=str(BIN_DIR), + env=env, + **kwargs, + ) + _tie_to_parent(self._process) + self._wait_for_health() + + def _wait_for_health(self, timeout: float = 180.0, interval: float = 1.0) -> None: + url = f"http://127.0.0.1:{self.port}/health" + deadline = time.time() + timeout + while time.time() < deadline: + if not self._alive(): + raise RuntimeError("llama-server exited during startup (bad model file or out of memory?).") + try: + with urlopen(Request(url, headers=_UA), timeout=2) as r: + if r.status == 200: + return + except Exception: + pass + time.sleep(interval) + self._terminate_locked() + raise TimeoutError(f"llama-server did not become ready within {timeout:.0f}s") + + def unload(self) -> None: + with self._lock: + self._terminate_locked() + + def _terminate_locked(self) -> None: + if self._process is None: + return + try: + self._process.terminate() + try: + self._process.wait(timeout=10) + except subprocess.TimeoutExpired: + self._process.kill() + except Exception: + pass + finally: + self._process = None + self._current_id = None + self._started_at = None + if self._log_file is not None: + self._log_file.close() + self._log_file = None + + def snapshot(self) -> dict: + alive = self._alive() + return { + "alive": alive, + "model_id": self._current_id if alive else None, + "port": self.port, + "uptime_seconds": round(time.monotonic() - self._started_at, 1) if alive and self._started_at else None, + "vram_mb": self.vram_mb if alive else None, + } + + +class LlamaPool: + """Pool of llama-server slots — one process per loaded model, each on its + own port, capped by BOTH resolve_max_models() and vram_budget_mb() + (LRU eviction). + + The idle reaper and unload_all() keep the old guarantees: models are + evicted after IDLE_TTL_SECONDS unused, and the 3D pipeline can reclaim all + VRAM at once. + """ + + def __init__(self) -> None: + self._lock = threading.Lock() + self._slots: dict[str, LlamaServerManager] = {} + # Ports held by a load in flight. A reserved slot has no process yet, so + # _alive() cannot tell it apart from a dead one; tracking the port rather + # than the model_id keeps the reservation valid even if unload() drops + # the slot from _slots while it is still starting. + self._loading_ports: set[int] = set() + # Signalled whenever a load lands, so callers held back by a concurrent + # cold start can re-run the limit check instead of loading on top of it. + self._cond = threading.Condition(self._lock) + self._reaper_started = False + + def ensure(self, model_id: str, spec: dict, hold: bool = False) -> LlamaServerManager: + """Blocking: return a ready slot serving `model_id`, starting it (and + evicting the least-recently-used slots) if needed. + + With `hold`, the slot is returned already claimed for one request and the + caller MUST call slot.release() when done (see busy()). Without it there + is a window between "loaded" and "in flight" in which a competing load + can evict the slot, and the caller then posts to a dead server: with + max_models = 1, two nodes asking for different models had one of them + answer 500 Internal Server Error. + + The slot is reserved under the pool lock but LOADED outside it. A cold + start blocks in _wait_for_health for up to 180 s, and holding the pool + lock that long froze every other caller: /llm/status, polled by an open + Settings → Agent page, hung for the whole load, and switch_model's + unload_all() — which exists to reclaim VRAM before a 3D model — queued + behind the very LLM competing for it. + + Loading outside the lock means two callers can be starting different + models at the same time (an agent turn plus a workflow LLM node). A slot + being started holds no process, so it cannot be evicted to make room: + when only such loads keep the pool over its limit, this waits for one to + land — it then becomes an ordinary eviction candidate — instead of + putting a second model on a card sized for one. + """ + incoming_mb = spec.get("vram_mb") or 0 + # A last resort, not a policy: past this the thing we are waiting for is + # stuck, and blocking forever would be worse than loading on top of it. + deadline = time.monotonic() + 300.0 + with self._lock: + while True: + slot = self._slots.get(model_id) + if slot is not None and slot._alive(): + break + if slot is not None and slot.port not in self._loading_ports: + # A server that crashed (OOM, CUDA error) leaves its slot + # behind, and its port is free for the next model. Respawning + # in place would then find that model's server healthy on + # "our" port and kill it. Start over on a fresh port instead. + self._drop_dead_locked(model_id, slot) + slot = None + # Also for a dead slot respawned in place, so reviving one + # cannot push the pool past the limit. Only alive slots that are + # not answering a request are eviction candidates, so this never + # evicts our own — nor anyone's live stream. + self._enforce_limit_locked(reserve=1, incoming_mb=incoming_mb) + remaining = deadline - time.monotonic() + if (not self._over_capacity_locked(incoming_mb) + or not self._transient_blockers_locked() + or remaining <= 0): + break + # Re-checked about once a second: capacity also frees up when a + # request finishes, which is not a load event and so notifies + # nothing. + self._cond.wait(min(1.0, remaining)) + if slot is None: + used_ports = {s.port for s in self._slots.values() if s._alive()} | self._loading_ports + port = next( + (SERVER_PORT + i for i in range(MAX_SLOT_PORTS) + if SERVER_PORT + i not in used_ports), + None, + ) + if port is None: + # A bare next() raised StopIteration here, which surfaced to + # the caller as "Could not start the local LLM: " — an empty + # message. Say what actually happened. + raise RuntimeError( + f"All {MAX_SLOT_PORTS} local-LLM slots are in use or starting up. " + "Wait for a run to finish, or lower the model limit in Settings → Agent." + ) + slot = LlamaServerManager(port) + self._slots[model_id] = slot + slot.vram_mb = incoming_mb + self._loading_ports.add(slot.port) + if hold: + # Claimed BEFORE the load, not after: a caller woken by this load + # landing would otherwise find the slot idle and evict it, and + # two racing callers spent their time killing each other's + # freshly loaded model until one gave up. + slot.hold() + self._start_reaper() + + # Outside the pool lock. LlamaServerManager holds its own, so concurrent + # callers for the same model serialise here and the later ones no-op. + try: + slot.ensure(model_id, spec) + except Exception: + if hold: + slot.release() # nothing to answer with — do not pin the slot + with self._lock: + if self._slots.get(model_id) is slot and not slot._alive(): + self._slots.pop(model_id, None) + raise + finally: + with self._lock: + self._loading_ports.discard(slot.port) + self._cond.notify_all() # waiters can re-run the limit check + return slot + + def _evictable_locked(self, slot: LlamaServerManager) -> bool: + """May this slot be unloaded to make room? Shared by the limit check and + the idle reaper, which both unload while holding the pool lock — a rule + applied to only one of them is a rule that does not hold. + + A slot answering a request is spared because terminating it truncates + the caller's stream mid-generation. A slot another thread is loading is + spared for a harder reason: from the moment Popen returns it is _alive() + (the long wait is _wait_for_health, up to 180 s) while its own lock is + held by the loader, so unload() would block on that lock and freeze the + whole pool with it — /llm/status, every ensure(), and the unload_all() + the 3D pipeline needs to reclaim VRAM. Its cost is still counted, it + just cannot be the victim.""" + return slot.busy_count == 0 and slot.port not in self._loading_ports + + def _enforce_limit_locked(self, reserve: int = 0, incoming_mb: int = 0) -> None: + """Evict LRU slots until the pool fits BOTH limits: the configured model + count, and the VRAM budget. + + The count alone was not enough. `max_models: 2` let any two models in, + so a 9 GB custom model plus a vision model on a 12 GB card oversubscribed + the card; Windows spills to shared memory instead of failing, so nothing + errored — a cold start just went from 5 s to 24 s with no explanation. + Budgeting on declared estimates keeps the pairs that actually fit (a 4B + and a 7B come to 10.4 of 11.5 GiB) and refuses the ones that never did. + + The incoming model is never rejected: if it does not fit even alone, + everything else is evicted and it loads by itself.""" + alive = [(mid, s) for mid, s in self._slots.items() if s._alive()] + alive.sort(key=lambda kv: kv[1]._last_used) # oldest first + # Slots reserved by a concurrent ensure() hold no process yet, so + # _alive() cannot see them. Counting only live slots let two concurrent + # ensure() calls — an agent turn plus a workflow LLM node — each + # conclude the pool was empty and both load, which on an 8 GB card + # (max_models = 1) is exactly the oversubscription this rule prevents. + loading = [s for s in self._slots.values() + if s.port in self._loading_ports and not s._alive()] + candidates = [kv for kv in alive if self._evictable_locked(kv[1])] + + def _evict_oldest() -> None: + mid, slot = candidates.pop(0) + alive.remove((mid, slot)) + slot.unload() + self._slots.pop(mid, None) + + max_n = resolve_max_models() + while len(alive) + len(loading) + reserve > max_n and candidates: + _evict_oldest() + + budget = vram_budget_mb() + if not budget or not incoming_mb: + return # no GPU detected, or an unknown estimate — count rule only + + def _committed_mb() -> int: + return sum(s.vram_mb for _mid, s in alive) + sum(s.vram_mb for s in loading) + + while candidates and _committed_mb() + incoming_mb > budget: + _evict_oldest() + + def _transient_blockers_locked(self) -> bool: + """Is the pool full only of things that end on their own — a load in + flight, or a slot answering a request? Those are worth waiting for; a + merely idle slot is not (it gets evicted instead).""" + return bool(self._loading_ports) or any( + s.busy_count for s in self._slots.values() if s._alive() + ) + + def _over_capacity_locked(self, incoming_mb: int) -> bool: + """Would loading one more model break either limit, counting the slots + another thread is starting right now?""" + alive = [s for s in self._slots.values() if s._alive()] + loading = [s for s in self._slots.values() + if s.port in self._loading_ports and not s._alive()] + if len(alive) + len(loading) + 1 > resolve_max_models(): + return True + budget = vram_budget_mb() + if not budget or not incoming_mb: + return False # no GPU detected, or an unknown estimate + committed = sum(s.vram_mb for s in alive) + sum(s.vram_mb for s in loading) + return committed + incoming_mb > budget + + def enforce_limit(self) -> None: + """Apply the configured limit right away (used when the user lowers it).""" + with self._lock: + self._enforce_limit_locked() + + def is_loaded(self, model_id: str) -> bool: + with self._lock: + slot = self._slots.get(model_id) + return slot is not None and slot._alive() + + def touch(self, model_id: str) -> None: + with self._lock: + slot = self._slots.get(model_id) + if slot is not None: + slot.touch() + + def unload(self, model_id: str) -> None: + with self._lock: + slot = self._slots.pop(model_id, None) + if slot is not None: + slot.unload() + + def unload_all(self, force: bool = False) -> None: + """Free the pool's VRAM. Slots answering a request or being started are + spared (see _evictable_locked): killing one cut a Text-to-CAD stream off + mid-answer when a 3D model loaded, and a cold start in another thread + handed its caller a dead slot. `force` is for shutdown only.""" + if force: + with self._lock: + slots = list(self._slots.values()) + self._slots.clear() + for slot in slots: + slot.unload() + return + with self._lock: + for mid, slot in list(self._slots.items()): + if self._evictable_locked(slot): + slot.unload() + self._slots.pop(mid, None) + + def _drop_dead_locked(self, model_id: str, slot: LlamaServerManager) -> None: + self._slots.pop(model_id, None) + slot.unload() # reaps the exited process and closes its log + + def _start_reaper(self) -> None: + if self._reaper_started or IDLE_TTL_SECONDS <= 0: + return + self._reaper_started = True + threading.Thread(target=self._reap_idle, daemon=True, name="llm-idle-reaper").start() + + def _reap_idle(self) -> None: + while True: + time.sleep(15) + self._reap_once() + + def _reap_once(self, now: Optional[float] = None) -> None: + now = time.monotonic() if now is None else now + with self._lock: + for mid, slot in list(self._slots.items()): + if not self._evictable_locked(slot): + continue # answering right now, or being loaded — never idle + if not slot._alive(): + self._drop_dead_locked(mid, slot) + elif now - slot._last_used > IDLE_TTL_SECONDS: + slot.unload() + self._slots.pop(mid, None) + + def snapshot(self) -> dict: + with self._lock: + slots = list(self._slots.values()) + servers = [s.snapshot() for s in slots if s._alive()] + first = servers[0] if servers else {"alive": False, "model_id": None, "port": SERVER_PORT, "uptime_seconds": None} + return { + **first, # legacy single-server shape + "servers": servers, + "max_models": resolve_max_models(), + "vram_gb": detect_vram_gb() or None, + "vram_budget_mb": vram_budget_mb() or None, + "vram_used_mb": sum(s.get("vram_mb") or 0 for s in servers) or None, + } + + +llama_pool = LlamaPool() diff --git a/api/services/mesh_ops/__init__.py b/api/services/mesh_ops/__init__.py new file mode 100644 index 00000000..40f9e61c --- /dev/null +++ b/api/services/mesh_ops/__init__.py @@ -0,0 +1,26 @@ +"""Unified mesh operation registry used by the API and workflow nodes.""" + +from .builtin import BUILTIN_MESH_OPS +from .registry import MeshOpsRegistry +from .types import ( + MeshOp, + MeshOpContext, + MeshOpExecutionError, + MeshOpNotFoundError, + MeshOpResult, + MeshOpUnavailableError, +) + + +mesh_ops_registry = MeshOpsRegistry(BUILTIN_MESH_OPS) + +__all__ = [ + "MeshOp", + "MeshOpContext", + "MeshOpExecutionError", + "MeshOpNotFoundError", + "MeshOpResult", + "MeshOpsRegistry", + "MeshOpUnavailableError", + "mesh_ops_registry", +] diff --git a/api/services/mesh_ops/builtin.py b/api/services/mesh_ops/builtin.py new file mode 100644 index 00000000..900bd193 --- /dev/null +++ b/api/services/mesh_ops/builtin.py @@ -0,0 +1,131 @@ +"""Built-in mesh operation definitions.""" + +from .operations import decimate_mesh, repair_mesh, smooth_mesh +from .types import MeshOp + + +REPAIR_PARAMS = ( + { + "id": "remove_duplicates", + "label": "Remove Duplicates", + "type": "boolean", + "default": True, + "tooltip": "Remove duplicate vertices and faces.", + }, + { + "id": "remove_degenerate", + "label": "Remove Degenerate Faces", + "type": "boolean", + "default": True, + "tooltip": "Remove zero-area faces and collapsed edges.", + }, + { + "id": "fix_non_manifold", + "label": "Fix Non-Manifold", + "type": "boolean", + "default": True, + "tooltip": "Detach faces causing non-manifold edges.", + }, + { + "id": "fill_holes", + "label": "Fill Holes", + "type": "boolean", + "default": True, + "tooltip": ( + "Fill simple boundary holes. Structural holes from AI generation " + "may not be fillable in post-processing." + ), + }, + { + "id": "max_hole_size", + "label": "Max Hole Size", + "type": "int", + "default": 2000, + "min": 10, + "max": 10000, + "tooltip": ( + "Maximum number of boundary edges of a hole to be filled. " + "Increase if large holes remain open." + ), + }, +) + +DECIMATE_PARAMS = ( + { + "id": "target_faces", + "label": "Target Triangles", + "type": "int", + "default": 10000, + "min": 100, + "max": 1000000, + "tooltip": "Target number of triangles after simplification.", + }, +) + +SMOOTH_PARAMS = ( + { + "id": "iterations", + "label": "Iterations", + "type": "int", + "default": 5, + "min": 1, + "max": 50, + "tooltip": ( + "Number of smoothing passes. More iterations = smoother result " + "but may lose fine details." + ), + }, + { + "id": "lambda_", + "label": "Smoothing Strength", + "type": "float", + "default": 0.5, + "min": 0.1, + "max": 1.0, + "step": 0.05, + "tooltip": ( + "Controls how far each vertex moves toward its neighbours per " + "iteration. Lower = more conservative." + ), + }, + { + "id": "mode", + "label": "Mode", + "type": "select", + "default": "taubin", + "options": [ + {"value": "taubin", "label": "Taubin (volume-preserving)"}, + {"value": "laplacian", "label": "Laplacian (stronger, may shrink)"}, + ], + "tooltip": ( + "Taubin alternates positive/negative steps to prevent mesh " + "shrinkage. Laplacian is simpler but tends to shrink the mesh over " + "many iterations." + ), + }, +) + + +BUILTIN_MESH_OPS = ( + MeshOp( + id="repair", + label="Repair Mesh", + params_schema=REPAIR_PARAMS, + fn=repair_mesh, + category="repair", + ), + MeshOp( + id="decimate", + label="Optimize Mesh", + params_schema=DECIMATE_PARAMS, + fn=decimate_mesh, + category="optimization", + ), + MeshOp( + id="smooth", + label="Smooth Mesh", + params_schema=SMOOTH_PARAMS, + fn=smooth_mesh, + category="optimization", + ), +) diff --git a/api/services/mesh_ops/meshopt_runner.cjs b/api/services/mesh_ops/meshopt_runner.cjs new file mode 100644 index 00000000..3436be44 --- /dev/null +++ b/api/services/mesh_ops/meshopt_runner.cjs @@ -0,0 +1,120 @@ +/** + * meshoptimizer backend for the Python mesh-op registry. + * + * Dependencies are resolved from the built-in mesh-optimizer extension so the + * packaged app keeps one copy of glTF Transform and meshoptimizer. + */ +const fs = require('fs') +const path = require('path') +const Module = require('module') + +function emit(message) { + process.stdout.write(`${JSON.stringify(message)}\n`) +} + +function progress(percent, label) { + emit({ type: 'progress', percent, label }) +} + +function log(message) { + emit({ type: 'log', message: String(message) }) +} + +function countTriangles(document) { + let count = 0 + for (const mesh of document.getRoot().listMeshes()) { + for (const primitive of mesh.listPrimitives()) { + const indices = primitive.getIndices() + if (indices) { + count += Math.round(indices.getCount() / 3) + } else { + const positions = primitive.getAttribute('POSITION') + if (positions) count += Math.round(positions.getCount() / 3) + } + } + } + return count +} + +async function run(payload) { + const requireExtension = Module.createRequire( + path.join(payload.dependencyDir, 'package.json'), + ) + const { NodeIO } = requireExtension('@gltf-transform/core') + const { ALL_EXTENSIONS } = requireExtension('@gltf-transform/extensions') + const { simplify, weld } = requireExtension('@gltf-transform/functions') + const { MeshoptSimplifier } = requireExtension('meshoptimizer') + + const targetFaces = Math.max( + 100, + Math.round(Number(payload.params?.target_faces ?? 10000)), + ) + log(`Target: ${targetFaces} triangles — input: ${payload.inputPath}`) + + await MeshoptSimplifier.ready + + progress(10, 'Loading mesh…') + const io = new NodeIO().registerExtensions(ALL_EXTENSIONS) + const document = await io.read(payload.inputPath) + const currentFaces = countTriangles(document) + log(`Current triangles: ${currentFaces}`) + + if (currentFaces <= targetFaces) { + log('Already within target — skipping simplification') + if (!payload.outputPath) { + progress(100, 'Done') + return { filePath: payload.inputPath, faceCount: currentFaces } + } + + fs.mkdirSync(path.dirname(payload.outputPath), { recursive: true }) + progress(85, 'Writing output…') + await io.write(payload.outputPath, document) + progress(100, 'Done') + log(`Output: ${payload.outputPath}`) + return { filePath: payload.outputPath, faceCount: currentFaces } + } + + const ratio = Math.min(1, targetFaces / currentFaces) + log( + `Simplification ratio: ${ratio.toFixed(4)} ` + + `(~${Math.round(currentFaces * ratio)} triangles)`, + ) + const error = Math.max(0.001, 1 - ratio) + + if (currentFaces < 500000) { + progress(25, 'Welding vertices…') + await document.transform(weld()) + } else { + log(`Skipping weld (${currentFaces} faces > 500k threshold)`) + } + + progress(55, 'Simplifying mesh…') + await document.transform( + simplify({ simplifier: MeshoptSimplifier, ratio, error, lockBorder: false }), + ) + + progress(85, 'Writing output…') + const outputPath = payload.outputPath || path.join( + payload.workspaceDir, + 'Workflows', + `mesh-optimizer-${Date.now()}.glb`, + ) + fs.mkdirSync(path.dirname(outputPath), { recursive: true }) + await io.write(outputPath, document) + + progress(100, 'Done') + log(`Output: ${outputPath}`) + return { filePath: outputPath, faceCount: countTriangles(document) } +} + +async function main() { + const raw = fs.readFileSync(0, 'utf8').trim() + if (!raw) throw new Error('mesh-optimizer: missing request payload') + const result = await run(JSON.parse(raw)) + emit({ type: 'done', result }) +} + +main().catch((error) => { + emit({ type: 'error', message: String(error) }) + process.exitCode = 1 +}) diff --git a/api/services/mesh_ops/operations.py b/api/services/mesh_ops/operations.py new file mode 100644 index 00000000..7fae3845 --- /dev/null +++ b/api/services/mesh_ops/operations.py @@ -0,0 +1,419 @@ +"""Canonical implementations for Modly's built-in mesh operations.""" + +import json +import os +import re +import shutil +import subprocess +import tempfile +import time +from pathlib import Path +from typing import Any, Mapping + +from .types import ( + MeshOpContext, + MeshOpExecutionError, + MeshOpResult, + MeshOpUnavailableError, +) + + +def _output_path(context: MeshOpContext, prefix: str) -> Path: + if context.output_path is not None: + output = Path(context.output_path) + else: + output = ( + context.workspace_dir + / "Workflows" + / f"{prefix}-{int(time.time() * 1000)}.glb" + ) + output.parent.mkdir(parents=True, exist_ok=True) + return output + + +def _load_single_mesh(input_path: Path, trimesh_module): + loaded = trimesh_module.load(str(input_path)) + if isinstance(loaded, trimesh_module.Scene): + geometries = list(loaded.geometry.values()) + return ( + trimesh_module.util.concatenate(geometries) + if len(geometries) > 1 + else geometries[0] + ) + return loaded + + +def _raw_geometry(mesh, trimesh_module): + loaded = trimesh_module.load(mesh, process=False) + if isinstance(loaded, trimesh_module.Scene): + geometries = list(loaded.geometry.values()) + loaded = ( + geometries[0] + if len(geometries) == 1 + else trimesh_module.util.concatenate(geometries) + ) + return trimesh_module.Trimesh( + vertices=loaded.vertices, + faces=loaded.faces, + process=False, + ) + + +def _face_count(mesh, trimesh_module) -> int: + if isinstance(mesh, trimesh_module.Scene): + return sum(len(geometry.faces) for geometry in mesh.geometry.values()) + return int(len(mesh.faces)) + + +def _has_texture(geometry, trimesh_module) -> bool: + if not isinstance(geometry.visual, trimesh_module.visual.TextureVisuals): + return False + material = geometry.visual.material + if material is None: + return False + return ( + getattr(material, "image", None) is not None + or getattr(material, "baseColorTexture", None) is not None + ) + + +def _texture_image(geometry): + material = geometry.visual.material + image = getattr(material, "image", None) + return image if image is not None else getattr(material, "baseColorTexture", None) + + +def _point_mtl_at_texture(mtl_path: str) -> None: + path = Path(mtl_path) + if not path.exists(): + return + contents = path.read_text(encoding="utf-8") + path.write_text( + re.sub(r"map_Kd\s+\S+", "map_Kd texture.png", contents), + encoding="utf-8", + ) + + +def _mesh_libraries(operation_name: str): + try: + import pymeshlab + except ImportError as exc: + raise MeshOpUnavailableError( + f"{operation_name}: pymeshlab is not available on this system" + ) from exc + + try: + import trimesh + except ImportError as exc: + raise MeshOpUnavailableError( + f"{operation_name}: trimesh is not available on this system" + ) from exc + + return pymeshlab, trimesh + + +def repair_mesh( + input_path: Path, + params: Mapping[str, Any], + context: MeshOpContext, +) -> MeshOpResult: + """Run the exact repair pipeline previously owned by mesh-repair.""" + pymeshlab, trimesh = _mesh_libraries("mesh-repair") + + remove_duplicates = bool(params.get("remove_duplicates", True)) + fix_non_manifold = bool(params.get("fix_non_manifold", True)) + remove_degenerate = bool(params.get("remove_degenerate", True)) + fill_holes = bool(params.get("fill_holes", True)) + max_hole_size = int(params.get("max_hole_size", 2000)) + output_path = _output_path(context, "mesh-repair") + + context.progress(10, "Loading mesh…") + geometry = _load_single_mesh(input_path, trimesh) + + temporary_dir = tempfile.mkdtemp() + try: + ply_input = os.path.join(temporary_dir, "input.ply") + ply_output = os.path.join(temporary_dir, "output.ply") + geometry.export(ply_input) + + mesh_set = pymeshlab.MeshSet() + mesh_set.load_new_mesh(ply_input) + + current = mesh_set.current_mesh() + context.log( + f"Input: {current.vertex_number()} verts, " + f"{current.face_number()} faces" + ) + + if remove_duplicates: + context.progress(20, "Removing duplicates…") + mesh_set.meshing_remove_duplicate_vertices() + mesh_set.meshing_remove_duplicate_faces() + + if remove_degenerate: + context.progress(40, "Removing degenerate faces…") + mesh_set.meshing_remove_null_faces() + mesh_set.meshing_remove_folded_faces() + + if fix_non_manifold: + context.progress(60, "Fixing non-manifold edges…") + try: + mesh_set.meshing_repair_non_manifold_edges(method=0) + except Exception as exc: + context.log(f"Non-manifold edge repair skipped: {exc}") + try: + mesh_set.meshing_repair_non_manifold_vertices() + except Exception as exc: + context.log(f"Non-manifold vertex repair skipped: {exc}") + + if fill_holes: + context.progress(75, "Filling holes…") + try: + mesh_set.meshing_close_holes( + maxholesize=max_hole_size, + newfaceselected=False, + selfintersection=False, + ) + except Exception as exc: + context.log( + "Hole fill skipped (mesh may still be non-manifold): " + f"{exc}" + ) + + current = mesh_set.current_mesh() + context.log( + f"Output: {current.vertex_number()} verts, " + f"{current.face_number()} faces" + ) + + context.progress(85, "Exporting…") + mesh_set.save_current_mesh(ply_output) + result = _raw_geometry(ply_output, trimesh) + finally: + shutil.rmtree(temporary_dir, ignore_errors=True) + + result.export(str(output_path)) + context.progress(100, "Done") + return MeshOpResult( + file_path=output_path, + details={"face_count": int(len(result.faces))}, + ) + + +def smooth_mesh( + input_path: Path, + params: Mapping[str, Any], + context: MeshOpContext, +) -> MeshOpResult: + """Run the exact Taubin/Laplacian pipeline previously owned by mesh-smoother.""" + pymeshlab, trimesh = _mesh_libraries("mesh-smoother") + + iterations = int(params.get("iterations", 5)) + strength = float(params.get("lambda_", 0.5)) + mode = str(params.get("mode", "taubin")) + output_path = _output_path(context, "mesh-smoother") + + context.log( + f"Mode: {mode}, iterations: {iterations}, strength: {strength}" + ) + context.progress(10, "Loading mesh…") + geometry = _load_single_mesh(input_path, trimesh) + + temporary_dir = tempfile.mkdtemp() + try: + mesh_set = pymeshlab.MeshSet() + if context.preserve_visuals: + context.progress(30, "Smoothing (laplacian)…") + if _has_texture(geometry, trimesh): + obj_input = os.path.join(temporary_dir, "input.obj") + texture_input = os.path.join(temporary_dir, "texture.png") + obj_output = os.path.join(temporary_dir, "output.obj") + + _texture_image(geometry).save(texture_input) + geometry.export(obj_input) + _point_mtl_at_texture(os.path.join(temporary_dir, "input.mtl")) + + mesh_set.load_new_mesh(obj_input) + mesh_set.apply_coord_laplacian_smoothing( + stepsmoothnum=iterations, + ) + context.progress(80, "Exporting…") + mesh_set.save_current_mesh(obj_output) + _point_mtl_at_texture(obj_output.replace(".obj", ".mtl")) + result = trimesh.load(obj_output) + else: + ply_input = os.path.join(temporary_dir, "input.ply") + ply_output = os.path.join(temporary_dir, "output.ply") + geometry.export(ply_input) + mesh_set.load_new_mesh(ply_input) + mesh_set.apply_coord_laplacian_smoothing( + stepsmoothnum=iterations, + ) + context.progress(80, "Exporting…") + mesh_set.save_current_mesh(ply_output) + result = trimesh.load(ply_output, force="mesh") + else: + ply_input = os.path.join(temporary_dir, "input.ply") + ply_output = os.path.join(temporary_dir, "output.ply") + geometry.export(ply_input) + + mesh_set.load_new_mesh(ply_input) + context.progress(30, f"Smoothing ({mode})…") + + if mode == "taubin": + mesh_set.apply_coord_taubin_smoothing( + lambda_=strength, + mu=-strength - 0.01, + stepsmoothnum=iterations, + ) + else: + mesh_set.apply_coord_laplacian_smoothing( + stepsmoothnum=iterations, + cotangentweight=False, + ) + + context.progress(80, "Exporting…") + mesh_set.save_current_mesh(ply_output) + result = _raw_geometry(ply_output, trimesh) + finally: + shutil.rmtree(temporary_dir, ignore_errors=True) + + result.export(str(output_path)) + face_count = _face_count(result, trimesh) + context.log(f"Output: {output_path} ({face_count} faces)") + context.progress(100, "Done") + return MeshOpResult( + file_path=output_path, + details={"face_count": face_count}, + ) + + +def _node_executable() -> tuple[str, bool]: + configured = os.environ.get("MODLY_NODE_EXECUTABLE") + if configured: + executable = Path(configured) + if not executable.is_file(): + raise MeshOpUnavailableError( + f"Configured Node runtime does not exist: {configured}" + ) + return str(executable), True + + executable = shutil.which("node") or shutil.which("nodejs") + if executable is None: + raise MeshOpUnavailableError( + "mesh-optimizer requires Node.js (or Modly's Electron runtime)" + ) + return executable, False + + +def _meshopt_dependency_dir() -> Path: + candidates: list[Path] = [] + extension_dir = os.environ.get("EXTENSION_DIR") + if extension_dir: + candidates.append(Path(extension_dir)) + + app_root = Path(__file__).resolve().parents[3] + candidates.extend( + [ + app_root / "builtin-extensions" / "mesh-optimizer", + app_root / "out" / "builtin-extensions" / "mesh-optimizer", + app_root / "src" / "areas" / "workflows" / "nodes" / "mesh-optimizer", + ] + ) + + for candidate in candidates: + if (candidate / "node_modules" / "meshoptimizer").exists(): + return candidate + + raise MeshOpUnavailableError( + "mesh-optimizer dependencies are unavailable; run `npm run build` " + "before starting Modly from source" + ) + + +def decimate_mesh( + input_path: Path, + params: Mapping[str, Any], + context: MeshOpContext, +) -> MeshOpResult: + """Run the existing glTF Transform + meshoptimizer implementation.""" + executable, electron_runtime = _node_executable() + dependency_dir = _meshopt_dependency_dir() + runner_path = Path(__file__).with_name("meshopt_runner.cjs") + + environment = os.environ.copy() + if electron_runtime: + environment["ELECTRON_RUN_AS_NODE"] = "1" + + payload = { + "inputPath": str(input_path), + "params": dict(params), + "workspaceDir": str(context.workspace_dir), + "dependencyDir": str(dependency_dir), + "outputPath": ( + str(context.output_path) if context.output_path is not None else None + ), + } + + process = subprocess.Popen( + [executable, str(runner_path)], + cwd=str(dependency_dir), + env=environment, + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + encoding="utf-8", + errors="replace", + bufsize=1, + ) + if process.stdin is None or process.stdout is None or process.stderr is None: + process.kill() + raise MeshOpExecutionError("mesh-optimizer failed to open its I/O pipes") + + process.stdin.write(json.dumps(payload) + "\n") + process.stdin.close() + + result_payload: dict[str, Any] | None = None + backend_error: str | None = None + for raw_line in process.stdout: + line = raw_line.strip() + if not line: + continue + try: + message = json.loads(line) + except json.JSONDecodeError: + context.log(line) + continue + + message_type = message.get("type") + if message_type == "progress": + context.progress( + int(message.get("percent", 0)), + str(message.get("label", "")), + ) + elif message_type == "log": + context.log(str(message.get("message", ""))) + elif message_type == "done": + result_payload = message.get("result") or {} + elif message_type == "error": + backend_error = str(message.get("message", "Unknown error")) + + stderr = process.stderr.read().strip() + return_code = process.wait() + if backend_error is not None: + raise MeshOpExecutionError(backend_error) + if return_code != 0: + raise MeshOpExecutionError( + stderr or f"mesh-optimizer exited with code {return_code}" + ) + if result_payload is None or not result_payload.get("filePath"): + raise MeshOpExecutionError("mesh-optimizer returned no output file") + + details: dict[str, Any] = {} + if result_payload.get("faceCount") is not None: + details["face_count"] = int(result_payload["faceCount"]) + return MeshOpResult( + file_path=Path(result_payload["filePath"]), + details=details, + ) diff --git a/api/services/mesh_ops/processor.py b/api/services/mesh_ops/processor.py new file mode 100644 index 00000000..9af04353 --- /dev/null +++ b/api/services/mesh_ops/processor.py @@ -0,0 +1,71 @@ +"""Adapter between workflow process-node NDJSON and the mesh-op registry.""" + +import json +import os +import sys +import tempfile +import traceback +from pathlib import Path + +from . import MeshOpContext, mesh_ops_registry + + +def _emit(message: dict) -> None: + print(json.dumps(message), flush=True) + + +def run_processor(operation_id: str, processor_id: str) -> None: + """Read one workflow request, run a registered op, and emit its result.""" + try: + raw = sys.stdin.readline() + if not raw: + raise ValueError(f"{processor_id}: missing request payload") + data = json.loads(raw) + input_data = data.get("input") or {} + input_path = input_data.get("filePath") + if not input_path or not Path(input_path).is_file(): + if processor_id == "mesh-optimizer": + raise FileNotFoundError( + "mesh-optimizer: input.filePath is required" + ) + raise FileNotFoundError( + f"{processor_id}: input file not found: {input_path}" + ) + + workspace_dir = Path( + data.get("workspaceDir") + or os.environ.get("WORKSPACE_DIR") + or Path.home() / ".modly" / "workspace" + ) + temp_dir = Path( + data.get("tempDir") + or os.environ.get("TEMP_DIR") + or tempfile.gettempdir() + ) + context = MeshOpContext( + workspace_dir=workspace_dir, + temp_dir=temp_dir, + progress_cb=lambda percent, label: _emit( + {"type": "progress", "percent": percent, "label": label} + ), + log_cb=lambda message: _emit({"type": "log", "message": message}), + ) + result = mesh_ops_registry.run( + operation_id, + Path(input_path), + data.get("params") or {}, + context, + ) + _emit( + { + "type": "done", + "result": {"filePath": str(result.file_path)}, + } + ) + except Exception as exc: + _emit( + { + "type": "error", + "message": f"{exc}\n{traceback.format_exc()}", + } + ) diff --git a/api/services/mesh_ops/registry.py b/api/services/mesh_ops/registry.py new file mode 100644 index 00000000..5d755340 --- /dev/null +++ b/api/services/mesh_ops/registry.py @@ -0,0 +1,116 @@ +"""Registry and dispatcher for mesh-editing operations.""" + +import re +from copy import deepcopy +from math import isfinite +from pathlib import Path +from typing import Any, Iterable, Mapping, Optional + +from .types import MeshOp, MeshOpContext, MeshOpNotFoundError, MeshOpResult + + +_OP_ID = re.compile(r"^[a-z][a-z0-9_-]*$") + + +def _apply_numeric_bounds( + parameter_id: str, + value: Any, + schema: Mapping[str, Any], +) -> Any: + """Clamp a supplied numeric value to the bounds declared by its schema.""" + minimum = schema.get("min") + maximum = schema.get("max") + if minimum is None and maximum is None: + return value + + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise TypeError(f"Parameter {parameter_id!r} must be numeric") + if isinstance(value, float) and not isfinite(value): + raise ValueError(f"Parameter {parameter_id!r} must be finite") + + if minimum is not None: + value = max(minimum, value) + if maximum is not None: + value = min(maximum, value) + return value + + +def _validate_param( + parameter_id: str, + value: Any, + schema: Mapping[str, Any], +) -> Any: + """Validate (and where applicable, clamp) a supplied param against its schema.""" + param_type = schema.get("type") + + if param_type == "boolean": + if not isinstance(value, bool): + raise TypeError(f"Parameter {parameter_id!r} must be a boolean") + return value + + if param_type == "select": + allowed = {option["value"] for option in schema.get("options", [])} + if allowed and value not in allowed: + raise ValueError( + f"Parameter {parameter_id!r} must be one of {sorted(allowed)}" + ) + return value + + return _apply_numeric_bounds(parameter_id, value, schema) + + +class MeshOpsRegistry: + """Stores mesh operations and provides one invocation path for every caller.""" + + def __init__(self, operations: Iterable[MeshOp] = ()) -> None: + self._operations: dict[str, MeshOp] = {} + for operation in operations: + self.register(operation) + + def register(self, operation: MeshOp) -> None: + if not _OP_ID.fullmatch(operation.id): + raise ValueError(f"Invalid mesh operation id: {operation.id!r}") + if operation.id in self._operations: + raise ValueError(f"Duplicate mesh operation id: {operation.id!r}") + self._operations[operation.id] = operation + + def get(self, operation_id: str) -> MeshOp: + try: + return self._operations[operation_id] + except KeyError as exc: + raise MeshOpNotFoundError(operation_id) from exc + + def describe(self) -> list[dict[str, Any]]: + return [operation.describe() for operation in self._operations.values()] + + def run( + self, + operation_id: str, + input_path: Path, + params: Optional[Mapping[str, Any]], + context: MeshOpContext, + ) -> MeshOpResult: + operation = self.get(operation_id) + path = Path(input_path) + if not path.is_file(): + raise FileNotFoundError(f"Input mesh not found: {path}") + + resolved_params = { + schema["id"]: deepcopy(schema["default"]) + for schema in operation.params_schema + if "id" in schema and "default" in schema + } + if params: + resolved_params.update(params) + + for schema in operation.params_schema: + parameter_id = schema.get("id") + if parameter_id not in resolved_params: + continue + resolved_params[parameter_id] = _validate_param( + parameter_id, + resolved_params[parameter_id], + schema, + ) + + return operation.fn(path, resolved_params, context) diff --git a/api/services/mesh_ops/types.py b/api/services/mesh_ops/types.py new file mode 100644 index 00000000..5951bfe1 --- /dev/null +++ b/api/services/mesh_ops/types.py @@ -0,0 +1,77 @@ +"""Shared types for Modly mesh operations.""" + +from copy import deepcopy +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Callable, Mapping, Optional + + +ProgressCallback = Callable[[int, str], None] +LogCallback = Callable[[str], None] + + +class MeshOpNotFoundError(LookupError): + """Raised when a caller requests an operation that is not registered.""" + + +class MeshOpUnavailableError(RuntimeError): + """Raised when an operation's runtime dependency is unavailable.""" + + +class MeshOpExecutionError(RuntimeError): + """Raised when an operation backend fails while processing a mesh.""" + + +@dataclass(frozen=True) +class MeshOpContext: + """Runtime paths and optional workflow-protocol callbacks for an operation.""" + + workspace_dir: Path + temp_dir: Path + output_path: Optional[Path] = None + preserve_visuals: bool = False + progress_cb: Optional[ProgressCallback] = None + log_cb: Optional[LogCallback] = None + + def progress(self, percent: int, label: str) -> None: + if self.progress_cb is not None: + self.progress_cb(percent, label) + + def log(self, message: str) -> None: + if self.log_cb is not None: + self.log_cb(message) + + +@dataclass(frozen=True) +class MeshOpResult: + """The file produced by an operation and optional JSON-safe measurements.""" + + file_path: Path + details: Mapping[str, Any] = field(default_factory=dict) + + +MeshOpFn = Callable[[Path, Mapping[str, Any], MeshOpContext], MeshOpResult] + + +@dataclass(frozen=True) +class MeshOp: + """One callable operation and the metadata consumed by the UI and agent.""" + + id: str + label: str + params_schema: tuple[Mapping[str, Any], ...] + fn: MeshOpFn + category: str + destructive: bool = False + undoable: bool = True + + def describe(self) -> dict[str, Any]: + """Return the public, serializable part of this registry entry.""" + return { + "id": self.id, + "label": self.label, + "params_schema": deepcopy(list(self.params_schema)), + "destructive": self.destructive, + "undoable": self.undoable, + "category": self.category, + } diff --git a/api/services/model_sources.py b/api/services/model_sources.py new file mode 100644 index 00000000..677a9bfc --- /dev/null +++ b/api/services/model_sources.py @@ -0,0 +1,592 @@ +"""Validation and readiness helpers for manifest-declared Hugging Face sources.""" + +from __future__ import annotations + +import math +import re +import unicodedata +from pathlib import Path +from typing import Any + + +_SAFE_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]*$") +_WINDOWS_DEVICE = re.compile( + r"^(?:con|prn|aux|nul|com[1-9]|lpt[1-9])(?:\..*)?$", re.IGNORECASE +) +_WINDOWS_UNSAFE = re.compile(r'[<>"|?*\x00-\x1f]') + + +def _portable_segment(value: str, field: str) -> str: + if ( + not value + or value in {".", ".."} + or value.endswith((".", " ")) + or ":" in value + or _WINDOWS_UNSAFE.search(value) + or _WINDOWS_DEVICE.fullmatch(value) + ): + raise ValueError(f'{field} contains unsafe path segment "{value}"') + return value + + +def safe_source_id(value: Any, field: str = "model source id") -> str: + if ( + not isinstance(value, str) + or not value + or value != value.strip() + or _SAFE_ID.fullmatch(value) is None + ): + raise ValueError(f"{field} must be a safe non-empty identifier") + return _portable_segment(value, field) + + +def safe_relative_path(value: Any, field: str, *, allow_dot: bool = False) -> str: + if not isinstance(value, str) or not value or value != value.strip(): + raise ValueError(f"{field} must be a non-empty relative path") + if allow_dot and value == ".": + return value + if value == "." or value.startswith("/") or "\\" in value: + raise ValueError(f"{field} must be a safe relative POSIX path") + for part in value.split("/"): + _portable_segment(part, field) + return value + + +def _safe_prefix(value: Any, field: str) -> str: + path = value[:-1] if isinstance(value, str) and value.endswith("/") else value + safe_relative_path(path, field) + return value + + +def _prefixes(value: Any, field: str) -> list[str] | None: + if not isinstance(value, list): + raise ValueError(f"{field} must be an array") + return [_safe_prefix(entry, f"{field}[{index}]") for index, entry in enumerate(value)] + + +def _safe_repo_id(value: Any, field: str) -> str: + if not isinstance(value, str) or not value or value != value.strip() or "\\" in value: + raise ValueError(f"{field} must be a non-empty Hugging Face repository id") + parts = value.split("/") + if len(parts) > 2 or any( + part in {"", ".", ".."} or _SAFE_ID.fullmatch(part) is None for part in parts + ): + raise ValueError(f"{field} is not a safe Hugging Face repository id") + return value + + +def _safe_revision(value: Any, field: str) -> str | None: + if value is None: + return None + if ( + not isinstance(value, str) + or not value + or value != value.strip() + or value.startswith("/") + or "\\" in value + or "\0" in value + or any(part in {"", ".", ".."} for part in value.split("/")) + ): + raise ValueError(f"{field} must be a safe non-empty revision") + return value + + +def normalize_model_sources( + node: dict[str, Any], *, 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(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"{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") + alias = unicodedata.normalize("NFC", source_id).casefold() + if alias in aliases: + raise ValueError( + f'model source ids "{aliases[alias]}" and "{source_id}" are not portable-unique' + ) + aliases[alias] = source_id + if raw.get("provider") != "huggingface": + raise ValueError(f'{field}.provider must be "huggingface"') + checks = raw.get("checks") + if not isinstance(checks, list) or not checks: + raise ValueError(f"{field}.checks must be a non-empty array") + + source: dict[str, Any] = { + "id": source_id, + "provider": "huggingface", + "repo_id": _safe_repo_id(raw.get("repo_id"), f"{field}.repo_id"), + "destination": safe_relative_path( + raw.get("destination"), f"{field}.destination", allow_dot=True + ), + "checks": [ + safe_relative_path(check, f"{field}.checks[{check_index}]") + for check_index, check in enumerate(checks) + ], + } + revision = ( + _safe_revision(raw["revision"], f"{field}.revision") + if "revision" in raw + else None + ) + if "revision" in raw and revision is None: + raise ValueError(f"{field}.revision must be a safe non-empty revision") + include = ( + _prefixes(raw["include_prefixes"], f"{field}.include_prefixes") + if "include_prefixes" in raw + else None + ) + skip = ( + _prefixes(raw["skip_prefixes"], f"{field}.skip_prefixes") + if "skip_prefixes" in raw + else None + ) + if revision is not None: + source["revision"] = revision + if include is not None: + source["include_prefixes"] = include + if skip is not None: + source["skip_prefixes"] = skip + sources.append(source) + return sources + + +def 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 + 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() + try: + relative = candidate.relative_to(root) + except ValueError: + return True + current = root + if current.exists() and current.is_symlink(): + return True + for part in relative.parts: + current /= part + if current.exists() and current.is_symlink(): + return True + return False + + +def resolve_model_root(models_dir: Path, model_id: str) -> Path: + if not isinstance(model_id, str): + raise ValueError("Model id must be a string") + parts = model_id.split("/") + if len(parts) != 2: + raise ValueError("Model id must identify one extension node") + extension_id = safe_source_id(parts[0], "extension id") + 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 + if _path_has_symlink(root, candidate): + raise ValueError("Model path resolves through a symlink") + try: + candidate.resolve().relative_to(root.resolve()) + except ValueError as exc: + raise ValueError("Model path escapes the models directory") from exc + return candidate + + +def resolve_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) + 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_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("/")) + if _path_has_symlink(destination, candidate): + raise ValueError("Download target resolves through a symlink") + return candidate + + +def model_sources_are_downloaded( + models_dir: Path, model_id: str, sources: list[dict[str, Any]] +) -> bool: + try: + model_root = resolve_model_root(models_dir, model_id) + 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_at_root( + model_root, source["destination"] + ) + if not destination.is_dir(): + return False + for check in source["checks"]: + candidate = resolve_download_path(destination, check) + if ( + not candidate.is_file() + or candidate.stat().st_size <= 0 + or _path_has_symlink(model_root, candidate) + ): + return False + return bool(sources) + except (KeyError, OSError, TypeError, ValueError): + 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: + """Reject cross-source aliases before the first file is written.""" + aliases: dict[str, tuple[str, str]] = {} + for source in sources: + source_id = source["id"] + destination = source["destination"] + source_files = { + safe_relative_path(filename, f'model source "{source_id}" file') + for filename in files_by_source[source_id] + } + missing_checks = [check for check in source["checks"] if check not in source_files] + if missing_checks: + raise ValueError( + f'Model source "{source_id}" checks files excluded from its download plan: ' + + ", ".join(missing_checks) + ) + for safe_filename in source_files: + target = safe_filename if destination == "." else f"{destination}/{safe_filename}" + for value in (target, f"{target}.part"): + alias = unicodedata.normalize("NFC", value).casefold() + for previous_alias, (previous_source, previous_target) in aliases.items(): + if previous_source == source_id: + continue + if ( + alias == previous_alias + or alias.startswith(f"{previous_alias}/") + or previous_alias.startswith(f"{alias}/") + ): + raise ValueError( + "Model sources have a portable target collision: " + f'"{previous_source}:{previous_target}" and "{source_id}:{value}"' + ) + aliases[alias] = (source_id, value) + + +def _prefixes_overlap(left: list[str], right: list[str]) -> bool: + return any( + a.lower().startswith(b.lower()) or b.lower().startswith(a.lower()) + for a in left + for b in right + ) + + +def _assert_param_offers_variants(param: str, params_schema: Any, ids: list[str]) -> None: + """Variant ids must be selectable, so the param they key has to exist and offer them.""" + entry = next( + (p for p in params_schema if isinstance(p, dict) and p.get("id") == param), + None, + ) if isinstance(params_schema, list) else None + if entry is None: + raise ValueError(f'weight_variants.param must name a params_schema entry ("{param}")') + options = entry.get("options") + if not isinstance(options, list): + return + values = { + str(option.get("value")) if isinstance(option, dict) else str(option) + for option in options + } + missing = [variant_id for variant_id in ids if variant_id not in values] + if missing: + raise ValueError( + f'the "{param}" param must offer every weight variant id ' + f'(missing: {", ".join(missing)})' + ) + + +def normalize_weight_variants( + node: dict[str, Any], params_schema: Any = None +) -> dict[str, Any] | None: + """Validate a node's separately installable weight variants (e.g. quantizations).""" + if "weight_variants" not in node: + return None + if "model_sources" in node: + raise ValueError("weight_variants cannot be combined with model_sources") + if "weight_groups" in node: + raise ValueError("weight_variants cannot be combined with weight_groups") + if not isinstance(node.get("hf_repo"), str) or not node["hf_repo"]: + raise ValueError("weight_variants requires hf_repo on the same node") + raw = node["weight_variants"] + if not isinstance(raw, dict): + raise ValueError("weight_variants must be an object") + param = safe_source_id(raw.get("param"), "weight_variants.param") + raw_options = raw.get("options") + if not isinstance(raw_options, list) or not raw_options: + raise ValueError("weight_variants.options must be a non-empty array") + + aliases: dict[str, str] = {} + options: list[dict[str, Any]] = [] + for index, raw_option in enumerate(raw_options): + field = f"weight_variants.options[{index}]" + if not isinstance(raw_option, dict): + raise ValueError(f"{field} must be an object") + variant_id = safe_source_id(raw_option.get("id"), f"{field}.id") + alias = unicodedata.normalize("NFC", variant_id).casefold() + if alias in aliases: + raise ValueError( + f'weight variant ids "{aliases[alias]}" and "{variant_id}" are not portable-unique' + ) + aliases[alias] = variant_id + label = raw_option.get("label", variant_id) + if not isinstance(label, str) or not label.strip(): + raise ValueError(f"{field}.label must be a non-empty string") + size_gb = raw_option.get("size_gb") + if size_gb is not None and ( + isinstance(size_gb, bool) + or not isinstance(size_gb, (int, float)) + or not math.isfinite(size_gb) + or size_gb <= 0 + ): + raise ValueError(f"{field}.size_gb must be a positive number") + vram_gb = raw_option.get("vram_gb") + if vram_gb is not None and ( + isinstance(vram_gb, bool) + or not isinstance(vram_gb, (int, float)) + or not math.isfinite(vram_gb) + or vram_gb <= 0 + ): + raise ValueError(f"{field}.vram_gb must be a positive number") + include = _prefixes(raw_option.get("include_prefixes"), f"{field}.include_prefixes") + if not include: + raise ValueError(f"{field}.include_prefixes must be a non-empty array") + checks = raw_option.get("checks") + if not isinstance(checks, list) or not checks: + raise ValueError(f"{field}.checks must be a non-empty array") + safe_checks: list[str] = [] + for check_index, check in enumerate(checks): + path = safe_relative_path(check, f"{field}.checks[{check_index}]") + if not any(path.startswith(prefix) for prefix in include): + raise ValueError( + f"{field}.checks[{check_index}] is not covered by its include_prefixes" + ) + safe_checks.append(path) + option: dict[str, Any] = { + "id": variant_id, + "label": label, + "include_prefixes": include, + "checks": safe_checks, + } + if size_gb is not None: + option["size_gb"] = size_gb + if vram_gb is not None: + option["vram_gb"] = vram_gb + options.append(option) + + for index, option in enumerate(options): + for other in options[index + 1:]: + if _prefixes_overlap(option["include_prefixes"], other["include_prefixes"]): + raise ValueError( + f'weight variants "{option["id"]}" and "{other["id"]}" share files' + ) + download_check = node.get("download_check") + if isinstance(download_check, str) and any( + _prefixes_overlap([download_check], option["include_prefixes"]) for option in options + ): + raise ValueError("download_check must name a file outside every weight variant") + default = raw.get("default", options[0]["id"]) + if not any(option["id"] == default for option in options): + raise ValueError("weight_variants.default must name one of its options") + _assert_param_offers_variants( + param, + node.get("params_schema") if params_schema is None else params_schema, + [option["id"] for option in options], + ) + return {"param": param, "default": default, "options": options} + + +def installed_weight_variants( + models_dir: Path, model_id: str, variants: dict[str, Any] +) -> list[str]: + try: + model_root = resolve_model_root(models_dir, model_id) + except ValueError: + return [] + + def _present(check: str) -> bool: + try: + candidate = resolve_download_path(model_root, check) + return candidate.is_file() and candidate.stat().st_size > 0 + except (OSError, ValueError): + return False + + return [ + option["id"] + for option in variants["options"] + if all(_present(check) for check in option["checks"]) + ] + + +def missing_weight_variant( + models_dir: Path, + model_id: str, + variants: dict[str, Any] | None, + params: dict[str, Any], +) -> dict[str, Any] | None: + """The variant selected by params when its files are not installed, else None.""" + if not variants: + return None + selected = params.get(variants["param"]) + selected_id = variants["default"] if selected is None else str(selected) + option = next((o for o in variants["options"] if o["id"] == selected_id), None) + if option is None or option["id"] in installed_weight_variants(models_dir, model_id, variants): + return None + return option diff --git a/api/services/scene_input.py b/api/services/scene_input.py new file mode 100644 index 00000000..c5395050 --- /dev/null +++ b/api/services/scene_input.py @@ -0,0 +1,138 @@ +"""Validation shared by the API and isolated model runner for scene inputs.""" +import json +import math +import re +import stat +from pathlib import Path, PurePosixPath, PureWindowsPath + +SCHEMA = "modly.scene-manifest.v1" +MANIFEST = "scene-manifest.json" +MAX_MANIFEST_BYTES = 1024 * 1024 +MAX_REFERENCES = 4096 +MAX_REFERENCED_FILE_BYTES = 16 * 1024**3 +MAX_REFERENCED_TOTAL_BYTES = 64 * 1024**3 + + +def _relative(value: str, *, allow_dot: bool = False) -> Path: + if not isinstance(value, str) or not value or value != value.strip() or "\x00" in value: + raise ValueError("Scene path must be a nonempty workspace-relative path") + value = value.replace("\\", "/") + if value == "." and allow_dot: + return Path(".") + if (PurePosixPath(value).is_absolute() or PureWindowsPath(value).is_absolute() + or re.match(r"^[A-Za-z][A-Za-z0-9+.-]*:", value) + or re.search(r"%(?:25|2e|2f|5c|00)", value, re.I) + or re.search(r"%(?![0-9a-f]{2})", value, re.I) + or any(part in ("", ".", "..") for part in value.split("/"))): + raise ValueError("Scene path must be a safe workspace-relative path") + return Path(*value.split("/")) + + +def _inside(path: Path, root: Path) -> Path: + try: + relative = path.relative_to(root) + except ValueError as exc: + raise ValueError("Scene path escapes its allowed root") from exc + current = root + for part in relative.parts: + current = current / part + try: + info = current.lstat() + except OSError as exc: + raise ValueError("Scene referenced path is missing or unreadable") from exc + is_reparse = bool(getattr(info, "st_file_attributes", 0) & getattr(stat, "FILE_ATTRIBUTE_REPARSE_POINT", 0)) + if current.is_symlink() or is_reparse: + raise ValueError("Scene referenced path must not use symlinks or reparse points") + try: + resolved = path.resolve(strict=True) + except OSError as exc: + raise ValueError("Scene referenced path is missing or unreadable") from exc + if not resolved.is_relative_to(root): + raise ValueError("Scene path escapes its allowed root") + return resolved + + +def validate_scene_input(workspace: Path, scene_path: str) -> Path: + """Return a canonical manifest Path, rejecting traversal and symlink escapes. + + V1 sceneRoot is workspace-relative, except '.' denotes the manifest directory. + An asset's path and preview references are sceneRoot-relative; workspacePath + is always workspace-relative. Existence never determines which base applies. + """ + root = workspace.resolve(strict=True) + relative = _relative(scene_path) + requested = root / relative + if requested.name != MANIFEST: + if requested.suffix.lower() == ".json": + raise ValueError(f"Scene input accepts a directory or {MANIFEST}") + requested = requested / MANIFEST + try: + manifest_path = _inside(requested, root) + if not manifest_path.is_file(): + raise ValueError("Scene manifest is not a file") + if manifest_path.stat().st_size > MAX_MANIFEST_BYTES: + raise ValueError("Scene manifest exceeds the 1 MiB limit") + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + except (OSError, UnicodeError, json.JSONDecodeError) as exc: + raise ValueError("Scene manifest is missing or invalid JSON") from exc + if not isinstance(manifest, dict) or manifest.get("schema") != SCHEMA: + raise ValueError(f"Scene manifest schema must be {SCHEMA}") + scene_root = _relative(manifest.get("sceneRoot"), allow_dot=True) + # v1 sceneRoot is workspace-relative; only '.' means the manifest's directory. + # Never select a base according to which unrelated path happens to exist. + scene_root_candidate = manifest_path.parent if scene_root == Path(".") else root / scene_root + root_path = _inside(scene_root_candidate, root) + if not root_path.is_dir(): + raise ValueError("Scene root is not a directory") + if not isinstance(manifest.get("assets"), list): + raise ValueError("Scene manifest assets must be an array") + if len(manifest["assets"]) > MAX_REFERENCES: + raise ValueError("Scene manifest contains too many asset references") + total_bytes = 0 + for asset in manifest["assets"]: + # Assets may be opaque extension metadata. Validate only references the + # host understands; never silently resolve a supplied path outside workspace. + if isinstance(asset, dict): + for field, base in (("workspacePath", root), ("path", root_path)): + if field not in asset: + continue + target = _inside(base / _relative(asset[field]), root) + if not target.is_file(): + raise ValueError(f"Scene asset {field} is not a file") + size = target.stat().st_size + if size > MAX_REFERENCED_FILE_BYTES: + raise ValueError("Scene referenced file exceeds the size limit") + total_bytes += size + if total_bytes > MAX_REFERENCED_TOTAL_BYTES: + raise ValueError("Scene referenced files exceed the total size limit") + preview = manifest.get("preview", {}) + if not isinstance(preview, dict): + raise ValueError("Scene preview must be an object") + for name in ("image", "video"): + if name in preview: + relative_preview = _relative(preview[name]) + target = _inside(root_path / relative_preview, root) + if not target.is_file(): + raise ValueError(f"Scene preview {name} is not a file") + view = manifest.get("initialView") + if view is not None: + if not isinstance(view, dict): + raise ValueError("Scene initialView must be an object") + for field in ("position", "target", "up"): + triple = view.get(field) + if triple is None and field == "up": + continue + if (not isinstance(triple, list) or len(triple) != 3 + or any(not isinstance(n, (int, float)) or isinstance(n, bool) + or not math.isfinite(n) for n in triple)): + raise ValueError(f"Scene initialView {field} must be a finite numeric triple") + if view["position"] == view["target"] or view.get("up") == [0, 0, 0]: + raise ValueError("Scene initialView has degenerate camera vectors") + return manifest_path + + +def revalidate_scene_manifest(workspace: Path, manifest_path: Path) -> Path: + """Recheck the file immediately before model invocation in the worker.""" + root = workspace.resolve(strict=True) + candidate = _inside(manifest_path, root) + return validate_scene_input(root, candidate.relative_to(root).as_posix()) diff --git a/api/tests/test_agent_router.py b/api/tests/test_agent_router.py new file mode 100644 index 00000000..ecd2ed7c --- /dev/null +++ b/api/tests/test_agent_router.py @@ -0,0 +1,53 @@ +import asyncio +import unittest +from unittest import mock + +import httpx + +import routers.agent as agent + + +class _MockClientFactory: + """Builds real AsyncClients wired to a MockTransport, so execute_tool talks to + a fake Modly API instead of the network.""" + + def __init__(self, handler) -> None: + self._handler = handler + self._real = httpx.AsyncClient + + def __call__(self, *args, **kwargs): + kwargs["transport"] = httpx.MockTransport(self._handler) + return self._real(*args, **kwargs) + + +def _run_tool(name: str, handler) -> tuple[str, object]: + factory = _MockClientFactory(handler) + with mock.patch.object(agent.httpx, "AsyncClient", factory): + return asyncio.run(agent.execute_tool(name, {}, {})) + + +class UnloadModelsErrorTests(unittest.TestCase): + """unload_models must report a failed unload, not claim success (like every + other POST tool and like the MCP server's modly_unload_models).""" + + def test_http_error_is_surfaced(self) -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(500, text="boom") + + text, payload = _run_tool("unload_models", handler) + # Before the fix the response was discarded and the success string was + # returned even on a 500; now the shared HTTPStatusError handler runs. + self.assertTrue(text.startswith("API error 500"), text) + self.assertIsNone(payload) + + def test_success_still_reports_unloaded(self) -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(200, json={"ok": True}) + + text, payload = _run_tool("unload_models", handler) + self.assertIn("unloaded", text.lower()) + self.assertIsNone(payload) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_agent_workflow_unload.py b/api/tests/test_agent_workflow_unload.py new file mode 100644 index 00000000..665ea90d --- /dev/null +++ b/api/tests/test_agent_workflow_unload.py @@ -0,0 +1,181 @@ +import asyncio +import json +import unittest +from unittest import mock + +import httpx + +import routers.agent as agent + + +class _ScriptedServer: + """MockTransport wiring: serves a scripted /chat/completions sequence and + records every request body agent_chat sends.""" + + def __init__(self, chat_bodies) -> None: + self._chat = list(chat_bodies) + self.requests: list[dict] = [] + self._real = httpx.AsyncClient + + def _handler(self, request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/chat/completions"): + self.requests.append({"headers": dict(request.headers), "body": json.loads(request.content)}) + return httpx.Response(200, json=self._chat.pop(0)) + return httpx.Response(404, json={}) + + def __call__(self, *args, **kwargs): + kwargs["transport"] = httpx.MockTransport(self._handler) + return self._real(*args, **kwargs) + + +class _FakeSlot: + base_url = "http://127.0.0.1:8791/v1" + + def __init__(self) -> None: + self.released = 0 + + def release(self) -> None: + self.released += 1 + + +def _assistant(content: str = "", tool_calls=None) -> dict: + msg: dict = {"role": "assistant", "content": content} + if tool_calls: + msg["tool_calls"] = tool_calls + return {"choices": [{"message": msg}]} + + +_RUN_WF1 = {"id": "call_1", "type": "function", "function": {"name": "run_workflow", "arguments": '{"workflow_id": "wf1"}'}} +_SPEC = {"vision": False} + + +def _run_local(chat_bodies, context): + scripted = _ScriptedServer(chat_bodies) + slot = _FakeSlot() + request = agent.AgentChatRequest( + messages=[agent.ChatMessage(role="user", content="make me a thing")], + model="qwen3-4b", + context=context, + ) + with mock.patch.object(agent.httpx, "AsyncClient", scripted), \ + mock.patch.object(agent.llm_server, "resolve_model", return_value=_SPEC), \ + mock.patch.object(agent.llama_pool, "ensure", return_value=slot), \ + mock.patch.object(agent.llama_pool, "unload_all") as unload_all: + response = asyncio.run(agent.agent_chat(request)) + return scripted, slot, unload_all, response + + +class WorkflowVramUnloadTests(unittest.TestCase): + def test_llm_unloaded_on_the_normal_return_after_a_workflow(self) -> None: + # Round 1 dispatches a workflow; round 2 is the final answer (no tools) and + # returns early. The VRAM unload must still fire on that path. + chat = [_assistant(tool_calls=[_RUN_WF1]), _assistant(content="Running your workflow now.")] + _, slot, unload_all, response = _run_local(chat, {"workflows": [{"id": "wf1", "name": "My Workflow"}]}) + unload_all.assert_called_once() + self.assertEqual(slot.released, 1) + self.assertEqual([a.tool for a in response.actions], ["run_workflow"]) + + def test_slot_is_released_before_the_unload(self) -> None: + # unload_all() spares held slots, so unloading before the release kept + # the agent's own model on the GPU through the whole workflow. + chat = [_assistant(tool_calls=[_RUN_WF1]), _assistant(content="ok")] + released_at_unload: list[int] = [] + slot = _FakeSlot() + request = agent.AgentChatRequest( + messages=[agent.ChatMessage(role="user", content="go")], + model="qwen3-4b", + context={"workflows": [{"id": "wf1", "name": "My Workflow"}]}, + ) + with mock.patch.object(agent.httpx, "AsyncClient", _ScriptedServer(chat)), \ + mock.patch.object(agent.llm_server, "resolve_model", return_value=_SPEC), \ + mock.patch.object(agent.llama_pool, "ensure", return_value=slot), \ + mock.patch.object(agent.llama_pool, "unload_all", lambda: released_at_unload.append(slot.released)): + asyncio.run(agent.agent_chat(request)) + self.assertEqual(released_at_unload, [1]) + + def test_no_unload_when_no_workflow_was_dispatched(self) -> None: + _, slot, unload_all, _ = _run_local([_assistant(content="Here is some info.")], {}) + unload_all.assert_not_called() + self.assertEqual(slot.released, 1) + + def test_a_single_system_message_leads_the_conversation(self) -> None: + # Qwen3.5's template rejects a system message anywhere but first (HTTP 400). + context = {"currentMeshPath": "/tmp/a.glb", "extensions": [{"id": "remesh", "name": "Remesh"}]} + scripted, *_ = _run_local([_assistant(content="ok")], context) + roles = [m["role"] for m in scripted.requests[0]["body"]["messages"]] + self.assertEqual(roles, ["system", "user"]) + system = scripted.requests[0]["body"]["messages"][0]["content"] + self.assertIn("Current mesh path: /tmp/a.glb", system) + self.assertIn("remesh", system) + + def test_tool_result_answers_its_call_id(self) -> None: + chat = [_assistant(tool_calls=[_RUN_WF1]), _assistant(content="done")] + scripted, *_ = _run_local(chat, {"workflows": [{"id": "wf1", "name": "My Workflow"}]}) + tool_msgs = [m for m in scripted.requests[1]["body"]["messages"] if m["role"] == "tool"] + self.assertEqual(tool_msgs[0]["tool_call_id"], "call_1") + + +class ExternalProviderTests(unittest.TestCase): + def test_external_provider_sends_the_key_and_never_touches_the_local_pool(self) -> None: + scripted = _ScriptedServer([_assistant(tool_calls=[_RUN_WF1]), _assistant(content="ok")]) + request = agent.AgentChatRequest( + messages=[agent.ChatMessage(role="user", content="hi")], + model="gpt-test", + provider=agent.ProviderConfig(type="external", base_url="https://llm.test/v1/", api_key="sk-test"), + context={"workflows": [{"id": "wf1", "name": "My Workflow"}]}, + ) + with mock.patch.object(agent.httpx, "AsyncClient", scripted), \ + mock.patch.object(agent.llama_pool, "ensure") as ensure, \ + mock.patch.object(agent.llama_pool, "unload_all") as unload_all: + response = asyncio.run(agent.agent_chat(request)) + self.assertEqual(response.message, "ok") + self.assertEqual(scripted.requests[0]["headers"]["authorization"], "Bearer sk-test") + ensure.assert_not_called() + unload_all.assert_not_called() + + def test_text_only_provider_gets_a_retry_without_the_image(self) -> None: + bodies: list[dict] = [] + + def handler(request: httpx.Request) -> httpx.Response: + body = json.loads(request.content) + bodies.append(body) + if isinstance(body["messages"][1]["content"], list): + return httpx.Response(400, json={"error": {"message": "image input not supported"}}) + return httpx.Response(200, json=_assistant(content="ok")) + + real = httpx.AsyncClient + request = agent.AgentChatRequest( + messages=[agent.ChatMessage(role="user", content="what is this", images=["data:image/png;base64,AAAA"])], + model="text-only", + provider=agent.ProviderConfig(type="external", base_url="https://llm.test/v1", api_key="k"), + ) + with mock.patch.object(agent.httpx, "AsyncClient", lambda *a, **k: real(*a, transport=httpx.MockTransport(handler), **k)): + response = asyncio.run(agent.agent_chat(request)) + self.assertEqual(response.message, "ok") + self.assertEqual(len(bodies), 2) + self.assertIn("what is this", bodies[1]["messages"][1]["content"]) + + def test_provider_error_shows_its_message_not_raw_json(self) -> None: + def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(401, json={"error": {"message": "Incorrect API key provided.", "type": "invalid_request_error"}}) + + real = httpx.AsyncClient + request = agent.AgentChatRequest( + messages=[agent.ChatMessage(role="user", content="hi")], + provider=agent.ProviderConfig(type="external", base_url="https://llm.test/v1", api_key="bad"), + ) + with mock.patch.object(agent.httpx, "AsyncClient", lambda *a, **k: real(*a, transport=httpx.MockTransport(handler), **k)): + response = asyncio.run(agent.agent_chat(request)) + self.assertEqual(response.message, "LLM error (401): Incorrect API key provided.") + + def test_missing_provider_url_is_reported(self) -> None: + request = agent.AgentChatRequest( + messages=[agent.ChatMessage(role="user", content="hi")], + provider=agent.ProviderConfig(type="external"), + ) + response = asyncio.run(agent.agent_chat(request)) + self.assertIn("No provider URL", response.message) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_base_generator.py b/api/tests/test_base_generator.py index 580b3d25..9adfe00c 100644 --- a/api/tests/test_base_generator.py +++ b/api/tests/test_base_generator.py @@ -43,6 +43,12 @@ def test_download_progress_print_is_ascii_safe(self) -> None: self.assertTrue(first_line.endswith("...")) self.assertTrue(all(ord(ch) < 128 for ch in printed)) + def test_typed_artifact_generation_requires_an_explicit_override(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + gen = _DownloadGen(Path(tmp) / "model", Path(tmp) / "out") + with self.assertRaisesRegex(NotImplementedError, "scene artifact generation"): + gen.generate_artifact("scene", Path(tmp) / "scene-manifest.json", {}) + if __name__ == "__main__": - unittest.main() \ No newline at end of file + unittest.main() diff --git a/api/tests/test_export_router.py b/api/tests/test_export_router.py new file mode 100644 index 00000000..2bed2a68 --- /dev/null +++ b/api/tests/test_export_router.py @@ -0,0 +1,208 @@ +import base64 +import io +import tempfile +import unittest +from pathlib import Path + +from fastapi import HTTPException + +# The export router imports trimesh at module load; skip the whole suite (rather +# than breaking `unittest discover`) in minimal environments without it. +try: + import numpy as np + import trimesh + + import routers.export as export_router + from services import imported_sources + + HAVE_TRIMESH = True +except Exception: # noqa: BLE001 + HAVE_TRIMESH = False + + +def _token(rel_path: str) -> str: + return base64.urlsafe_b64encode(rel_path.encode("utf-8")).decode("ascii").rstrip("=") + + +def _load_stl(resp) -> "trimesh.Trimesh": + return trimesh.load(io.BytesIO(resp.body), file_type="stl") + + +@unittest.skipUnless(HAVE_TRIMESH, "trimesh not installed") +class ExportForSlicerTests(unittest.TestCase): + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self.workspace = Path(self._tmp.name).resolve() + self._orig_workspace = export_router.registry.WORKSPACE_DIR + export_router.registry.WORKSPACE_DIR = self.workspace + # A box that is tallest along Y (glTF up-axis) and unit-sized, matching + # what image-to-3D generators emit. Exported to GLB, it reloads as a + # Scene so the flatten path is exercised too. + box = trimesh.creation.box(extents=[0.3, 1.0, 0.3]) + self.rel = "Workflows/hero.glb" + (self.workspace / "Workflows").mkdir(parents=True, exist_ok=True) + box.export(str(self.workspace / self.rel)) + + def tearDown(self) -> None: + export_router.registry.WORKSPACE_DIR = self._orig_workspace + self._tmp.cleanup() + + def test_converts_glb_to_stl_with_download_filename(self) -> None: + resp = export_router.export_for_slicer("stl", _token(self.rel), "model.stl") + self.assertEqual(resp.media_type, "model/stl") + self.assertIn('filename="model.stl"', resp.headers["content-disposition"]) + mesh = _load_stl(resp) + self.assertGreater(len(mesh.faces), 0) + + def test_reorients_y_up_to_z_up(self) -> None: + # The box is tallest in Y; after the Y->Z rotation it must be tallest in + # Z so it imports standing upright on the slicer bed. + resp = export_router.export_for_slicer("stl", _token(self.rel), "model.stl") + ex = _load_stl(resp).extents + self.assertEqual(int(np.argmax(ex)), 2, f"expected Z to be the tallest axis, got extents {ex}") + + def test_normalizes_longest_edge_to_default_print_size(self) -> None: + resp = export_router.export_for_slicer("stl", _token(self.rel), "model.stl") + longest = float(max(_load_stl(resp).extents)) + self.assertAlmostEqual(longest, export_router.DEFAULT_PRINT_LONGEST_MM, places=3) + + def test_leaves_a_real_world_sized_mesh_alone(self) -> None: + # A mesh that already carries a physical size is the user's own: silently + # shrinking a 180 mm part to 50 mm would waste a print. + big = trimesh.creation.box(extents=[40.0, 180.0, 40.0]) + rel = "Workflows/part.glb" + big.export(str(self.workspace / rel)) + resp = export_router.export_for_slicer("stl", _token(rel), "model.stl") + self.assertAlmostEqual(float(max(_load_stl(resp).extents)), 180.0, places=2) + + def test_does_not_rotate_a_z_up_stl_source(self) -> None: + # STL is conventionally Z-up already; the glTF Y->Z rotation would lay an + # upright model on its side. + rel = "Workflows/upright.stl" + trimesh.creation.box(extents=[10.0, 10.0, 30.0]).export(str(self.workspace / rel)) + resp = export_router.export_for_slicer("stl", _token(rel), "model.stl") + ex = _load_stl(resp).extents + self.assertEqual(int(np.argmax(ex)), 2, f"expected Z to stay the tallest axis, got extents {ex}") + + def test_rejects_unsupported_format(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("glb", _token(self.rel), "model.glb") + self.assertEqual(ctx.exception.status_code, 400) + + def test_rejects_filename_extension_mismatch(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token(self.rel), "model.obj") + self.assertEqual(ctx.exception.status_code, 400) + + def test_rejects_malformed_token(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", "!!!not-base64!!!", "model.stl") + self.assertEqual(ctx.exception.status_code, 400) + + def test_rejects_path_traversal(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token("../escape.glb"), "model.stl") + self.assertEqual(ctx.exception.status_code, 400) + + def test_rejects_sibling_prefix_escape(self) -> None: + # A sibling dir whose name starts with the workspace dir name must not be + # reachable — the old str.startswith containment guard would allow it. + sibling = self.workspace.parent / (self.workspace.name + "-secret") + sibling.mkdir(parents=True, exist_ok=True) + (sibling / "x.glb").write_bytes(b"nope") + rel = f"../{self.workspace.name}-secret/x.glb" + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token(rel), "model.stl") + self.assertEqual(ctx.exception.status_code, 400) + + def test_missing_file_is_404(self) -> None: + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token("Workflows/nope.glb"), "model.stl") + self.assertEqual(ctx.exception.status_code, 404) + + +@unittest.skipUnless(HAVE_TRIMESH, "trimesh not installed") +class FlattenAndScaleHelperTests(unittest.TestCase): + def test_flatten_bakes_scene_node_transforms(self) -> None: + # Two boxes placed at different positions via scene-graph transforms. + # util.concatenate(geometry.values()) would ignore the transforms; the + # scene-level flatten must reflect them in the combined bounds. + scene = trimesh.Scene() + scene.add_geometry(trimesh.creation.box(extents=[2, 2, 2]), transform=trimesh.transformations.translation_matrix([0, 0, 0])) + scene.add_geometry(trimesh.creation.box(extents=[2, 2, 2]), transform=trimesh.transformations.translation_matrix([100, 0, 0])) + mesh = export_router._to_single_mesh(scene) + self.assertIsInstance(mesh, trimesh.Trimesh) + # Combined X extent spans both boxes: ~101 (from -1 to 101). + self.assertGreater(mesh.extents[0], 100.0) + + def test_scale_to_print_size(self) -> None: + mesh = trimesh.creation.box(extents=[1.0, 2.0, 4.0]) + export_router._scale_to_print_size(mesh, longest_mm=80.0) + self.assertAlmostEqual(float(max(mesh.extents)), 80.0, places=3) + + def test_scale_ignores_degenerate_mesh(self) -> None: + # A single point cloud has zero extent; scaling must not divide by zero. + mesh = trimesh.Trimesh(vertices=[[0, 0, 0]], faces=[]) + export_router._scale_to_print_size(mesh) # must not raise + + +@unittest.skipUnless(HAVE_TRIMESH, "trimesh not installed") +class ImportedSourceSlicerTests(unittest.TestCase): + """Meshes imported from outside the workspace are sliceable — but only the + exact files the user picked, and never with the wrong up-axis.""" + + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self.outside = Path(self._tmp.name).resolve() + self._orig_workspace = export_router.registry.WORKSPACE_DIR + # A workspace elsewhere, so nothing here is reachable as a relative path. + self._ws_tmp = tempfile.TemporaryDirectory() + export_router.registry.WORKSPACE_DIR = Path(self._ws_tmp.name).resolve() + imported_sources.clear() + # Unit-sized and tallest along Y, as a glTF export would be. + self.mesh_path = self.outside / "imported.glb" + trimesh.creation.box(extents=[0.3, 1.0, 0.3]).export(str(self.mesh_path)) + + def tearDown(self) -> None: + export_router.registry.WORKSPACE_DIR = self._orig_workspace + imported_sources.clear() + self._tmp.cleanup() + self._ws_tmp.cleanup() + + def test_rejects_an_absolute_path_that_was_never_imported(self) -> None: + # The whole point of the registry: an absolute path alone buys nothing. + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + self.assertEqual(ctx.exception.status_code, 400) + + def test_serves_a_registered_import(self) -> None: + imported_sources.register(self.mesh_path, self.mesh_path) + resp = export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + self.assertEqual(resp.media_type, "model/stl") + self.assertGreater(len(_load_stl(resp).faces), 0) + + def test_rotates_a_registered_gltf_import(self) -> None: + imported_sources.register(self.mesh_path, self.mesh_path) + resp = export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + ex = _load_stl(resp).extents + self.assertEqual(int(np.argmax(ex)), 2, f"expected Z to be the tallest axis, got extents {ex}") + + def test_does_not_rotate_an_import_that_was_stl_before_conversion(self) -> None: + # `import-by-path` converts STL/OBJ/PLY to GLB without touching the axes, + # so the .glb container here still holds Z-up data. Rotating it would be + # exactly the bug the rotation exists to prevent. + imported_sources.register(self.mesh_path, self.outside / "original.stl") + resp = export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + ex = _load_stl(resp).extents + self.assertEqual(int(np.argmax(ex)), 1, f"expected Y to stay the tallest axis, got extents {ex}") + + def test_registered_but_deleted_file_is_404(self) -> None: + imported_sources.register(self.mesh_path, self.mesh_path) + self.mesh_path.unlink() + with self.assertRaises(HTTPException) as ctx: + export_router.export_for_slicer("stl", _token(str(self.mesh_path)), "model.stl") + self.assertEqual(ctx.exception.status_code, 404) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_extension_process.py b/api/tests/test_extension_process.py index b77317ef..29602cd6 100644 --- a/api/tests/test_extension_process.py +++ b/api/tests/test_extension_process.py @@ -1,7 +1,11 @@ import io +import json import platform import queue import unittest +import json +import tempfile +from unittest.mock import patch from pathlib import Path from services.extension_process import ExtensionProcess, _venv_python @@ -12,6 +16,31 @@ def _make_proc() -> ExtensionProcess: class ExtensionProcessTests(unittest.TestCase): + def test_generation_envelope_pins_worker_model_id(self) -> None: + proc = _make_proc() + sent = [] + proc._send = sent.append + proc._receive_generation = lambda *args: Path("result.glb") + proc._generate_request({"image_b64": ""}, {}, None, None) + self.assertEqual(sent[0]["model_id"], "demo") + + def test_generate_artifact_sends_typed_scene_without_image_bytes(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + workspace = Path(tmp) / "workspace" + scene = workspace / "Workflows" / "room" + scene.mkdir(parents=True) + manifest = scene / "scene-manifest.json" + manifest.write_text(json.dumps({ + "schema": "modly.scene-manifest.v1", "sceneRoot": ".", "assets": [], + })) + proc = _make_proc() + calls = [] + proc._generate_request = lambda payload, params, progress, cancel: calls.append((payload, params)) or manifest + with patch("services.generator_registry.WORKSPACE_DIR", workspace): + result = proc.generate_artifact("scene", manifest, {"quality": "high"}) + self.assertEqual(result, manifest) + self.assertEqual(calls, [({"input": {"kind": "scene", "path": str(manifest.resolve())}}, {"quality": "high"})]) + def test_read_loop_writes_sentinel_to_own_queue_only(self) -> None: proc = _make_proc() @@ -102,6 +131,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": str(Path("/tmp/models/ext/_shared/base"))}, + ) + class MissingModuleExtractionTests(unittest.TestCase): def test_extracts_module_name_from_message(self) -> None: @@ -137,6 +183,24 @@ def test_returns_none_for_unknown_module(self) -> None: self.assertIsNone(proc._resolve_auto_repair_package("totally_unknown_pkg")) +class StderrDecodingTests(unittest.TestCase): + """tqdm draws partial blocks with U+258D/U+258F, whose UTF-8 bytes are + undefined in cp1252. Decoding the child's pipes with the locale codec killed + _stderr_loop mid-run; nothing then drained stderr and the child blocked + forever on write once the pipe filled (a 3D generation froze at 80%).""" + + def test_tqdm_partial_blocks_survive_the_loop(self) -> None: + proc = _make_proc() + bar = "Volume Decoding: 1%|█▍▏| 123/13827\r" + fake_proc = type("FakeProc", (), {"stderr": io.StringIO(bar)})() + + proc._stderr_loop(fake_proc) # must not raise + + def test_child_is_told_to_write_utf8(self) -> None: + proc = _make_proc() + self.assertEqual(proc._build_env().get("PYTHONIOENCODING"), "utf-8") + + class RecvTests(unittest.TestCase): def test_returns_message_from_queue(self) -> None: proc = _make_proc() @@ -155,5 +219,50 @@ def test_empty_queue_raises_timeout_error(self) -> None: proc._recv(timeout=0.05) +class GenerateErrorLoadedFlagTests(unittest.TestCase): + """ + Issue #239: a generation that fails during lazy texture setup leaves the + worker without a model. If _loaded stays True, GeneratorRegistry.get_active() + skips load() forever and every later run reuses the broken worker. + """ + + def _failing_generate(self, error_msg: dict) -> ExtensionProcess: + proc = _make_proc() + proc._loaded = True + proc._send = lambda msg: None # type: ignore[assignment] + proc._queue.put(error_msg) + with self.assertRaises(RuntimeError): + proc.generate(b"", {}) + return proc + + def test_clears_loaded_when_worker_reports_model_lost(self) -> None: + proc = self._failing_generate( + {"type": "error", "message": "No module named 'xatlas'", "loaded": False} + ) + self.assertFalse(proc._loaded) + + def test_keeps_loaded_when_worker_still_has_its_model(self) -> None: + proc = self._failing_generate( + {"type": "error", "message": "bad input image", "loaded": True} + ) + self.assertTrue(proc._loaded) + + def test_keeps_loaded_when_worker_reports_no_state(self) -> None: + proc = self._failing_generate({"type": "error", "message": "boom"}) + self.assertTrue(proc._loaded) + + def test_error_still_propagates_the_original_cause(self) -> None: + proc = _make_proc() + proc._loaded = True + proc._send = lambda msg: None # type: ignore[assignment] + proc._queue.put( + {"type": "error", "message": "short", "traceback": "full traceback here", + "loaded": False} + ) + with self.assertRaises(RuntimeError) as ctx: + proc.generate(b"", {}) + self.assertIn("full traceback here", str(ctx.exception)) + + if __name__ == "__main__": unittest.main() diff --git a/api/tests/test_generation_router.py b/api/tests/test_generation_router.py new file mode 100644 index 00000000..3d69118c --- /dev/null +++ b/api/tests/test_generation_router.py @@ -0,0 +1,174 @@ +import asyncio +import tempfile +import threading +import unittest +from pathlib import Path + +from fastapi import BackgroundTasks + +import services.generator_registry as registry +import routers.generation as generation +from schemas.generation import JobStatus + + +class _FakeGenerator: + """Writes its output into whatever directory generation assigns it.""" + + def __init__(self) -> None: + self.outputs_dir: Path | None = None + + def generate(self, image_bytes, params, progress_cb, cancel_event=None) -> Path: + out = Path(self.outputs_dir) / "model.glb" + out.write_bytes(b"glb") + return out + + +class _FakeUpload: + """Minimal UploadFile stand-in: an image content-type and readable bytes.""" + + def __init__(self, content_type: str = "image/png", data: bytes = b"\x89PNG\r\n") -> None: + self.content_type = content_type + self._data = data + + async def read(self) -> bytes: + return self._data + + +class _FakeRegistry: + def __init__(self, gen: _FakeGenerator) -> None: + self._gen = gen + + def active_status(self) -> dict: + # Report loaded so _run_generation skips the download/load thread. + return {"loaded": True, "name": "fake", "downloaded": True} + + def assert_weight_variant_installed(self, params: dict, model_id=None) -> None: + pass + + def get_active(self) -> _FakeGenerator: + return self._gen + + # generate_from_image looks the model up and switches to it before filing the job. + def get_generator(self, model_id: str) -> _FakeGenerator: + return self._gen + + def get_manifest(self, model_id: str) -> dict: + return {"output": "mesh"} + + def switch_model(self, model_id: str) -> None: + pass + + +class RunGenerationWorkspaceTests(unittest.TestCase): + """A generation started after the workspace path is relocated at runtime + (POST /settings/paths) must file its output under the *current* workspace, + not the one captured when the module was imported.""" + + def setUp(self) -> None: + self._prev_registry = generation.generator_registry + self._prev_ws = registry.WORKSPACE_DIR + self._tmp = tempfile.TemporaryDirectory() + # The user relocated the workspace: the registry global now points here. + registry.WORKSPACE_DIR = Path(self._tmp.name) / "new_workspace" + # Keep the test hermetic against the module's import-time binding: if the + # stale name still exists (before the fix) redirect it into the temp tree + # so the assertion — not a stray write to the real workspace — is what + # catches the bug. + self._had_stale = hasattr(generation, "WORKSPACE_DIR") + if self._had_stale: + generation.WORKSPACE_DIR = Path(self._tmp.name) / "old_workspace" + + def tearDown(self) -> None: + generation.generator_registry = self._prev_registry + registry.WORKSPACE_DIR = self._prev_ws + if self._had_stale: + generation.WORKSPACE_DIR = self._prev_ws + for store in ( + generation._jobs, + generation._cancel_events, + generation._cancelled, + generation._completed_at, + ): + store.clear() + self._tmp.cleanup() + + def _run(self, collection: str) -> tuple[_FakeGenerator, JobStatus]: + gen = _FakeGenerator() + generation.generator_registry = _FakeRegistry(gen) + job_id = "job-test" + generation._jobs[job_id] = JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = threading.Event() + asyncio.run(generation._run_generation(job_id, b"img", {}, collection)) + return gen, generation._jobs[job_id] + + def test_output_lands_under_the_current_workspace(self) -> None: + gen, job = self._run("MyColl") + self.assertEqual(Path(gen.outputs_dir), registry.WORKSPACE_DIR / "MyColl") + self.assertEqual(job.status, "done") + self.assertEqual(job.output_url, "/workspace/MyColl/model.glb") + + def test_missing_weight_variant_fails_the_job_before_generation(self) -> None: + class _MissingVariantRegistry(_FakeRegistry): + def assert_weight_variant_installed(self, params: dict, model_id=None) -> None: + raise RuntimeError(f'{params["gguf_quant"]} weights for trellis2/generate are not installed.') + + gen = _FakeGenerator() + generation.generator_registry = _MissingVariantRegistry(gen) + job_id = "job-missing-variant" + generation._jobs[job_id] = JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = threading.Event() + asyncio.run(generation._run_generation(job_id, b"img", {"gguf_quant": "Q6_K"}, "MyColl")) + + job = generation._jobs[job_id] + self.assertEqual(job.status, "error") + self.assertIn("Q6_K weights for trellis2/generate are not installed", job.error) + self.assertIsNone(gen.outputs_dir) + + +class GenerateFromImageWorkspaceTests(unittest.TestCase): + """The request path must survive the relocation too, not just the worker: + sanitize_collection() checks the name's containment against the workspace + root before the job is filed, so it has to read the same live binding -- a + stale (or missing) module-level name there fails every POST + /generate/from-image, whatever _run_generation does afterwards.""" + + def setUp(self) -> None: + self._prev_registry = generation.generator_registry + self._prev_ws = registry.WORKSPACE_DIR + self._tmp = tempfile.TemporaryDirectory() + registry.WORKSPACE_DIR = Path(self._tmp.name) / "new_workspace" + generation.generator_registry = _FakeRegistry(_FakeGenerator()) + + def tearDown(self) -> None: + generation.generator_registry = self._prev_registry + registry.WORKSPACE_DIR = self._prev_ws + for store in ( + generation._jobs, + generation._cancel_events, + generation._cancelled, + generation._completed_at, + ): + store.clear() + self._tmp.cleanup() + + def test_request_is_filed_after_the_workspace_moves(self) -> None: + background = BackgroundTasks() + response = asyncio.run( + generation.generate_from_image( + background, + image=_FakeUpload(), + model_id="sf3d", + collection="MyColl", + remesh="quad", + enable_texture=False, + texture_resolution=1024, + params="{}", + ) + ) + self.assertIn(response["job_id"], generation._jobs) + # add_task(_run_generation, job_id, image_bytes, full_params, collection) + self.assertEqual(background.tasks[0].args[3], "MyColl") + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_generator_registry.py b/api/tests/test_generator_registry.py index 7340b72c..b21253b7 100644 --- a/api/tests/test_generator_registry.py +++ b/api/tests/test_generator_registry.py @@ -4,6 +4,7 @@ import os import sys import tempfile +import threading import unittest from pathlib import Path @@ -155,6 +156,277 @@ def test_legacy_generator_supports_eager_and_lazy_sibling_imports(self) -> None: self.registry.reload() self.assertNotIn(str(extension.resolve()), sys.path) + def test_scene_and_existing_custom_io_types_are_registered(self) -> None: + for extension_id, input_kind in (("scene-io", "scene"), ("capture-io", "capture"), ("video-io", "video")): + extension = self._make_extension(extension_id) + manifest = { + "id": extension_id, "name": extension_id, "type": "model", + "generator_class": "TestGenerator", + "nodes": [{"id": "generate", "input": input_kind, "output": "scene"}], + } + (extension / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + (extension / "generator.py").write_text( + "from services.generators.base import BaseGenerator\n" + "class TestGenerator(BaseGenerator):\n" + " def load(self): self._model = object()\n" + " def generate(self, value, params, progress_cb=None, cancel_event=None): return self.outputs_dir\n", + encoding="utf-8", + ) + + self.registry.initialize() + self.assertEqual(self.registry.get_manifest("scene-io/generate")["input"], "scene") + self.assertEqual(self.registry.get_manifest("capture-io/generate")["input"], "capture") + self.assertEqual(self.registry.get_manifest("video-io/generate")["input"], "video") + + def test_scene_input_rejects_multi_input_shapes_but_image_multi_can_output_scene(self) -> None: + cases = { + "scene-mixed": {"input": "scene", "inputs": ["scene", "text"], "output": "mesh"}, + "scene-array": {"input": "scene", "inputs": ["scene"], "output": "mesh"}, + "images-scene": {"input": "image", "inputs": ["image", "image"], "output": "scene"}, + } + for extension_id, node in cases.items(): + extension = self._make_extension(extension_id) + (extension / "manifest.json").write_text(json.dumps({ + "id": extension_id, "name": extension_id, "type": "model", + "generator_class": "TestGenerator", + "nodes": [{"id": "generate", **node}], + }), encoding="utf-8") + (extension / "generator.py").write_text( + "from services.generators.base import BaseGenerator\n" + "class TestGenerator(BaseGenerator):\n" + " def load(self): self._model = object()\n" + " def generate(self, value, params, progress_cb=None, cancel_event=None): return self.outputs_dir\n", + encoding="utf-8", + ) + self.registry.initialize() + self.assertIn("scene-mixed/generate", self.registry.load_errors()) + self.assertIn("scene-array/generate", self.registry.load_errors()) + self.assertIn("images-scene/generate", self.registry._generators) + + def test_declared_sources_block_generation_even_when_generator_overrides_readiness(self) -> None: + extension = self._make_extension("multi-source") + manifest = { + "id": "multi-source", + "name": "multi-source", + "type": "model", + "generator_class": "TestGenerator", + "nodes": [{ + "id": "generate", + "model_sources": [ + { + "id": "primary", + "provider": "huggingface", + "repo_id": "org/main", + "destination": ".", + "checks": ["main.bin"], + }, + { + "id": "encoder", + "provider": "huggingface", + "repo_id": "org/encoder", + "destination": "auxiliary/encoder", + "checks": ["encoder.bin"], + }, + ], + }], + } + (extension / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + (extension / "generator.py").write_text( + "\n".join([ + "from services.generators.base import BaseGenerator", + "class TestGenerator(BaseGenerator):", + " def is_downloaded(self): return True", + " def load(self): self._model = object()", + " def generate(self, image_bytes, params, progress_cb=None, cancel_event=None):", + " return self.outputs_dir / 'result.glb'", + ]), + encoding="utf-8", + ) + + self.registry.initialize() + self.registry._active_id = "multi-source/generate" + with self.assertRaisesRegex(RuntimeError, "Model sources are incomplete"): + self.registry.get_active() + self.assertFalse(self.registry.all_status()[0]["downloaded"]) + + model_root = self.models_dir / "multi-source" / "generate" + (model_root / "auxiliary" / "encoder").mkdir(parents=True) + (model_root / "main.bin").write_bytes(b"main") + (model_root / "auxiliary" / "encoder" / "encoder.bin").write_bytes(b"encoder") + self.assertIsNotNone(self.registry.get_active()) + self.assertTrue(self.registry.all_status()[0]["downloaded"]) + (model_root / "main.bin").unlink() + with self.assertRaisesRegex(RuntimeError, "Model sources are incomplete"): + self.registry.get_active() + + def test_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.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)) + + 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)) + 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_selected_weight_variant_must_be_installed_before_generation(self) -> None: + def variant(quant: str) -> dict: + return { + "id": quant, + "include_prefixes": [f"dit_{quant}.gguf"], + "checks": [f"dit_{quant}.gguf"], + } + + extension = self._make_extension("quantized") + manifest = { + "id": "quantized", + "name": "quantized", + "type": "model", + "generator_class": "TestGenerator", + "params_schema": [ + {"id": "quant", "type": "select", "options": [{"value": "Q4"}, {"value": "Q5"}]} + ], + "nodes": [{ + "id": "generate", + "hf_repo": "org/model-gguf", + "download_check": "pipeline.json", + "weight_variants": { + "param": "quant", + "default": "Q5", + "options": [variant("Q4"), variant("Q5")], + }, + }], + } + (extension / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") + (extension / "generator.py").write_text( + "\n".join([ + "from services.generators.base import BaseGenerator", + "class TestGenerator(BaseGenerator):", + " def is_downloaded(self): return True", + " def load(self): self._model = object()", + " def generate(self, image_bytes, params, progress_cb=None, cancel_event=None):", + " return self.outputs_dir / 'result.glb'", + ]), + encoding="utf-8", + ) + + self.registry.initialize() + manifest_variants = self.registry.get_manifest("quantized/generate")["weight_variants"] + self.assertEqual([option["id"] for option in manifest_variants["options"]], ["Q4", "Q5"]) + + model_root = self.models_dir / "quantized" / "generate" + model_root.mkdir(parents=True) + (model_root / "dit_Q5.gguf").write_bytes(b"q5") + # The job's model is checked, not whichever model is still active: the + # switch to the job's model only happens once the job starts running. + self.registry._active_id = "other/generate" + model_id = "quantized/generate" + self.registry.assert_weight_variant_installed({}, model_id) + self.registry.assert_weight_variant_installed({"quant": "Q5"}, model_id) + self.registry.assert_weight_variant_installed({"quant": "fp16"}, model_id) + with self.assertRaisesRegex(RuntimeError, "Q4 weights for quantized/generate are not installed"): + self.registry.assert_weight_variant_installed({"quant": "Q4"}, model_id) + + # Without an explicit model id the active one is used (legacy callers). + self.registry._active_id = model_id + with self.assertRaisesRegex(RuntimeError, "Q4 weights"): + self.registry.assert_weight_variant_installed({"quant": "Q4"}) + + + def test_activate_ready_generator_switches_before_loading_exact_model(self) -> None: + class Generator: + DISPLAY_NAME = "test" + def __init__(self): + self.loaded = False + self.unloads = 0 + def is_downloaded(self): return True + def is_loaded(self): return self.loaded + def load(self): self.loaded = True + def unload(self): + self.loaded = False + self.unloads += 1 + + first = Generator() + second = Generator() + first.loaded = True + self.registry._generators = {"demo/a": first, "demo/b": second} + self.registry._manifests = { + "demo/a": {"name": "A"}, + "demo/b": {"name": "B"}, + } + self.registry._active_id = "demo/a" + + selected = self.registry.activate_ready_generator("demo/b") + + self.assertIs(selected, second) + self.assertEqual(self.registry._active_id, "demo/b") + self.assertFalse(first.loaded) + self.assertEqual(first.unloads, 1) + self.assertTrue(second.loaded) + def test_reload_preserves_legacy_path_owned_by_the_host(self) -> None: extension = self._make_extension("host-owned-path") self._write_manifest(extension, extension_id="host-owned-path") @@ -414,6 +686,27 @@ def test_valid_capability_authorizes_exact_pending_extension_once(self) -> None: self.assertNotIn("pending-update/generate", self.registry._generators) self.assertIn("pending-update/generate", self.registry.load_errors()) + @unittest.skipUnless(sys.platform == "win32", "8.3 short paths are Windows-only") + def test_valid_capability_authorizes_extension_under_a_short_path(self) -> None: + # GitHub's Windows runners use an 8.3 TEMP (C:\Users\RUNNER~1\...): the + # capability destination is resolved (long form) while discovery walks + # the configured, short-form EXTENSIONS_DIR. + import ctypes + + capability = self._make_loadable_pending_extension("pending-short") + buffer = ctypes.create_unicode_buffer(32768) + if not ctypes.windll.kernel32.GetShortPathNameW(str(self.extensions_dir), buffer, len(buffer)): + self.skipTest("short path unavailable") + short_dir = Path(buffer.value) + if str(short_dir) == str(self.extensions_dir): + self.skipTest("8.3 names are disabled on this volume") + registry_module.EXTENSIONS_DIR = short_dir + + self.registry.reload(capability) + + self.assertIn("pending-short/generate", self.registry._generators) + self.assertEqual(self.registry.load_errors(), {}) + def test_public_reload_and_predictable_id_cannot_bypass_pending_state(self) -> None: self._make_loadable_pending_extension("pending-public") @@ -517,5 +810,87 @@ def test_process_extension_is_skipped_without_model_errors(self) -> None: self.assertEqual(self.registry.load_errors(), {}) +class _StatusOnlyGenerator: + DISPLAY_NAME = "Fake" + VRAM_GB = 1 + + def is_downloaded(self) -> bool: + return True + + def is_loaded(self) -> bool: + return False + + def params_schema(self) -> list: + return [{"id": "steps"}] + + +class _DirectGenerator: + def __init__(self, sticky: bool = False) -> None: + self.loaded = True + self.sticky = sticky + + def unload(self) -> None: + if not self.sticky: + self.loaded = False + + def is_loaded(self) -> bool: + return self.loaded + + +class GeneratorRegistryUnloadTests(unittest.TestCase): + def test_unload_all_unloads_every_model_before_reporting_a_stuck_one(self): + registry = GeneratorRegistry() + stuck = _DirectGenerator(sticky=True) + other = _DirectGenerator() + registry._generators = {"demo/stuck": stuck, "demo/other": other} + + with self.assertRaisesRegex(RuntimeError, "demo/stuck"): + registry.unload_all() + + self.assertFalse(other.loaded) + + +class GeneratorRegistryLockTests(unittest.TestCase): + def test_status_reads_do_not_wait_for_an_in_progress_load(self): + # A load holds the lifecycle lock for its whole duration (first-run + # downloads included); status endpoints must keep answering meanwhile. + registry = GeneratorRegistry() + registry._generators["demo/generate"] = _StatusOnlyGenerator() + registry._manifests["demo/generate"] = {"name": "Demo"} + registry._active_id = "demo/generate" + + lock_held = threading.Event() + release = threading.Event() + + def hold_lock() -> None: + with registry._lifecycle_lock: + lock_held.set() + release.wait(5) + + holder = threading.Thread(target=hold_lock) + holder.start() + self.addCleanup(holder.join) + self.addCleanup(release.set) + self.assertTrue(lock_held.wait(5)) + + results = {} + + def read_status() -> None: + results["active"] = registry.active_status() + results["all"] = registry.all_status() + results["params"] = registry.params_schema("demo/generate") + results["model"] = registry.model_status("demo/generate") + + reader = threading.Thread(target=read_status) + reader.start() + reader.join(2) + + self.assertFalse(reader.is_alive(), "status reads blocked on the lifecycle lock") + self.assertEqual(results["active"]["id"], "demo/generate") + self.assertEqual([m["id"] for m in results["all"]], ["demo/generate"]) + self.assertEqual(results["params"], [{"id": "steps"}]) + self.assertFalse(results["model"]["loaded"]) + + if __name__ == "__main__": unittest.main() diff --git a/api/tests/test_llm_downloads.py b/api/tests/test_llm_downloads.py new file mode 100644 index 00000000..31d29915 --- /dev/null +++ b/api/tests/test_llm_downloads.py @@ -0,0 +1,121 @@ +import hashlib +import importlib +import tempfile +import unittest +from pathlib import Path +from unittest import mock + +llm_server = importlib.import_module("services.llm_server") +llm_router = importlib.import_module("routers.llm") + +_VISION = { + "id": "vl", + "hf_filename": "weights.gguf", + "hf_mmproj_filename": "mmproj-F16.gguf", +} + + +class DiscardIncompleteTests(unittest.TestCase): + """Cancelling a download has to leave nothing behind. A vision model fetches + weights then projector, so a cancel during the second one used to leave the + finished weights on disk under a model still reported `downloaded: false` — + the UI offers no trash button for those, so the space was unreclaimable.""" + + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self._orig = llm_server.LLM_MODELS_DIR + llm_server.LLM_MODELS_DIR = Path(self._tmp.name) + + def tearDown(self) -> None: + llm_server.LLM_MODELS_DIR = self._orig + self._tmp.cleanup() + + def _touch(self, name: str) -> Path: + p = llm_server.LLM_MODELS_DIR / name + p.write_bytes(b"x") + return p + + def test_removes_a_part_file(self): + self._touch("weights.gguf.part") + removed = llm_router._discard_incomplete(_VISION) + self.assertEqual(removed, ["weights.gguf.part"]) + self.assertEqual(list(llm_server.LLM_MODELS_DIR.iterdir()), []) + + def test_removes_a_sibling_that_finished_before_the_cancel(self): + self._touch("weights.gguf") # done + self._touch("mmproj-vl.gguf.part") # in flight + removed = llm_router._discard_incomplete(_VISION) + self.assertIn("weights.gguf", removed) + self.assertIn("mmproj-vl.gguf.part", removed) + self.assertEqual(list(llm_server.LLM_MODELS_DIR.iterdir()), []) + + def test_leaves_a_complete_model_alone(self): + self._touch("weights.gguf") + self._touch("mmproj-vl.gguf") + self.assertEqual(llm_router._discard_incomplete(_VISION), []) + self.assertEqual(len(list(llm_server.LLM_MODELS_DIR.iterdir())), 2) + + def test_nothing_on_disk_is_not_an_error(self): + self.assertEqual(llm_router._discard_incomplete(_VISION), []) + + +class _FakeResponse: + def __init__(self, body: bytes) -> None: + self._body = body + self.headers = {"Content-Length": str(len(body))} + + def read(self, _size: int) -> bytes: + chunk, self._body = self._body, b"" + return chunk + + def __enter__(self): + return self + + def __exit__(self, *_exc) -> None: + return None + + +class EngineDigestTests(unittest.TestCase): + """The engine archive's files are executed, so a download whose bytes differ + from the digest GitHub published for the asset must be refused — and must + not be left behind in the temp folder.""" + + BODY = b"llama-server archive bytes" + + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + tmp_dir = self._tmp.name + real_mkstemp = tempfile.mkstemp # llm_server.tempfile is this same module + patches = [ + mock.patch.object(llm_server, "urlopen", lambda *_a, **_k: _FakeResponse(self.BODY)), + mock.patch.object(llm_server.tempfile, "mkstemp", lambda suffix="": real_mkstemp(suffix=suffix, dir=tmp_dir)), + ] + for p in patches: + p.start() + self.addCleanup(p.stop) + + def tearDown(self) -> None: + self._tmp.cleanup() + + def _download(self, digest): + asset = {"name": "llama-bin-win-x64.zip", "browser_download_url": "https://example.invalid/a.zip"} + if digest is not None: + asset["digest"] = digest + return llm_server._download_asset(asset, lambda _msg: None, lambda: None, "engine") + + def test_matching_digest_keeps_the_archive(self): + path = self._download("sha256:" + hashlib.sha256(self.BODY).hexdigest()) + self.assertEqual(path.read_bytes(), self.BODY) + + def test_mismatching_digest_is_refused_and_removed(self): + with self.assertRaises(RuntimeError): + self._download("sha256:" + "0" * 64) + self.assertEqual(list(Path(self._tmp.name).iterdir()), []) + + def test_asset_without_digest_is_still_accepted(self): + path = self._download(None) + self.assertEqual(path.read_bytes(), self.BODY) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_llm_pool_budget.py b/api/tests/test_llm_pool_budget.py new file mode 100644 index 00000000..9bc34a95 --- /dev/null +++ b/api/tests/test_llm_pool_budget.py @@ -0,0 +1,279 @@ +import importlib +import unittest + +llm_server = importlib.import_module("services.llm_server") + + +class _FakeSlot: + """Stands in for a loaded LlamaServerManager: alive, with a last-used stamp + and a VRAM estimate.""" + + def __init__(self, last_used: float, vram_mb: int, busy: int = 0, port: int = 0) -> None: + self._last_used = last_used + self.vram_mb = vram_mb + self.port = port + self.busy_count = busy + self.unloaded = False + + def _alive(self) -> bool: + return not self.unloaded + + def unload(self) -> None: + self.unloaded = True + + +class EstimateVramTests(unittest.TestCase): + def test_declared_estimate_wins(self): + self.assertEqual(llm_server.estimate_vram_mb({"vram_estimate_mb": 4200}), 4200) + + def test_custom_model_is_derived_from_file_size(self): + # 9 GB of weights: enough on its own to fill a 12 GB card. + mb = llm_server.estimate_vram_mb({"size_bytes": 9 * 1024**3}) + self.assertGreater(mb, 11000) + self.assertLess(mb, 12500) + + def test_unknown_size_yields_zero(self): + self.assertEqual(llm_server.estimate_vram_mb({}), 0) + + +class PoolBudgetTests(unittest.TestCase): + """Numbers are the ones measured on the reference machine: an RTX 5070 + reporting 11.9 GiB, so a 768 MiB reserve leaves a 11.4 GiB budget.""" + + def setUp(self) -> None: + self.pool = llm_server.LlamaPool() + self._orig_max = llm_server.resolve_max_models + self._orig_budget = llm_server.vram_budget_mb + llm_server.resolve_max_models = lambda: 2 + llm_server.vram_budget_mb = lambda: int(11.9 * 1024) - 768 # 11417 + + def tearDown(self) -> None: + llm_server.resolve_max_models = self._orig_max + llm_server.vram_budget_mb = self._orig_budget + + def _load(self, **slots) -> None: + for i, (mid, vram) in enumerate(slots.items()): + self.pool._slots[mid] = _FakeSlot(last_used=float(i), vram_mb=vram) + + def test_a_pair_that_fits_is_kept(self): + # qwen3-4b (4200) already loaded, cadquery-coder-7b (6200) incoming. + self._load(qwen4b=4200) + self.pool._enforce_limit_locked(reserve=1, incoming_mb=6200) + self.assertEqual(list(self.pool._slots), ["qwen4b"]) + + def test_a_pair_that_does_not_fit_evicts_the_oldest(self): + # qwen3-4b (4200) + qwen3-vl-8b (7800) = 12000 > 11417. + self._load(qwen4b=4200) + self.pool._enforce_limit_locked(reserve=1, incoming_mb=7800) + self.assertEqual(list(self.pool._slots), []) + + def test_an_oversized_model_gets_the_card_to_itself(self): + # The custom 14B alone exceeds the budget: it still loads, alone. + self._load(qwen4b=4200, coder7b=6200) + self.pool._enforce_limit_locked(reserve=1, incoming_mb=11700) + self.assertEqual(list(self.pool._slots), []) + + def test_count_limit_still_applies_under_the_budget(self): + # Three tiny models fit the VRAM budget but not `max_models: 2`. + self._load(a=500, b=500) + self.pool._enforce_limit_locked(reserve=1, incoming_mb=500) + self.assertEqual(list(self.pool._slots), ["b"]) + + def test_evicts_least_recently_used_first(self): + self._load(oldest=4200, newest=4200) + llm_server.resolve_max_models = lambda: 3 + self.pool._enforce_limit_locked(reserve=1, incoming_mb=7800) + self.assertNotIn("oldest", self.pool._slots) + + def test_no_gpu_falls_back_to_the_count_rule(self): + llm_server.vram_budget_mb = lambda: 0 + self._load(a=9000) + self.pool._enforce_limit_locked(reserve=1, incoming_mb=9000) + self.assertEqual(list(self.pool._slots), ["a"]) + + def test_unknown_estimate_never_evicts(self): + # A model with no size on disk must not push everything out. + self._load(a=4200) + self.pool._enforce_limit_locked(reserve=1, incoming_mb=0) + self.assertEqual(list(self.pool._slots), ["a"]) + + def test_a_concurrent_load_counts_against_the_budget(self): + # ensure() reserves under the lock but loads outside it, so a slot being + # started has no process yet. Counting only live slots let two callers + # (agent turn + workflow LLM node) each load 7.8 GB on an 11.4 GB card. + loading = _FakeSlot(last_used=0.0, vram_mb=7800, port=llm_server.SERVER_PORT) + loading.unloaded = True # reserved: not alive yet + self.pool._slots["incoming-a"] = loading + self.pool._loading_ports.add(llm_server.SERVER_PORT) + self._load(alive_one=4200) + + self.pool._enforce_limit_locked(reserve=1, incoming_mb=7800) + # 7800 (in flight) + 7800 (incoming) already blows the budget, so the + # live 4.2 GB model goes. + self.assertNotIn("alive_one", self.pool._slots) + + def test_a_concurrent_load_counts_against_the_model_limit(self): + loading = _FakeSlot(last_used=0.0, vram_mb=500, port=llm_server.SERVER_PORT) + loading.unloaded = True + self.pool._slots["incoming-a"] = loading + self.pool._loading_ports.add(llm_server.SERVER_PORT) + self._load(a=500) # 1 live + 1 loading + 1 incoming > max_models (2) + + self.pool._enforce_limit_locked(reserve=1, incoming_mb=500) + self.assertNotIn("a", self.pool._slots) + + def test_a_slot_answering_a_request_is_never_evicted(self): + # Evicting it would terminate llama-server mid-generation: the caller + # gets a truncated stream and the node fails with no message. + self.pool._slots["busy"] = _FakeSlot(last_used=0.0, vram_mb=7800, busy=1) + self.pool._enforce_limit_locked(reserve=1, incoming_mb=7800) + self.assertIn("busy", self.pool._slots) + + def test_a_slot_being_started_is_never_evicted(self): + # From the moment Popen returns, a slot is _alive() while the loading + # thread still holds its lock for the whole of _wait_for_health (up to + # 180 s). unload() runs UNDER the pool lock, so evicting it there would + # block every other pool operation — /llm/status, ensure(), and the + # unload_all() the 3D pipeline calls to reclaim VRAM — for that long. + starting = _FakeSlot(last_used=0.0, vram_mb=7800, port=llm_server.SERVER_PORT) + self.pool._slots["starting"] = starting + self.pool._loading_ports.add(llm_server.SERVER_PORT) + + self.pool._enforce_limit_locked(reserve=1, incoming_mb=7800) + self.assertIn("starting", self.pool._slots) + self.assertFalse(starting.unloaded) + + def test_over_capacity_counts_loads_in_flight(self): + loading = _FakeSlot(last_used=0.0, vram_mb=7800, port=llm_server.SERVER_PORT) + loading.unloaded = True # reserved, no process yet + self.pool._slots["incoming-a"] = loading + self.pool._loading_ports.add(llm_server.SERVER_PORT) + # Nothing is alive, yet the card is already spoken for. + self.assertTrue(self.pool._over_capacity_locked(7800)) + self.pool._loading_ports.clear() + self.assertFalse(self.pool._over_capacity_locked(7800)) + + def test_only_loads_and_live_requests_are_worth_waiting_for(self): + # ensure() waits for these to end instead of loading a second model on a + # one-model card; a merely idle slot is evicted rather than waited on. + idle = _FakeSlot(last_used=0.0, vram_mb=4200, port=llm_server.SERVER_PORT) + self.pool._slots["idle"] = idle + self.assertFalse(self.pool._transient_blockers_locked()) + + idle.busy_count = 1 + self.assertTrue(self.pool._transient_blockers_locked()) + + idle.busy_count = 0 + self.pool._loading_ports.add(llm_server.SERVER_PORT + 1) + self.assertTrue(self.pool._transient_blockers_locked()) + + def test_a_crashed_slot_is_restarted_on_a_fresh_port(self): + # Its port was handed to another model meanwhile; respawning in place + # killed that model's server through _kill_stale_server. + dead = _FakeSlot(last_used=0.0, vram_mb=500, port=llm_server.SERVER_PORT) + dead.unloaded = True + self.pool._slots["x"] = dead + self.pool._slots["y"] = _FakeSlot(last_used=1.0, vram_mb=500, port=llm_server.SERVER_PORT) + orig = llm_server.LlamaServerManager.ensure + llm_server.LlamaServerManager.ensure = lambda self, mid, spec: None + try: + slot = self.pool.ensure("x", {"vram_mb": 500}) + finally: + llm_server.LlamaServerManager.ensure = orig + self.assertIsNot(slot, dead) + self.assertNotEqual(slot.port, llm_server.SERVER_PORT) + self.assertFalse(self.pool._slots["y"].unloaded) + + def test_unload_all_spares_requests_and_cold_starts(self): + idle = _FakeSlot(last_used=0.0, vram_mb=4200, port=llm_server.SERVER_PORT) + busy = _FakeSlot(last_used=0.0, vram_mb=4200, busy=1, port=llm_server.SERVER_PORT + 1) + starting = _FakeSlot(last_used=0.0, vram_mb=4200, port=llm_server.SERVER_PORT + 2) + self.pool._slots.update(idle=idle, busy=busy, starting=starting) + self.pool._loading_ports.add(starting.port) + + self.pool.unload_all() + self.assertEqual(sorted(self.pool._slots), ["busy", "starting"]) + self.assertTrue(idle.unloaded) + + self.pool.unload_all(force=True) + self.assertEqual(list(self.pool._slots), []) + self.assertTrue(busy.unloaded and starting.unloaded) + + def test_no_free_port_raises_a_readable_error(self): + for i in range(llm_server.MAX_SLOT_PORTS): + self.pool._loading_ports.add(llm_server.SERVER_PORT + i) + with self.assertRaises(RuntimeError) as ctx: + self.pool.ensure("whatever", {"vram_mb": 500}) + self.assertIn("slots", str(ctx.exception)) + + +class ReaperTests(unittest.TestCase): + def setUp(self) -> None: + self.pool = llm_server.LlamaPool() + + def test_an_idle_slot_is_unloaded(self): + old = _FakeSlot(last_used=0.0, vram_mb=4200) + self.pool._slots["old"] = old + self.pool._reap_once(now=llm_server.IDLE_TTL_SECONDS + 1) + self.assertEqual(list(self.pool._slots), []) + self.assertTrue(old.unloaded) + + def test_a_slot_answering_a_long_request_survives(self): + # A generation longer than the TTL (Text-to-CAD, a 14B on CPU) used to + # be killed mid-answer: _last_used only moved once a call had finished. + busy = _FakeSlot(last_used=0.0, vram_mb=4200, busy=1) + self.pool._slots["busy"] = busy + self.pool._reap_once(now=llm_server.IDLE_TTL_SECONDS * 10) + self.assertEqual(list(self.pool._slots), ["busy"]) + self.assertFalse(busy.unloaded) + + def test_a_crashed_slot_is_dropped(self): + dead = _FakeSlot(last_used=0.0, vram_mb=4200) + dead.unloaded = True + self.pool._slots["dead"] = dead + self.pool._reap_once(now=1.0) + self.assertEqual(list(self.pool._slots), []) + + + def test_a_slot_being_started_is_not_idle(self): + # Same reason as the eviction case: the reaper also unloads under the + # pool lock, so a slot mid-cold-start must not be its victim. + starting = _FakeSlot(last_used=0.0, vram_mb=4200, port=llm_server.SERVER_PORT) + self.pool._slots["starting"] = starting + self.pool._loading_ports.add(llm_server.SERVER_PORT) + self.pool._reap_once(now=llm_server.IDLE_TTL_SECONDS * 10) + self.assertEqual(list(self.pool._slots), ["starting"]) + self.assertFalse(starting.unloaded) + + +class BusyContextTests(unittest.TestCase): + def test_busy_counts_nest_and_stamp_last_used(self): + slot = llm_server.LlamaServerManager(port=1) + self.assertEqual(slot.busy_count, 0) + with slot.busy(): + self.assertEqual(slot.busy_count, 1) + with slot.busy(): + self.assertEqual(slot.busy_count, 2) + self.assertEqual(slot.busy_count, 1) + self.assertEqual(slot.busy_count, 0) + self.assertGreater(slot._last_used, 0.0) + + def test_hold_and_release_pair_up(self): + # ensure(hold=True) claims the slot before returning it; the caller + # releases it once the request is over. + slot = llm_server.LlamaServerManager(port=1) + slot.hold() + self.assertEqual(slot.busy_count, 1) + slot.release() + self.assertEqual(slot.busy_count, 0) + + def test_busy_is_released_when_the_request_raises(self): + slot = llm_server.LlamaServerManager(port=1) + with self.assertRaises(ValueError): + with slot.busy(): + raise ValueError("client disconnected") + self.assertEqual(slot.busy_count, 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_mesh_ops_operations.py b/api/tests/test_mesh_ops_operations.py new file mode 100644 index 00000000..b7ffeff0 --- /dev/null +++ b/api/tests/test_mesh_ops_operations.py @@ -0,0 +1,323 @@ +import io +import json +import shutil +import subprocess +import tempfile +import unittest +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import patch + +from services.mesh_ops import MeshOpContext +from services.mesh_ops import operations + + +class _FakeGeometry: + def __init__(self, faces=None) -> None: + self.faces = faces if faces is not None else [1, 2, 3] + self.vertices = [1, 2, 3, 4] + self.exports = [] + + def export(self, path) -> None: + self.exports.append(str(path)) + Path(path).touch() + + +class _FakeScene: + pass + + +class _FakeMesh: + def vertex_number(self) -> int: + return 4 + + def face_number(self) -> int: + return 3 + + +class _FakeMeshSet: + def __init__(self) -> None: + self.calls = [] + + def current_mesh(self): + return _FakeMesh() + + def __getattr__(self, name): + def call(*args, **kwargs): + self.calls.append((name, args, kwargs)) + if name == "save_current_mesh": + Path(args[0]).touch() + + return call + + +class _InputCapture: + def __init__(self) -> None: + self.value = "" + + def write(self, value) -> None: + self.value += value + + def close(self) -> None: + pass + + +class _FakeProcess: + def __init__(self, messages, return_code=0, stderr="") -> None: + self.stdin = _InputCapture() + self.stdout = iter(f"{json.dumps(message)}\n" for message in messages) + self.stderr = io.StringIO(stderr) + self.return_code = return_code + self.killed = False + + def wait(self) -> int: + return self.return_code + + def kill(self) -> None: + self.killed = True + + +class MeshOpOperationRegressionTests(unittest.TestCase): + def test_meshopt_runner_is_valid_javascript(self) -> None: + node = shutil.which("node") or shutil.which("nodejs") + if node is None: + self.skipTest("Node.js is unavailable") + runner = Path(operations.__file__).with_name("meshopt_runner.cjs") + result = subprocess.run( + [node, "--check", str(runner)], + capture_output=True, + text=True, + check=False, + ) + self.assertEqual(result.returncode, 0, result.stderr) + + def test_repair_keeps_the_original_filter_order_and_arguments(self) -> None: + mesh_set = _FakeMeshSet() + source = _FakeGeometry() + result_geometry = _FakeGeometry(faces=[1, 2]) + events = [] + + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + input_path = root / "input.glb" + output_path = root / "output.glb" + input_path.touch() + context = MeshOpContext( + workspace_dir=root, + temp_dir=root, + output_path=output_path, + progress_cb=lambda percent, label: events.append( + ("progress", percent, label) + ), + log_cb=lambda message: events.append(("log", message)), + ) + + with ( + patch.object( + operations, + "_mesh_libraries", + return_value=( + SimpleNamespace(MeshSet=lambda: mesh_set), + object(), + ), + ), + patch.object(operations, "_load_single_mesh", return_value=source), + patch.object( + operations, + "_raw_geometry", + return_value=result_geometry, + ), + ): + result = operations.repair_mesh(input_path, {}, context) + + self.assertEqual(result.file_path, output_path) + self.assertEqual(result.details, {"face_count": 2}) + self.assertEqual( + [call[0] for call in mesh_set.calls], + [ + "load_new_mesh", + "meshing_remove_duplicate_vertices", + "meshing_remove_duplicate_faces", + "meshing_remove_null_faces", + "meshing_remove_folded_faces", + "meshing_repair_non_manifold_edges", + "meshing_repair_non_manifold_vertices", + "meshing_close_holes", + "save_current_mesh", + ], + ) + self.assertEqual(mesh_set.calls[5][2], {"method": 0}) + self.assertEqual( + mesh_set.calls[7][2], + { + "maxholesize": 2000, + "newfaceselected": False, + "selfintersection": False, + }, + ) + self.assertIn(("progress", 100, "Done"), events) + + def test_smooth_keeps_taubin_and_laplacian_parameter_semantics(self) -> None: + for mode, expected_method, expected_arguments in ( + ( + "taubin", + "apply_coord_taubin_smoothing", + {"lambda_": 0.4, "mu": -0.41000000000000003, "stepsmoothnum": 7}, + ), + ( + "laplacian", + "apply_coord_laplacian_smoothing", + {"stepsmoothnum": 7, "cotangentweight": False}, + ), + ): + with self.subTest(mode=mode), tempfile.TemporaryDirectory() as directory: + mesh_set = _FakeMeshSet() + geometry = _FakeGeometry() + root = Path(directory) + input_path = root / "input.glb" + input_path.touch() + context = MeshOpContext( + workspace_dir=root, + temp_dir=root, + output_path=root / "output.glb", + ) + + with ( + patch.object( + operations, + "_mesh_libraries", + return_value=( + SimpleNamespace(MeshSet=lambda: mesh_set), + SimpleNamespace(Scene=_FakeScene), + ), + ), + patch.object( + operations, + "_load_single_mesh", + return_value=geometry, + ), + patch.object( + operations, + "_raw_geometry", + return_value=geometry, + ), + ): + operations.smooth_mesh( + input_path, + {"iterations": 7, "lambda_": 0.4, "mode": mode}, + context, + ) + + smoothing_call = next( + call for call in mesh_set.calls if call[0] == expected_method + ) + self.assertEqual(smoothing_call[2], expected_arguments) + + def test_legacy_smooth_keeps_its_original_laplacian_arguments(self) -> None: + mesh_set = _FakeMeshSet() + geometry = _FakeGeometry() + fake_trimesh = SimpleNamespace( + Scene=_FakeScene, + load=lambda path, **kwargs: geometry, + ) + + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + input_path = root / "input.glb" + input_path.touch() + context = MeshOpContext( + workspace_dir=root, + temp_dir=root, + output_path=root / "output.glb", + preserve_visuals=True, + ) + with ( + patch.object( + operations, + "_mesh_libraries", + return_value=( + SimpleNamespace(MeshSet=lambda: mesh_set), + fake_trimesh, + ), + ), + patch.object( + operations, + "_load_single_mesh", + return_value=geometry, + ), + patch.object(operations, "_has_texture", return_value=False), + ): + operations.smooth_mesh( + input_path, + {"iterations": 9, "lambda_": 0.5, "mode": "laplacian"}, + context, + ) + + smoothing_call = next( + call + for call in mesh_set.calls + if call[0] == "apply_coord_laplacian_smoothing" + ) + self.assertEqual(smoothing_call[2], {"stepsmoothnum": 9}) + + def test_decimate_forwards_meshopt_progress_logs_and_result(self) -> None: + messages = [ + {"type": "log", "message": "Current triangles: 12"}, + {"type": "progress", "percent": 55, "label": "Simplifying mesh…"}, + { + "type": "done", + "result": {"filePath": "/workspace/result.glb", "faceCount": 5}, + }, + ] + fake_process = _FakeProcess(messages) + events = [] + + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + input_path = root / "input.glb" + input_path.touch() + context = MeshOpContext( + workspace_dir=root, + temp_dir=root, + progress_cb=lambda percent, label: events.append( + ("progress", percent, label) + ), + log_cb=lambda message: events.append(("log", message)), + ) + with ( + patch.object( + operations, + "_node_executable", + return_value=("node", False), + ), + patch.object( + operations, + "_meshopt_dependency_dir", + return_value=root, + ), + patch.object( + operations.subprocess, + "Popen", + return_value=fake_process, + ) as popen, + ): + result = operations.decimate_mesh( + input_path, + {"target_faces": 5}, + context, + ) + + self.assertEqual(result.file_path, Path("/workspace/result.glb")) + self.assertEqual(result.details, {"face_count": 5}) + self.assertIn(("log", "Current triangles: 12"), events) + self.assertIn(("progress", 55, "Simplifying mesh…"), events) + command = popen.call_args.args[0] + self.assertEqual(command[0], "node") + self.assertTrue(command[1].endswith("meshopt_runner.cjs")) + payload = json.loads(fake_process.stdin.value) + self.assertEqual(payload["inputPath"], str(input_path)) + self.assertEqual(payload["params"], {"target_faces": 5}) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_mesh_ops_processor.py b/api/tests/test_mesh_ops_processor.py new file mode 100644 index 00000000..2f8a6659 --- /dev/null +++ b/api/tests/test_mesh_ops_processor.py @@ -0,0 +1,76 @@ +import io +import json +import tempfile +import unittest +from contextlib import redirect_stdout +from pathlib import Path +from unittest.mock import patch + +from services.mesh_ops import MeshOpResult +from services.mesh_ops import processor + + +class _FakeRegistry: + def __init__(self, output_path: Path) -> None: + self.output_path = output_path + self.calls = [] + + def run(self, operation_id, input_path, params, context): + self.calls.append((operation_id, input_path, params, context)) + context.progress(35, "Working…") + context.log("shared implementation") + return MeshOpResult(self.output_path) + + +class MeshOpProcessorTests(unittest.TestCase): + def test_workflow_protocol_forwards_to_registry_and_preserves_events(self) -> None: + with tempfile.TemporaryDirectory() as directory: + root = Path(directory) + input_path = root / "input.glb" + output_path = root / "output.glb" + input_path.touch() + registry = _FakeRegistry(output_path) + request = { + "input": {"filePath": str(input_path)}, + "params": {"iterations": 7}, + "workspaceDir": str(root), + "tempDir": str(root / "tmp"), + } + + stdout = io.StringIO() + with ( + patch.object(processor, "mesh_ops_registry", registry), + patch.object(processor.sys, "stdin", io.StringIO(json.dumps(request))), + redirect_stdout(stdout), + ): + processor.run_processor("smooth", "mesh-smoother") + + messages = [json.loads(line) for line in stdout.getvalue().splitlines()] + self.assertEqual( + [message["type"] for message in messages], + ["progress", "log", "done"], + ) + self.assertEqual(messages[-1]["result"]["filePath"], str(output_path)) + self.assertEqual(registry.calls[0][0:3], ( + "smooth", + input_path, + {"iterations": 7}, + )) + self.assertEqual(registry.calls[0][3].workspace_dir, root) + + def test_missing_input_uses_existing_node_error_contract(self) -> None: + stdout = io.StringIO() + request = {"input": {}, "params": {}} + with ( + patch.object(processor.sys, "stdin", io.StringIO(json.dumps(request))), + redirect_stdout(stdout), + ): + processor.run_processor("repair", "mesh-repair") + + message = json.loads(stdout.getvalue()) + self.assertEqual(message["type"], "error") + self.assertIn("mesh-repair: input file not found: None", message["message"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_mesh_ops_registry.py b/api/tests/test_mesh_ops_registry.py new file mode 100644 index 00000000..31a777a2 --- /dev/null +++ b/api/tests/test_mesh_ops_registry.py @@ -0,0 +1,194 @@ +import json +import tempfile +import unittest +from pathlib import Path + +from services.mesh_ops import ( + MeshOp, + MeshOpContext, + MeshOpNotFoundError, + MeshOpResult, + MeshOpsRegistry, + mesh_ops_registry, +) + + +class MeshOpsRegistryTests(unittest.TestCase): + def test_builtin_metadata_is_serializable_and_complete(self) -> None: + descriptions = mesh_ops_registry.describe() + + self.assertEqual( + [description["id"] for description in descriptions], + ["repair", "decimate", "smooth"], + ) + for description in descriptions: + self.assertIn(description["category"], {"repair", "optimization"}) + self.assertFalse(description["destructive"]) + self.assertTrue(description["undoable"]) + self.assertIsInstance(description["params_schema"], list) + json.dumps(descriptions) + + def test_run_applies_schema_defaults_without_mutating_metadata(self) -> None: + calls = [] + + def operation(input_path, params, context): + calls.append((input_path, params, context)) + return MeshOpResult(input_path) + + registry = MeshOpsRegistry( + [ + MeshOp( + id="example", + label="Example", + params_schema=( + {"id": "amount", "type": "int", "default": 3}, + ), + fn=operation, + category="test", + ) + ] + ) + + with tempfile.TemporaryDirectory() as directory: + input_path = Path(directory) / "mesh.glb" + input_path.touch() + context = MeshOpContext(Path(directory), Path(directory)) + registry.run("example", input_path, {"extra": True}, context) + + self.assertEqual(calls[0][1], {"amount": 3, "extra": True}) + calls[0][1]["amount"] = 99 + self.assertEqual( + registry.describe()[0]["params_schema"][0]["default"], + 3, + ) + + def test_run_clamps_supplied_values_to_schema_bounds(self) -> None: + calls = [] + + def operation(input_path, params, context): + calls.append(params) + return MeshOpResult(input_path) + + registry = MeshOpsRegistry( + [ + MeshOp( + id="bounded", + label="Bounded", + params_schema=( + { + "id": "count", + "type": "int", + "default": 3, + "min": 1, + "max": 5, + }, + { + "id": "strength", + "type": "float", + "default": 0.5, + "min": 0.1, + "max": 1.0, + }, + ), + fn=operation, + category="test", + ) + ] + ) + + with tempfile.TemporaryDirectory() as directory: + input_path = Path(directory) / "mesh.glb" + input_path.touch() + context = MeshOpContext(Path(directory), Path(directory)) + registry.run( + "bounded", + input_path, + {"count": 99, "strength": -2.0}, + context, + ) + + self.assertEqual(calls[0], {"count": 5, "strength": 0.1}) + + def test_run_rejects_non_numeric_values_for_bounded_params(self) -> None: + operation = MeshOp( + id="bounded", + label="Bounded", + params_schema=( + {"id": "amount", "type": "int", "min": 1, "max": 5}, + ), + fn=lambda path, params, context: MeshOpResult(path), + category="test", + ) + registry = MeshOpsRegistry([operation]) + + with tempfile.TemporaryDirectory() as directory: + input_path = Path(directory) / "mesh.glb" + input_path.touch() + context = MeshOpContext(Path(directory), Path(directory)) + with self.assertRaisesRegex(TypeError, "'amount' must be numeric"): + registry.run( + "bounded", + input_path, + {"amount": "a lot"}, + context, + ) + + def test_invalid_duplicate_and_unknown_ids_are_rejected(self) -> None: + operation = MeshOp( + id="valid", + label="Valid", + params_schema=(), + fn=lambda path, params, context: MeshOpResult(path), + category="test", + ) + registry = MeshOpsRegistry([operation]) + + with self.assertRaisesRegex(ValueError, "Duplicate"): + registry.register(operation) + with self.assertRaisesRegex(ValueError, "Invalid"): + registry.register( + MeshOp( + id="Not Valid", + label="Invalid", + params_schema=(), + fn=operation.fn, + category="test", + ) + ) + with self.assertRaises(MeshOpNotFoundError): + registry.get("missing") + + def test_workflow_manifests_share_registry_schemas_and_thin_adapters(self) -> None: + repository_root = Path(__file__).resolve().parents[2] + nodes_root = repository_root / "src" / "areas" / "workflows" / "nodes" + cases = { + "repair": ("mesh-repair", "repair"), + "decimate": ("mesh-optimizer", "decimate"), + "smooth": ("mesh-smoother", "smooth"), + } + + descriptions = { + description["id"]: description + for description in mesh_ops_registry.describe() + } + for operation_id, (extension_id, wrapper_operation_id) in cases.items(): + extension_dir = nodes_root / extension_id + manifest = json.loads( + (extension_dir / "manifest.json").read_text(encoding="utf-8") + ) + self.assertEqual(manifest["entry"], "processor.py") + self.assertEqual( + manifest["nodes"][0]["params_schema"], + descriptions[operation_id]["params_schema"], + ) + + wrapper = (extension_dir / "processor.py").read_text(encoding="utf-8") + self.assertIn( + f'run_processor("{wrapper_operation_id}", "{extension_id}")', + wrapper, + ) + self.assertLess(len(wrapper.splitlines()), 30) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_mesh_routers_workspace.py b/api/tests/test_mesh_routers_workspace.py new file mode 100644 index 00000000..f03e86df --- /dev/null +++ b/api/tests/test_mesh_routers_workspace.py @@ -0,0 +1,144 @@ +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +import trimesh +from fastapi import HTTPException + +import routers.export as export_router +import routers.optimize as optimize_router +import services.generator_registry as registry +from services.mesh_ops import MeshOpResult + +IDENTITY = [ + [1.0, 0.0, 0.0, 0.0], + [0.0, 1.0, 0.0, 0.0], + [0.0, 0.0, 1.0, 0.0], + [0.0, 0.0, 0.0, 1.0], +] + + +class MeshRoutersAfterWorkspaceMoveTests(unittest.TestCase): + """Export and mesh-edit endpoints must resolve paths against the workspace as + it is *now*. POST /settings/paths rebinds registry.WORKSPACE_DIR when the user + moves the workspace; a name captured at import keeps pointing at the old + folder, so a model generated after the move can't be exported or edited.""" + + def setUp(self) -> None: + self._tmp = tempfile.TemporaryDirectory() + self.root = Path(self._tmp.name) + self._prev_ws = registry.WORKSPACE_DIR + # The user moved the workspace: the registry global now points here. + registry.WORKSPACE_DIR = self.root / "new_workspace" + # Keep the test hermetic against import-time bindings: if a router still + # holds its own WORKSPACE_DIR name, point it at an empty "old" folder in + # the temp tree so the assertions -- not the real workspace -- catch it. + (self.root / "old_workspace").mkdir() + self._stale = [] + for module in (export_router, optimize_router): + if hasattr(module, "WORKSPACE_DIR"): + self._stale.append((module, module.WORKSPACE_DIR)) + module.WORKSPACE_DIR = self.root / "old_workspace" + + mesh_dir = registry.WORKSPACE_DIR / "MyColl" + mesh_dir.mkdir(parents=True) + trimesh.creation.box().export(mesh_dir / "mesh.glb") + + def tearDown(self) -> None: + registry.WORKSPACE_DIR = self._prev_ws + for module, value in self._stale: + module.WORKSPACE_DIR = value + self._tmp.cleanup() + + def test_export_router_converts_a_mesh_in_the_moved_workspace(self) -> None: + response = export_router.export_mesh("stl", "MyColl/mesh.glb") + self.assertEqual(response.status_code, 200) + self.assertGreater(len(response.body), 0) + + def test_optimize_export_converts_a_mesh_in_the_moved_workspace(self) -> None: + response = optimize_router.export_mesh(path="MyColl/mesh.glb", format="obj") + self.assertEqual(response.status_code, 200) + self.assertIn(b"v ", response.body) + + def test_transform_writes_its_result_into_the_moved_workspace(self) -> None: + result = optimize_router.transform_mesh( + optimize_router.TransformRequest(path="MyColl/mesh.glb", matrix=IDENTITY) + ) + self.assertTrue(result["url"].startswith("/workspace/MyColl/mesh_xf_")) + written = registry.WORKSPACE_DIR / result["url"].removeprefix("/workspace/") + self.assertTrue(written.is_file()) + + def test_decimate_and_smooth_read_their_input_from_the_moved_workspace(self) -> None: + # /optimize/mesh and /optimize/smooth resolve their input through this helper. + resolved = optimize_router._resolve_input_path("MyColl/mesh.glb") + self.assertEqual(resolved, (registry.WORKSPACE_DIR / "MyColl" / "mesh.glb").resolve()) + + def test_decimate_and_smooth_write_their_result_into_the_moved_workspace(self) -> None: + # The backends (meshoptimizer, pymeshlab) aren't available here, so stand in + # for the mesh-ops registry and check the router's own path handling: the + # workspace it hands the operation, where the output goes, and the URL. + class _RecordingRegistry: + def __init__(self) -> None: + self.contexts = [] + + def run(self, operation_id, input_path, params, context): + self.contexts.append(context) + context.output_path.touch() + return MeshOpResult(context.output_path, {"face_count": 12}) + + ops = _RecordingRegistry() + with patch.object(optimize_router, "mesh_ops_registry", ops): + decimated = optimize_router.optimize_mesh( + optimize_router.OptimizeRequest(path="MyColl/mesh.glb", target_faces=500) + ) + smoothed = optimize_router.smooth_mesh( + optimize_router.SmoothRequest(path="MyColl/mesh.glb", iterations=2) + ) + + self.assertEqual(decimated["url"], "/workspace/MyColl/mesh_opt500.glb") + self.assertEqual(smoothed["url"], "/workspace/MyColl/mesh_smooth2.glb") + for context in ops.contexts: + self.assertEqual(context.workspace_dir, registry.WORKSPACE_DIR) + self.assertEqual(context.output_path.parent, registry.WORKSPACE_DIR / "MyColl") + + def test_a_path_leaving_the_workspace_is_still_refused(self) -> None: + # Reading the live workspace must not loosen the containment check. + trimesh.creation.box().export(self.root / "outside.glb") + calls = ( + lambda: export_router.export_mesh("stl", "../outside.glb"), + lambda: optimize_router.export_mesh(path="../outside.glb", format="obj"), + ) + for call in calls: + with self.assertRaises(HTTPException) as raised: + call() + self.assertEqual(raised.exception.status_code, 400) + + def test_a_sibling_folder_sharing_the_workspace_prefix_is_refused(self) -> None: + # "-secret" starts with the workspace's path string, so a string + # prefix check let it through; the containment check must compare ancestry. + sibling = self.root / "new_workspace-secret" + sibling.mkdir() + trimesh.creation.box().export(sibling / "private.glb") + escaping = "../new_workspace-secret/private.glb" + calls = { + "/export/{fmt}": lambda: export_router.export_mesh("stl", escaping), + "/optimize/export": lambda: optimize_router.export_mesh(path=escaping, format="obj"), + "/optimize/mesh, /smooth, /transform": lambda: optimize_router._resolve_input_path(escaping), + "/optimize/ply-to-splat": lambda: optimize_router.ply_to_splat(escaping), + } + for route, call in calls.items(): + with self.subTest(route=route), self.assertRaises(HTTPException) as raised: + call() + self.assertEqual(raised.exception.status_code, 400) + + def test_is_within_workspace_compares_ancestry(self) -> None: + workspace = registry.WORKSPACE_DIR.resolve() + self.assertTrue(registry.is_within_workspace(workspace)) + self.assertTrue(registry.is_within_workspace(workspace / "MyColl" / "mesh.glb")) + self.assertFalse(registry.is_within_workspace(workspace.parent / "new_workspace-secret" / "x.glb")) + self.assertFalse(registry.is_within_workspace(workspace.parent)) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_model_router.py b/api/tests/test_model_router.py new file mode 100644 index 00000000..4090a0ea --- /dev/null +++ b/api/tests/test_model_router.py @@ -0,0 +1,313 @@ +import asyncio +import json +import sys +import tempfile +import threading +import types +import unittest +from pathlib import Path +from unittest.mock import patch + +from starlette.requests import Request + +import routers.model as model_router + + +SOURCES = [ + { + "id": "primary", + "provider": "huggingface", + "repo_id": "org/main", + "destination": ".", + "checks": ["main.bin"], + }, + { + "id": "encoder", + "provider": "huggingface", + "repo_id": "org/encoder", + "destination": "auxiliary/encoder", + "checks": ["encoder.bin"], + }, +] + + +def request_for(sources: list[dict]) -> Request: + body = json.dumps({"sources": sources}).encode() + sent = False + + async def receive(): + nonlocal sent + if sent: + return {"type": "http.disconnect"} + sent = True + return {"type": "http.request", "body": body, "more_body": False} + + return Request({ + "type": "http", + "method": "POST", + "path": "/model/hf-download-sources", + "headers": [(b"authorization", b"Bearer test-token")], + "query_string": b"", + "server": ("test", 80), + "client": ("test", 1), + "scheme": "http", + }, receive) + + +async def collect_events(response) -> list[dict]: + payload = "" + async for chunk in response.body_iterator: + payload += chunk.decode() if isinstance(chunk, bytes) else chunk + return [ + json.loads(block[6:]) + for block in payload.strip().split("\n\n") + if block.startswith("data: ") + ] + + +class MultiSourceRouterTests(unittest.TestCase): + def setUp(self) -> None: + self.tempdir = tempfile.TemporaryDirectory(prefix="modly-model-router-") + self.models_dir = Path(self.tempdir.name) / "models" + self.models_dir.mkdir() + self.old_models_dir = model_router.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.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) + else: + sys.modules["huggingface_hub"] = self.old_hf_module + self.tempdir.cleanup() + + def install_hf_stub(self, files: dict[str, list[str]], calls: list[str]) -> None: + module = types.ModuleType("huggingface_hub") + + def list_repo_files(repo_id, revision=None, token=None): + calls.append(f"list:{repo_id}:{revision}:{token}") + return files[repo_id] + + def hf_hub_url(repo_id, filename, revision=None): + return f"https://example.invalid/{repo_id}/{revision or 'main'}/{filename}" + + module.list_repo_files = list_repo_files + module.hf_hub_url = hf_hub_url + sys.modules["huggingface_hub"] = module + + def test_lists_every_source_before_sequential_download_with_monotonic_progress(self) -> None: + calls: list[str] = [] + controls: list[int] = [] + self.install_hf_stub({"org/main": ["main.bin"], "org/encoder": ["encoder.bin"]}, calls) + + def fake_download(**kwargs): + calls.append(f"download:{kwargs['dest_dir']}:{kwargs['filename']}") + controls.append(id(kwargs["control"])) + target = Path(kwargs["dest_dir"]) / kwargs["filename"] + target.parent.mkdir(parents=True, exist_ok=True) + target.write_bytes(b"data") + kwargs["progress_cb"]({ + "percent": kwargs["base_percent"], + "file": kwargs["filename"], + "fileIndex": kwargs["file_index"], + "totalFiles": kwargs["total_files"], + "status": "Downloading...", + "bytesDownloaded": 4, + "stalledSeconds": 0, + }) + return 4 + + async def run(): + with patch.object(model_router, "_download_file_streamed", fake_download): + response = await model_router.hf_download_sources( + request_for(SOURCES), "pixal3d/generate" + ) + return await collect_events(response) + + events = asyncio.run(run()) + first_download = next(index for index, value in enumerate(calls) if value.startswith("download:")) + self.assertTrue(all(value.startswith("list:") for value in calls[:first_download])) + self.assertEqual(len(set(controls)), 1) + self.assertEqual([event["percent"] for event in events if "percent" in event], sorted( + event["percent"] for event in events if "percent" in event + )) + self.assertEqual(events[-1], {"percent": 100, "status": "done"}) + self.assertTrue((self.models_dir / "pixal3d/generate/main.bin").is_file()) + self.assertTrue((self.models_dir / "pixal3d/generate/auxiliary/encoder/encoder.bin").is_file()) + + def test_pause_cancel_and_resume_reuse_one_model_control(self) -> None: + calls: list[str] = [] + self.install_hf_stub({"org/main": ["main.bin"]}, calls) + source = [SOURCES[0]] + mode = "pause" + + def controlled_download(**kwargs): + target = Path(kwargs["dest_dir"]) / kwargs["filename"] + target.parent.mkdir(parents=True, exist_ok=True) + part = Path(f"{target}.part") + part.write_bytes(b"partial") + if mode == "pause": + kwargs["control"]["pause"].set() + model_router._check_download_control(kwargs["control"]) + if mode == "cancel": + kwargs["control"]["cancel"].set() + model_router._check_download_control(kwargs["control"]) + part.replace(target) + return target.stat().st_size + + async def one_run(): + with patch.object(model_router, "_download_file_streamed", controlled_download): + response = await model_router.hf_download_sources( + request_for(source), "pixal3d/generate" + ) + return await collect_events(response) + + paused = asyncio.run(one_run()) + self.assertTrue(paused[-1]["paused"]) + self.assertTrue((self.models_dir / "pixal3d/generate/main.bin.part").is_file()) + + mode = "cancel" + cancelled = asyncio.run(one_run()) + self.assertTrue(cancelled[-1]["cancelled"]) + self.assertFalse((self.models_dir / "pixal3d/generate/main.bin.part").exists()) + + mode = "resume" + resumed = asyncio.run(one_run()) + self.assertEqual(resumed[-1], {"percent": 100, "status": "done"}) + self.assertTrue((self.models_dir / "pixal3d/generate/main.bin").is_file()) + + def test_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) + + async def run(): + response = await model_router.hf_download_sources( + request_for([SOURCES[0]]), "pixal3d/generate" + ) + return await collect_events(response) + + events = asyncio.run(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) + self.assertEqual(model_router.Request.__module__, "urllib.request") + + +class ModelLifecycleEndpointTests(unittest.TestCase): + def test_status_and_switch_do_not_block_the_event_loop_while_lifecycle_is_contended(self) -> None: + class ContendedRegistry: + def __init__(self): + self.lock = threading.Lock() + def active_status(self): + with self.lock: + return {"id": "demo/a", "loaded": True} + def switch_model(self, model_id): + with self.lock: + return None + + registry = ContendedRegistry() + previous = model_router.generator_registry + model_router.generator_registry = registry + + async def assert_responsive(call): + registry.lock.acquire() + release = threading.Timer(0.25, registry.lock.release) + release.start() + try: + start = asyncio.get_running_loop().time() + task = asyncio.create_task(call()) + await asyncio.sleep(0.02) + elapsed = asyncio.get_running_loop().time() - start + self.assertLess(elapsed, 0.15) + return await task + finally: + release.join() + + async def run(): + status = await assert_responsive(model_router.model_status) + switched = await assert_responsive(lambda: model_router.switch_model("demo/b")) + return status, switched + + try: + status, switched = asyncio.run(run()) + finally: + model_router.generator_registry = previous + + self.assertEqual(status["id"], "demo/a") + self.assertEqual(switched, {"active": "demo/b"}) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_model_sources.py b/api/tests/test_model_sources.py new file mode 100644 index 00000000..54ecbf06 --- /dev/null +++ b/api/tests/test_model_sources.py @@ -0,0 +1,249 @@ +import os +import tempfile +import unittest +from pathlib import Path + +from services.model_sources import ( + installed_weight_variants, + missing_weight_variant, + model_sources_are_downloaded, + normalize_model_sources, + normalize_weight_group_references, + normalize_weight_groups, + normalize_weight_variants, + resolve_model_root, + resolve_weight_group_root, + resolve_weight_storage_root, + validate_source_file_plan, + validate_model_node_ids, + weight_group_sources_are_downloaded, +) + + +def valid_node() -> dict: + return { + "model_sources": [ + { + "id": "primary", + "provider": "huggingface", + "repo_id": "org/main", + "destination": ".", + "checks": ["pipeline.json"], + }, + { + "id": "encoder", + "provider": "huggingface", + "repo_id": "org/encoder", + "revision": "refs/pr/1", + "destination": "auxiliary/encoder", + "include_prefixes": ["config.json", "weights/"], + "checks": ["config.json", "model.safetensors"], + }, + ] + } + + +class ModelSourcesTests(unittest.TestCase): + def test_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"]) + self.assertIsNone(normalize_model_sources({ + "hf_repo": "legacy/repo", + "download_check": "../generate/model.safetensors", + "hf_skip_prefixes": ["weights/**"], + })) + + def test_rejects_unsafe_and_non_portable_declarations(self) -> None: + source = valid_node()["model_sources"][0] + for destination in ("../outside", "aux/CON", "aux/name.", "C:/models"): + with self.subTest(destination=destination), self.assertRaises(ValueError): + normalize_model_sources({ + "model_sources": [{**source, "destination": destination}] + }) + with self.assertRaisesRegex(ValueError, "provider"): + normalize_model_sources({ + "model_sources": [{**source, "provider": "url"}] + }) + with self.assertRaisesRegex(ValueError, "portable-unique"): + normalize_model_sources({ + "model_sources": [source, {**source, "id": "PRIMARY"}] + }) + with self.assertRaisesRegex(ValueError, "checks"): + normalize_model_sources({ + "model_sources": [{**source, "checks": []}] + }) + + def test_rejects_portable_cross_source_file_collisions(self) -> None: + sources = normalize_model_sources(valid_node()) or [] + with self.assertRaisesRegex(ValueError, "portable target collision"): + validate_source_file_plan(sources, { + "primary": ["pipeline.json", "Auxiliary/Encoder/model.safetensors"], + "encoder": ["config.json", "model.safetensors"], + }) + + def test_rejects_checks_excluded_from_the_download_plan(self) -> None: + sources = normalize_model_sources(valid_node()) or [] + with self.assertRaisesRegex(ValueError, "excluded from its download plan"): + validate_source_file_plan(sources, { + "primary": ["pipeline.json"], + "encoder": ["config.json"], + }) + + def test_requires_all_checks_and_rejects_symlinked_extension_ancestry(self) -> None: + sources = normalize_model_sources(valid_node()) or [] + with tempfile.TemporaryDirectory(prefix="modly-model-sources-") as tmp: + models = Path(tmp) / "models" + model_root = models / "pixal3d" / "generate" + encoder = model_root / "auxiliary" / "encoder" + encoder.mkdir(parents=True) + (model_root / "pipeline.json").write_text("{}", encoding="utf-8") + (encoder / "config.json").write_text("{}", encoding="utf-8") + self.assertFalse(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + (encoder / "model.safetensors").write_bytes(b"x") + self.assertTrue(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + + (encoder / "model.safetensors").write_bytes(b"") + self.assertFalse(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + (encoder / "model.safetensors").unlink() + (encoder / "model.safetensors").mkdir() + self.assertFalse(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + + for child in sorted((models / "pixal3d").rglob("*"), reverse=True): + child.unlink() if child.is_file() else child.rmdir() + (models / "pixal3d").rmdir() + outside = Path(tmp) / "outside" + (outside / "generate").mkdir(parents=True) + try: + os.symlink(outside, models / "pixal3d", target_is_directory=True) + except (NotImplementedError, OSError) as exc: + self.skipTest(f"Symlinks unavailable: {exc}") + with self.assertRaisesRegex(ValueError, "symlink"): + resolve_model_root(models, "pixal3d/generate") + self.assertFalse(model_sources_are_downloaded(models, "pixal3d/generate", sources)) + + 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)) + + def test_weight_variants_validate_and_report_the_missing_selection(self) -> None: + def variant(quant: str) -> dict: + return { + "id": quant, + "include_prefixes": [f"dit/model_{quant}.gguf"], + "checks": [f"dit/model_{quant}.gguf"], + } + + node = { + "hf_repo": "org/model-gguf", + "download_check": "pipeline.json", + "params_schema": [ + {"id": "quant", "type": "select", "options": [{"value": "Q4"}, {"value": "Q5"}]} + ], + "weight_variants": {"param": "quant", "options": [variant("Q4"), variant("Q5")]}, + } + variants = normalize_weight_variants(node) or {} + self.assertEqual(variants["default"], "Q4") + self.assertIsNone(normalize_weight_variants({"hf_repo": "org/model"})) + + def with_options(options: list, **extra) -> dict: + return {**node, "weight_variants": {"param": "quant", "options": options, **extra}} + + overlapping = {**variant("Q5"), "include_prefixes": ["dit/"]} + for broken, message in ( + ({**node, "hf_repo": ""}, "requires hf_repo"), + ({**node, "model_sources": []}, "cannot be combined"), + ({**node, "weight_groups": ["base"]}, "cannot be combined with weight_groups"), + ({**node, "download_check": "dit/model_Q4.gguf"}, "download_check"), + (with_options([variant("Q4"), overlapping]), "share files"), + (with_options([variant("Q4")], default="Q8"), "default"), + (with_options([{**variant("Q4"), "checks": ["x.gguf"]}]), "not covered"), + (with_options([{**variant("Q4"), "size_gb": True}]), "size_gb"), + (with_options([{**variant("Q4"), "vram_gb": 0}]), "vram_gb"), + (with_options([{**variant("Q4"), "vram_gb": "6"}]), "vram_gb"), + ({**node, "params_schema": []}, "must name a params_schema entry"), + ( + {**node, "params_schema": [{"id": "quant", "options": [{"value": "Q5"}]}]}, + "must offer every weight variant id", + ), + ): + with self.subTest(message=message), self.assertRaisesRegex(ValueError, message): + normalize_weight_variants(broken) + + # A param that declares no options is left to the node to interpret. + self.assertIsNotNone( + normalize_weight_variants({**node, "params_schema": [{"id": "quant", "type": "string"}]}) + ) + self.assertNotIn("vram_gb", variants["options"][0]) + with_vram = normalize_weight_variants(with_options([{**variant("Q4"), "vram_gb": 6.5}])) or {} + self.assertEqual(with_vram["options"][0]["vram_gb"], 6.5) + + with tempfile.TemporaryDirectory(prefix="modly-weight-variants-") as tmp: + models = Path(tmp) / "models" + dit = models / "trellis" / "generate" / "dit" + dit.mkdir(parents=True) + (dit / "model_Q5.gguf").write_bytes(b"q5") + (dit / "model_Q4.gguf").write_bytes(b"") + self.assertEqual(installed_weight_variants(models, "trellis/generate", variants), ["Q5"]) + missing = missing_weight_variant(models, "trellis/generate", variants, {}) + self.assertEqual((missing or {}).get("id"), "Q4") + for params in ({"quant": "Q5"}, {"quant": "fp16"}): + with self.subTest(params=params): + self.assertIsNone(missing_weight_variant(models, "trellis/generate", variants, params)) + self.assertIsNone(missing_weight_variant(models, "trellis/generate", None, {"quant": "Q4"})) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_optimize_mesh_ops.py b/api/tests/test_optimize_mesh_ops.py new file mode 100644 index 00000000..6c702ddb --- /dev/null +++ b/api/tests/test_optimize_mesh_ops.py @@ -0,0 +1,207 @@ +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +from fastapi import HTTPException + +from routers import optimize +from services.mesh_ops import ( + MeshOp, + MeshOpExecutionError, + MeshOpNotFoundError, + MeshOpResult, + MeshOpsRegistry, +) + + +class _FakeRegistry: + def __init__(self, output_path: Path) -> None: + self.output_path = output_path + self.calls = [] + + def describe(self): + return [{"id": "repair", "category": "repair", "params_schema": []}] + + def run(self, operation_id, input_path, params, context): + self.calls.append((operation_id, input_path, params, context)) + output_path = context.output_path or self.output_path + return MeshOpResult(output_path, {"face_count": 42}) + + +class OptimizeMeshOpsRouteTests(unittest.TestCase): + def test_generic_list_and_run_routes_use_the_shared_registry(self) -> None: + with tempfile.TemporaryDirectory() as directory: + workspace = Path(directory) + input_path = workspace / "input.glb" + output_path = workspace / "Workflows" / "output.glb" + input_path.touch() + registry = _FakeRegistry(output_path) + + with ( + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), + patch.object(optimize, "mesh_ops_registry", registry), + ): + descriptions = optimize.list_mesh_operations() + response = optimize.run_mesh_operation( + "repair", + optimize.MeshOpRequest( + path="input.glb", + params={"fill_holes": False}, + ), + ) + + self.assertEqual(descriptions[0]["id"], "repair") + self.assertEqual( + response, + { + "path": "Workflows/output.glb", + "url": "/workspace/Workflows/output.glb", + "face_count": 42, + }, + ) + self.assertEqual(registry.calls[0][0:3], ( + "repair", + input_path, + {"fill_holes": False}, + )) + + def test_legacy_routes_delegate_with_their_existing_clamps_and_names(self) -> None: + with tempfile.TemporaryDirectory() as directory: + workspace = Path(directory) + input_path = workspace / "model.glb" + fallback_output = workspace / "Workflows" / "unused.glb" + input_path.touch() + registry = _FakeRegistry(fallback_output) + + with ( + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), + patch.object(optimize, "mesh_ops_registry", registry), + ): + optimize_response = optimize.optimize_mesh( + optimize.OptimizeRequest(path="model.glb", target_faces=2) + ) + smooth_response = optimize.smooth_mesh( + optimize.SmoothRequest(path="model.glb", iterations=99) + ) + + optimize_call, smooth_call = registry.calls + self.assertEqual(optimize_call[0], "decimate") + self.assertEqual(optimize_call[2], {"target_faces": 100}) + self.assertEqual( + optimize_call[3].output_path, + workspace / "model_opt100.glb", + ) + self.assertEqual(smooth_call[0], "smooth") + self.assertEqual( + smooth_call[2], + {"iterations": 20, "lambda_": 0.5, "mode": "laplacian"}, + ) + self.assertEqual( + smooth_call[3].output_path, + workspace / "model_smooth20.glb", + ) + self.assertTrue(smooth_call[3].preserve_visuals) + self.assertEqual(optimize_response["face_count"], 42) + self.assertEqual(smooth_response["url"], "/workspace/model_smooth20.glb") + + def test_unknown_generic_operation_is_a_404(self) -> None: + class MissingRegistry: + def run(self, operation_id, input_path, params, context): + raise MeshOpNotFoundError(operation_id) + + with tempfile.TemporaryDirectory() as directory: + workspace = Path(directory) + input_path = workspace / "input.glb" + input_path.touch() + with ( + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), + patch.object(optimize, "mesh_ops_registry", MissingRegistry()), + self.assertRaises(HTTPException) as raised, + ): + optimize.run_mesh_operation( + "missing", + optimize.MeshOpRequest(path="input.glb"), + ) + + self.assertEqual(raised.exception.status_code, 404) + + def test_generic_run_route_enforces_registry_bounds(self) -> None: + calls = [] + + def operation(input_path, params, context): + calls.append(params) + return MeshOpResult(input_path) + + registry = MeshOpsRegistry( + [ + MeshOp( + id="bounded", + label="Bounded", + params_schema=( + { + "id": "amount", + "type": "int", + "default": 3, + "min": 1, + "max": 5, + }, + ), + fn=operation, + category="test", + ) + ] + ) + + with tempfile.TemporaryDirectory() as directory: + workspace = Path(directory) + input_path = workspace / "input.glb" + input_path.touch() + with ( + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), + patch.object(optimize, "mesh_ops_registry", registry), + ): + optimize.run_mesh_operation( + "bounded", + optimize.MeshOpRequest( + path="input.glb", + params={"amount": 99}, + ), + ) + + self.assertEqual(calls, [{"amount": 5}]) + + def test_missing_input_during_execution_is_a_404(self) -> None: + class MissingInputRegistry: + def run(self, operation_id, input_path, params, context): + raise FileNotFoundError(f"Input mesh not found: {input_path}") + + self._assert_operation_error(MissingInputRegistry(), 404) + + def test_backend_execution_failure_is_a_503(self) -> None: + class FailedBackendRegistry: + def run(self, operation_id, input_path, params, context): + raise MeshOpExecutionError("mesh-optimizer exited with code 1") + + self._assert_operation_error(FailedBackendRegistry(), 503) + + def _assert_operation_error(self, registry, expected_status: int) -> None: + with tempfile.TemporaryDirectory() as directory: + workspace = Path(directory) + input_path = workspace / "input.glb" + input_path.touch() + with ( + patch.object(optimize.registry, "WORKSPACE_DIR", workspace), + patch.object(optimize, "mesh_ops_registry", registry), + self.assertRaises(HTTPException) as raised, + ): + optimize.run_mesh_operation( + "decimate", + optimize.MeshOpRequest(path="input.glb"), + ) + + self.assertEqual(raised.exception.status_code, expected_status) + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_runner.py b/api/tests/test_runner.py index f11a9c22..084e3e7b 100644 --- a/api/tests/test_runner.py +++ b/api/tests/test_runner.py @@ -5,6 +5,7 @@ import json import tempfile import importlib +from unittest.mock import patch from contextlib import redirect_stdout from pathlib import Path @@ -20,6 +21,30 @@ class RunnerTests(unittest.TestCase): + def test_decode_typed_scene_revalidates_worker_workspace_and_keeps_legacy_image(self) -> None: + with tempfile.TemporaryDirectory() as tmp: + workspace = Path(tmp) / "workspace" + scene = workspace / "Workflows" / "room" + scene.mkdir(parents=True) + manifest = scene / "scene-manifest.json" + manifest.write_text(json.dumps({ + "schema": "modly.scene-manifest.v1", "sceneRoot": ".", "assets": [], + })) + with patch.object(runner, "WORKSPACE_DIR", workspace): + typed = runner.decode_model_input({"input": {"kind": "scene", "path": str(manifest)}}) + self.assertEqual(typed.kind, "scene") + self.assertEqual(typed.path, manifest.resolve()) + self.assertEqual(runner.decode_model_input({"image_b64": "aW1hZ2U="}), b"image") + with self.assertRaises(ValueError): + runner.decode_model_input({"input": {"kind": "video", "path": str(manifest)}}) + + def test_runner_model_envelope_rejects_cross_node_dispatch(self) -> None: + manifest = {"id": "pixal3d"} + node = {"id": "worldsculpt"} + runner.validate_requested_model({"model_id": "pixal3d/worldsculpt"}, manifest, node) + with self.assertRaisesRegex(ValueError, "does not match"): + runner.validate_requested_model({"model_id": "pixal3d/generate"}, manifest, node) + def test_select_node_uses_model_dir_override(self) -> None: manifest = { "nodes": [ @@ -32,6 +57,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 @@ -154,5 +190,226 @@ def test_send_writes_single_json_line(self) -> None: self.assertEqual(json.loads(written), {"type": "ready", "params_schema": []}) +_FAKE_TEXGEN_GENERATOR = ''' +from pathlib import Path + +INSTANCES = [] + + +class FakeTexGen: + """ + Mimics an extension whose texture pipeline is built lazily on first use and + frees the shape pipeline first to make room for it. + + The first texture setup fails (missing `xatlas`), which is what issue #239 + reports; a later attempt succeeds once the dependency is available. + """ + + def __init__(self, model_dir, outputs_dir): + self.model_dir = model_dir + self.outputs_dir = outputs_dir + self._model = None + self.load_calls = 0 + self.texgen_attempts = 0 + INSTANCES.append(self) + + def is_loaded(self): + return self._model is not None + + def load(self): + self.load_calls += 1 + self._model = lambda image: "mesh" + + def unload(self): + self._model = None + + def _setup_texgen(self): + # Free the shape pipeline before building the texture pipeline. + self._model = None + self.texgen_attempts += 1 + if self.texgen_attempts == 1: + raise RuntimeError("No module named 'xatlas'") + # Texture setup succeeded: shape pipeline comes back. + self.load() + + def generate(self, image_bytes, params, progress_cb=None, cancel_event=None): + # Stands in for `self._model(image)` on a generator whose model is gone. + if self._model is None: + raise TypeError("'NoneType' object is not callable") + self._model(image_bytes) + if params.get("enable_texture"): + self._setup_texgen() + return Path("out.glb") +''' + + +class _RunnerDriver: + """Runs runner.main() against a throwaway extension dir.""" + + def __init__(self, generator_src: str, generator_class: str) -> None: + self.ext_dir = Path(tempfile.mkdtemp(prefix="modly-texgen-test-")) + (self.ext_dir / "generator.py").write_text(generator_src, encoding="utf-8") + (self.ext_dir / "manifest.json").write_text( + json.dumps({"id": "demo-ext", "generator_class": generator_class}), + encoding="utf-8", + ) + + def run(self, actions: list) -> list: + """Feeds actions on stdin, returns the parsed messages runner emitted.""" + original_ext_dir = runner.EXT_DIR + original_stdin = sys.stdin + original_module = sys.modules.pop("generator", None) + runner.EXT_DIR = self.ext_dir + sys.stdin = io.StringIO("".join(json.dumps(a) + "\n" for a in actions)) + out = io.StringIO() + try: + with redirect_stdout(out): + runner.main() + self.generator_module = sys.modules["generator"] + finally: + runner.EXT_DIR = original_ext_dir + sys.stdin = original_stdin + sys.modules.pop("generator", None) + if original_module is not None: + sys.modules["generator"] = original_module + return [json.loads(line) for line in out.getvalue().splitlines() if line.strip()] + + +class 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})() + self.assertTrue(runner._generator_is_loaded(gen)) + + def test_treats_raising_is_loaded_as_not_loaded(self) -> None: + class Gen: + def is_loaded(self): + raise RuntimeError("half-initialised") + + self.assertFalse(runner._generator_is_loaded(Gen())) + + def test_ensure_model_loaded_is_a_noop_when_already_loaded(self) -> None: + class Gen: + load_calls = 0 + + def is_loaded(self): + return True + + def load(self): + self.load_calls += 1 + + gen = Gen() + self.assertFalse(runner._ensure_model_loaded(gen)) + self.assertEqual(gen.load_calls, 0) + + def test_ensure_model_loaded_reloads_when_model_is_gone(self) -> None: + class Gen: + def __init__(self): + self._model = None + self.load_calls = 0 + + def is_loaded(self): + return self._model is not None + + def load(self): + self.load_calls += 1 + self._model = object() + + gen = Gen() + self.assertTrue(runner._ensure_model_loaded(gen)) + self.assertEqual(gen.load_calls, 1) + self.assertIsNotNone(gen._model) + + +class TextureSetupRecoveryTests(unittest.TestCase): + """ + Regression tests for issue #239: a failed lazy texture setup left the worker + with _model = None and no recovery path, so every later generation raised + "TypeError: 'NoneType' object is not callable" until the process was killed. + """ + + def setUp(self) -> None: + self.driver = _RunnerDriver(_FAKE_TEXGEN_GENERATOR, "FakeTexGen") + self.texture_run = { + "action": "generate", + "image_b64": "", + "params": {"enable_texture": True}, + } + + def test_worker_recovers_after_failed_texture_setup(self) -> None: + messages = self.driver.run([ + {"action": "load"}, + dict(self.texture_run, id="run-1"), + dict(self.texture_run, id="run-2"), + ]) + + by_id = {m.get("id"): m for m in messages if m.get("type") in ("done", "error")} + + # First run fails during texture setup and surfaces the real cause. + self.assertEqual(by_id["run-1"]["type"], "error") + self.assertIn("xatlas", by_id["run-1"]["message"]) + + # The retry must succeed instead of dying on a None model. + self.assertEqual( + by_id["run-2"]["type"], "done", + msg=f"retry did not recover: {by_id['run-2']}", + ) + self.assertNotIn("NoneType", json.dumps(by_id["run-2"])) + + # …and the worker's model is genuinely back. + gen = self.driver.generator_module.INSTANCES[0] + self.assertIsNotNone(gen._model) + self.assertTrue(gen.is_loaded()) + + def test_failed_run_reports_that_the_model_was_lost(self) -> None: + messages = self.driver.run([ + {"action": "load"}, + dict(self.texture_run, id="run-1"), + ]) + + error = next(m for m in messages if m.get("type") == "error") + self.assertIs(error["loaded"], False) + + def test_reload_before_generate_is_logged(self) -> None: + messages = self.driver.run([ + {"action": "load"}, + dict(self.texture_run, id="run-1"), + dict(self.texture_run, id="run-2"), + ]) + + logs = [m for m in messages if m.get("type") == "log"] + self.assertTrue( + any("reloaded before generating" in m.get("message", "") for m in logs), + msg=f"expected a reload log, got {logs}", + ) + + def test_successful_run_does_not_reload_the_model(self) -> None: + messages = self.driver.run([ + {"action": "load"}, + {"action": "generate", "id": "run-1", "image_b64": "", "params": {}}, + ]) + + self.assertEqual( + [m["type"] for m in messages if m.get("id") == "run-1"], ["done"] + ) + self.assertEqual(self.driver.generator_module.INSTANCES[0].load_calls, 1) + + if __name__ == "__main__": unittest.main() diff --git a/api/tests/test_scene_generation.py b/api/tests/test_scene_generation.py new file mode 100644 index 00000000..90333ca5 --- /dev/null +++ b/api/tests/test_scene_generation.py @@ -0,0 +1,339 @@ +import asyncio +import json +import tempfile +import threading +import unittest +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from unittest.mock import patch + +from fastapi import BackgroundTasks, HTTPException +from pydantic import ValidationError + +import routers.generation as generation +import services.generator_registry as registry +from schemas.generation import GenerateFromArtifactRequest + + +class _Registry: + def __init__(self): + self.switched = False + def get_generator(self, model_id): return object() + def get_manifest(self, model_id): return {"input": "scene"} + def switch_model(self, model_id): self.switched = True + def assert_weight_variant_installed(self, params, model_id=None): pass + + +class SceneGenerationTests(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + self.workspace = Path(self.tmp.name) / "workspace" + self.scene = self.workspace / "Workflows" / "room" + self.scene.mkdir(parents=True) + self.manifest = self.scene / "scene-manifest.json" + self.manifest.write_text(json.dumps({"schema": "modly.scene-manifest.v1", "sceneRoot": ".", "assets": []})) + self.registry = _Registry() + self.patches = [patch.object(generation, "generator_registry", self.registry), patch.object(registry, "WORKSPACE_DIR", self.workspace)] + for item in self.patches: item.start() + + def tearDown(self): + for item in reversed(self.patches): item.stop() + generation._jobs.clear(); generation._cancel_events.clear(); generation._cancelled.clear(); generation._completed_at.clear(); generation._job_generators.clear() + self.tmp.cleanup() + + def test_generic_route_queues_typed_scene_and_strips_reserved_params(self): + tasks = BackgroundTasks() + result = asyncio.run(generation.generate_from_artifact(GenerateFromArtifactRequest( + input_kind="scene", input_path="Workflows/room", model_id="demo/scene", + params={"artifact_path": "/etc/passwd", "input_kind": "image", "quality": "high"}, + ), tasks)) + queued = tasks.tasks[0] + self.assertEqual(queued.args[1].kind, "scene") + self.assertEqual(queued.args[1].path, self.manifest.resolve()) + self.assertEqual(queued.args[2]["scene_manifest_path"], str(self.manifest.resolve())) + self.assertNotIn("artifact_path", queued.args[2]) + self.assertNotIn("input_kind", queued.args[2]) + self.assertEqual(queued.args[5], "demo/scene") + self.assertEqual(result["job_id"], queued.args[0]) + self.assertFalse(self.registry.switched) + + def test_generic_route_rejects_unsupported_kind_and_model_mismatch(self): + for kind in ("video", "capture", "image"): + with self.subTest(kind=kind), self.assertRaises(ValidationError): + GenerateFromArtifactRequest( + input_kind=kind, input_path="Workflows/room", model_id="demo/scene") + self.registry.get_manifest = lambda _model_id: {"input": "image"} + with self.assertRaises(HTTPException) as caught: + asyncio.run(generation.generate_from_artifact(GenerateFromArtifactRequest( + input_kind="scene", input_path="Workflows/room", model_id="demo/image"), BackgroundTasks())) + self.assertEqual(caught.exception.status_code, 400) + + def test_rejects_traversal_before_switch_or_queue(self): + with self.assertRaises(HTTPException): + asyncio.run(generation.generate_from_artifact(GenerateFromArtifactRequest( + input_kind="scene", input_path="../outside", model_id="demo/scene"), BackgroundTasks())) + self.assertFalse(self.registry.switched) + + def test_queued_scene_job_is_pinned_to_requested_model(self): + calls = [] + + class Generator: + outputs_dir = None + def is_loaded(self): return True + def generate_artifact(self, kind, path, params, progress_cb, cancel_event=None): + calls.append(("model-a", kind, path)) + output = Path(self.outputs_dir) / "result.glb" + output.write_bytes(b"glb") + return output + + generator = Generator() + registry_stub = type("Registry", (), { + "assert_weight_variant_installed": lambda self, params, model_id=None: None, + "model_status": lambda self, model_id: {"name": model_id, "downloaded": True, "loaded": True}, + "get_generator": lambda self, model_id: generator if model_id == "demo/a" else (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), + "activate_ready_generator": lambda self, model_id: generator if model_id == "demo/a" else (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), + "get_active": lambda self: (_ for _ in ()).throw(AssertionError("mutable active model must not be used")), + })() + job_id = "pinned-scene" + generation._jobs[job_id] = generation.JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = __import__("threading").Event() + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(generation._run_generation( + job_id, generation.TypedArtifactInput("scene", self.manifest.resolve()), {}, + "Workflows", "mesh", "demo/a", + )) + self.assertEqual(calls[0][0], "model-a") + self.assertEqual(generation._jobs[job_id].status, "done") + + def test_missing_pinned_model_fails_actionably(self): + registry_stub = type("Registry", (), { + "assert_weight_variant_installed": lambda self, params, model_id=None: None, + "model_status": lambda self, model_id: {"name": model_id, "downloaded": True, "loaded": False}, + "get_generator": lambda self, model_id: object(), + "activate_ready_generator": lambda self, model_id: (_ for _ in ()).throw(ValueError(f"Unknown model ID: {model_id}")), + "get_active": lambda self: (_ for _ in ()).throw(AssertionError("must not use active model")), + })() + job_id = "missing-scene" + generation._jobs[job_id] = generation.JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = __import__("threading").Event() + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(generation._run_generation( + job_id, generation.TypedArtifactInput("scene", self.manifest.resolve()), {}, + "Workflows", "mesh", "demo/missing", + )) + self.assertEqual(generation._jobs[job_id].status, "error") + self.assertIn("Unknown model ID: demo/missing", generation._jobs[job_id].error) + + def test_interleaved_pinned_jobs_serialize_model_lifecycle_and_cancel_exact_job(self): + class Proc: + def __init__(self, owner): + self.owner = owner + self.alive = True + def poll(self): + return None if self.alive else 0 + def kill(self): + self.alive = False + self.owner.killed = True + + class Generator: + def __init__(self, model_id): + self.model_id = model_id + self.outputs_dir = None + self.loaded = False + self.started = threading.Event() + self.killed = False + self._loaded = True + self._proc = Proc(self) + def is_loaded(self): return self.loaded + def is_downloaded(self): return True + def load(self): self.loaded = True + def unload(self): self.loaded = False + def generate_artifact(self, kind, path, params, progress_cb, cancel_event=None): + self.started.set() + if self.model_id == "demo/a": + if cancel_event is None or not cancel_event.wait(2): + raise AssertionError("first job was not cancelled") + raise generation.GenerationCancelled() + output = Path(self.outputs_dir) / "result.glb" + output.write_bytes(b"glb") + return output + + class Registry: + def __init__(self): + self.generators = {model_id: Generator(model_id) for model_id in ("demo/a", "demo/b")} + self.active_id = "demo/a" + def assert_weight_variant_installed(self, params, model_id=None): pass + def get_generator(self, model_id): return self.generators[model_id] + def model_status(self, model_id): + gen = self.generators[model_id] + return {"name": model_id, "downloaded": True, "loaded": gen.loaded} + def switch_model(self, model_id): + if model_id != self.active_id: + self.generators[self.active_id].unload() + self.active_id = model_id + def activate_ready_generator(self, model_id): + self.switch_model(model_id) + gen = self.generators[model_id] + gen.load() + if sum(candidate.loaded for candidate in self.generators.values()) != 1: + raise AssertionError("more than one generator is resident") + return gen + def get_active(self): raise AssertionError("pinned jobs must not consult mutable active state") + + registry_stub = Registry() + # Reproduce the request interleaving: A switches first, then B switches + # before either background task has started. + registry_stub.switch_model("demo/a") + registry_stub.switch_model("demo/b") + for job_id in ("job-a", "job-b"): + generation._jobs[job_id] = generation.JobStatus(job_id=job_id, status="pending", progress=0) + generation._cancel_events[job_id] = threading.Event() + + async def run_both(): + first = asyncio.create_task(generation._run_generation( + "job-a", generation.TypedArtifactInput("scene", self.manifest.resolve()), {}, + "Workflows", "mesh", "demo/a", + )) + second = asyncio.create_task(generation._run_generation( + "job-b", generation.TypedArtifactInput("scene", self.manifest.resolve()), {}, + "Workflows", "mesh", "demo/b", + )) + started = await asyncio.to_thread(registry_stub.generators["demo/a"].started.wait, 1) + self.assertTrue(started) + await asyncio.sleep(0.05) + self.assertFalse(registry_stub.generators["demo/b"].started.is_set()) + await generation.cancel_job("job-a") + await asyncio.gather(first, second) + + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(run_both()) + + self.assertTrue(registry_stub.generators["demo/a"].killed) + self.assertFalse(registry_stub.generators["demo/b"].killed) + self.assertEqual(generation._jobs["job-a"].status, "cancelled") + self.assertEqual(generation._jobs["job-b"].status, "done") + self.assertEqual(registry_stub.active_id, "demo/b") + self.assertFalse(registry_stub.generators["demo/a"].loaded) + self.assertTrue(registry_stub.generators["demo/b"].loaded) + + def test_cancelling_queued_generation_does_not_wedge_later_jobs(self): + release_first = threading.Event() + first_started = threading.Event() + + class Generator: + outputs_dir = None + + def __init__(self): + self.calls = 0 + + def generate(self, image_bytes, params, progress_cb, cancel_event=None): + self.calls += 1 + call_number = self.calls + if call_number == 1: + first_started.set() + if not release_first.wait(2): + raise AssertionError("first generation was not released") + output = Path(self.outputs_dir) / f"result-{call_number}.glb" + output.write_bytes(b"glb") + return output + + generator = Generator() + registry_stub = type("Registry", (), { + "assert_weight_variant_installed": lambda self, params, model_id=None: None, + "get_generator": lambda self, model_id: generator, + "model_status": lambda self, model_id: { + "name": model_id, "downloaded": True, "loaded": True, + }, + "activate_ready_generator": lambda self, model_id: generator, + "get_active": lambda self: (_ for _ in ()).throw( + AssertionError("pinned jobs must not use mutable active state") + ), + })() + for job_id in ("blocking", "cancelled-waiter", "after-cancel"): + generation._jobs[job_id] = generation.JobStatus( + job_id=job_id, status="pending", progress=0, + ) + generation._cancel_events[job_id] = threading.Event() + + async def run_scenario(): + first = asyncio.create_task(generation._run_generation( + "blocking", b"image", {}, "Workflows", "mesh", "demo/a", + )) + started = await asyncio.to_thread(first_started.wait, 1) + self.assertTrue(started) + + queued = asyncio.create_task(generation._run_generation( + "cancelled-waiter", b"image", {}, "Workflows", "mesh", "demo/b", + )) + await asyncio.sleep(0) + queued.cancel() + with self.assertRaises(asyncio.CancelledError): + await queued + + release_first.set() + await asyncio.wait_for(first, 1) + await asyncio.wait_for(generation._run_generation( + "after-cancel", b"image", {}, "Workflows", "mesh", "demo/c", + ), 1) + + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(run_scenario()) + + self.assertEqual(generation._jobs["cancelled-waiter"].status, "cancelled") + self.assertEqual(generation._jobs["after-cancel"].status, "done") + self.assertEqual(generator.calls, 2) + + def test_pinned_jobs_do_not_starve_small_default_executor(self): + class Generator: + outputs_dir = None + + def __init__(self): + self.calls = 0 + + def generate(self, image_bytes, params, progress_cb, cancel_event=None): + self.calls += 1 + output = Path(self.outputs_dir) / f"result-{self.calls}.glb" + output.write_bytes(b"glb") + return output + + generator = Generator() + registry_stub = type("Registry", (), { + "assert_weight_variant_installed": lambda self, params, model_id=None: None, + "get_generator": lambda self, model_id: generator, + "model_status": lambda self, model_id: { + "name": model_id, "downloaded": True, "loaded": True, + }, + "activate_ready_generator": lambda self, model_id: generator, + "get_active": lambda self: (_ for _ in ()).throw( + AssertionError("pinned jobs must not use mutable active state") + ), + })() + job_ids = ("small-pool-a", "small-pool-b", "small-pool-c") + for job_id in job_ids: + generation._jobs[job_id] = generation.JobStatus( + job_id=job_id, status="pending", progress=0, + ) + generation._cancel_events[job_id] = threading.Event() + + default_executor = ThreadPoolExecutor(max_workers=2) + + async def run_scenario(): + asyncio.get_running_loop().set_default_executor(default_executor) + await asyncio.wait_for(asyncio.gather(*( + generation._run_generation( + job_id, b"image", {}, "Workflows", "mesh", f"demo/{job_id}", + ) + for job_id in job_ids + )), 2) + + try: + with patch.object(generation, "generator_registry", registry_stub): + asyncio.run(run_scenario()) + finally: + default_executor.shutdown(wait=True) + + self.assertEqual(generator.calls, 3) + self.assertEqual( + [generation._jobs[job_id].status for job_id in job_ids], + ["done", "done", "done"], + ) diff --git a/api/tests/test_scene_input.py b/api/tests/test_scene_input.py new file mode 100644 index 00000000..aec4bc5e --- /dev/null +++ b/api/tests/test_scene_input.py @@ -0,0 +1,93 @@ +import json +import sys +import tempfile +import unittest +from pathlib import Path + +from services.scene_input import validate_scene_input, revalidate_scene_manifest + + +class SceneInputTests(unittest.TestCase): + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + self.workspace = Path(self.tmp.name) / "workspace" + self.scene = self.workspace / "Workflows" / "room" + self.scene.mkdir(parents=True) + (self.scene / "model.glb").write_bytes(b"mesh") + self.manifest = self.scene / "scene-manifest.json" + self.manifest.write_text(json.dumps({ + "schema": "modly.scene-manifest.v1", + "sceneRoot": ".", + "assets": [{"path": "model.glb"}], + })) + + def tearDown(self): + self.tmp.cleanup() + + def test_directory_and_manifest_are_canonical_paths(self): + expected = self.manifest.resolve() + self.assertEqual(validate_scene_input(self.workspace, "Workflows/room"), expected) + self.assertEqual(validate_scene_input(self.workspace, "Workflows/room/scene-manifest.json"), expected) + self.assertEqual(revalidate_scene_manifest(self.workspace, expected), expected) + + def test_rejects_traversal_absolute_encoded_and_non_manifest_json(self): + for path in ("../outside", "/etc/passwd", "C:/outside", "Workflows/room/../room", "Workflows/room/other.json", "Workflows/%2e%2e"): + with self.subTest(path=path), self.assertRaises(ValueError): + validate_scene_input(self.workspace, path) + + def test_rejects_symlink_escape_when_supported(self): + outside = Path(self.tmp.name) / "outside" + outside.mkdir() + (outside / "scene-manifest.json").write_text(self.manifest.read_text()) + try: + (self.workspace / "Workflows" / "link").symlink_to(outside, target_is_directory=True) + except OSError as exc: + self.skipTest(f"directory symlinks are unavailable on this host: {exc}") + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/link") + + @unittest.skipUnless(sys.platform == "win32", "directory junctions are Windows-only") + def test_rejects_junction_escape_on_windows(self): + # Junctions need no privilege (unlike symlinks), so they are the + # realistic escape on Windows and keep the reparse-point check covered. + import _winapi + + outside = Path(self.tmp.name) / "outside" + outside.mkdir() + (outside / "scene-manifest.json").write_text(self.manifest.read_text()) + (outside / "model.glb").write_bytes(b"mesh") + _winapi.CreateJunction(str(outside), str(self.workspace / "Workflows" / "link")) + + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/link") + + data = json.loads(self.manifest.read_text()) + data["assets"] = [{"workspacePath": "Workflows/link/model.glb"}] + self.manifest.write_text(json.dumps(data)) + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/room") + + def test_rejects_missing_assets_and_oversized_manifest(self): + data = json.loads(self.manifest.read_text()) + data["assets"] = [{"path": "missing.glb"}] + self.manifest.write_text(json.dumps(data)) + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/room") + + self.manifest.write_text(" " * (1024 * 1024 + 1)) + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/room") + + def test_rejects_malformed_manifest_and_asset_paths(self): + original = json.loads(self.manifest.read_text()) + for patch in ( + {"schema": "wrong"}, + {"assets": "bad"}, + {"assets": [{"path": "../escape.glb"}]}, + {"preview": {"image": None}}, + {"initialView": {"position": [0, 0, 0], "target": [0, 0, 0]}}, + ): + with self.subTest(patch=patch): + self.manifest.write_text(json.dumps({**original, **patch})) + with self.assertRaises(ValueError): + validate_scene_input(self.workspace, "Workflows/room") diff --git a/api/tests/test_workflow_runs_lifecycle.py b/api/tests/test_workflow_runs_lifecycle.py new file mode 100644 index 00000000..47433573 --- /dev/null +++ b/api/tests/test_workflow_runs_lifecycle.py @@ -0,0 +1,204 @@ +import asyncio +import tempfile +import threading +import time +import unittest +from pathlib import Path + +from fastapi import BackgroundTasks + +import routers.generation as generation +import routers.workflow_runs as workflow_runs +import services.generator_registry as registry_module +from schemas.generation import JobStatus + + +class _FakeUpload: + """Minimal UploadFile stand-in: an image content-type and readable bytes.""" + + def __init__(self, content_type: str = "image/png", data: bytes = b"\x89PNG\r\n") -> None: + self.content_type = content_type + self._data = data + + async def read(self) -> bytes: + return self._data + + +class _FakeRegistry: + """Accepts any model id and exposes the attrs cancel_run pokes at.""" + + _generators: dict = {} + _active_id = None + + def get_generator(self, model_id: str) -> object: + return object() + + def get_manifest(self, model_id: str) -> dict: + return {"output": "mesh"} + + def switch_model(self, model_id: str) -> None: + pass + + +def _clear_job_stores() -> None: + for store in ( + generation._jobs, + generation._cancel_events, + generation._cancelled, + generation._completed_at, + generation._job_generators, + ): + store.clear() + + +class WorkflowRunJobLifecycleTests(unittest.TestCase): + """The headless /workflow-runs surface shares the job dicts with /generate, + so it must take part in the same TTL purge — otherwise long-running + automation leaks a JobStatus + Event per run forever.""" + + def setUp(self) -> None: + self._prev = workflow_runs.generator_registry + workflow_runs.generator_registry = _FakeRegistry() + _clear_job_stores() + + def tearDown(self) -> None: + workflow_runs.generator_registry = self._prev + _clear_job_stores() + + def test_create_run_purges_terminal_jobs_past_ttl(self) -> None: + stale = "stale-run" + generation._jobs[stale] = JobStatus(job_id=stale, status="done", progress=100) + generation._cancel_events[stale] = threading.Event() + generation._completed_at[stale] = time.monotonic() - generation._JOB_TTL - 1 + + background = BackgroundTasks() + asyncio.run( + workflow_runs.create_run_from_image( + background, + image=_FakeUpload(), + model_id="sf3d", + collection="Default", + params="{}", + ) + ) + + # Before the fix create_run_from_image never purged, so the stale job lingered. + self.assertNotIn(stale, generation._jobs) + self.assertNotIn(stale, generation._completed_at) + self.assertNotIn(stale, generation._cancel_events) + + def test_cancel_run_records_completion_so_it_can_be_purged(self) -> None: + run_id = "run-1" + generation._jobs[run_id] = JobStatus(job_id=run_id, status="running", progress=10) + generation._cancel_events[run_id] = threading.Event() + + asyncio.run(workflow_runs.cancel_run(run_id)) + + self.assertEqual(generation._jobs[run_id].status, "cancelled") + # Without a _completed_at stamp the purge sweep can never evict a cancelled run. + self.assertIn(run_id, generation._completed_at) + + def test_workflow_and_generate_routes_share_exact_model_queue_and_cancellation(self) -> None: + class Proc: + def __init__(self, owner): + self.owner = owner + self.alive = True + def poll(self): return None if self.alive else 0 + def kill(self): + self.alive = False + self.owner.killed = True + + class Generator: + def __init__(self, model_id): + self.model_id = model_id + self.outputs_dir = None + self.loaded = False + self.started = threading.Event() + self.killed = False + self._loaded = True + self._proc = Proc(self) + def is_loaded(self): return self.loaded + def load(self): self.loaded = True + def unload(self): self.loaded = False + def generate(self, image_bytes, params, progress_cb, cancel_event=None): + self.started.set() + if self.model_id == "demo/workflow": + if cancel_event is None or not cancel_event.wait(2): + raise AssertionError("workflow job was not cancelled") + raise generation.GenerationCancelled() + output = Path(self.outputs_dir) / "result.glb" + output.write_bytes(b"glb") + return output + + class Registry: + def __init__(self): + ids = ("demo/workflow", "demo/generate") + self.generators = {model_id: Generator(model_id) for model_id in ids} + self.active_id = "demo/generate" + def assert_weight_variant_installed(self, params, model_id=None): pass + def get_generator(self, model_id): return self.generators[model_id] + def get_manifest(self, model_id): return {"output": "mesh"} + def model_status(self, model_id): + gen = self.generators[model_id] + return {"name": model_id, "downloaded": True, "loaded": gen.loaded} + def activate_ready_generator(self, model_id): + if model_id != self.active_id: + self.generators[self.active_id].unload() + self.active_id = model_id + gen = self.generators[model_id] + gen.load() + if sum(candidate.loaded for candidate in self.generators.values()) != 1: + raise AssertionError("more than one generator is resident") + return gen + def get_active(self): raise AssertionError("queued routes must use their pinned model id") + + fake_registry = Registry() + previous_generation_registry = generation.generator_registry + previous_workflow_registry = workflow_runs.generator_registry + previous_workspace = registry_module.WORKSPACE_DIR + with tempfile.TemporaryDirectory() as tmp: + registry_module.WORKSPACE_DIR = Path(tmp) + generation.generator_registry = fake_registry + workflow_runs.generator_registry = fake_registry + + async def run_cross_route_jobs(): + workflow_background = BackgroundTasks() + workflow_response = await workflow_runs.create_run_from_image( + workflow_background, _FakeUpload(), "demo/workflow", "Workflows", "{}", + ) + generate_background = BackgroundTasks() + generate_response = await generation.generate_from_image( + generate_background, _FakeUpload(), "demo/generate", "Workflows", + "quad", False, 1024, "{}", + ) + self.assertEqual(workflow_background.tasks[0].args[5], "demo/workflow") + self.assertEqual(generate_background.tasks[0].args[5], "demo/generate") + + first = asyncio.create_task(workflow_background.tasks[0]()) + second = asyncio.create_task(generate_background.tasks[0]()) + started = await asyncio.to_thread( + fake_registry.generators["demo/workflow"].started.wait, 1, + ) + self.assertTrue(started) + await asyncio.sleep(0.05) + self.assertFalse(fake_registry.generators["demo/generate"].started.is_set()) + await workflow_runs.cancel_run(workflow_response["run_id"]) + await asyncio.gather(first, second) + return workflow_response["run_id"], generate_response["job_id"] + + try: + workflow_job_id, generate_job_id = asyncio.run(run_cross_route_jobs()) + finally: + generation.generator_registry = previous_generation_registry + workflow_runs.generator_registry = previous_workflow_registry + registry_module.WORKSPACE_DIR = previous_workspace + + self.assertTrue(fake_registry.generators["demo/workflow"].killed) + self.assertFalse(fake_registry.generators["demo/generate"].killed) + self.assertEqual(generation._jobs[workflow_job_id].status, "cancelled") + self.assertEqual(generation._jobs[generate_job_id].status, "done") + self.assertEqual(fake_registry.active_id, "demo/generate") + + +if __name__ == "__main__": + unittest.main() diff --git a/api/tests/test_workflow_runs_router.py b/api/tests/test_workflow_runs_router.py index f045bef8..7537a484 100644 --- a/api/tests/test_workflow_runs_router.py +++ b/api/tests/test_workflow_runs_router.py @@ -24,6 +24,9 @@ class _FakeRegistry: def get_generator(self, model_id: str) -> object: return object() + def get_manifest(self, model_id: str) -> dict: + return {"output": "mesh"} + def switch_model(self, model_id: str) -> None: pass @@ -79,7 +82,7 @@ def _collection_forwarded(self, collection: str) -> str: params="{}", ) ) - # add_task(_run_generation, job_id, image_bytes, full_params, collection) + # add_task(_run_generation, ..., collection, output_kind, model_id) return background.tasks[0].args[3] def test_collection_is_forwarded_to_the_run(self) -> None: diff --git a/arch/decisions/AMD-ROCM-SUPPORT.md b/arch/decisions/AMD-ROCM-SUPPORT.md new file mode 100644 index 00000000..09671fc7 --- /dev/null +++ b/arch/decisions/AMD-ROCM-SUPPORT.md @@ -0,0 +1,100 @@ +# AMD-ROCM-SUPPORT + +- Status: proposed +- Date: 2026-08-17 + +## Decision + +Modly supports AMD Radeon GPUs through ROCm on Linux and Windows. Detection is +automatic and requires no ROCm installation on the user's machine. + +Scope and operating rules: + +- Detection is centralised in `electron/main/gpu-detect.ts` and produces, for AMD + machines, a compute target plus the pip index and requirements an extension's + torch install must end up using. +- NVIDIA keeps detection priority. On a machine with both vendors the existing + CUDA behaviour is unchanged. +- AMD machines report `gpu_sm = 0` and `cuda_version = 0`, never a synthesised + compute capability. +- Extensions are told the flavour via a `torch_flavor` setup argument. Extensions + that ignore it are corrected by a rewrite shim in `electron/main/setup-launcher.ts`. +- The ROCm wheel source differs by platform: `download.pytorch.org/whl/rocm7.2` + on Linux, `repo.amd.com/rocm/whl-multi-arch/` on Windows. +- Every automatic choice has an environment-variable override. See + `docs/running-on-amd-rocm.md`. + +## Context + +Modly never installs PyTorch itself. Each extension ships a `setup.py` that +creates its own venv and installs torch from an index it hardcodes — and those +scripts are third-party code in separate GitHub repositories that Modly cannot +edit. Before this work `detectGpuInfo()` only probed `nvidia-smi`, so an AMD +machine was reported as `accelerator: 'cpu'` with `gpu_sm: 0`, which sent every +extension down its legacy CUDA 11.8 branch and installed a torch that cannot see +the GPU at all. + +That leaves two distinct problems, and both have to be solved: + +- Extensions that *do* understand AMD were never told. The official + `modly-hunyuan3d-mini-extension` has accepted a `torch_flavor: "rocm"` argument + for some time; Modly simply never sent it. +- Extensions that don't understand AMD — `triposg`, `trellis2`, and the rest — + hardcode `--index-url .../whl/cu124` and have no branch to select. Passing an + argument achieves nothing for them. + +The wheel sources are also not symmetric across platforms. `download.pytorch.org` +publishes no ROCm wheels for Windows at all; AMD's own multi-arch index does, but +there the compute target is selected by a pip extra (`torch[device-gfx1200]`) +rather than by the index URL, which means Windows needs the compute target +*before* the install, not after. + +## Consequences + +- **A rewrite shim is unavoidable.** Correcting extensions we cannot edit means + intercepting their pip calls. The launcher already patched `subprocess` for two + other compatibility fixes, so ROCm redirection joins those rather than + introducing a new mechanism. +- **The decision is made in TypeScript, applied in Python.** The launcher is an + inline Python string that cannot be unit-tested in isolation, so index and + requirement resolution lives in `gpu-detect.ts` and reaches the launcher as + environment variables. The launcher is separately exercised end-to-end by + `setup-launcher.test.mjs`, which runs it against the command shapes the + official extensions actually use. +- **`gpu_sm = 0` is load-bearing, not a placeholder.** Extensions written before + `torch_flavor` branch on that number, and 0 selects their most conservative + path. It also keeps them off `rembg[gpu]`, whose `onnxruntime-gpu` is + CUDA-only. Reporting a synthesised capability instead would break both. + PyTorch's HIP build answers the whole `torch.cuda` API, so + `get_device_capability()` reports `(12, 0)` for a gfx1200 Radeon — + indistinguishable from an sm_120 Blackwell. `api/routers/extensions.py` asks + `nvidia-smi` rather than torch for this reason (and because Modly's main venv + carries no torch at all). +- **Compute-target discovery is platform-specific.** Linux reads + `gfx_target_version` from the kernel's KFD topology, which needs no ROCm + install and no external binary. Windows has no equivalent, so it maps PCI + device ids from `Win32_VideoController` through a table keyed by silicon. That + table is a maintenance surface: new AMD silicon needs an entry, and an unmapped + AMD card falls back to CPU with an actionable message rather than guessing a + wheel. +- **The Linux and Windows torch versions diverge.** Linux gets unpinned wheels + from the pytorch.org ROCm index (currently torch 2.11+); Windows gets a pinned + pair from AMD's index. Extension code written against torch 2.6/2.7 may not + survive that jump, which is why `MODLY_ROCM_INDEX` and `MODLY_ROCM_TORCH_SPEC` + exist as first-class escape hatches rather than debug affordances. +- **Linux is verified, Windows is not.** On a Radeon RX 9060 XT (gfx1200), + `torch 2.13.0+rocm7.2` loads, rocBLAS and MIOpen kernels execute, 14 GB of the + card's 16 GB allocates and reads back cleanly ([ROCm #6295](https://github.com/ROCm/ROCm/issues/6295), + which reports this card capped near 8 GB, did not reproduce), and a full + image-to-3D generation completes through the normal `ExtensionProcess` path in + 221 s. For Windows the wheel URLs, `cp311` availability and index layout were + checked, but no end-to-end run has been performed. +- **This work also required fixing an unrelated AppImage bug** to be verifiable + at all. `ensureStableEmbeddedPython()` copied the bundled runtime with + `fs.cp`, which rewrites relative symlinks into absolute paths pointing back at + the ephemeral `/tmp/.mount_Modly-XXXXXX/` mount — so the "stable" copy was not + stable, and every extension venv built from it died on the next launch with a + misleading `No module named 'PIL'`. See `electron/main/copy-runtime.ts`. +- **Texture generation is out of scope.** `api/texture_baker` already carries a + HIP build path, but those native extensions are not built as part of standard + extension setup, so texture generation is not covered by this ADR. diff --git a/arch/decisions/README.md b/arch/decisions/README.md index 063df577..925441ea 100644 --- a/arch/decisions/README.md +++ b/arch/decisions/README.md @@ -8,3 +8,4 @@ single reviewable document. Current ADRs: - [APPLE-SILICON-SUPPORT](./APPLE-SILICON-SUPPORT.md) +- [AMD-ROCM-SUPPORT](./AMD-ROCM-SUPPORT.md) diff --git a/docs/running-on-amd-rocm.md b/docs/running-on-amd-rocm.md new file mode 100644 index 00000000..5fe60f8a --- /dev/null +++ b/docs/running-on-amd-rocm.md @@ -0,0 +1,193 @@ +# Running Modly on an AMD GPU (ROCm) + +Modly's default GPU path is NVIDIA/CUDA (plus Metal/MPS on Apple Silicon). This page +covers the AMD path: how Modly detects a Radeon card, which PyTorch wheels it steers +extensions to, and what to do when the automatic choice is wrong. + +Nothing here needs a manual setup: install the app, install an extension, and the AMD +path is taken automatically when an AMD GPU is present. + +--- + +## 1. Requirements + +| | Linux | Windows | +|---|---|---| +| GPU | RDNA 2 or newer discrete Radeon (see the table below) | same | +| Driver | in-tree `amdgpu` kernel driver — any recent distro kernel | Adrenalin with the ROCm runtime (26.2.2 or newer) | +| ROCm install | **not required** — the PyTorch ROCm wheels bundle their own runtime | not required | +| Device access | `/dev/kfd` and `/dev/dri/renderD*` must be readable by your user | n/a | + +On most distributions `/dev/kfd` is world-accessible. If it is not, add yourself to the +`render` group and log back in: + +```bash +ls -l /dev/kfd # crw-rw-rw- → nothing to do +sudo usermod -aG render "$USER" +``` + +Integrated (APU) graphics are not mapped by the Windows detection table. They work on +Linux whenever the kernel publishes a compute target, but they are not a target Modly +tests against. + +--- + +## 2. How detection works + +Detection lives in `electron/main/gpu-detect.ts` and runs before an extension's +`setup.py`, in this order: + +1. `MODLY_TORCH_FLAVOR` (`cuda` / `rocm` / `cpu`) — an explicit override wins over everything. +2. Apple Silicon → MPS. +3. `nvidia-smi` → CUDA. **NVIDIA keeps priority**: on a machine with both vendors, + nothing about the existing CUDA behaviour changes. +4. AMD: + - **Linux** — reads `gfx_target_version` from the kernel's KFD topology + (`/sys/class/kfd/kfd/topology/nodes/*/properties`). No ROCm install and no external + binary involved; the `amdgpu` driver publishes this on its own. + - **Windows** — reads the PCI device id from `Win32_VideoController` via PowerShell + and maps it to a compute target. +5. Otherwise → CPU. + +You can check what was detected in the logs (Settings → Logs), on the line beginning +`[ext-setup] accelerator=`: + +``` +[ext-setup] accelerator=rocm gfx=gfx1200 torch_index=https://download.pytorch.org/whl/rocm7.2 +``` + +### Supported compute targets + +| Silicon | Compute target | Cards | +|---|---|---| +| Navi 44 | `gfx1200` | RX 9060 XT | +| Navi 48 | `gfx1201` | RX 9070, RX 9070 XT, RX 9070 GRE, AI PRO R9700 | +| Navi 31 | `gfx1100` | RX 7900 XT/XTX/GRE, PRO W7800/W7900 | +| Navi 32 | `gfx1101` | RX 7700 XT, RX 7800 XT, PRO W7700 | +| Navi 33 | `gfx1102` | RX 7600 series, PRO W7500/W7600 | +| Navi 21 | `gfx1030` | RX 6800/6800 XT/6900 XT/6950 XT, PRO W6800 | +| Navi 22/23/24 | `gfx1031` / `gfx1032` / `gfx1034` | RX 6700 / 6600 / 6400 series | + +On Linux the target is read from the kernel, so any AMD GPU the kernel knows about is +picked up — the table above only bounds the *Windows* mapping. AMD officially supports +`gfx1030`, `gfx110x` and `gfx120x`; the RDNA 2 mid-range entries work in practice but +are not part of AMD's supported matrix. + +--- + +## 3. Which wheels get installed + +Modly does not install PyTorch itself — each extension's `setup.py` does, from an index +it hardcodes. On an AMD machine Modly corrects that choice two ways: + +- It passes `torch_flavor: "rocm"` in the setup arguments. Extensions that know about + AMD (the official `hunyuan3d-mini` one does) branch on it themselves. +- For every extension that doesn't, a compatibility shim in the setup launcher rewrites + the pip command: the CUDA (or CPU) index is swapped for the ROCm one, and the pinned + `torch==…` requirements are replaced. Non-torch installs are left untouched. + +| Platform | Index | Requirements | +|---|---|---| +| Linux | `https://download.pytorch.org/whl/rocm7.2` | `torch`, `torchvision`, unpinned | +| Windows | `https://repo.amd.com/rocm/whl-multi-arch/` | `torch[device-gfxNNNN]==2.11.0+rocm7.14.0` and matching `torchvision` | + +The two differ because `download.pytorch.org` publishes no ROCm wheels for Windows at +all. AMD's multi-arch index does, and there the compute target is selected by a pip +extra rather than by the index URL. + +`sm` and `cuda_version` are deliberately reported as `0` for AMD. Extensions written +before `torch_flavor` existed branch on those numbers, and `0` sends them down their +most conservative path — which also keeps them off `rembg[gpu]`, whose `onnxruntime-gpu` +is CUDA-only. + +--- + +## 4. Verifying an install + +After installing an extension (or running **Repair** on it from the Models page), check +what actually landed in its venv: + +```bash +# Linux; adjust to your extensions directory +EXT=~/Documents/Modly/extensions/hunyuan3d-mini +"$EXT/venv/bin/python" -c "import torch; print(torch.__version__, '| hip', torch.version.hip, '| avail', torch.cuda.is_available(), '|', torch.cuda.get_device_name(0))" +``` + +Expected: a `+rocm…` version, a non-null `hip`, `True`, and your card's name. A version +ending in `+cu118` or `+cu124` means the AMD path was not taken — check the detection +line in the logs. + +Then confirm real VRAM is usable, not just that the device opens: + +```bash +"$EXT/venv/bin/python" -c "import torch; x = torch.empty(5_000_000_000, dtype=torch.float16, device='cuda'); torch.cuda.synchronize(); print('10 GB allocated OK')" +``` + +--- + +## 5. Escape hatches + +All of these are environment variables read at detection time — set them before +launching Modly. + +| Variable | Effect | +|---|---| +| `MODLY_TORCH_FLAVOR` | `cuda` / `rocm` / `cpu`. Forces the path, skipping detection entirely. | +| `MODLY_ROCM_GFX` | Forces the compute target (e.g. `gfx1201`). Needed for an AMD card the Windows table doesn't map. | +| `MODLY_ROCM_INDEX` | Overrides the pip index, e.g. `https://download.pytorch.org/whl/rocm6.4` to fall back to torch 2.8/2.9. | +| `MODLY_ROCM_TORCH_SPEC` | Overrides the requirements entirely, space-separated: `"torch==2.9.1 torchvision==0.24.1"`. | +| `HSA_OVERRIDE_GFX_VERSION` | ROCm's own override, for cards without native kernels (e.g. `10.3.0` on an unsupported RDNA 2 part). Not needed on RDNA 3/4. | + +After changing any of these, run **Repair** on the extension so its venv is rebuilt. + +--- + +## 6. Troubleshooting + +| Symptom | Cause | Fix | +|---|---|---| +| Logs say `accelerator=cpu` on an AMD machine | `/dev/kfd` missing or unreadable (Linux), or an unmapped PCI id (Windows) | Check `ls -l /dev/kfd`; on Windows set `MODLY_ROCM_GFX` | +| `torch.__version__` ends in `+cu118` | Extension venv predates AMD support | Run **Repair** on the extension | +| `torch.cuda.is_available()` is `False` with a ROCm build | Device nodes not accessible from the process | Add your user to the `render` group, log back in | +| `HIP error: invalid device function` | Wheel has no kernels for your card | Set `HSA_OVERRIDE_GFX_VERSION` to a supported nearby target | +| Extension imports fail after install (`diffusers`/`transformers`) | The ROCm index only carries recent torch (2.11+), newer than some extensions expect | `MODLY_ROCM_INDEX=https://download.pytorch.org/whl/rocm6.4`, then **Repair** | +| Allocations above ~8 GB segfault on a 16 GB RX 9060 XT | [ROCm issue #6295](https://github.com/ROCm/ROCm/issues/6295) — did not reproduce on torch 2.13.0+rocm7.2 (see below) | If you hit it, pin an older stack via `MODLY_ROCM_INDEX` | + +--- + +## 7. Verified configuration and limitations + +The Linux path was measured end to end on a Radeon RX 9060 XT (Navi 44, gfx1200, +16 GB), CachyOS kernel 7.1.8, with the stack this page installs by default: + +- `torch 2.13.0+rocm7.2` / `torchvision 0.28.0+rocm7.2`, HIP runtime 7.2.53211 +- Detection reported `gfx1200`; `torch.cuda.is_available()` is `True` and + `get_device_properties(0).gcnArchName` is `gfx1200` +- rocBLAS (fp16 matmul) and MIOpen (conv2d) kernels both execute — no + `no kernel image is available` failures +- Allocations of 4/6/8/10/12/14 GB all succeeded, filled and read back. + [ROCm #6295](https://github.com/ROCm/ROCm/issues/6295), which reports this exact + card capped near 8 GB with a segfault, **did not reproduce** on this stack. + It remains open upstream, so it may still affect older ROCm builds. +- A full image-to-3D generation with `hunyuan3d-mini` completed through the normal + `ExtensionProcess` subprocess path: model load 18 s, generation 221 s, 15.8 MB GLB. + +### Attention kernels + +During generation PyTorch emits: + +> Mem Efficient attention on Current AMD GPU is still experimental. Enable it with +> `TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL=1` + +Generation works without it — `scaled_dot_product_attention` falls back to a +slower path. Setting `TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL=1` before launching +Modly enables the memory-efficient kernels, at the cost of running code AMD still +labels experimental on RDNA 3/4. Modly does not set it for you. + +Limitations: + +- **Windows is untested** by the Modly maintainers. The wheel URLs and the `cp311` + availability were verified, but no end-to-end run has been done on that path. +- **Texture generation** relies on optional native extensions (`texture_baker`, + `uv_unwrapper`) whose CUDA kernels have a HIP path but are not built as part of the + standard extension setup. diff --git a/electron/main/artifact-registry-service.test.ts b/electron/main/artifact-registry-service.test.ts index 882120d4..276e90a1 100644 --- a/electron/main/artifact-registry-service.test.ts +++ b/electron/main/artifact-registry-service.test.ts @@ -64,6 +64,18 @@ test('lists Workflows and Exports assets while skipping hidden, cache, and inter assert.equal(result.success && result.entries.find((entry) => entry.workspacePath.endsWith('exported.ply'))?.openable, false) })) +test('registers a generated scene directory through its canonical manifest artifact', () => withWorkspace(async (workspaceDir) => { + await mkdir(path.join(workspaceDir, 'Workflows/world'), { recursive: true }) + await writeFile(path.join(workspaceDir, 'Workflows/world/scene-manifest.json'), JSON.stringify({ + schema: 'modly.scene-manifest.v1', sceneRoot: '.', assets: [], + })) + const result = await listWorkspaceAssetLibrary({ workspaceDir }) + assert.equal(result.success, true) + const scene = result.success && result.entries.find((entry) => entry.workspacePath === 'Workflows/world/scene-manifest.json') + assert.equal(scene && scene.capability, 'scene-manifest') + assert.equal(scene && scene.state, 'ready') +})) + test('reads and opens only safe GLB/GLTF workspace assets', () => withWorkspace(async (workspaceDir) => { await mkdir(path.join(workspaceDir, 'Workflows/checkpoints'), { recursive: true }) await mkdir(path.join(workspaceDir, 'Exports'), { recursive: true }) diff --git a/electron/main/artifact-registry-service.ts b/electron/main/artifact-registry-service.ts index 3c26de73..2acb657f 100644 --- a/electron/main/artifact-registry-service.ts +++ b/electron/main/artifact-registry-service.ts @@ -124,7 +124,7 @@ export function classifyAssetLibraryCandidate(candidate: AssetLibraryClassificat if (candidate.workspacePath.endsWith('.world.json')) { return { capability: 'generated-world', state: 'ready', previewKind: 'text', openable: false, nonOpenableReason: 'Generated worlds are list-only in this release.' } } - if (candidate.workspacePath.endsWith('.scene.json')) { + if (candidate.workspacePath.endsWith('.scene.json') || candidate.workspacePath.endsWith('/scene-manifest.json')) { return { capability: 'scene-manifest', state: 'ready', previewKind: 'text', openable: false, nonOpenableReason: 'Scene manifests are list-only in this release.' } } if (INTRINSIC_MOTION_EXTENSIONS.has(extension)) { @@ -153,6 +153,7 @@ function objectField(value: unknown): Record | undefined { function manifestCapabilityFor(workspacePath: string): 'generated-world' | 'scene-manifest' | undefined { if (workspacePath.endsWith('.world.json')) return 'generated-world' if (workspacePath.endsWith('.scene.json')) return 'scene-manifest' + if (workspacePath.endsWith('/scene-manifest.json')) return 'scene-manifest' return undefined } diff --git a/electron/main/bounded-file-reader.test.mjs b/electron/main/bounded-file-reader.test.mjs new file mode 100644 index 00000000..03071180 --- /dev/null +++ b/electron/main/bounded-file-reader.test.mjs @@ -0,0 +1,48 @@ +import assert from 'node:assert/strict' +import { readFileSync } from 'node:fs' +import { resolve } from 'node:path' +import test from 'node:test' + +const { + MAX_SCENE_MANIFEST_BYTES, + readLocalFileBase64, +} = await import(new URL('./bounded-file-reader.ts', import.meta.url).href) + +test('the authoritative readFileBase64 IPC delegates to the bounded reader', () => { + const handlers = readFileSync(resolve('electron/main/ipc-handlers.ts'), 'utf8') + assert.match( + handlers, + /ipcMain\.handle\('fs:readFileBase64',[\s\S]*?readLocalFileBase64\(filePath\)/, + ) +}) + +test('oversized scene manifests are rejected before their bytes are read', async () => { + let reads = 0 + await assert.rejects( + readLocalFileBase64('/workspace/room/scene-manifest.json', { + statFile: async () => ({ size: MAX_SCENE_MANIFEST_BYTES + 1, isFile: () => true }), + readBytes: async () => { reads++; return new Uint8Array() }, + }), + /exceeds the 1 MiB limit/, + ) + assert.equal(reads, 0) +}) + +test('scene manifests at the bound are read and encoded', async () => { + const bytes = new TextEncoder().encode('{"schema":"modly.scene-manifest.v1"}') + const encoded = await readLocalFileBase64('/workspace/room/scene-manifest.json', { + statFile: async () => ({ size: MAX_SCENE_MANIFEST_BYTES, isFile: () => true }), + readBytes: async () => bytes, + }) + assert.equal(encoded, Buffer.from(bytes).toString('base64')) +}) + +test('a manifest that grows after stat is rejected before base64 transfer', async () => { + await assert.rejects( + readLocalFileBase64('/workspace/room/scene-manifest.json', { + statFile: async () => ({ size: 1, isFile: () => true }), + readBytes: async () => new Uint8Array(MAX_SCENE_MANIFEST_BYTES + 1), + }), + /exceeds the 1 MiB limit/, + ) +}) diff --git a/electron/main/bounded-file-reader.ts b/electron/main/bounded-file-reader.ts new file mode 100644 index 00000000..f7b18ba4 --- /dev/null +++ b/electron/main/bounded-file-reader.ts @@ -0,0 +1,40 @@ +import { basename } from 'node:path' +import { readFile, stat } from 'node:fs/promises' + +export const SCENE_MANIFEST_FILE_NAME = 'scene-manifest.json' +export const MAX_SCENE_MANIFEST_BYTES = 1024 * 1024 + +type FileReadDependencies = { + statFile: (filePath: string) => Promise<{ size: number; isFile: () => boolean }> + readBytes: (filePath: string) => Promise +} + +const DEFAULT_DEPENDENCIES: FileReadDependencies = { + statFile: stat, + readBytes: readFile, +} + +export async function readLocalFileBase64( + filePath: string, + dependencies: FileReadDependencies = DEFAULT_DEPENDENCIES, +): Promise { + if (typeof filePath !== 'string' || filePath.trim().length === 0) { + throw new Error('fs:readFileBase64 requires a non-empty file path') + } + + const isSceneManifest = basename(filePath) === SCENE_MANIFEST_FILE_NAME + if (isSceneManifest) { + const fileInfo = await dependencies.statFile(filePath) + if (!fileInfo.isFile()) throw new Error('Scene manifest is not a file') + if (fileInfo.size > MAX_SCENE_MANIFEST_BYTES) { + throw new Error('Scene manifest exceeds the 1 MiB limit') + } + } + + const bytes = await dependencies.readBytes(filePath) + // Recheck the transferred bytes in case the file grew between stat and read. + if (isSceneManifest && bytes.byteLength > MAX_SCENE_MANIFEST_BYTES) { + throw new Error('Scene manifest exceeds the 1 MiB limit') + } + return Buffer.from(bytes).toString('base64') +} diff --git a/electron/main/copy-runtime.test.mjs b/electron/main/copy-runtime.test.mjs new file mode 100644 index 00000000..b5dcbdb0 --- /dev/null +++ b/electron/main/copy-runtime.test.mjs @@ -0,0 +1,100 @@ +/** + * Guards the symlink property that makes the AppImage's "stable" Python copy + * actually stable. Getting this wrong is silent at copy time and only breaks on + * the *next* launch, once the source mount is gone — so it needs a test that + * checks the links rather than the copy succeeding. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { + mkdtempSync, mkdirSync, writeFileSync, symlinkSync, readlinkSync, rmSync, existsSync, +} from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve, isAbsolute } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-copy-test-')), 'copy-runtime.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/copy-runtime.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { copyRuntimeTree } = loadModule() + +// makeRuntime() below needs symlinkSync, which on Windows requires Developer +// Mode (or elevation) — and the AppImage failure mode this file guards is +// Linux-only anyway. +const symlinkOpts = { + skip: process.platform === 'win32' && 'symlinks require Developer Mode on Windows', +} + +/** Builds a miniature of the bundled runtime: a real binary plus relative links. */ +function makeRuntime(root) { + mkdirSync(join(root, 'bin'), { recursive: true }) + mkdirSync(join(root, 'lib'), { recursive: true }) + writeFileSync(join(root, 'bin', 'python3.11'), '#!/bin/sh\n', 'utf8') + symlinkSync('python3.11', join(root, 'bin', 'python3')) + symlinkSync('python3.11', join(root, 'bin', 'python')) + writeFileSync(join(root, 'lib', 'libpython3.11.so.1.0'), '', 'utf8') + symlinkSync('libpython3.11.so.1.0', join(root, 'lib', 'libpython3.11.so')) +} + +test('copyRuntimeTree keeps relative symlinks relative', symlinkOpts, async () => { + const dir = mkdtempSync(join(tmpdir(), 'modly-runtime-')) + const source = join(dir, 'mount', 'python-embed') + const dest = join(dir, 'stable', 'python-embed') + makeRuntime(source) + + await copyRuntimeTree(source, dest) + + for (const link of ['bin/python3', 'bin/python', 'lib/libpython3.11.so']) { + const target = readlinkSync(join(dest, link)) + assert.ok( + !isAbsolute(target), + `${link} was rewritten to an absolute path (${target}); it would point back at the source mount`, + ) + assert.ok(!target.includes(source), `${link} still references the source tree`) + } + + rmSync(dir, { recursive: true, force: true }) +}) + +test('the copy survives the source being deleted', symlinkOpts, async () => { + // This is the actual failure mode: the AppImage mount disappears between + // launches, and every venv built from the copy dies with it. + const dir = mkdtempSync(join(tmpdir(), 'modly-runtime-')) + const source = join(dir, 'mount', 'python-embed') + const dest = join(dir, 'stable', 'python-embed') + makeRuntime(source) + + await copyRuntimeTree(source, dest) + rmSync(join(dir, 'mount'), { recursive: true, force: true }) + + // existsSync follows symlinks, so this is false for a dangling link — exactly + // the check generator_registry.py's _venv_python(...).exists() performs. + assert.ok(existsSync(join(dest, 'bin', 'python3')), 'bin/python3 is dangling after the source went away') + assert.ok(existsSync(join(dest, 'lib', 'libpython3.11.so')), 'libpython3.11.so is dangling') +}) + +test('copyRuntimeTree copies regular files and directory structure', symlinkOpts, async () => { + const dir = mkdtempSync(join(tmpdir(), 'modly-runtime-')) + const source = join(dir, 'mount', 'python-embed') + const dest = join(dir, 'stable', 'python-embed') + makeRuntime(source) + + await copyRuntimeTree(source, dest) + + assert.ok(existsSync(join(dest, 'bin', 'python3.11'))) + assert.ok(existsSync(join(dest, 'lib', 'libpython3.11.so.1.0'))) + + rmSync(dir, { recursive: true, force: true }) +}) diff --git a/electron/main/copy-runtime.ts b/electron/main/copy-runtime.ts new file mode 100644 index 00000000..eb062f5c --- /dev/null +++ b/electron/main/copy-runtime.ts @@ -0,0 +1,35 @@ +/** + * Copying the bundled Python runtime out of an ephemeral AppImage mount. + * + * Kept apart from python-setup.ts (which imports electron) so the symlink + * behaviour this depends on can be tested — it is subtle, silent when wrong, + * and only shows up on the *next* launch. + */ + +import { cp } from 'fs/promises' + +/** + * Copies a self-contained runtime tree, keeping relative symlinks relative. + * + * `verbatimSymlinks` is the whole point. Without it `fs.cp` resolves a relative + * link (`bin/python3 -> python3.11`) into an absolute path pointing back at the + * *source* tree. When the source is an AppImage mount at + * /tmp/.mount_Modly-XXXXXX/ — which is a different path on every launch — the + * copy silently keeps a hard dependency on a directory that is about to vanish: + * + * - `bin/python3` in the "stable" copy points into the old mount + * - `sys._base_executable` of any venv made from it inherits that path + * - every extension venv records it in pyvenv.cfg and as a bin/python symlink + * - on the next launch the mount is gone, so every extension venv is dead + * + * The user-visible symptom is remote from the cause: the extension registry + * finds no usable venv, falls back to importing generator.py in the main API + * process, and reports a missing third-party module such as `No module named 'PIL'`. + */ +export async function copyRuntimeTree(source: string, destination: string): Promise { + await cp(source, destination, { + recursive: true, + preserveTimestamps: true, + verbatimSymlinks: true, + }) +} diff --git a/electron/main/extension-install-utils.test.mjs b/electron/main/extension-install-utils.test.mjs index d5d1a389..ec192d71 100644 --- a/electron/main/extension-install-utils.test.mjs +++ b/electron/main/extension-install-utils.test.mjs @@ -53,6 +53,156 @@ test('validateInstallManifest still rejects missing process entry files', () => ) }) +test('validateInstallManifest accepts multi-source nodes and preserves legacy shapes', () => { + const mod = loadModule() + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'multi-model', + generator_class: 'Generator', + nodes: [{ + id: 'generate', + model_sources: [ + { + id: 'primary', provider: 'huggingface', repo_id: 'org/main', + destination: '.', checks: ['pipeline.json'], + }, + { + id: 'encoder', provider: 'huggingface', repo_id: 'org/encoder', + destination: 'auxiliary/encoder', checks: ['model.safetensors'], + }, + ], + }], + }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) + + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'legacy', + generator_class: 'Generator', + nodes: [{ + id: 'projection', + hf_repo: 'org/legacy', + download_check: '../generate/model.safetensors', + hf_skip_prefixes: ['weights/**'], + }], + }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository')) + + // 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 accepts scene IO without rejecting third-party artifact kinds', () => { + const mod = loadModule() + const files = { hasEntryFile: () => false, hasGeneratorFile: () => true } + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'scene-model', generator_class: 'Generator', + nodes: [{ id: 'normalize', input: 'scene', output: 'scene' }], + }, files, 'repository')) + for (const input of ['capture', 'video']) { + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'future-model', generator_class: 'Generator', + nodes: [{ id: 'future', input, output: 'scene' }], + }, files, 'repository')) + } +}) + +test('scene is model-only, single-input, while image-multi to scene stays valid', () => { + const mod = loadModule() + const modelFiles = { hasEntryFile: () => false, hasGeneratorFile: () => true } + const processFiles = { hasEntryFile: () => true, hasGeneratorFile: () => false } + for (const node of [ + { id: 'mixed', input: 'scene', inputs: ['scene', 'text'], output: 'mesh' }, + { id: 'duplicate', input: 'scene', inputs: ['scene', 'scene'], output: 'mesh' }, + { id: 'hidden', input: 'image', inputs: ['scene'], output: 'mesh' }, + ]) { + assert.throws(() => mod.validateInstallManifest({ id: 'bad', generator_class: 'Generator', nodes: [node] }, modelFiles, 'repository'), /scene.*single|single.*scene/i) + } + for (const node of [ + { id: 'input', input: 'scene', output: 'mesh' }, + { id: 'output', input: 'image', output: 'scene' }, + ]) { + assert.throws(() => mod.validateInstallManifest({ id: 'proc', type: 'process', entry: 'processor.js', nodes: [node] }, processFiles, 'repository'), /scene.*model|model.*scene/i) + } + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'images-to-scene', generator_class: 'Generator', + nodes: [{ id: 'prepare', input: 'image', inputs: ['image', 'image'], output: 'scene' }], + }, modelFiles, 'repository')) + for (const output of ['scene', 'mesh']) { + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: `scene-to-${output}`, generator_class: 'Generator', + nodes: [{ id: 'generate', input: 'scene', output }], + }, modelFiles, 'repository')) + } +}) + +test('validateInstallManifest rejects malformed or process model_sources', () => { + const mod = loadModule() + const source = { + id: 'weights', provider: 'huggingface', repo_id: 'org/model', + destination: '../outside', checks: ['model.safetensors'], + } + assert.throws(() => mod.validateInstallManifest({ + id: 'unsafe', generator_class: 'Generator', + nodes: [{ id: 'generate', model_sources: [source] }], + }, { hasEntryFile: () => false, hasGeneratorFile: () => true }, 'repository'), /destination/i) + + assert.throws(() => mod.validateInstallManifest({ + id: 'process', type: 'process', entry: 'processor.js', + nodes: [{ id: 'run', model_sources: [{ ...source, destination: '.' }] }], + }, { hasEntryFile: () => true, hasGeneratorFile: () => false }, 'repository'), /only for model nodes/i) +}) + +test('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() @@ -317,3 +467,43 @@ 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/) + } +}) + +test('validateInstallManifest validates weight variants and keeps them off process nodes', () => { + const mod = loadModule() + const files = { hasEntryFile: () => true, hasGeneratorFile: () => true } + const weightVariants = { + param: 'quant', + options: [{ id: 'Q4', include_prefixes: ['dit/model_Q4.gguf'], checks: ['dit/model_Q4.gguf'] }], + } + const paramsSchema = [{ id: 'quant', type: 'select', options: [{ value: 'Q4' }] }] + assert.doesNotThrow(() => mod.validateInstallManifest({ + id: 'quantized', generator_class: 'Generator', params_schema: paramsSchema, + nodes: [{ id: 'generate', hf_repo: 'org/model', weight_variants: weightVariants }], + }, files, 'repository')) + assert.throws(() => mod.validateInstallManifest({ + id: 'quantized', generator_class: 'Generator', + nodes: [{ id: 'generate', hf_repo: 'org/model', weight_variants: weightVariants }], + }, files, 'repository'), /must name a params_schema entry/) + assert.throws(() => mod.validateInstallManifest({ + id: 'quantized', generator_class: 'Generator', params_schema: paramsSchema, + nodes: [{ id: 'generate', weight_variants: weightVariants }], + }, files, 'repository'), /requires hf_repo/) + assert.throws(() => mod.validateInstallManifest({ + id: 'proc', type: 'process', entry: 'processor.js', + nodes: [{ id: 'run', hf_repo: 'org/model', weight_variants: weightVariants }], + }, files, 'repository'), /weight_variants is supported only for model nodes/) +}) + diff --git a/electron/main/extension-install-utils.ts b/electron/main/extension-install-utils.ts index 4195cdca..2677f1ec 100644 --- a/electron/main/extension-install-utils.ts +++ b/electron/main/extension-install-utils.ts @@ -1,9 +1,32 @@ +import { + normalizeModelSources, + normalizeWeightGroupReferences, + normalizeWeightGroups, + normalizeWeightVariants, + validateModelNodeIds, + safeModelSourceId, + type ModelWeightNode, + type WeightVariantNode, +} from './model-sources' + export interface InstallManifest { id?: string type?: 'model' | 'process' entry?: string generator_class?: string - nodes?: Array<{ id?: string }> + model_sources?: unknown + params_schema?: unknown + weight_groups?: unknown + nodes?: Array<{ + id?: string + input?: unknown + inputs?: unknown + output?: unknown + hf_repo?: unknown + model_sources?: unknown + weight_groups?: unknown + weight_variants?: unknown + } & ModelWeightNode & WeightVariantNode> } export interface ValidatedInstallManifest { @@ -20,6 +43,21 @@ export interface ExtensionReloadPayload { errors: Record } +export function assertSupportedSceneNodeShape( + kind: 'model' | 'process', + node: { id?: string; input?: unknown; inputs?: unknown; output?: unknown }, + declaredInputs: unknown[], + output: unknown, +): void { + const usesSceneInput = declaredInputs.includes('scene') + if (kind === 'process' && (usesSceneInput || output === 'scene')) { + throw new Error('manifest.json: scene input and output are supported only for model nodes') + } + if (kind === 'model' && usesSceneInput && (node.inputs !== undefined || node.input !== 'scene')) { + throw new Error(`manifest.json: ${node.id ?? 'node'} must declare scene as its single input field`) + } +} + export type IncompleteInstallRecoveryAction = | 'none' | 'remove-incomplete' @@ -38,6 +76,43 @@ export function validateInstallManifest( const isProcess = manifest.type === 'process' const entryFile = manifest.entry ?? 'processor.js' const nodes = Array.isArray(manifest.nodes) ? manifest.nodes.filter((node) => node?.id) : [] + 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) + 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 declaredInputs = Array.isArray(node.inputs) ? node.inputs : [node.input ?? 'image'] + const output = node.output ?? 'mesh' + assertSupportedSceneNodeShape(isProcess ? 'process' : 'model', node, declaredInputs, output) + 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.weight_variants !== undefined) { + throw new Error('manifest.json: weight_variants is supported only for model nodes') + } + 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 && node.weight_variants === undefined) continue + const nodeId = safeModelSourceId(node.id, 'model node id') + if (node.model_sources !== undefined) normalizeModelSources(node) + // Before the weight_groups/hf_repo check, so a variants + groups node gets + // the explicit "cannot be combined" error. + normalizeWeightVariants(node, node.params_schema ?? manifest.params_schema) + 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`, + ) + } + } if (isProcess) { if (!opts.hasEntryFile(entryFile)) { diff --git a/electron/main/gpu-detect.test.mjs b/electron/main/gpu-detect.test.mjs new file mode 100644 index 00000000..3746da5c --- /dev/null +++ b/electron/main/gpu-detect.test.mjs @@ -0,0 +1,253 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-gpu-test-')), 'gpu-detect.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/gpu-detect.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const mod = loadModule() + +// ─── KFD topology parsing ───────────────────────────────────────────────────── + +test('formatGfxTarget decodes gfx_target_version across GPU generations', () => { + // Encoding is major*10000 + minor*100 + step, minor and step read as hex + // digits. 120000 is the value this machine's RX 9060 XT actually reports. + assert.equal(mod.formatGfxTarget(120000), 'gfx1200') // Navi 44 / RX 9060 XT + assert.equal(mod.formatGfxTarget(120001), 'gfx1201') // Navi 48 + assert.equal(mod.formatGfxTarget(110000), 'gfx1100') // Navi 31 + assert.equal(mod.formatGfxTarget(110002), 'gfx1102') // Navi 33 + assert.equal(mod.formatGfxTarget(100300), 'gfx1030') // Navi 21 + assert.equal(mod.formatGfxTarget(90402), 'gfx942') // MI300, hex step + assert.equal(mod.formatGfxTarget(90010), 'gfx90a') // MI200, hex step +}) + +test('formatGfxTarget rejects the "no GPU" and malformed encodings', () => { + assert.equal(mod.formatGfxTarget(0), null) + assert.equal(mod.formatGfxTarget(-1), null) + assert.equal(mod.formatGfxTarget(1.5), null) +}) + +test('parseKfdGfxTarget skips the CPU node and reads the first GPU', () => { + // Verbatim shape of /sys/class/kfd/kfd/topology/nodes/*/properties + const cpuNode = 'cpu_cores_count 32\nsimd_count 0\ngfx_target_version 0\n' + const gpuNode = 'cpu_cores_count 0\nsimd_count 64\ngfx_target_version 120000\n' + + assert.equal(mod.parseKfdGfxTarget([cpuNode, gpuNode]), 'gfx1200') +}) + +test('parseKfdGfxTarget prefers the discrete GPU over an APU', () => { + // Ryzen 7840 (gfx1103, 24 SIMDs) + RX 7900 XTX (gfx1100, 384 SIMDs): the APU + // gets the lower node number, but gfx1103 has no published wheels — the card + // with more SIMDs is the one the user bought for this. + const cpuNode = 'cpu_cores_count 16\nsimd_count 0\ngfx_target_version 0\n' + const apuNode = 'cpu_cores_count 0\nsimd_count 24\ngfx_target_version 110003\n' + const dgpuNode = 'cpu_cores_count 0\nsimd_count 384\ngfx_target_version 110000\n' + + assert.equal(mod.parseKfdGfxTarget([cpuNode, apuNode, dgpuNode]), 'gfx1100') + // Order must not matter + assert.equal(mod.parseKfdGfxTarget([cpuNode, dgpuNode, apuNode]), 'gfx1100') +}) + +test('parseKfdGfxTarget returns null without a usable GPU node', () => { + assert.equal(mod.parseKfdGfxTarget([]), null) + assert.equal(mod.parseKfdGfxTarget(['simd_count 0\ngfx_target_version 0\n']), null) + // A node advertising SIMDs but no compute target is not something we can target + assert.equal(mod.parseKfdGfxTarget(['simd_count 64\ngfx_target_version 0\n']), null) +}) + +test('parseKfdGfxTarget does not confuse a key with its prefix', () => { + // simd_count must not be satisfied by e.g. "max_simd_count" + const node = 'max_simd_count 999\nsimd_count 64\ngfx_target_version 110002\n' + assert.equal(mod.parseKfdGfxTarget([node]), 'gfx1102') +}) + +// ─── nvidia-smi parsing (non-regression) ────────────────────────────────────── + +test('parseNvidiaSmi maps compute cap and driver version to CUDA version', () => { + assert.deepEqual(mod.parseNvidiaSmi('8.6, 551.61\n'), { sm: 86, cudaVersion: 124 }) + assert.deepEqual(mod.parseNvidiaSmi('12.0, 572.16\n'), { sm: 120, cudaVersion: 128 }) + assert.deepEqual(mod.parseNvidiaSmi('6.1, 470.82\n'), { sm: 61, cudaVersion: 118 }) +}) + +test('parseNvidiaSmi falls back to sm 86 on an unparseable compute cap', () => { + assert.deepEqual(mod.parseNvidiaSmi('N/A, 551.61\n'), { sm: 86, cudaVersion: 124 }) +}) + +test('parseNvidiaSmi returns null on empty output', () => { + assert.equal(mod.parseNvidiaSmi(''), null) + assert.equal(mod.parseNvidiaSmi(' \n'), null) +}) + +// ─── Windows adapter mapping ────────────────────────────────────────────────── + +test('parseWindowsVideoControllers accepts both the object and array JSON shapes', () => { + const single = mod.parseWindowsVideoControllers( + '{"Name":"AMD Radeon RX 9060 XT","PNPDeviceID":"PCI\\\\VEN_1002&DEV_7590&SUBSYS_06391043&REV_C0\\\\4&1"}', + ) + assert.deepEqual(single, [ + { name: 'AMD Radeon RX 9060 XT', pnpDeviceId: 'PCI\\VEN_1002&DEV_7590&SUBSYS_06391043&REV_C0\\4&1' }, + ]) + + const many = mod.parseWindowsVideoControllers( + '[{"Name":"A","PNPDeviceID":"PCI\\\\VEN_1002&DEV_7550"},{"Name":"B","PNPDeviceID":"PCI\\\\VEN_8086&DEV_1234"}]', + ) + assert.equal(many.length, 2) +}) + +test('parseWindowsVideoControllers survives malformed PowerShell output', () => { + assert.deepEqual(mod.parseWindowsVideoControllers('not json'), []) + assert.deepEqual(mod.parseWindowsVideoControllers('{"Name":"No id"}'), []) +}) + +test('parseAmdPciDeviceId only matches AMD vendor ids', () => { + assert.equal(mod.parseAmdPciDeviceId('PCI\\VEN_1002&DEV_7590&SUBSYS_0'), '7590') + assert.equal(mod.parseAmdPciDeviceId('PCI\\VEN_10DE&DEV_2684'), null) +}) + +test('resolveWindowsGfxTarget maps device ids by silicon, not by marketing range', () => { + const target = (deviceId) => + mod.resolveWindowsGfxTarget([{ name: 'x', pnpDeviceId: `PCI\\VEN_1002&DEV_${deviceId}` }]).gfxTarget + + assert.equal(target('7590'), 'gfx1200') // Navi 44 + assert.equal(target('7550'), 'gfx1201') // Navi 48 + assert.equal(target('744C'), 'gfx1100') // Navi 31, uppercase from WMI + assert.equal(target('747e'), 'gfx1101') // Navi 32 + // 0x73f0 sells as "RX 7600M XT" but is Navi 33 — it must not land with its + // 0x73xx RDNA2 neighbours. + assert.equal(target('73f0'), 'gfx1102') + assert.equal(target('73bf'), 'gfx1030') // Navi 21 +}) + +test('resolveWindowsGfxTarget reports unmapped AMD cards instead of guessing', () => { + const result = mod.resolveWindowsGfxTarget([ + { name: 'AMD Radeon RX 9999', pnpDeviceId: 'PCI\\VEN_1002&DEV_FFFF' }, + ]) + assert.equal(result.gfxTarget, null) + assert.deepEqual(result.amdAdapters, ['AMD Radeon RX 9999']) +}) + +test('resolveWindowsGfxTarget ignores non-AMD adapters', () => { + const result = mod.resolveWindowsGfxTarget([ + { name: 'NVIDIA RTX 4090', pnpDeviceId: 'PCI\\VEN_10DE&DEV_2684' }, + ]) + assert.equal(result.gfxTarget, null) + assert.deepEqual(result.amdAdapters, []) +}) + +// ─── ROCm wheel source resolution ───────────────────────────────────────────── + +test('resolveRocmTorchSpec uses the unpinned pytorch.org index on Linux', () => { + const { indexUrl, specs } = mod.resolveRocmTorchSpec('linux', 'gfx1200', {}) + assert.equal(indexUrl, 'https://download.pytorch.org/whl/rocm7.2') + assert.deepEqual(specs, ['torch', 'torchvision']) +}) + +test('resolveRocmTorchSpec uses AMD\'s index with a device extra on Windows', () => { + // download.pytorch.org publishes no ROCm wheels for Windows at all. + const { indexUrl, specs } = mod.resolveRocmTorchSpec('win32', 'gfx1200', {}) + assert.equal(indexUrl, 'https://repo.amd.com/rocm/whl-multi-arch/') + assert.deepEqual(specs, [ + 'torch[device-gfx1200]==2.11.0+rocm7.14.0', + 'torchvision[device-gfx1200]==0.26.0+rocm7.14.0', + ]) +}) + +test('resolveRocmTorchSpec honours MODLY_ROCM_INDEX and MODLY_ROCM_TORCH_SPEC', () => { + const rolledBack = mod.resolveRocmTorchSpec('linux', 'gfx1200', { + MODLY_ROCM_INDEX: 'https://download.pytorch.org/whl/rocm6.4', + MODLY_ROCM_TORCH_SPEC: 'torch==2.8.0 torchvision==0.23.0', + }) + assert.equal(rolledBack.indexUrl, 'https://download.pytorch.org/whl/rocm6.4') + assert.deepEqual(rolledBack.specs, ['torch==2.8.0', 'torchvision==0.23.0']) +}) + +test('resolveRocmTorchSpec stays unpinned on Windows without a compute target', () => { + // No target means no device extra to ask for; better an install that fails + // loudly than one silently pinned to the wrong architecture. + const { specs } = mod.resolveRocmTorchSpec('win32', null, {}) + assert.deepEqual(specs, ['torch', 'torchvision']) +}) + +test('torchFlavorFor maps accelerators to the setup.py argument', () => { + assert.equal(mod.torchFlavorFor('rocm'), 'rocm') + assert.equal(mod.torchFlavorFor('cuda'), 'cuda') + assert.equal(mod.torchFlavorFor('cpu'), 'cpu') + assert.equal(mod.torchFlavorFor('mps'), 'cpu') +}) + +// ─── Detection precedence ───────────────────────────────────────────────────── + +test('detectGpuInfo keeps Apple Silicon on MPS', async () => { + const info = await mod.detectGpuInfo({ env: {}, platform: 'darwin', arch: 'arm64' }) + assert.equal(info.accelerator, 'mps') +}) + +test('detectGpuInfo forced to rocm resolves wheels without probing hardware', async () => { + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'rocm', MODLY_ROCM_GFX: 'gfx1201' }, + platform: 'linux', + arch: 'x64', + }) + assert.equal(info.accelerator, 'rocm') + assert.equal(info.gfxTarget, 'gfx1201') + assert.equal(info.torchIndexUrl, 'https://download.pytorch.org/whl/rocm7.2') + // sm/cudaVersion stay at 0 so CUDA-era extensions take their most + // conservative branch (and keep off rembg[gpu], which is CUDA-only). + assert.equal(info.sm, 0) + assert.equal(info.cudaVersion, 0) +}) + +test('detectGpuInfo forced to rocm overrides MPS on Apple Silicon', async () => { + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'rocm', MODLY_ROCM_GFX: 'gfx1200' }, + platform: 'darwin', + arch: 'arm64', + }) + assert.equal(info.accelerator, 'rocm') +}) + +test('detectGpuInfo forced to cuda overrides MPS on Apple Silicon', async () => { + // The override has to beat every probe, the darwin/arm64 default included. + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'cuda' }, + platform: 'darwin', + arch: 'arm64', + }) + assert.equal(info.accelerator, 'cuda') +}) + +test('detectGpuInfo forced to cpu short-circuits everything', async () => { + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'cpu', MODLY_ROCM_GFX: 'gfx1200' }, + platform: 'linux', + arch: 'x64', + }) + assert.deepEqual(info, { sm: 0, cudaVersion: 0, accelerator: 'cpu' }) +}) + +test('detectGpuInfo falls back to CPU when forced to rocm on Windows with no target', async () => { + const logs = [] + const info = await mod.detectGpuInfo({ + env: { MODLY_TORCH_FLAVOR: 'rocm' }, + platform: 'win32', + arch: 'x64', + onLog: (line) => logs.push(line), + }) + assert.equal(info.accelerator, 'cpu') + assert.ok(logs.some((l) => l.includes('MODLY_ROCM_GFX'))) +}) diff --git a/electron/main/gpu-detect.ts b/electron/main/gpu-detect.ts new file mode 100644 index 00000000..44ca4531 --- /dev/null +++ b/electron/main/gpu-detect.ts @@ -0,0 +1,407 @@ +/** + * GPU detection and PyTorch flavour resolution. + * + * Deliberately free of electron imports: the parsing and resolution helpers are + * pure and get unit-tested by bundling this file directly (gpu-detect.test.mjs). + * + * Modly never installs PyTorch itself — each extension's setup.py does, from an + * index it picks on its own. What we produce here is the information that lets + * that choice land on the right wheels: the accelerator, and for AMD the ROCm + * compute target plus the pip index/requirements the setup must end up using. + */ + +import { spawn } from 'child_process' +import { existsSync, readdirSync, readFileSync } from 'fs' +import { join } from 'path' + +export type Accelerator = 'cuda' | 'rocm' | 'mps' | 'cpu' +export type TorchFlavor = 'cuda' | 'rocm' | 'cpu' + +export interface GpuInfo { + sm: number + cudaVersion: number + accelerator: Accelerator + /** ROCm compute target ("gfx1200"). Only set when accelerator is 'rocm'. */ + gfxTarget?: string + /** pip --index-url the torch install has to come from (ROCm only). */ + torchIndexUrl?: string + /** pip requirements replacing whatever torch pins an extension hardcodes. */ + torchSpecs?: string[] +} + +// ─── ROCm wheel sources ─────────────────────────────────────────────────────── +// +// Linux and Windows need different indexes. download.pytorch.org publishes no +// ROCm wheels for Windows at all, so Windows has to go through AMD's multi-arch +// index, where the compute target is selected by a pip extra +// (torch[device-gfx1200]) instead of by the index URL. + +const ROCM_LINUX_INDEX = 'https://download.pytorch.org/whl/rocm7.2' +const ROCM_WINDOWS_INDEX = 'https://repo.amd.com/rocm/whl-multi-arch/' + +// Pinned because AMD's index carries several ROCm builds side by side; this is +// the newest pair published for cp311 (Modly's embedded Python) on Windows. +const ROCM_WINDOWS_TORCH = '2.11.0+rocm7.14.0' +const ROCM_WINDOWS_TORCHVISION = '0.26.0+rocm7.14.0' + +const KFD_TOPOLOGY_DIR = '/sys/class/kfd/kfd/topology/nodes' + +/** + * AMD PCI device id → ROCm compute target, for Windows where there is no KFD + * topology to read. Keyed by silicon rather than by marketing name: 0x73f0 is + * sold as an "RX 7600M XT" but is Navi 33, so it belongs with gfx1102, not with + * its 0x73xx neighbours. Device ids come from the pci.ids database. + */ +const WINDOWS_PCI_GFX_TARGETS: Record = { + // Navi 21 (RDNA 2) + '73a1': 'gfx1030', '73a2': 'gfx1030', '73a3': 'gfx1030', '73a5': 'gfx1030', + '73ab': 'gfx1030', '73ae': 'gfx1030', '73af': 'gfx1030', '73bf': 'gfx1030', + // Navi 22 (RDNA 2) + '73c3': 'gfx1031', '73ce': 'gfx1031', '73df': 'gfx1031', + // Navi 23 (RDNA 2) + '73e0': 'gfx1032', '73e1': 'gfx1032', '73e3': 'gfx1032', '73ef': 'gfx1032', + '73ff': 'gfx1032', + // Navi 24 (RDNA 2) + '7421': 'gfx1034', '7422': 'gfx1034', '7423': 'gfx1034', '7424': 'gfx1034', + '743f': 'gfx1034', + // Navi 31 (RDNA 3) + '7448': 'gfx1100', '7449': 'gfx1100', '744a': 'gfx1100', '744b': 'gfx1100', + '744c': 'gfx1100', '745e': 'gfx1100', + // Navi 32 (RDNA 3) + '7460': 'gfx1101', '7461': 'gfx1101', '7470': 'gfx1101', '747e': 'gfx1101', + // Navi 33 (RDNA 3) + '73f0': 'gfx1102', '7480': 'gfx1102', '7481': 'gfx1102', '7483': 'gfx1102', + '7487': 'gfx1102', '7489': 'gfx1102', '748b': 'gfx1102', '7499': 'gfx1102', + '749f': 'gfx1102', + // Navi 44 / Navi 48 (RDNA 4) + '7590': 'gfx1200', + '7550': 'gfx1201', '7551': 'gfx1201', +} + +// ─── Pure parsing helpers ───────────────────────────────────────────────────── + +/** + * Parses `nvidia-smi --query-gpu=compute_cap,driver_version --format=csv,noheader`. + * Returns null when the output carries no usable line. + */ +export function parseNvidiaSmi(stdout: string): { sm: number; cudaVersion: number } | null { + const line = stdout.trim().split('\n')[0]?.trim() // e.g. "8.6, 551.61" + if (!line) return null + + const parts = line.split(',').map((s) => s.trim()) + const sm = Math.round(parseFloat(parts[0] ?? '') * 10) // → 86 + + // Derive max supported CUDA version from driver version + // Driver ≥ 520 → CUDA 11.8, ≥ 525 → 12.0, ≥ 530 → 12.1, ≥ 535 → 12.2, + // ≥ 545 → 12.3, ≥ 550 → 12.4, ≥ 555 → 12.5, ≥ 560 → 12.6 + const driverMajor = parseInt((parts[1] ?? '').split('.')[0] ?? '0', 10) + let cudaVersion = 118 // safe minimum + if (driverMajor >= 570) cudaVersion = 128 // Blackwell (RTX 50xx, sm_120) + else if (driverMajor >= 560) cudaVersion = 126 + else if (driverMajor >= 555) cudaVersion = 125 + else if (driverMajor >= 550) cudaVersion = 124 + else if (driverMajor >= 545) cudaVersion = 123 + else if (driverMajor >= 535) cudaVersion = 122 + else if (driverMajor >= 530) cudaVersion = 121 + else if (driverMajor >= 525) cudaVersion = 120 + else if (driverMajor >= 520) cudaVersion = 118 + + return { sm: isNaN(sm) ? 86 : sm, cudaVersion } +} + +/** + * Decodes the KFD `gfx_target_version` integer (major*10000 + minor*100 + step, + * with minor and step read as hex digits) into a compute target name: + * 120000 → gfx1200, 90402 → gfx942, 90010 → gfx90a. + */ +export function formatGfxTarget(version: number): string | null { + if (!Number.isInteger(version) || version <= 0) return null + const major = Math.floor(version / 10000) + const minor = Math.floor((version % 10000) / 100) + const step = version % 100 + if (major <= 0 || minor > 15 || step > 15) return null + return `gfx${major}${minor.toString(16)}${step.toString(16)}` +} + +/** + * Picks the compute target out of the KFD topology node `properties` files. + * Node 0 is the CPU node (simd_count 0) and is skipped. Among GPU nodes the + * largest simd_count wins: on an APU + dGPU machine the APU commonly gets the + * lower node number, and its gfx target (e.g. gfx1103) often has no published + * wheels — the discrete card is always the one with more SIMDs. + */ +export function parseKfdGfxTarget(nodeProperties: string[]): string | null { + let best: { target: string; simdCount: number } | null = null + for (const text of nodeProperties) { + const simdCount = readKfdProperty(text, 'simd_count') + if (simdCount <= 0) continue + const target = formatGfxTarget(readKfdProperty(text, 'gfx_target_version')) + if (!target) continue + if (!best || simdCount > best.simdCount) best = { target, simdCount } + } + return best?.target ?? null +} + +function readKfdProperty(text: string, key: string): number { + const match = new RegExp(`^${key}\\s+(\\d+)\\s*$`, 'm').exec(text) + return match ? parseInt(match[1], 10) : 0 +} + +export interface VideoController { + name: string + pnpDeviceId: string +} + +/** + * Parses the JSON emitted by `Get-CimInstance Win32_VideoController | ConvertTo-Json`. + * PowerShell emits a bare object rather than an array when there is one adapter. + */ +export function parseWindowsVideoControllers(stdout: string): VideoController[] { + let parsed: unknown + try { + parsed = JSON.parse(stdout) + } catch { + return [] + } + const list = Array.isArray(parsed) ? parsed : [parsed] + return list + .filter((entry): entry is Record => !!entry && typeof entry === 'object') + .map((entry) => ({ + name: typeof entry['Name'] === 'string' ? entry['Name'] : '', + pnpDeviceId: typeof entry['PNPDeviceID'] === 'string' ? entry['PNPDeviceID'] : '', + })) + .filter((c) => c.pnpDeviceId !== '') +} + +/** Extracts the PCI device id from a PNPDeviceID, e.g. `PCI\VEN_1002&DEV_7590&…` → "7590". */ +export function parseAmdPciDeviceId(pnpDeviceId: string): string | null { + const match = /VEN_1002&DEV_([0-9A-F]{4})/i.exec(pnpDeviceId) + return match ? match[1].toLowerCase() : null +} + +/** + * Resolves a compute target from the installed display adapters. Returns the + * AMD adapters it saw as well, so an unmapped card can be reported by name + * instead of silently falling back to CPU. + */ +export function resolveWindowsGfxTarget( + controllers: VideoController[], +): { gfxTarget: string | null; amdAdapters: string[] } { + const amdAdapters: string[] = [] + let gfxTarget: string | null = null + + for (const controller of controllers) { + const deviceId = parseAmdPciDeviceId(controller.pnpDeviceId) + if (!deviceId) continue + amdAdapters.push(controller.name || `PCI 1002:${deviceId}`) + gfxTarget ??= WINDOWS_PCI_GFX_TARGETS[deviceId] ?? null + } + + return { gfxTarget, amdAdapters } +} + +/** + * The pip index and requirements an extension's torch install has to end up + * using on this platform. `MODLY_ROCM_INDEX` and `MODLY_ROCM_TORCH_SPEC` + * (space-separated requirements) override either half — the escape hatch when a + * newer torch breaks an extension and you need to drop back to, say, rocm6.4. + */ +export function resolveRocmTorchSpec( + platform: string, + gfxTarget: string | null, + env: NodeJS.ProcessEnv = process.env, +): { indexUrl: string; specs: string[] } { + const indexOverride = env['MODLY_ROCM_INDEX']?.trim() + const specOverride = env['MODLY_ROCM_TORCH_SPEC']?.trim() + + const isWindows = platform === 'win32' + const indexUrl = indexOverride || (isWindows ? ROCM_WINDOWS_INDEX : ROCM_LINUX_INDEX) + + if (specOverride) return { indexUrl, specs: specOverride.split(/\s+/).filter(Boolean) } + + // AMD's multi-arch index ships one torch per compute target, selected by a + // pip extra. The pytorch.org ROCm index bakes the targets into a single + // wheel, so there is nothing to select and nothing to pin. + if (isWindows && gfxTarget) { + return { + indexUrl, + specs: [ + `torch[device-${gfxTarget}]==${ROCM_WINDOWS_TORCH}`, + `torchvision[device-${gfxTarget}]==${ROCM_WINDOWS_TORCHVISION}`, + ], + } + } + return { indexUrl, specs: ['torch', 'torchvision'] } +} + +/** The `torch_flavor` value extension setup.py scripts branch on. */ +export function torchFlavorFor(accelerator: Accelerator): TorchFlavor { + if (accelerator === 'rocm') return 'rocm' + if (accelerator === 'cuda') return 'cuda' + return 'cpu' +} + +export function describeGpuInfo(info: GpuInfo): string { + const bits = [`accelerator=${info.accelerator}`] + if (info.accelerator === 'cuda') bits.push(`sm=${info.sm}`, `cuda=${info.cudaVersion}`) + if (info.gfxTarget) bits.push(`gfx=${info.gfxTarget}`) + if (info.torchIndexUrl) bits.push(`torch_index=${info.torchIndexUrl}`) + return bits.join(' ') +} + +// ─── Detection ──────────────────────────────────────────────────────────────── + +function cpuInfo(): GpuInfo { + return { sm: 0, cudaVersion: 0, accelerator: 'cpu' } +} + +/** + * AMD keeps sm/cudaVersion at 0 on purpose. Extensions that predate `torch_flavor` + * branch on those two numbers, and 0 sends them down their most conservative + * path — which also keeps them off `rembg[gpu]`, whose onnxruntime-gpu is + * CUDA-only. Their torch install is then corrected by the ROCm setup shim. + */ +function rocmInfo( + gfxTarget: string | null, + platform: string, + env: NodeJS.ProcessEnv, +): GpuInfo { + const { indexUrl, specs } = resolveRocmTorchSpec(platform, gfxTarget, env) + return { + sm: 0, + cudaVersion: 0, + accelerator: 'rocm', + ...(gfxTarget ? { gfxTarget } : {}), + torchIndexUrl: indexUrl, + torchSpecs: specs, + } +} + +function readFlavorOverride(env: NodeJS.ProcessEnv): TorchFlavor | null { + const raw = env['MODLY_TORCH_FLAVOR']?.trim().toLowerCase() + return raw === 'cuda' || raw === 'rocm' || raw === 'cpu' ? raw : null +} + +function queryNvidiaSmi(): Promise<{ sm: number; cudaVersion: number } | null> { + return new Promise((resolve) => { + // Query compute cap + driver version in one call + const proc = spawn('nvidia-smi', ['--query-gpu=compute_cap,driver_version', '--format=csv,noheader'], { + stdio: ['ignore', 'pipe', 'ignore'], + }) + let out = '' + proc.stdout?.on('data', (d: Buffer) => { out += d.toString() }) + proc.on('close', (code) => resolve(code === 0 ? parseNvidiaSmi(out) : null)) + proc.on('error', () => resolve(null)) + }) +} + +/** + * Reads the compute target straight out of the kernel's KFD topology. This + * needs no ROCm installation and no external binary — the amdgpu driver alone + * publishes it, which is exactly the state of a machine that has only ever run + * PyTorch ROCm wheels (they bundle their own runtime). + */ +function readKfdGfxTarget(): string | null { + if (!existsSync('/dev/kfd')) return null + try { + const nodes = readdirSync(KFD_TOPOLOGY_DIR).sort((a, b) => Number(a) - Number(b)) + const properties = nodes.map((node) => { + try { + return readFileSync(join(KFD_TOPOLOGY_DIR, node, 'properties'), 'utf-8') + } catch { + return '' + } + }) + return parseKfdGfxTarget(properties) + } catch { + return null + } +} + +function queryWindowsVideoControllers(): Promise { + return new Promise((resolve) => { + const proc = spawn('powershell', [ + '-NoProfile', '-NonInteractive', '-Command', + 'Get-CimInstance Win32_VideoController | Select-Object Name,PNPDeviceID | ConvertTo-Json -Compress', + ], { stdio: ['ignore', 'pipe', 'ignore'] }) + let out = '' + proc.stdout?.on('data', (d: Buffer) => { out += d.toString() }) + proc.on('close', (code) => resolve(code === 0 ? parseWindowsVideoControllers(out) : [])) + proc.on('error', () => resolve([])) + }) +} + +export interface DetectOptions { + env?: NodeJS.ProcessEnv + platform?: string + arch?: string + onLog?: (line: string) => void +} + +export async function detectGpuInfo(options: DetectOptions = {}): Promise { + const env = options.env ?? process.env + const platform = options.platform ?? process.platform + const arch = options.arch ?? process.arch + const log = options.onLog ?? (() => {}) + + const forced = readFlavorOverride(env) + if (forced) log(`[gpu-detect] MODLY_TORCH_FLAVOR=${forced} — skipping auto-detection`) + + if (forced === 'cpu') return cpuInfo() + + if (forced === 'rocm') { + const { gfxTarget } = await resolveGfxTarget(env, platform) + if (!gfxTarget && platform === 'win32') { + log('[gpu-detect] ROCm forced on Windows but no compute target found — set MODLY_ROCM_GFX (e.g. gfx1200)') + return cpuInfo() + } + return rocmInfo(gfxTarget, platform, env) + } + + // A forced flavour wins over every probe, the Apple-Silicon MPS default + // included; nvidia-smi still runs so the CUDA path learns its real + // sm/cudaVersion instead of the conservative zeros. + if (forced === 'cuda') { + const nvidia = await queryNvidiaSmi() + return nvidia ? { ...nvidia, accelerator: 'cuda' } : { sm: 0, cudaVersion: 0, accelerator: 'cuda' } + } + + if (platform === 'darwin' && arch === 'arm64') { + return { sm: 0, cudaVersion: 0, accelerator: 'mps' } + } + + // NVIDIA keeps priority: on a machine with both, nothing about the existing + // CUDA behaviour changes. + const nvidia = await queryNvidiaSmi() + if (nvidia) return { ...nvidia, accelerator: 'cuda' } + + const { gfxTarget, amdAdapters } = await resolveGfxTarget(env, platform) + if (gfxTarget) { + log(`[gpu-detect] AMD GPU detected — compute target ${gfxTarget}`) + return rocmInfo(gfxTarget, platform, env) + } + + if (amdAdapters.length > 0) { + log( + `[gpu-detect] AMD GPU found (${amdAdapters.join(', ')}) but its ROCm compute target is unknown. ` + + 'Falling back to CPU — set MODLY_ROCM_GFX (e.g. gfx1201) to force one.', + ) + } + + return cpuInfo() +} + +/** + * Also hands back the AMD adapters the Windows probe saw, so the caller can + * report an unmapped card without spawning a second PowerShell round-trip. + */ +async function resolveGfxTarget( + env: NodeJS.ProcessEnv, + platform: string, +): Promise<{ gfxTarget: string | null; amdAdapters: string[] }> { + const override = env['MODLY_ROCM_GFX']?.trim() + if (override) return { gfxTarget: override, amdAdapters: [] } + if (platform === 'linux') return { gfxTarget: readKfdGfxTarget(), amdAdapters: [] } + if (platform === 'win32') return resolveWindowsGfxTarget(await queryWindowsVideoControllers()) + return { gfxTarget: null, amdAdapters: [] } +} diff --git a/electron/main/hf-token.ts b/electron/main/hf-token.ts new file mode 100644 index 00000000..2ea9061a --- /dev/null +++ b/electron/main/hf-token.ts @@ -0,0 +1,61 @@ +import { getSettings, setSettings } from './settings-store' +import { encryptSecretSync, decryptSecretSync } from './secure-store' +import { logger } from './logger' + +/** + * Hugging Face token, encrypted at rest in settings.json the same way the + * provider API keys are (see secure-store.ts). It used to sit there in plain + * text while every other secret in the app was encrypted. + * + * The plaintext lives only in this module's cache: `getHfToken()` is sync + * because the callers that need it are sync (child-process env building), while + * encrypt/decrypt are async. + */ +let cached = '' + +export function getHfToken(): string { + return cached +} + +/** + * Decrypt the stored token into the cache, re-encrypting a legacy plaintext one. + * Synchronous on purpose: the FastAPI bridge reads the token while building its + * child-process env at startup, so an async init could lose a race with it. + */ +export function initHfToken(userData: string): string { + const stored = getSettings(userData).hfToken ?? '' + if (!stored) { + cached = '' + return cached + } + + const decrypted = decryptSecretSync(stored) + if (decrypted === null) { + // One of our blobs, but not decryptable here (different OS user/machine). + // Leave settings.json alone — rewriting would turn an unreadable-but-intact + // blob into a permanently lost one — and expose no token rather than + // handing the ciphertext out as a credential. + cached = '' + logger.error('[hf-token] the stored Hugging Face token could not be decrypted — re-enter it in Settings → Integrations') + return cached + } + + cached = decrypted + // An unchanged value means it was never encrypted (saved before encryption + // existed) — upgrade it in place. + if (cached === stored) { + try { + setSettings(userData, { hfToken: encryptSecretSync(cached) }) + logger.info('[hf-token] migrated a plaintext token to encrypted storage') + } catch (err) { + logger.error(`[hf-token] could not migrate the stored token: ${err}`) + } + } + return cached +} + +/** Persist the token encrypted and refresh the cache. */ +export function setHfToken(userData: string, token: string): void { + cached = token + setSettings(userData, { hfToken: token ? encryptSecretSync(token) : '' }) +} diff --git a/electron/main/index.ts b/electron/main/index.ts index 0cbd5e5a..2f3f5b83 100644 --- a/electron/main/index.ts +++ b/electron/main/index.ts @@ -33,7 +33,8 @@ function createWindow(): void { preload: join(__dirname, '../preload/index.js'), sandbox: false, contextIsolation: true, - nodeIntegration: false + nodeIntegration: false, + backgroundThrottling: false } }) diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 005f1f78..e9cac040 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -1,8 +1,8 @@ import { ipcMain, BrowserWindow, Notification, dialog, app, shell } from 'electron' import { buildSync } from 'esbuild' import { autoUpdater } from 'electron-updater' -import { join } from 'path' -import { rm as rmAsync, readFile, writeFile, mkdir, readdir, rename, cp, symlink, lstat } from 'fs/promises' +import { basename, join } from 'path' +import { rm as rmAsync, readFile, writeFile, mkdir, readdir, rename, cp, symlink, lstat, copyFile } from 'fs/promises' import { existsSync, mkdirSync, readdirSync, statSync } from 'fs' import axios from 'axios' import * as tar from 'tar' @@ -13,7 +13,34 @@ import { isModelDownloaded, listDownloadedModels, downloadModelFromHF, + downloadModelSourcesFromHF, + type DownloadProgress, } from './model-downloader' +import { + legacyDownloadSteps, + resolveInstalledExtensionSharedWeightGroups, + resolveInstalledModelDownloadPlan, +} from './model-download-plan' +import { + areModelSourcesDownloaded, + areModelSourcesDownloadedAtRoot, + areWeightGroupSourcesDownloaded, + installedWeightVariants, + listWeightVariantFiles, + modelHasLocalData, + normalizeModelSources, + normalizeWeightGroupReferences, + normalizeWeightGroups, + normalizeWeightVariants, + validateModelNodeIds, + 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' import { logger } from './logger' @@ -30,6 +57,8 @@ import { isInternalExtensionDirName, resolveExtensionPathWithinRoot, } from './extension-path-guard' +import { detectGpuInfo, describeGpuInfo, torchFlavorFor, type GpuInfo } from './gpu-detect' +import { SETUP_LAUNCHER_SOURCE } from './setup-launcher' import { assertCompatibleExtensionUpdateType, expectedModelIds, @@ -38,6 +67,7 @@ import { validateExtensionReloadPayload, validateExistingExtensionReplacement, validateInstallManifest, + assertSupportedSceneNodeShape, } from './extension-install-utils' import { beginExtensionRegistrationTransaction, @@ -56,65 +86,20 @@ import { } from './extension-install-recovery' import { registerWorkspaceAssetLibraryIpcHandlers } from './artifact-registry-service' import { updatesSupported } from './updater' +import { ModelWeightOperations } from './model-weight-operations' +import { readLocalFileBase64 } from './bounded-file-reader' +import { encryptSecret, decryptSecret } from './secure-store' +import { getHfToken, initHfToken, setHfToken } from './hf-token' type WindowGetter = () => BrowserWindow | null const pExecFile = promisify(execFile) -// ─── GPU detect (best-effort, no Python required) ───────────────────────────── - -interface GpuInfo { - sm: number - cudaVersion: number - accelerator: 'cuda' | 'mps' | 'cpu' -} - -function detectGpuInfo(): Promise { - if (process.platform === 'darwin' && process.arch === 'arm64') { - return Promise.resolve({ sm: 0, cudaVersion: 0, accelerator: 'mps' }) - } - - return new Promise((resolve) => { - // Query compute cap + driver version in one call - const proc = spawn('nvidia-smi', ['--query-gpu=compute_cap,driver_version', '--format=csv,noheader'], { - stdio: ['ignore', 'pipe', 'ignore'], - }) - let out = '' - proc.stdout?.on('data', (d: Buffer) => { out += d.toString() }) - proc.on('close', (code) => { - if (code === 0) { - const line = out.trim().split('\n')[0].trim() // e.g. "8.6, 551.61" - const parts = line.split(',').map(s => s.trim()) - const sm = Math.round(parseFloat(parts[0] ?? '') * 10) // → 86 - // Derive max supported CUDA version from driver version - // Driver ≥ 520 → CUDA 11.8, ≥ 525 → 12.0, ≥ 530 → 12.1, ≥ 535 → 12.2, - // ≥ 545 → 12.3, ≥ 550 → 12.4, ≥ 555 → 12.5, ≥ 560 → 12.6 - const driverMajor = parseInt((parts[1] ?? '').split('.')[0] ?? '0', 10) - let cudaVersion = 118 // safe minimum - if (driverMajor >= 570) cudaVersion = 128 // Blackwell (RTX 50xx, sm_120) - else if (driverMajor >= 560) cudaVersion = 126 - else if (driverMajor >= 555) cudaVersion = 125 - else if (driverMajor >= 550) cudaVersion = 124 - else if (driverMajor >= 545) cudaVersion = 123 - else if (driverMajor >= 535) cudaVersion = 122 - else if (driverMajor >= 530) cudaVersion = 121 - else if (driverMajor >= 525) cudaVersion = 120 - else if (driverMajor >= 520) cudaVersion = 118 - resolve({ sm: isNaN(sm) ? 86 : sm, cudaVersion, accelerator: 'cuda' }) - } else { - resolve({ sm: 0, cudaVersion: 0, accelerator: 'cpu' }) - } - }) - proc.on('error', () => resolve({ sm: 0, cudaVersion: 0, accelerator: 'cpu' })) - }) -} - // ─── Run an extension's setup.py directly (no FastAPI needed) ───────────────── function runExtensionSetup( - extDir: string, - gpuSm: number, - cudaVersion: number, - onLog?: (line: string) => void, + extDir: string, + gpu: GpuInfo, + onLog?: (line: string) => void, ): Promise { return new Promise((resolve, reject) => { const userData = app.getPath('userData') @@ -128,113 +113,34 @@ function runExtensionSetup( const pipCacheDir = join(getSettings(userData).dependenciesDir, 'pip-cache') try { mkdirSync(pipCacheDir, { recursive: true }) } catch { /* pip creates it too */ } - const accelerator = process.platform === 'darwin' && process.arch === 'arm64' ? 'mps' : gpuSm > 0 ? 'cuda' : 'cpu' + const torchFlavor = torchFlavorFor(gpu.accelerator) const args = JSON.stringify({ python_exe: pythonExe, ext_dir: extDir, - gpu_sm: gpuSm, - cuda_version: cudaVersion, - accelerator, + gpu_sm: gpu.sm, + cuda_version: gpu.cudaVersion, + accelerator: gpu.accelerator, + // Extensions that know about AMD branch on torch_flavor (the official + // hunyuan3d-mini one does). Those that don't get corrected by the ROCm + // shim in setup-launcher.ts instead. + torch_flavor: torchFlavor, + gfx_target: gpu.gfxTarget ?? '', + torch_index_url: gpu.torchIndexUrl ?? '', platform: process.platform, arch: process.arch, }) - const launcher = ` -import runpy -import subprocess -import sys - -setup_py = sys.argv[1] -setup_args = sys.argv[2:] - -_original_run = subprocess.run -_original_check_call = subprocess.check_call -_original_check_output = subprocess.check_output - -def _is_cuda_torch_index(value): - return isinstance(value, str) and value.startswith("https://download.pytorch.org/whl/cu") - -def _mentions_torch(command): - if not isinstance(command, (list, tuple)): - return False - return any(str(part).startswith(("torch==", "torchvision==", "torchaudio==")) for part in command) - -def _rewrite_command(command): - if sys.platform != "darwin" or not _mentions_torch(command): - return command - if not isinstance(command, (list, tuple)): - return command - - rewritten = [] - changed = False - i = 0 - while i < len(command): - part = command[i] - text = str(part) - if text in ("--index-url", "-i", "--extra-index-url") and i + 1 < len(command) and _is_cuda_torch_index(str(command[i + 1])): - changed = True - i += 2 - continue - if text.startswith("--index-url=") or text.startswith("--extra-index-url="): - value = text.split("=", 1)[1] - if _is_cuda_torch_index(value): - changed = True - i += 1 - continue - rewritten.append(part) - i += 1 - - if changed: - print("[Modly setup compat] Removed CUDA-only PyTorch index on macOS; pip will use macOS wheels.", file=sys.stderr) - return rewritten - return command - -def _is_pip_command(command): - if not isinstance(command, (list, tuple)): - return False - return any("pip" in str(part).lower() for part in command[:3]) - -def _strip_no_cache(command): - # Extension setup scripts often hardcode --no-cache-dir, which forces pip to - # re-download multi-GB wheels on every retry. Modly provides a shared cache - # via PIP_CACHE_DIR, so drop the flag and let pip use it. - if not _is_pip_command(command): - return command - if not any(str(part) == "--no-cache-dir" for part in command): - return command - print("[Modly setup compat] Removed --no-cache-dir so pip reuses the shared wheel cache.", file=sys.stderr) - return [part for part in command if str(part) != "--no-cache-dir"] - -def _transform_command(command): - return _strip_no_cache(_rewrite_command(command)) - -def _patched_run(*args, **kwargs): - args = list(args) - if args: - args[0] = _transform_command(args[0]) - return _original_run(*args, **kwargs) - -def _patched_check_call(*args, **kwargs): - args = list(args) - if args: - args[0] = _transform_command(args[0]) - return _original_check_call(*args, **kwargs) - -def _patched_check_output(*args, **kwargs): - args = list(args) - if args: - args[0] = _transform_command(args[0]) - return _original_check_output(*args, **kwargs) - -subprocess.run = _patched_run -subprocess.check_call = _patched_check_call -subprocess.check_output = _patched_check_output - -sys.argv = [setup_py] + setup_args -runpy.run_path(setup_py, run_name="__main__") -` + const launcher = SETUP_LAUNCHER_SOURCE + // The rewrite decision itself is made in gpu-detect.ts (and unit-tested + // there); the launcher above only applies what these carry. const proc = spawn(pythonExe, ['-c', launcher, setupPy, args], { stdio: ['ignore', 'pipe', 'pipe'], - env: { ...process.env, PIP_CACHE_DIR: pipCacheDir }, + env: { + ...process.env, + PIP_CACHE_DIR: pipCacheDir, + MODLY_TORCH_FLAVOR: torchFlavor, + MODLY_TORCH_INDEX_URL: gpu.torchIndexUrl ?? '', + MODLY_TORCH_SPECS: JSON.stringify(gpu.torchSpecs ?? []), + }, }) const handleLine = (line: string) => { if (line) onLog?.(line) } @@ -265,7 +171,60 @@ const renameWithRetry = (from: string, to: string, label: string) => renameExtensionWithRetry(from, to, label, logger) export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGetter): void { - const activeDownloads = new Map() + type ActiveDownload = { + progress: DownloadProgress & { variantId?: string } + done: Promise + finish: () => void + targetRoots: string[] + currentTargetId?: string + stopRequested?: 'pause' | 'cancel' + } + const activeDownloads = 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') + } + // No backend listening means no Python process can hold the weight files + // open, so deletion is safe. A backend that answers (or times out) is not. + const backendUnreachable = (err: unknown): boolean => { + if (!axios.isAxiosError(err) || err.response) return false + const cause = (err as { cause?: { code?: string } }).cause + return err.code === 'ECONNREFUSED' || cause?.code === 'ECONNREFUSED' + } + async function unloadForRemoval(modelIds: string[]) { + for (const id of modelIds) { + try { + const response = await axios.post( + `${API_BASE_URL}/model/unload/${encodeURIComponent(id)}`, {}, { timeout: 40_000 }, + ) + if (response.data?.unloaded !== true) throw new Error('Model unload was not confirmed; weights were preserved') + } catch (err) { + if (backendUnreachable(err)) return + throw err + } + } + } + // Unload only this extension's generators, not every model in the app. + async function unloadExtensionForRemoval(extensionId: string) { + let modelIds: string[] + try { + const { data } = await axios.get<{ id: string }[]>(`${API_BASE_URL}/model/all`, { timeout: 10_000 }) + modelIds = data.map((model) => model.id).filter((id) => id.startsWith(`${extensionId}/`)) + } catch (err) { + if (backendUnreachable(err)) return + throw err + } + await unloadForRemoval(modelIds) + } + const resolveModelPlan = (modelId: unknown) => resolveInstalledModelDownloadPlan({ + modelId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + const LOCKED_MODEL_FILES_ERROR = 'Model files are still locked after several attempts. Close any programs using the model and try again.' // 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')) @@ -297,6 +256,11 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe } }) + // Secure storage — OS-level encryption (Keychain/DPAPI/libsecret) for secrets + // the renderer would otherwise have to keep in plain-text localStorage (API keys). + ipcMain.handle('secure:encrypt', (_, plainText: string) => encryptSecret(plainText)) + ipcMain.handle('secure:decrypt', (_, stored: string) => decryptSecret(stored)) + // Window controls (frameless window) ipcMain.on('window:minimize', () => getWindow()?.minimize()) ipcMain.on('window:maximize', () => { @@ -350,6 +314,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe workflowsDir: join(baseDir, 'workflows'), extensionsDir: join(baseDir, 'extensions'), dependenciesDir: join(baseDir, 'dependencies'), + agentDir: join(baseDir, 'agent'), }) }) @@ -446,25 +411,59 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }) ipcMain.handle('model:delete', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { - const modelDir = join(getSettings(app.getPath('userData')).modelsDir, modelId) + if (activeDownloads.has(modelId)) { + return { success: false, error: 'Cannot remove model weights while their download is active' } + } + let modelDir: string + try { + await resolveModelPlan(modelId) + modelDir = resolveModelRoot(getSettings(app.getPath('userData')).modelsDir, modelId) + } catch (err) { + return { success: false, error: String(err) } + } - // Unload the model and wait for confirmation so file handles are released try { - await axios.post(`${API_BASE_URL}/model/unload/${encodeURIComponent(modelId)}`, {}, { timeout: 10_000 }) - // Give the OS a moment to release file locks (Windows holds handles briefly after close) - await new Promise(resolve => setTimeout(resolve, 1_500)) - } catch { - // Unload failed (model may not be loaded) — still attempt deletion + 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) } } + }) - // 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), + ipcMain.handle('model:deleteWeightVariant', async (_, modelId: string, variantId: string): Promise<{ success: boolean; error?: string }> => { + if (activeDownloads.has(modelId)) { + return { success: false, error: 'Cannot remove model weights while their download is active' } + } + try { + const plan = await resolveModelPlan(modelId) + const variant = plan.kind === 'legacy' + ? plan.weightVariants?.options.find((option) => option.id === variantId) + : undefined + if (!variant) throw new Error(`Model node "${modelId}" has no weight variant "${String(variantId)}"`) + const modelsDir = getSettings(app.getPath('userData')).modelsDir + // Reserve the node root (blocks a concurrent download of any of its variants), + // unload with confirmation, then list and remove only this variant's files. + const removed = await weightOperations.remove( + [resolveModelRoot(modelsDir, modelId)], + () => unloadForRemoval([modelId]), + async (): Promise>> => { + for (const file of await listWeightVariantFiles(modelsDir, modelId, variant)) { + const result = await rmWithRetry(file, 'model-variant-delete') + if (!result.ok) return result + } + return { ok: true } + }, + ) + notifyWeightChange() + if (removed.ok) return { success: true } + return { success: false, error: removed.locked ? LOCKED_MODEL_FILES_ERROR : String(removed.error) } + } catch (err) { + return { success: false, error: String(err) } } }) @@ -476,13 +475,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }) // Read local file → base64 (bypasses file:// restrictions in the renderer) - ipcMain.handle('fs:readFileBase64', async (_, filePath: string) => { - if (typeof filePath !== 'string' || filePath.trim().length === 0) { - throw new Error('fs:readFileBase64 requires a non-empty file path') - } - const buffer = await readFile(filePath) - return buffer.toString('base64') - }) + ipcMain.handle('fs:readFileBase64', (_, filePath: string) => readLocalFileBase64(filePath)) ipcMain.handle('fs:readScreenshotDataUrl', async (_, filename: string) => { const filePath = app.isPackaged @@ -498,49 +491,273 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe return listDownloadedModels(modelsDir) }) - ipcMain.handle('model:isDownloaded', (_, modelId: string, downloadCheck?: string): boolean => { + ipcMain.handle('model:isDownloaded', async (_, modelId: string): Promise => { const modelsDir = getSettings(app.getPath('userData')).modelsDir - return isModelDownloaded(modelsDir, modelId, downloadCheck) + try { + const plan = await resolveInstalledModelDownloadPlan({ + modelId, + userExtensionsDir: getSettings(app.getPath('userData')).extensionsDir, + builtinExtensionsDir: getBuiltinExtensionsDir(), + blockedExtensionIds: activeExtensionInstalls, + }) + 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) + && (!plan.weightVariants || installedWeightVariants(modelsDir, modelId, plan.weightVariants).length > 0) + } catch { + return false + } + }) + + // null means "unknown" (unreadable plan, or a node without variants) — the renderer + // must not read an empty array as "no variant installed". + ipcMain.handle('model:installedWeightVariants', async (_, modelId: string): Promise => { + try { + const plan = await resolveModelPlan(modelId) + return plan.kind === 'legacy' && plan.weightVariants + ? installedWeightVariants(getSettings(app.getPath('userData')).modelsDir, modelId, plan.weightVariants) + : null + } catch { + return null + } + }) + + ipcMain.handle('model:hasLocalData', async (_, modelId: string): Promise => { + try { + await resolveModelPlan(modelId) + return modelHasLocalData(getSettings(app.getPath('userData')).modelsDir, modelId) + } catch { + return false + } + }) + + 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, + ) + const removed = await weightOperations.remove( + [groupRoot], + () => unloadForRemoval(group.dependentModelIds), + () => rmWithRetry(groupRoot, 'shared-model-delete'), + ) + notifyWeightChange() + 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, + ) + const removed = await weightOperations.remove( + [extensionRoot], + () => unloadExtensionForRemoval(safeExtensionId), + () => rmWithRetry(extensionRoot, 'extension-model-delete'), + ) + notifyWeightChange() + 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, progress]) => ({ modelId, ...progress })) + [...activeDownloads.entries()].map(([modelId, active]) => ({ modelId, ...active.progress })) ) ipcMain.handle('model:download', async ( event, - { repoId, modelId, skipPrefixes, includePrefixes }: { repoId: string; modelId: string; skipPrefixes?: string[]; includePrefixes?: string[] }, + modelId: string, + requestedVariantId?: string | null, ) => { if (activeDownloads.has(modelId)) { return { success: false, error: 'Download already in progress' } } - activeDownloads.set(modelId, { percent: 0 }) + const variantId = requestedVariantId ?? undefined + let plan: Awaited> + let legacySteps: ReturnType = [] try { - await downloadModelFromHF(repoId, modelId, (progress) => { - activeDownloads.set(modelId, progress) - event.sender.send('model:downloadProgress', { modelId, ...progress }) - }, skipPrefixes, includePrefixes) + plan = await resolveModelPlan(modelId) + if (plan.kind === 'multi-source') { + if (variantId !== undefined) throw new Error(`Model node "${modelId}" does not declare weight variants`) + } else { + // Shared files first (every variant excluded), then the requested variant. + legacySteps = legacyDownloadSteps(plan, variantId) + } + } 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 allTargets = 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, + }] : []), + ] + : [] + const targetRoots = plan.kind === 'multi-source' + ? allTargets.map((target) => resolveWeightStorageRoot(modelsDir, target.targetId)) + : [resolveModelRoot(modelsDir, modelId)] + 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, variantId }, done, finish, targetRoots } + activeDownloads.set(modelId, active) + interruptedTargets.set(modelId, targetRoots) + try { + const onProgress = (progress: DownloadProgress) => { + active.progress = { ...progress, variantId } + event.sender.send('model:downloadProgress', { modelId, variantId, ...progress }) + } + if (plan.kind === 'multi-source') { + if (managedTargets.length === 0) { + 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( + 99, + Math.round(((index + progress.percent / 100) / managedTargets.length) * 100), + ) + onProgress({ + ...progress, + percent: aggregatePercent, + 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 + // Shared files and the variant are separate passes sharing one 0-100 bar. + for (const [index, step] of legacySteps.entries()) { + if (active.stopRequested) throw new Error(`Model download ${active.stopRequested === 'pause' ? 'paused' : 'cancelled'}`) + await downloadModelFromHF( + plan.repoId, + modelId, + (progress) => onProgress({ ...progress, percent: Math.round((index * 100 + progress.percent) / legacySteps.length) }), + step.skipPrefixes, + step.includePrefixes, + ) + } + } + interruptedTargets.delete(modelId) return { success: true } } catch (err: any) { const message = err?.message ?? String(err) if (message.includes('paused')) { - event.sender.send('model:downloadProgress', { modelId, percent: 0, status: 'paused', paused: true }) + event.sender.send('model:downloadProgress', { modelId, variantId, percent: 0, status: 'paused', paused: true }) return { success: false, paused: true } } if (message.includes('cancelled')) { - event.sender.send('model:downloadProgress', { modelId, percent: 0, status: 'cancelled', cancelled: true }) + event.sender.send('model:downloadProgress', { modelId, variantId, percent: 0, status: 'cancelled', cancelled: true }) return { success: false, cancelled: true } } return { success: false, error: String(err) } } finally { - activeDownloads.delete(modelId) + if (activeDownloads.get(modelId) === active) activeDownloads.delete(modelId) + release() + active.finish() + notifyWeightChange() } }) ipcMain.handle('model:pauseDownload', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { try { + 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: modelId }, + params: { model_id: targetId }, timeout: 5000, }) return { success: true } @@ -551,17 +768,42 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe ipcMain.handle('model:cancelDownload', async (_, modelId: string): Promise<{ success: boolean; error?: string }> => { try { - await axios.post(`${API_BASE_URL}/model/hf-download/cancel`, null, { - params: { model_id: modelId }, - timeout: 5000, - }) - const modelDir = join(getSettings(app.getPath('userData')).modelsDir, modelId) - await rmAsync(modelDir, { recursive: true, force: true }) + 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 }, + timeout: 5000, + }) + } + if (active) { + 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() + } return { success: true } } catch (err) { return { success: false, error: String(err) } - } finally { - activeDownloads.delete(modelId) } }) @@ -595,6 +837,26 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // Shell ipcMain.handle('shell:openExternal', (_, url: string) => shell.openExternal(url)) + // Open a model in OrcaSlicer via its orcaslicer://open?file= deeplink. + // + // The returned error only covers the shell refusing the call outright. It is + // NOT an install check: on Windows an unregistered scheme still makes + // ShellExecuteEx succeed — the OS shows its own "You'll need a new app to open + // this orcaslicer link" dialog and this resolves with success. Detecting a + // missing OrcaSlicer would take a per-platform handler probe (registry on + // Windows), so the renderer must not promise the user that it knows. + ipcMain.handle('slicer:open', async (_, url: string): Promise<{ success: boolean; error?: string }> => { + if (typeof url !== 'string' || !url.startsWith('orcaslicer://')) { + return { success: false, error: 'slicer:open requires an orcaslicer:// URL' } + } + try { + await shell.openExternal(url) + return { success: true } + } catch (err) { + return { success: false, error: err instanceof Error ? err.message : String(err) } + } + }) + // App info // System memory (used/available/total bytes). // On macOS, matches Activity Monitor's "Memory Used": @@ -638,33 +900,40 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe arch: process.arch, })) - // Settings — seed HF token into main-process env at startup + // Settings — decrypt the HF token (migrating a legacy plaintext one) and seed + // it into the main-process env at startup. { - const initialToken = getSettings(app.getPath('userData')).hfToken ?? '' - if (initialToken) { - process.env['HUGGING_FACE_HUB_TOKEN'] = initialToken - process.env['HF_TOKEN'] = initialToken + const token = initHfToken(app.getPath('userData')) + if (token) { + process.env['HUGGING_FACE_HUB_TOKEN'] = token + process.env['HF_TOKEN'] = token } } ipcMain.handle('settings:get', () => { - return getSettings(app.getPath('userData')) + // hfToken is stored encrypted — hand the renderer the usable value. + return { ...getSettings(app.getPath('userData')), hfToken: getHfToken() } }) ipcMain.handle('settings:set', async (_event, patch: { modelsDir?: string; workspaceDir?: string; extensionsDir?: string; hfToken?: string }) => { - const updated = setSettings(app.getPath('userData'), patch) + if (patch.modelsDir !== undefined && weightOperations.busy) { + throw new Error('Cannot change model storage while model weights are busy') + } + const { hfToken, ...dirs } = patch + setSettings(app.getPath('userData'), dirs) // Keep main-process env in sync so child processes spawned after token change inherit it - if (patch.hfToken !== undefined) { - process.env['HUGGING_FACE_HUB_TOKEN'] = patch.hfToken - process.env['HF_TOKEN'] = patch.hfToken + if (hfToken !== undefined) { + setHfToken(app.getPath('userData'), hfToken) + process.env['HUGGING_FACE_HUB_TOKEN'] = hfToken + process.env['HF_TOKEN'] = hfToken // Also push the token into the live FastAPI process env so extension // subprocesses spawned by ExtensionProcess._build_env() pick it up // without requiring a full app restart. try { - await axios.post(`${API_BASE_URL}/settings/hf-token`, { token: patch.hfToken }, { timeout: 3000 }) + await axios.post(`${API_BASE_URL}/settings/hf-token`, { token: hfToken }, { timeout: 3000 }) } catch { /* FastAPI may not be running yet — ignore */ } } - return updated + return { ...getSettings(app.getPath('userData')), hfToken: getHfToken() } }) // Directory picker @@ -786,6 +1055,40 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe return result.canceled ? null : result.filePaths[0] }) + // Add a local GGUF to the agent's models folder. Picking and copying both + // happen here, so the renderer never hands main an arbitrary path to copy. + ipcMain.handle('agent:addModel', async (): Promise<{ success: boolean; cancelled?: boolean; fileName?: string; error?: string }> => { + const win = getWindow() + if (!win) return { success: false, error: 'No window available' } + const result = await dialog.showOpenDialog(win, { + title: 'Add a model', + filters: [{ name: 'GGUF model', extensions: ['gguf'] }], + properties: ['openFile'], + }) + if (result.canceled || !result.filePaths[0]) return { success: false, cancelled: true } + + const src = result.filePaths[0] + const fileName = basename(src) + if (!fileName.toLowerCase().endsWith('.gguf')) return { success: false, error: 'Only .gguf files can be added.' } + + const modelsDir = join(getSettings(app.getPath('userData')).agentDir, 'models') + const dest = join(modelsDir, fileName) + if (existsSync(dest)) return { success: false, error: `"${fileName}" is already in your models.` } + + // Copied under a temporary name first: the API lists every *.gguf in the + // folder, so a multi-GB copy in progress would otherwise show up as a model. + const partial = `${dest}.part` + try { + await mkdir(modelsDir, { recursive: true }) + await copyFile(src, partial) + await rename(partial, dest) + return { success: true, fileName } + } catch (err) { + await rmAsync(partial, { force: true }).catch(() => {}) + return { success: false, error: String(err) } + } + }) + ipcMain.handle('fs:moveDirectory', async (_, { src, dest }: { src: string; dest: string }) => { try { await mkdir(dest, { recursive: true }) @@ -858,22 +1161,27 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // extension type 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 nodes?: { id: string name?: string - input?: 'mesh' | 'image' | 'text' | 'audio' - inputs?: ('mesh' | 'image' | 'text' | 'audio')[] + input?: string + inputs?: string[] input_labels?: string[] - output?: 'mesh' | 'image' | 'text' | 'audio' + output?: string params_schema?: unknown[] param_defaults?: Record hf_repo?: string download_check?: string hf_skip_prefixes?: string[] hf_include_prefixes?: string[] + model_sources?: unknown + weight_groups?: unknown + weight_variants?: unknown }[] } @@ -889,26 +1197,88 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe builtin, } - const nodes = (parsed.nodes ?? []).map(n => ({ - id: n.id, - name: n.name ?? n.id, - input: n.input ?? 'image' as const, - inputs: n.inputs, - inputLabels: n.input_labels, - output: n.output ?? 'mesh' as const, - paramsSchema: n.params_schema ?? parsed.params_schema ?? [], - paramDefaults: { ...(parsed.param_defaults ?? {}), ...(n.param_defaults ?? {}) }, - hfRepo: n.hf_repo, - downloadCheck: n.download_check, - hfSkipPrefixes: n.hf_skip_prefixes, - hfIncludePrefixes: n.hf_include_prefixes, - })) + if (parsed.model_sources !== undefined) { + throw new Error('manifest.json: model_sources must be declared on a model node') + } + 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) + if (weightGroups || parsed.nodes?.some((node) => node.model_sources !== undefined || node.weight_groups !== undefined)) { + validateModelNodeIds(parsed.nodes ?? []) + } + const nodes = (parsed.nodes ?? []).map(n => { + const declaredInputs = Array.isArray(n.inputs) ? n.inputs : [n.input ?? 'image'] + const output = n.output ?? 'mesh' + assertSupportedSceneNodeShape(parsed.type === 'process' ? 'process' : 'model', n, declaredInputs, output) + 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') + } + if (parsed.type === 'process' && n.weight_variants !== undefined) { + throw new Error('manifest.json: weight_variants is supported only for model nodes') + } + const modelSources = normalizeModelSources(n) + const weightVariants = normalizeWeightVariants(n, n.params_schema ?? parsed.params_schema) + 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: nodeId, + name: n.name ?? n.id, + input: n.input ?? 'image' as const, + inputs: n.inputs, + inputLabels: n.input_labels, + output: n.output ?? 'mesh' as const, + paramsSchema: n.params_schema ?? parsed.params_schema ?? [], + paramDefaults: { ...(parsed.param_defaults ?? {}), ...(n.param_defaults ?? {}) }, + hfRepo: n.hf_repo, + downloadCheck: n.download_check, + hfSkipPrefixes: n.hf_skip_prefixes, + hfIncludePrefixes: n.hf_include_prefixes, + hasModelSources: modelSources !== undefined, + weightGroups: groupRefs, + weightVariants: weightVariants && { + param: weightVariants.param, + default: weightVariants.default, + options: weightVariants.options.map((option) => ({ + id: option.id, + label: option.label, + sizeGb: option.size_gb, + vramGb: option.vram_gb, + })), + }, + } + }) if (parsed.type === 'process') { 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( @@ -1307,8 +1677,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // 7a. Python process extension: run setup.py if present (same as model extensions) if (existsSync(join(destDir, 'setup.py'))) { emit({ step: 'setting_up', message: 'Setting up Python environment…' }) - const { sm: gpuSm, cudaVersion } = await detectGpuInfo() - await runExtensionSetup(destDir, gpuSm, cudaVersion, (line) => { + const gpu = await detectGpuInfo({ onLog: (line) => logger.info(line) }) + logger.info(`[ext-setup] ${describeGpuInfo(gpu)}`) + await runExtensionSetup(destDir, gpu, (line) => { logger.info(`[ext-setup] ${line}`) emit({ step: 'setting_up', message: line }) }) @@ -1343,8 +1714,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // 7c. Model extension: run setup.py directly (no FastAPI required) if (existsSync(join(destDir, 'setup.py'))) { emit({ step: 'setting_up', message: 'Setting up Python environment…' }) - const { sm: gpuSm, cudaVersion } = await detectGpuInfo() - await runExtensionSetup(destDir, gpuSm, cudaVersion, (line) => { + const gpu = await detectGpuInfo({ onLog: (line) => logger.info(line) }) + logger.info(`[ext-setup] ${describeGpuInfo(gpu)}`) + await runExtensionSetup(destDir, gpu, (line) => { logger.info(`[ext-setup] ${line}`) emit({ step: 'setting_up', message: line }) }) @@ -1473,6 +1845,9 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe // Uninstall an extension — built-ins cannot be uninstalled ipcMain.handle('extensions:uninstall', async (_, extensionId: string) => { try { + if ([...activeDownloads.keys()].some((modelId) => modelId.split('/', 1)[0] === extensionId)) { + return { success: false, error: 'Cannot uninstall an extension while its model download is active' } + } // Corrupted folders can carry arbitrary names (manual copies, failed // unzips), so only enforce root confinement for the deletion path. The // strict id pattern still guards the built-in check — a non-conforming @@ -1533,7 +1908,8 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }, 'extension folder', ) - const { sm: gpuSm, cudaVersion } = await detectGpuInfo() + const gpu = await detectGpuInfo({ onLog: (line) => logger.info(line) }) + logger.info(`[ext-repair] ${describeGpuInfo(gpu)}`) await runExtensionRepairTransaction( { extensionsDir, @@ -1548,8 +1924,7 @@ export function setupIpcHandlers(pythonBridge: PythonBridge, getWindow: WindowGe }, setup: () => runExtensionSetup( extDir, - gpuSm, - cudaVersion, + gpu, (line) => logger.info(`[ext-repair] ${line}`), ), validate: async (validationCapability) => { @@ -1807,6 +2182,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.test.mjs b/electron/main/model-download-plan.test.mjs new file mode 100644 index 00000000..1a927cbc --- /dev/null +++ b/electron/main/model-download-plan.test.mjs @@ -0,0 +1,235 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, rmSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-model-plan-module-')), 'model-plan.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/model-download-plan.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +function setupExtension(manifest) { + const root = mkdtempSync(join(tmpdir(), 'modly-action-plan-')) + const user = join(root, 'user') + const builtin = join(root, 'builtin') + const extension = join(user, manifest.id) + mkdirSync(extension, { recursive: true }) + mkdirSync(builtin) + const manifestPath = join(extension, 'manifest.json') + writeFileSync(manifestPath, JSON.stringify(manifest)) + return { root, user, builtin, manifestPath } +} + +test('re-reads the installed manifest for each action and resolves only node-owned sources', async () => { + const { resolveInstalledModelDownloadPlan } = loadModule() + const manifest = { + id: 'pixal3d', + type: 'model', + nodes: [{ + id: 'generate', + model_sources: [{ + id: 'primary', provider: 'huggingface', repo_id: 'org/old', + destination: '.', checks: ['model.safetensors'], + }], + }], + } + const fixture = setupExtension(manifest) + const args = { + modelId: 'pixal3d/generate', + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + } + try { + const first = await resolveInstalledModelDownloadPlan(args) + assert.equal(first.kind, 'multi-source') + assert.equal(first.sources[0].repo_id, 'org/old') + + manifest.nodes[0].model_sources[0].repo_id = 'org/new' + writeFileSync(fixture.manifestPath, JSON.stringify(manifest)) + const second = await resolveInstalledModelDownloadPlan(args) + assert.equal(second.sources[0].repo_id, 'org/new') + } finally { + rmSync(fixture.root, { recursive: true, force: true }) + } +}) + +test('keeps legacy sibling checks and wildcard filters unchanged', async () => { + const { resolveInstalledModelDownloadPlan } = loadModule() + const fixture = setupExtension({ + id: 'triposplat', + type: 'model', + nodes: [{ + id: 'projection', + hf_repo: 'VAST-AI/TripoSplat', + download_check: '../generate/diffusion_models/triposplat_fp16.safetensors', + hf_skip_prefixes: ['weights/**', 'assets/*'], + }], + }) + try { + const plan = await resolveInstalledModelDownloadPlan({ + modelId: 'triposplat/projection', + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + }) + assert.equal(plan.kind, 'legacy') + assert.equal(plan.downloadCheck, '../generate/diffusion_models/triposplat_fp16.safetensors') + assert.deepEqual(plan.skipPrefixes, ['weights/**', 'assets/*']) + } finally { + rmSync(fixture.root, { recursive: true, force: true }) + } +}) + +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 }) + } + } +}) + +test('downloads a weight-variant node as shared files first, then one variant at a time', async () => { + const { legacyDownloadSteps, resolveInstalledModelDownloadPlan } = loadModule() + const fixture = setupExtension({ + id: 'trellis', + type: 'model', + nodes: [{ + id: 'generate', + hf_repo: 'org/model-gguf', + download_check: 'pipeline.json', + hf_include_prefixes: ['pipeline.json', 'dit/'], + hf_skip_prefixes: ['README.md'], + params_schema: [{ id: 'quant', type: 'select', options: [{ value: 'Q4' }, { value: 'Q5' }] }], + weight_variants: { + param: 'quant', + default: 'Q5', + options: ['Q4', 'Q5'].map((quant) => ({ + id: quant, + include_prefixes: [`dit/model_${quant}.gguf`], + checks: [`dit/model_${quant}.gguf`], + })), + }, + }], + }) + try { + const plan = await resolveInstalledModelDownloadPlan({ + modelId: 'trellis/generate', + userExtensionsDir: fixture.user, + builtinExtensionsDir: fixture.builtin, + }) + assert.equal(plan.kind, 'legacy') + assert.deepEqual(legacyDownloadSteps(plan), [ + { includePrefixes: ['pipeline.json', 'dit/'], skipPrefixes: ['README.md', 'dit/model_Q4.gguf', 'dit/model_Q5.gguf'] }, + { includePrefixes: ['dit/model_Q5.gguf'], skipPrefixes: ['README.md'] }, + ]) + // Picking one variant still fetches the shared files first, so a node installed + // variant-first is never left without its pipeline files. + assert.deepEqual(legacyDownloadSteps(plan, 'Q4'), [ + { includePrefixes: ['pipeline.json', 'dit/'], skipPrefixes: ['README.md', 'dit/model_Q4.gguf', 'dit/model_Q5.gguf'] }, + { includePrefixes: ['dit/model_Q4.gguf'], skipPrefixes: ['README.md'] }, + ]) + assert.throws(() => legacyDownloadSteps(plan, 'Q8'), /no weight variant "Q8"/) + + const withoutVariants = { ...plan, weightVariants: undefined } + assert.deepEqual(legacyDownloadSteps(withoutVariants), [ + { includePrefixes: ['pipeline.json', 'dit/'], skipPrefixes: ['README.md'] }, + ]) + assert.throws(() => legacyDownloadSteps(withoutVariants, 'Q4'), /no weight variant/) + } finally { + rmSync(fixture.root, { recursive: true, force: true }) + } +}) diff --git a/electron/main/model-download-plan.ts b/electron/main/model-download-plan.ts new file mode 100644 index 00000000..6dc78668 --- /dev/null +++ b/electron/main/model-download-plan.ts @@ -0,0 +1,301 @@ +import { existsSync } from 'node:fs' +import { readFile, readdir } from 'node:fs/promises' +import { join } from 'node:path' + +import { + EXT_INCOMPLETE_MARKER, + EXT_REGISTRATION_PENDING_MARKER, + assertSafeExtensionId, + resolveExtensionPathWithinRoot, +} from './extension-path-guard' +import { + normalizeModelSources, + normalizeWeightGroupReferences, + normalizeWeightGroups, + normalizeWeightVariants, + validateModelNodeIds, + safeModelSourceId, + weightGroupTargetId, + type ModelSource, + type ModelWeightGroup, + type WeightVariants, +} from './model-sources' + +interface InstalledNode { + id?: unknown + hf_repo?: unknown + download_check?: unknown + hf_skip_prefixes?: unknown + hf_include_prefixes?: unknown + model_sources?: unknown + weight_groups?: unknown + weight_variants?: unknown + params_schema?: unknown +} + +interface InstalledManifest { + id?: unknown + type?: unknown + model_sources?: unknown + weight_groups?: unknown + params_schema?: unknown + nodes?: unknown +} + +export interface InstalledSharedWeightGroup extends ModelWeightGroup { + targetId: string + dependentModelIds: string[] +} + +export type InstalledModelDownloadPlan = { + kind: 'legacy' + modelId: string + extensionId: string + nodeId: string + repoId: string + downloadCheck?: string + skipPrefixes?: string[] + includePrefixes?: string[] + weightVariants?: WeightVariants +} | { + kind: 'multi-source' + modelId: string + extensionId: string + nodeId: string + sources: ModelSource[] + sharedGroups: InstalledSharedWeightGroup[] +} + +async function hasPendingRegistration(root: string, extensionId: string): Promise { + try { + const prefix = `${EXT_REGISTRATION_PENDING_MARKER}-${extensionId}-` + return (await readdir(root)).some((name) => ( + name.startsWith(prefix) && /^\d+$/.test(name.slice(prefix.length)) + )) + } catch { + return false + } +} + +function installedSharedGroups( + manifest: InstalledManifest, + extensionId: string, + 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) { + 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 { + parsed = JSON.parse(raw) + } catch { + throw new Error(`Extension "${extensionId}" has an invalid manifest.json`) + } + if (typeof parsed !== 'object' || parsed === null || Array.isArray(parsed)) { + throw new Error(`Extension "${extensionId}" has an invalid manifest.json`) + } + const manifest = parsed as InstalledManifest + if (manifest.id !== extensionId) throw new Error(`Installed manifest id does not match extension "${extensionId}"`) + if (manifest.type !== undefined && manifest.type !== 'model') { + throw new Error(`Extension "${extensionId}" is not a model extension`) + } + if (manifest.model_sources !== undefined) { + throw new Error('manifest.json: model_sources must be declared on a model node') + } + if (!Array.isArray(manifest.nodes)) throw new Error(`Extension "${extensionId}" does not declare model nodes`) + + const 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`) + } + + const node = matches[0] + const modelId = `${extensionId}/${nodeId}` + const weightVariants = normalizeWeightVariants(node, node.params_schema ?? manifest.params_schema) + const sources = normalizeModelSources(node) + 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`) + } + return { + kind: 'legacy', + modelId, + extensionId, + nodeId, + repoId: node.hf_repo, + downloadCheck: typeof node.download_check === 'string' ? node.download_check : undefined, + skipPrefixes: node.hf_skip_prefixes as string[] | undefined, + includePrefixes: node.hf_include_prefixes as string[] | undefined, + ...(weightVariants ? { weightVariants } : {}), + } +} + + +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) +} + +export interface LegacyDownloadStep { + includePrefixes?: string[] + skipPrefixes?: string[] +} + +/** + * Hugging Face filter passes for one download action. A node that declares weight + * variants always fetches its shared files first (every variant excluded), then the + * requested variant — its default one when no variant id is given. The shared pass + * skips files already complete on disk, so it stays cheap on a resume or a second + * variant. + */ +export function legacyDownloadSteps( + plan: Extract, + variantId?: unknown, +): LegacyDownloadStep[] { + const variants = plan.weightVariants + if (!variants) { + if (variantId !== undefined) { + throw new Error(`Model node "${plan.modelId}" has no weight variant "${String(variantId)}"`) + } + return [{ includePrefixes: plan.includePrefixes, skipPrefixes: plan.skipPrefixes }] + } + const selected = variantId === undefined ? variants.default : variantId + const variant = variants.options.find((option) => option.id === selected) + if (!variant) throw new Error(`Model node "${plan.modelId}" has no weight variant "${String(variantId)}"`) + return [ + { + includePrefixes: plan.includePrefixes, + skipPrefixes: [...(plan.skipPrefixes ?? []), ...variants.options.flatMap((option) => option.include_prefixes)], + }, + { includePrefixes: variant.include_prefixes, skipPrefixes: plan.skipPrefixes }, + ] +} + +/** Re-read the installed manifest for every model action; renderer metadata is never trusted. */ +export async function resolveInstalledModelDownloadPlan(args: { + modelId: unknown + userExtensionsDir: string + builtinExtensionsDir: string + blockedExtensionIds?: ReadonlySet +}): Promise { + if (typeof args.modelId !== 'string') throw new Error('Model id must be a string') + const parts = args.modelId.split('/') + if (parts.length !== 2) throw new Error('Model id must identify one extension node') + const extensionId = assertSafeExtensionId(parts[0]) + 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`) + } + + const userPath = resolveExtensionPathWithinRoot(args.userExtensionsDir, extensionId) + const builtinPath = resolveExtensionPathWithinRoot(args.builtinExtensionsDir, extensionId) + const extensionPath = existsSync(userPath) ? userPath : existsSync(builtinPath) ? builtinPath : undefined + if (!extensionPath) throw new Error(`Extension "${extensionId}" is not installed`) + const extensionRoot = extensionPath === userPath ? args.userExtensionsDir : args.builtinExtensionsDir + if ( + existsSync(join(extensionPath, EXT_INCOMPLETE_MARKER)) + || existsSync(join(extensionPath, EXT_REGISTRATION_PENDING_MARKER)) + || await hasPendingRegistration(extensionRoot, extensionId) + ) { + throw new Error(`Extension "${extensionId}" has an incomplete installation`) + } + + const manifestPath = join(extensionPath, 'manifest.json') + if (!existsSync(manifestPath)) throw new Error(`Extension "${extensionId}" has no manifest.json`) + return parseManifest(await readFile(manifestPath, 'utf-8'), extensionId, nodeId) +} diff --git a/electron/main/model-download-preload.test.mjs b/electron/main/model-download-preload.test.mjs new file mode 100644 index 00000000..c88ebf6b --- /dev/null +++ b/electron/main/model-download-preload.test.mjs @@ -0,0 +1,64 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, readFileSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-preload-download-')), 'electron-api.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/preload/electron-api.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +test('renderer model actions keep node and shared-weight identities explicit', async () => { + const { createElectronApi } = loadModule() + const calls = [] + const ipc = { + invoke: async (...args) => { calls.push(args); return { success: true } }, + send: () => {}, + on: () => {}, + removeAllListeners: () => {}, + } + const api = createElectronApi( + ipc, + { setZoomFactor: () => {} }, + { getPathForFile: () => '' }, + ) + + 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'], + ]) +}) + +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(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-downloader.ts b/electron/main/model-downloader.ts index 8c5571d9..abad7b81 100644 --- a/electron/main/model-downloader.ts +++ b/electron/main/model-downloader.ts @@ -4,8 +4,8 @@ */ import { existsSync, readdirSync, statSync, readFileSync } from 'fs' import { join } from 'path' -import { getSettings } from './settings-store' -import { app } from 'electron' +import { getHfToken } from './hf-token' +import type { ModelSource } from './model-sources' export interface DownloadProgress { percent: number @@ -121,7 +121,6 @@ export async function downloadModelFromHF( includePrefixes?: string[], ): Promise { const { net } = require('electron') - const STALL_TIMEOUT_MS = 120_000 let url = `${PYTHON_API_URL}/model/hf-download?repo_id=${encodeURIComponent(repoId)}&model_id=${encodeURIComponent(modelId)}` if (skipPrefixes && skipPrefixes.length > 0) { url += `&skip_prefixes=${encodeURIComponent(JSON.stringify(skipPrefixes))}` @@ -129,26 +128,63 @@ export async function downloadModelFromHF( if (includePrefixes && includePrefixes.length > 0) { url += `&include_prefixes=${encodeURIComponent(JSON.stringify(includePrefixes))}` } - const hfToken = getSettings(app.getPath('userData')).hfToken - if (hfToken) { - url += `&token=${encodeURIComponent(hfToken)}` - } - - const res = await net.fetch(url) + // Header, never a query param: uvicorn logs the full request line to stdout, + // python-bridge pipes that into runtime.log, and `log:readAll` hands that file + // to the user for bug reports. The token used to ride all the way there. + const hfToken = getHfToken() // decrypted cache — settings.json holds the ciphertext + const res = await net.fetch(url, hfToken ? { headers: { 'X-HF-Token': hfToken } } : undefined) if (!res.ok) throw new Error(`HuggingFace download failed: HTTP ${res.status}`) + await consumeDownloadStream(res, onProgress) +} + +/** Download a validated node-level source plan through one aggregate SSE stream. */ +export async function downloadModelSourcesFromHF( + modelId: string, + sources: ModelSource[], + onProgress: ProgressCallback, +): Promise { + const { net } = require('electron') + const headers: Record = { 'Content-Type': 'application/json' } + const hfToken = getHfToken() + if (hfToken) headers.Authorization = `Bearer ${hfToken}` + const url = `${PYTHON_API_URL}/model/hf-download-sources?model_id=${encodeURIComponent(modelId)}` + const res = await net.fetch(url, { + method: 'POST', + headers, + body: JSON.stringify({ sources }), + }) + if (!res.ok) throw new Error(`HuggingFace multi-source download failed: HTTP ${res.status}`) + await consumeDownloadStream(res, onProgress) +} + +async function consumeDownloadStream( + res: Response, + onProgress: ProgressCallback, +): Promise { if (!res.body) throw new Error('No response body from HF download stream') const decoder = new TextDecoder() const reader = res.body.getReader() let buffer = '' + const STALL_TIMEOUT_MS = 120_000 async function readWithTimeout() { - return await Promise.race([ - reader.read(), - new Promise((_, reject) => { - setTimeout(() => reject(new Error(`Model download stalled for ${Math.round(STALL_TIMEOUT_MS / 1000)}s`)), STALL_TIMEOUT_MS) - }), - ]) + // The timer has to be cleared: an SSE stream emitting ~10 events/s otherwise + // keeps every timer created in the last 120 s alive in the main process. + let timer: NodeJS.Timeout | undefined + try { + return await Promise.race([ + reader.read(), + new Promise((_, reject) => { + timer = setTimeout( + () => reject(new Error(`Model download stalled for ${Math.round(STALL_TIMEOUT_MS / 1000)}s`)), + STALL_TIMEOUT_MS, + ) + }), + ]) + } finally { + if (timer) clearTimeout(timer) + } } while (true) { diff --git a/electron/main/model-sources.test.mjs b/electron/main/model-sources.test.mjs new file mode 100644 index 00000000..751c4f56 --- /dev/null +++ b/electron/main/model-sources.test.mjs @@ -0,0 +1,261 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, rmSync, symlinkSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-model-sources-module-')), 'model-sources.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/model-sources.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const validNode = () => ({ + model_sources: [ + { + id: 'primary', + provider: 'huggingface', + repo_id: 'org/main', + destination: '.', + checks: ['pipeline.json'], + }, + { + id: 'encoder', + provider: 'huggingface', + repo_id: 'org/encoder', + revision: 'refs/pr/1', + destination: 'auxiliary/encoder', + include_prefixes: ['config.json', 'weights/'], + skip_prefixes: ['README.md'], + checks: ['config.json', 'model.safetensors'], + }, + ], +}) + +test('validates the new model_sources contract without reinterpreting legacy fields', () => { + const { normalizeModelSources } = loadModule() + const sources = normalizeModelSources(validNode()) + assert.equal(sources.length, 2) + assert.equal(sources[1].destination, 'auxiliary/encoder') + + assert.equal(normalizeModelSources({ + hf_repo: 'legacy/repo', + download_check: '../generate/model.safetensors', + hf_skip_prefixes: ['weights/**'], + }), undefined) +}) + +test('rejects unsafe destinations, unsupported providers, and non-portable source aliases', () => { + const { normalizeModelSources } = loadModule() + const source = validNode().model_sources[0] + for (const destination of ['../outside', 'aux/CON', 'aux/name.', 'C:/models']) { + assert.throws( + () => normalizeModelSources({ model_sources: [{ ...source, destination }] }), + /destination|unsafe/i, + ) + } + assert.throws( + () => normalizeModelSources({ model_sources: [{ ...source, provider: 'url' }] }), + /provider.*huggingface/i, + ) + assert.throws( + () => normalizeModelSources({ model_sources: [source, { ...source, id: 'PRIMARY' }] }), + /portable-unique/i, + ) + assert.throws( + () => normalizeModelSources({ model_sources: [{ ...source, checks: [] }] }), + /checks.*non-empty/i, + ) +}) + +test('requires every declared check and rejects symlinked extension-root ancestry', (t) => { + const { areModelSourcesDownloaded, normalizeModelSources } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-model-readiness-')) + const models = join(root, 'models') + const modelRoot = join(models, 'pixal3d', 'generate') + const sources = normalizeModelSources(validNode()) + mkdirSync(join(modelRoot, 'auxiliary', 'encoder'), { recursive: true }) + writeFileSync(join(modelRoot, 'pipeline.json'), '{}') + writeFileSync(join(modelRoot, 'auxiliary', 'encoder', 'config.json'), '{}') + + try { + assert.equal(areModelSourcesDownloaded(models, 'pixal3d/generate', sources), false) + writeFileSync(join(modelRoot, 'auxiliary', 'encoder', 'model.safetensors'), 'x') + assert.equal(areModelSourcesDownloaded(models, 'pixal3d/generate', sources), true) + + writeFileSync(join(modelRoot, 'auxiliary', 'encoder', 'model.safetensors'), '') + assert.equal(areModelSourcesDownloaded(models, 'pixal3d/generate', sources), false) + rmSync(join(modelRoot, 'auxiliary', 'encoder', 'model.safetensors')) + mkdirSync(join(modelRoot, 'auxiliary', 'encoder', 'model.safetensors')) + assert.equal(areModelSourcesDownloaded(models, 'pixal3d/generate', sources), false) + + rmSync(join(models, 'pixal3d'), { recursive: true, force: true }) + const outside = join(root, 'outside') + mkdirSync(join(outside, 'generate'), { recursive: true }) + try { + symlinkSync(outside, join(models, 'pixal3d'), 'dir') + } catch (error) { + t.skip(`Symlinks unavailable: ${error}`) + return + } + assert.equal(areModelSourcesDownloaded(models, 'pixal3d/generate', sources), false) + } finally { + rmSync(root, { recursive: true, force: true }) + } +}) + +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 }) + } +}) + + +const quantNode = () => ({ + hf_repo: 'org/model-gguf', + download_check: 'pipeline.json', + params_schema: [{ + id: 'quant', + label: 'Quantization', + type: 'select', + default: 'Q5', + options: [{ value: 'Q4', label: 'Q4' }, { value: 'Q5', label: 'Q5' }], + }], + weight_variants: { + param: 'quant', + default: 'Q5', + options: ['Q4', 'Q5'].map((quant) => ({ + id: quant, + include_prefixes: [`dit/model_${quant}.gguf`], + checks: [`dit/model_${quant}.gguf`], + })), + }, +}) + +test('validates weight variants and rejects ambiguous or unsafe declarations', () => { + const { normalizeWeightVariants } = loadModule() + const variants = normalizeWeightVariants(quantNode()) + assert.equal(variants.param, 'quant') + assert.equal(variants.default, 'Q5') + assert.deepEqual(variants.options.map((option) => option.label), ['Q4', 'Q5']) + assert.equal(normalizeWeightVariants({ hf_repo: 'org/model' }), undefined) + + const node = quantNode() + const [q4, q5] = node.weight_variants.options + const withOptions = (options, extra = {}) => ({ ...node, weight_variants: { ...node.weight_variants, options, ...extra } }) + const cases = [ + [{ ...node, hf_repo: undefined }, /requires hf_repo/], + [{ ...node, model_sources: [] }, /cannot be combined/], + [{ ...node, weight_groups: ['base'] }, /cannot be combined with weight_groups/], + [{ ...node, download_check: 'dit/model_Q4.gguf' }, /download_check/], + [withOptions([q4, q5], { default: 'Q8' }), /default/], + [withOptions([q4, { ...q5, include_prefixes: ['dit/'] }]), /share files/], + [withOptions([q4, { ...q5, id: 'q4' }]), /portable-unique/], + [withOptions([{ ...q4, checks: ['other.gguf'] }]), /not covered/], + [withOptions([{ ...q4, include_prefixes: ['../outside'] }]), /unsafe/], + [withOptions([{ ...q4, size_gb: -1 }]), /size_gb/], + [withOptions([{ ...q4, vram_gb: 0 }]), /vram_gb/], + [withOptions([{ ...q4, vram_gb: '6' }]), /vram_gb/], + [{ ...node, params_schema: undefined }, /must name a params_schema entry/], + [{ ...node, params_schema: [{ id: 'steps', type: 'int' }] }, /must name a params_schema entry/], + [ + { ...node, params_schema: [{ id: 'quant', type: 'select', options: [{ value: 'Q5' }] }] }, + /must offer every weight variant id \(missing: Q4\)/, + ], + ] + for (const [candidate, error] of cases) assert.throws(() => normalizeWeightVariants(candidate), error) + + // A param that declares no options (or an unknown shape) is left to the node. + assert.doesNotThrow(() => normalizeWeightVariants({ ...node, params_schema: [{ id: 'quant', type: 'string' }] })) + assert.equal('vram_gb' in variants.options[0], false) + assert.equal(normalizeWeightVariants(withOptions([{ ...q4, vram_gb: 6.5 }, q5])).options[0].vram_gb, 6.5) +}) + +test('reports installed variants and lists only the files of the variant being removed', async () => { + const { installedWeightVariants, listWeightVariantFiles, normalizeWeightVariants } = loadModule() + const variants = normalizeWeightVariants(quantNode()) + const root = mkdtempSync(join(tmpdir(), 'modly-weight-variants-')) + const models = join(root, 'models') + const nodeRoot = join(models, 'trellis', 'generate') + mkdirSync(join(nodeRoot, 'dit'), { recursive: true }) + writeFileSync(join(nodeRoot, 'pipeline.json'), '{}') + writeFileSync(join(nodeRoot, 'dit', 'model_Q5.gguf'), 'q5') + writeFileSync(join(nodeRoot, 'dit', 'model_Q4.gguf'), '') + writeFileSync(join(nodeRoot, 'dit', 'model_Q4.gguf.part'), 'partial') + + try { + assert.deepEqual(installedWeightVariants(models, 'trellis/generate', variants), ['Q5']) + assert.deepEqual(installedWeightVariants(models, '../escape', variants), []) + const files = await listWeightVariantFiles(models, 'trellis/generate', variants.options[0]) + assert.deepEqual( + files.map((file) => file.slice(nodeRoot.length + 1).replaceAll('\\', '/')).sort(), + ['dit/model_Q4.gguf', 'dit/model_Q4.gguf.part'], + ) + } finally { + rmSync(root, { recursive: true, force: true }) + } +}) diff --git a/electron/main/model-sources.ts b/electron/main/model-sources.ts new file mode 100644 index 00000000..194e916e --- /dev/null +++ b/electron/main/model-sources.ts @@ -0,0 +1,535 @@ +import { existsSync, lstatSync, readdirSync, statSync } from 'node:fs' +import { readdir, rm } from 'node:fs/promises' +import { isAbsolute, relative, resolve, sep } from 'node:path' + +export interface ModelSource { + id: string + provider: 'huggingface' + repo_id: string + revision?: string + destination: string + include_prefixes?: string[] + skip_prefixes?: string[] + checks: string[] +} + +export interface ModelSourceNode { + model_sources?: unknown +} + +export interface ModelWeightGroup { + id: string + sources: ModelSource[] +} + +export interface ModelWeightManifest { + weight_groups?: unknown +} + +export interface ModelWeightNode extends ModelSourceNode { + weight_groups?: unknown +} + +/** One separately installable set of files inside a node's model directory (e.g. a quantization). */ +export interface WeightVariant { + id: string + label: string + size_gb?: number + vram_gb?: number + include_prefixes: string[] + checks: string[] +} + +export interface WeightVariants { + /** params_schema id whose value selects the variant at generation time */ + param: string + default: string + options: WeightVariant[] +} + +export interface WeightVariantNode { + weight_variants?: unknown + hf_repo?: unknown + download_check?: unknown + model_sources?: unknown + weight_groups?: unknown + params_schema?: unknown +} + +const SAFE_ID = /^[A-Za-z0-9][A-Za-z0-9._-]*$/ +const WINDOWS_DEVICE = /^(?:con|prn|aux|nul|com[1-9]|lpt[1-9])(?:\..*)?$/i +const WINDOWS_UNSAFE = /[<>"|?*\u0000-\u001f]/ + +function portableSegment(value: string, field: string): string { + if ( + !value + || value === '.' + || value === '..' + || value.endsWith('.') + || value.endsWith(' ') + || value.includes(':') + || WINDOWS_UNSAFE.test(value) + || WINDOWS_DEVICE.test(value) + ) { + throw new Error(`${field} contains unsafe path segment "${value}"`) + } + return value +} + +export function safeModelSourceId(value: unknown, field = 'model source id'): string { + if (typeof value !== 'string' || !value || value !== value.trim() || !SAFE_ID.test(value)) { + throw new Error(`${field} must be a safe non-empty identifier`) + } + return portableSegment(value, field) +} + +export function safeModelRelativePath(value: unknown, field: string, allowDot = false): string { + if (typeof value !== 'string' || !value || value !== value.trim()) { + throw new Error(`${field} must be a non-empty relative path`) + } + if (allowDot && value === '.') return value + if (value === '.' || value.startsWith('/') || value.includes('\\') || isAbsolute(value)) { + throw new Error(`${field} must be a safe relative POSIX path`) + } + for (const part of value.split('/')) portableSegment(part, field) + return value +} + +function safePrefix(value: unknown, field: string): string { + if (typeof value !== 'string') return safeModelRelativePath(value, field) + const path = value.endsWith('/') ? value.slice(0, -1) : value + safeModelRelativePath(path, field) + return value +} + +function optionalPrefixes(value: unknown, field: string): string[] | undefined { + if (value === undefined) return undefined + if (!Array.isArray(value)) throw new Error(`${field} must be an array`) + return value.map((entry, index) => safePrefix(entry, `${field}[${index}]`)) +} + +function safeRepoId(value: unknown, field: string): string { + if (typeof value !== 'string' || !value || value !== value.trim() || value.includes('\\')) { + throw new Error(`${field} must be a non-empty Hugging Face repository id`) + } + const parts = value.split('/') + if (parts.length > 2 || parts.some((part) => !SAFE_ID.test(part) || part === '.' || part === '..')) { + throw new Error(`${field} is not a safe Hugging Face repository id`) + } + return value +} + +function safeRevision(value: unknown, field: string): string | undefined { + if (value === undefined) return undefined + if (typeof value !== 'string' || !value || value !== value.trim() || value.startsWith('/') || value.includes('\\') || value.includes('\0')) { + throw new Error(`${field} must be a safe non-empty revision`) + } + if (value.split('/').some((part) => !part || part === '.' || part === '..')) { + throw new Error(`${field} must be a safe revision`) + } + return value +} + +export function normalizeModelSources( + node: ModelSourceNode, + 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(`${fieldName} must be a non-empty array`) + } + + const seen = new Map() + return node.model_sources.map((raw, index) => { + const field = `${fieldName}[${index}]` + if (typeof raw !== 'object' || raw === null || Array.isArray(raw)) { + throw new Error(`${field} must be an object`) + } + const value = raw as Record + const id = safeModelSourceId(value.id, `${field}.id`) + const alias = id.normalize('NFC').toLowerCase() + const previous = seen.get(alias) + if (previous) throw new Error(`model source ids "${previous}" and "${id}" are not portable-unique`) + seen.set(alias, id) + if (value.provider !== 'huggingface') throw new Error(`${field}.provider must be "huggingface"`) + if (!Array.isArray(value.checks) || value.checks.length === 0) { + throw new Error(`${field}.checks must be a non-empty array`) + } + + const source: ModelSource = { + id, + provider: 'huggingface', + repo_id: safeRepoId(value.repo_id, `${field}.repo_id`), + destination: safeModelRelativePath(value.destination, `${field}.destination`, true), + checks: value.checks.map((check, checkIndex) => ( + safeModelRelativePath(check, `${field}.checks[${checkIndex}]`) + )), + } + const revision = safeRevision(value.revision, `${field}.revision`) + const include = optionalPrefixes(value.include_prefixes, `${field}.include_prefixes`) + const skip = optionalPrefixes(value.skip_prefixes, `${field}.skip_prefixes`) + if (revision !== undefined) source.revision = revision + if (include !== undefined) source.include_prefixes = include + if (skip !== undefined) source.skip_prefixes = skip + return source + }) +} + +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) { + 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 prefixesOverlap(left: string[], right: string[]): boolean { + return left.some((a) => right.some((b) => { + const lowerA = a.toLowerCase() + const lowerB = b.toLowerCase() + return lowerA.startsWith(lowerB) || lowerB.startsWith(lowerA) + })) +} + +/** Variant ids must be selectable, so the param they key has to exist and offer them. */ +function assertParamOffersVariants(param: string, paramsSchema: unknown, ids: string[]): void { + const entry = Array.isArray(paramsSchema) + ? paramsSchema.find((p) => typeof p === 'object' && p !== null && (p as { id?: unknown }).id === param) + : undefined + if (!entry) throw new Error(`weight_variants.param must name a params_schema entry ("${param}")`) + const options = (entry as { options?: unknown }).options + if (!Array.isArray(options)) return + const values = new Set(options.map((option) => + typeof option === 'object' && option !== null ? String((option as { value?: unknown }).value) : String(option))) + const missing = ids.filter((id) => !values.has(id)) + if (missing.length > 0) { + throw new Error(`the "${param}" param must offer every weight variant id (missing: ${missing.join(', ')})`) + } +} + +export function normalizeWeightVariants( + node: WeightVariantNode, + paramsSchema?: unknown, +): WeightVariants | undefined { + if (!Object.prototype.hasOwnProperty.call(node, 'weight_variants')) return undefined + if (Object.prototype.hasOwnProperty.call(node, 'model_sources')) { + throw new Error('weight_variants cannot be combined with model_sources') + } + if (Object.prototype.hasOwnProperty.call(node, 'weight_groups')) { + throw new Error('weight_variants cannot be combined with weight_groups') + } + if (typeof node.hf_repo !== 'string' || !node.hf_repo) { + throw new Error('weight_variants requires hf_repo on the same node') + } + const raw = node.weight_variants + if (typeof raw !== 'object' || raw === null || Array.isArray(raw)) { + throw new Error('weight_variants must be an object') + } + const value = raw as Record + const param = safeModelSourceId(value.param, 'weight_variants.param') + if (!Array.isArray(value.options) || value.options.length === 0) { + throw new Error('weight_variants.options must be a non-empty array') + } + + const seen = new Map() + const options = value.options.map((rawOption, index): WeightVariant => { + const field = `weight_variants.options[${index}]` + if (typeof rawOption !== 'object' || rawOption === null || Array.isArray(rawOption)) { + throw new Error(`${field} must be an object`) + } + const option = rawOption as Record + const id = safeModelSourceId(option.id, `${field}.id`) + const alias = id.normalize('NFC').toLowerCase() + const previous = seen.get(alias) + if (previous) throw new Error(`weight variant ids "${previous}" and "${id}" are not portable-unique`) + seen.set(alias, id) + const label = option.label ?? id + if (typeof label !== 'string' || !label.trim()) throw new Error(`${field}.label must be a non-empty string`) + const sizeGb = option.size_gb + if (sizeGb !== undefined && (typeof sizeGb !== 'number' || !Number.isFinite(sizeGb) || sizeGb <= 0)) { + throw new Error(`${field}.size_gb must be a positive number`) + } + const vramGb = option.vram_gb + if (vramGb !== undefined && (typeof vramGb !== 'number' || !Number.isFinite(vramGb) || vramGb <= 0)) { + throw new Error(`${field}.vram_gb must be a positive number`) + } + const include = optionalPrefixes(option.include_prefixes, `${field}.include_prefixes`) + if (!include?.length) throw new Error(`${field}.include_prefixes must be a non-empty array`) + if (!Array.isArray(option.checks) || option.checks.length === 0) { + throw new Error(`${field}.checks must be a non-empty array`) + } + const checks = option.checks.map((check, checkIndex) => { + const path = safeModelRelativePath(check, `${field}.checks[${checkIndex}]`) + if (!include.some((prefix) => path.startsWith(prefix))) { + throw new Error(`${field}.checks[${checkIndex}] is not covered by its include_prefixes`) + } + return path + }) + return { + id, + label, + ...(sizeGb === undefined ? {} : { size_gb: sizeGb }), + ...(vramGb === undefined ? {} : { vram_gb: vramGb }), + include_prefixes: include, + checks, + } + }) + + options.forEach((variant, index) => { + for (const other of options.slice(index + 1)) { + if (prefixesOverlap(variant.include_prefixes, other.include_prefixes)) { + throw new Error(`weight variants "${variant.id}" and "${other.id}" share files`) + } + } + }) + const downloadCheck = node.download_check + if (typeof downloadCheck === 'string' && options.some((variant) => prefixesOverlap([downloadCheck], variant.include_prefixes))) { + throw new Error('download_check must name a file outside every weight variant') + } + const defaultId = value.default ?? options[0].id + if (!options.some((variant) => variant.id === defaultId)) { + throw new Error('weight_variants.default must name one of its options') + } + assertParamOffersVariants(param, paramsSchema ?? node.params_schema, options.map((variant) => variant.id)) + return { param, default: defaultId as string, options } +} + +function pathHasSymlink(root: string, candidate: string): boolean { + const rootPath = resolve(root) + const rel = relative(rootPath, resolve(candidate)) + if (rel === '..' || rel.startsWith('../') || rel.startsWith('..\\') || isAbsolute(rel)) return true + let current = rootPath + try { + if (existsSync(current) && lstatSync(current).isSymbolicLink()) return true + for (const part of rel.split(/[/\\]/).filter(Boolean)) { + current = resolve(current, part) + if (existsSync(current) && lstatSync(current).isSymbolicLink()) return true + } + } catch { + return true + } + return false +} + +export function resolveModelRoot(modelsDir: string, modelId: string): string { + if (typeof modelId !== 'string') throw new Error('Model id must be a string') + const parts = modelId.split('/') + if (parts.length !== 2) throw new Error('Model id must identify one extension node') + const extensionId = safeModelSourceId(parts[0], 'extension id') + 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 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 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 { + if (!existsSync(modelRoot) || pathHasSymlink(modelRoot, modelRoot)) return false + return sources.every((source) => { + const destination = source.destination === '.' + ? modelRoot + : resolve(modelRoot, ...source.destination.split('/')) + if (!existsSync(destination) || pathHasSymlink(modelRoot, destination)) return false + return source.checks.every((check) => { + const candidate = resolve(destination, ...check.split('/')) + if (!existsSync(candidate) || pathHasSymlink(modelRoot, candidate)) return false + try { + const stat = statSync(candidate) + return stat.isFile() && stat.size > 0 + } catch { + return false + } + }) + }) + } catch { + return false + } +} + +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) + return existsSync(modelRoot) && readdirSync(modelRoot).length > 0 + } catch { + return false + } +} + +export function weightStorageHasLocalData(modelsDir: string, targetId: string): boolean { + try { + const root = resolveWeightStorageRoot(modelsDir, targetId) + return existsSync(root) && readdirSync(root).length > 0 + } catch { + return false + } +} + +function isDownloadedFile(root: string, relativePath: string): boolean { + const candidate = resolve(root, ...relativePath.split('/')) + if (!existsSync(candidate) || pathHasSymlink(root, candidate)) return false + try { + const stat = statSync(candidate) + return stat.isFile() && stat.size > 0 + } catch { + return false + } +} + +export function installedWeightVariants(modelsDir: string, modelId: string, variants: WeightVariants): string[] { + try { + const modelRoot = resolveModelRoot(modelsDir, modelId) + return variants.options + .filter((variant) => variant.checks.every((check) => isDownloadedFile(modelRoot, check))) + .map((variant) => variant.id) + } catch { + return [] + } +} + +/** Files of one variant on disk (in-progress `.part` files included), never through a symlink. */ +export async function listWeightVariantFiles(modelsDir: string, modelId: string, variant: WeightVariant): Promise { + const modelRoot = resolveModelRoot(modelsDir, modelId) + if (!existsSync(modelRoot)) return [] + const entries = await readdir(modelRoot, { recursive: true, withFileTypes: true }) + return entries + .filter((entry) => entry.isFile()) + .map((entry) => resolve(entry.parentPath ?? modelRoot, entry.name)) + .filter((path) => { + const relativePath = relative(modelRoot, path).split(sep).join('/') + return variant.include_prefixes.some((prefix) => relativePath.startsWith(prefix)) + && !pathHasSymlink(modelRoot, path) + }) +} + +// Mirrors the backend's cancel cleanup (api/routers/model.py): only the in-progress +// `.part` files are removed, so completed sources already on disk survive a cancel. +export async function removePartialDownloadArtifacts(modelRoot: string): Promise { + if (!existsSync(modelRoot)) return + const entries = await readdir(modelRoot, { recursive: true, withFileTypes: true }) + await Promise.all( + entries + .filter((entry) => entry.isFile() && entry.name.endsWith('.part')) + .map((entry) => rm(resolve(entry.parentPath ?? modelRoot, entry.name), { force: true })), + ) +} diff --git a/electron/main/model-weight-ipc.test.mjs b/electron/main/model-weight-ipc.test.mjs new file mode 100644 index 00000000..b77d75bc --- /dev/null +++ b/electron/main/model-weight-ipc.test.mjs @@ -0,0 +1,218 @@ +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 } }), + listModels: async () => ({ data: [{ id: 'demo/a' }, { id: 'demo/b' }, { id: 'other/generate' }] }), + 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), + get: (...args) => hooks.listModels(...args), + isAxiosError: (err) => Boolean(err?.isAxiosError), + } + 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) + }) +} + +function connectionRefused() { + return Object.assign(new Error('connect ECONNREFUSED 127.0.0.1:8765'), { isAxiosError: true, code: 'ECONNREFUSED' }) +} + +for (const action of ['deleteSharedGroup', 'deleteExtensionWeights']) { + test(`${action} still removes weights when the backend is not running`, async (t) => { + // No backend means no Python process can hold the files open. + const f = fixture(t) + await f.invoke('download', 'demo/a') + f.hooks.unload = async () => { throw connectionRefused() } + f.hooks.listModels = async () => { throw connectionRefused() } + assert.equal((await f.invoke(action, 'demo', 'base')).success, true) + assert.equal(f.removed.length, 1) + }) +} + +test('deleteExtensionWeights unloads only that extension models', async (t) => { + const f = fixture(t), unloaded = [] + f.hooks.unload = async (url) => { unloaded.push(url); return { data: { unloaded: true } } } + assert.equal((await f.invoke('deleteExtensionWeights', 'demo')).success, true) + assert.deepEqual(unloaded, ['http://test/model/unload/demo%2Fa', 'http://test/model/unload/demo%2Fb']) +}) + +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/main/process-runner-worker-exit.test.mjs b/electron/main/process-runner-worker-exit.test.mjs new file mode 100644 index 00000000..de1a11f1 --- /dev/null +++ b/electron/main/process-runner-worker-exit.test.mjs @@ -0,0 +1,106 @@ +/** + * A JS process extension runs in a worker thread that is kept warm between + * runs. If that worker dies mid-run -- an uncaught error outside the awaited + * processor call, running out of memory on a large mesh, process.exit() -- it + * never posts 'done' or 'error'. The run must settle with an error instead of + * leaving the workflow waiting forever, and the next run must get a fresh + * worker rather than posting into the dead one -- including when the worker + * died while idle, between two runs. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' +import vm from 'node:vm' + +function loadModule() { + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/process-runner.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + external: ['electron'], + }) + // process-runner imports electron's `app` (only the Python runner uses it); + // the real package cannot load outside Electron, so hand it a stub. + const module = { exports: {} } + const dependencies = (name) => (name === 'electron' ? { app: {} } : require(name)) + vm.runInNewContext(result.outputFiles[0].text, { + module, exports: module.exports, require: dependencies, process, console, Buffer, setTimeout, clearTimeout, + }) + return module.exports +} + +// Behaves according to params.mode; `runs` counts runs served by this worker, +// so a fresh worker starts again from 1. +function makeRunner() { + const { ProcessRunner } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-worker-exit-')) + const extDir = join(root, 'ext') + mkdirSync(extDir, { recursive: true }) + writeFileSync(join(extDir, 'processor.js'), [ + 'let runs = 0', + 'module.exports = async (input, params) => {', + ' runs += 1', + " if (params.mode === 'throw') throw new Error('bad input')", + " if (params.mode === 'crash') {", + " setTimeout(() => { throw new Error('worker blew up') }, 0)", + ' return new Promise(() => {})', + ' }', + " if (params.mode === 'exit') process.exit(3)", + // Returns normally, then a timer it left behind kills the idle worker. + " if (params.mode === 'late-crash') setTimeout(() => { throw new Error('late failure') }, 10)", + ' return { text: String(runs) }', + '}', + '', + ].join('\n')) + return new ProcessRunner(extDir, 'processor.js', join(root, 'workspace'), root) +} + +test('a run whose worker crashes rejects instead of hanging', { timeout: 5000 }, async () => { + const runner = makeRunner() + try { + await assert.rejects(runner.run({}, { mode: 'crash' }), /worker blew up/) + } finally { + runner.terminate() + } +}) + +test('after its worker exits, the runner starts a fresh one for the next run', { timeout: 5000 }, async () => { + const runner = makeRunner() + try { + await assert.rejects(runner.run({}, { mode: 'exit' }), /exited with code 3/) + assert.deepEqual(await runner.run({}, { mode: 'ok' }), { text: '1' }) + } finally { + runner.terminate() + } +}) + +test('an error thrown by the processor still rejects with its message and keeps the warm worker', { timeout: 5000 }, async () => { + const runner = makeRunner() + try { + await assert.rejects(runner.run({}, { mode: 'throw' }), { message: 'Error: bad input' }) + // Same worker thread: its run counter carried over instead of restarting. + assert.deepEqual(await runner.run({}, { mode: 'ok' }), { text: '2' }) + } finally { + runner.terminate() + } +}) + +test('a worker that dies between runs is replaced for the next run', { timeout: 5000 }, async () => { + const runner = makeRunner() + try { + // The run itself succeeds; the timer it leaves behind kills the idle worker. + assert.deepEqual(await runner.run({}, { mode: 'late-crash' }), { text: '1' }) + await new Promise((settle) => setTimeout(settle, 200)) + // A fresh worker serves the next run (counter restarts) instead of hanging. + assert.deepEqual(await runner.run({}, { mode: 'ok' }), { text: '1' }) + } finally { + runner.terminate() + } +}) diff --git a/electron/main/process-runner.test.mjs b/electron/main/process-runner.test.mjs new file mode 100644 index 00000000..1d72b866 --- /dev/null +++ b/electron/main/process-runner.test.mjs @@ -0,0 +1,122 @@ +/** + * Process-extension runners are cached per extension id and reused across + * workflow runs. The cache must not hand back a runner that was built for + * different arguments: the workspace folder is baked into each runner, so after + * the user moves the workspace in Settings (which updates the backend at + * runtime, without a restart) a stale runner keeps writing node output into the + * old folder, where the viewer can no longer find it. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, mkdirSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' +import vm from 'node:vm' + +function loadModule() { + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/process-runner.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + external: ['electron'], + }) + // process-runner imports electron's `app` (the Python runner reads it when it + // spawns); the real package cannot load outside Electron, so hand it a stub. + const app = { isPackaged: false, getAppPath: () => process.cwd() } + const module = { exports: {} } + const dependencies = (name) => (name === 'electron' ? { app } : require(name)) + vm.runInNewContext(result.outputFiles[0].text, { + module, exports: module.exports, require: dependencies, process, console, Buffer, setTimeout, clearTimeout, + }) + return module.exports +} + +// A JS process extension that reports the workspace it was given, plus how many +// runs this worker has served (module state survives only while it is reused). +function makeJsExtension(root) { + const extDir = join(root, 'js-ext') + mkdirSync(extDir, { recursive: true }) + writeFileSync(join(extDir, 'processor.js'), [ + 'let runs = 0', + 'module.exports = async (input, params, context) => {', + ' runs += 1', + ' return { filePath: context.workspaceDir, text: String(runs) }', + '}', + '', + ].join('\n')) + return extDir +} + +// A "Python" process extension driven through the same stdin/stdout protocol. +// Node stands in for the interpreter so the test needs no Python install. +function makeStdioExtension(root) { + const extDir = join(root, 'py-ext') + mkdirSync(extDir, { recursive: true }) + writeFileSync(join(extDir, 'processor.cjs'), [ + "let raw = ''", + "process.stdin.on('data', (chunk) => { raw += chunk })", + "process.stdin.on('end', () => {", + ' const data = JSON.parse(raw)', + " process.stdout.write(JSON.stringify({ type: 'done', result: { filePath: data.workspaceDir } }) + '\\n')", + '})', + '', + ].join('\n')) + return extDir +} + +test('a JS process runner follows the workspace after it moves', async () => { + const { getProcessRunner, terminateAllProcessRunners } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-runner-js-')) + const extDir = makeJsExtension(root) + const oldWorkspace = join(root, 'workspace-old') + const newWorkspace = join(root, 'workspace-new') + try { + const before = await getProcessRunner('js-ext', extDir, 'processor.js', oldWorkspace, root).run({}, {}) + assert.equal(before.filePath, oldWorkspace) + + const after = await getProcessRunner('js-ext', extDir, 'processor.js', newWorkspace, root).run({}, {}) + assert.equal(after.filePath, newWorkspace) + } finally { + terminateAllProcessRunners() + } +}) + +test('a Python process runner follows the workspace after it moves', async () => { + const { getPythonProcessRunner, terminateAllProcessRunners } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-runner-py-')) + const extDir = makeStdioExtension(root) + const oldWorkspace = join(root, 'workspace-old') + const newWorkspace = join(root, 'workspace-new') + try { + const before = await getPythonProcessRunner('py-ext', process.execPath, extDir, 'processor.cjs', oldWorkspace, root).run({}, {}) + assert.equal(before.filePath, oldWorkspace) + + const after = await getPythonProcessRunner('py-ext', process.execPath, extDir, 'processor.cjs', newWorkspace, root).run({}, {}) + assert.equal(after.filePath, newWorkspace) + } finally { + terminateAllProcessRunners() + } +}) + +test('unchanged arguments keep reusing the same warm runner', async () => { + const { getProcessRunner, terminateAllProcessRunners } = loadModule() + const root = mkdtempSync(join(tmpdir(), 'modly-runner-reuse-')) + const extDir = makeJsExtension(root) + const workspace = join(root, 'workspace') + try { + const first = getProcessRunner('js-ext', extDir, 'processor.js', workspace, root) + assert.equal((await first.run({}, {})).text, '1') + + const second = getProcessRunner('js-ext', extDir, 'processor.js', workspace, root) + assert.equal(second, first) + // Same worker thread: its module state carried over instead of reloading. + assert.equal((await second.run({}, {})).text, '2') + } finally { + terminateAllProcessRunners() + } +}) diff --git a/electron/main/process-runner.ts b/electron/main/process-runner.ts index 758f62d2..8b012c82 100644 --- a/electron/main/process-runner.ts +++ b/electron/main/process-runner.ts @@ -2,6 +2,7 @@ import { Worker } from 'worker_threads' import { spawn } from 'child_process' import { existsSync } from 'fs' import { join } from 'path' +import { app } from 'electron' // ─── Worker code for JS process extensions ──────────────────────────────────── @@ -114,6 +115,11 @@ export class ProcessRunner implements IProcessRunner { worker.once('error', (err) => { reject(err) }) + + // A worker can also die between runs (e.g. a timer the processor left + // behind throws after it returned). Forget it whenever it exits, so the + // next run starts a fresh one instead of posting into a dead thread. + worker.once('exit', () => this.discardWorker(worker)) }) } @@ -127,25 +133,52 @@ export class ProcessRunner implements IProcessRunner { const worker = this.worker! return new Promise((resolve, reject) => { + const settle = () => { + worker.off('message', handler) + worker.off('error', onError) + worker.off('exit', onExit) + } const handler = (msg: { type: string; result?: ProcessResult; message?: string; percent?: number; label?: string }) => { if (msg.type === 'progress') { onProgress?.(msg.percent ?? 0, msg.label ?? '') } else if (msg.type === 'log') { onLog?.(msg.message ?? '') } else if (msg.type === 'done') { - worker.off('message', handler) + settle() resolve(msg.result ?? {}) } else if (msg.type === 'error') { - worker.off('message', handler) + settle() reject(new Error(msg.message)) } } + // A worker that dies mid-run (an uncaught error, out of memory, + // process.exit) never posts 'done' or 'error'. Settle the run instead of + // waiting forever, and drop the dead worker so the next run starts a + // fresh one rather than posting into it. + const onError = (err: Error) => { + settle() + this.discardWorker(worker) + reject(err) + } + const onExit = (code: number) => { + settle() + this.discardWorker(worker) + reject(new Error(`Process extension worker exited with code ${code}`)) + } worker.on('message', handler) + worker.on('error', onError) + worker.on('exit', onExit) worker.postMessage({ action: 'run', input, params }) }) } + private discardWorker(worker: Worker): void { + if (this.worker !== worker) return + this.worker = null + this.ready = false + } + terminate(): void { this.worker?.terminate() this.worker = null @@ -160,12 +193,14 @@ export class ProcessRunner implements IProcessRunner { export class PythonProcessRunner implements IProcessRunner { private pythonExe: string + private extDir: string private scriptPath: string private workspaceDir: string private tempDir: string constructor(pythonExe: string, extDir: string, entry: string, workspaceDir: string, tempDir: string) { this.pythonExe = pythonExe + this.extDir = extDir this.scriptPath = join(extDir, entry) this.workspaceDir = workspaceDir this.tempDir = tempDir @@ -180,9 +215,22 @@ export class PythonProcessRunner implements IProcessRunner { return new Promise((resolve, reject) => { const proc = spawn(this.pythonExe, [this.scriptPath], { stdio: ['pipe', 'pipe', 'pipe'], - // Force UTF-8 stdio so Unicode prints from process extensions do not - // crash under legacy Windows codepages (cp1252/cp932). - env: { ...process.env, PYTHONUTF8: '1' }, + env: { + ...process.env, + // Force UTF-8 stdio so Unicode prints from process extensions do not + // crash under legacy Windows codepages (cp1252/cp932). + PYTHONUTF8: '1', + // Built-in process nodes may import shared services from the backend. + MODLY_API_DIR: app.isPackaged + ? join(process.resourcesPath, 'api') + : join(app.getAppPath(), 'api'), + // Electron is also the packaged Node runtime. meshopt_runner.cjs uses + // ELECTRON_RUN_AS_NODE when it launches this executable. + MODLY_NODE_EXECUTABLE: process.execPath, + EXTENSION_DIR: this.extDir, + WORKSPACE_DIR: this.workspaceDir, + TEMP_DIR: this.tempDir, + }, }) // Send input as a single JSON line on stdin @@ -270,6 +318,19 @@ export function getExtPythonExe(extDir: string): string | null { // ─── Registry (one runner per extension id, reused across calls) ────────────── const registry = new Map() +// The arguments each cached runner was built with. A runner bakes them in at +// construction, so a call with different ones — e.g. after the workspace is +// moved in Settings, which updates paths without a restart — must not get the +// old runner back, or node output keeps landing in the previous folder. +const registryArgs = new Map() + +function canReuseRunner(extensionId: string, args: string[]): boolean { + const key = JSON.stringify(args) + if (registry.has(extensionId) && registryArgs.get(extensionId) === key) return true + terminateProcessRunner(extensionId) + registryArgs.set(extensionId, key) + return false +} export function getProcessRunner( extensionId: string, @@ -278,7 +339,7 @@ export function getProcessRunner( workspaceDir: string, tempDir: string, ): ProcessRunner { - if (!registry.has(extensionId)) { + if (!canReuseRunner(extensionId, [extDir, entry, workspaceDir, tempDir])) { registry.set(extensionId, new ProcessRunner(extDir, entry, workspaceDir, tempDir)) } return registry.get(extensionId)! as ProcessRunner @@ -292,7 +353,7 @@ export function getPythonProcessRunner( workspaceDir: string, tempDir: string, ): PythonProcessRunner { - if (!registry.has(extensionId)) { + if (!canReuseRunner(extensionId, [pythonExe, extDir, entry, workspaceDir, tempDir])) { registry.set(extensionId, new PythonProcessRunner(pythonExe, extDir, entry, workspaceDir, tempDir)) } return registry.get(extensionId)! as PythonProcessRunner diff --git a/electron/main/python-bridge.ts b/electron/main/python-bridge.ts index 94dcfb63..87aa2b6b 100644 --- a/electron/main/python-bridge.ts +++ b/electron/main/python-bridge.ts @@ -3,7 +3,8 @@ import { join } from 'path' import { app, BrowserWindow } from 'electron' import { existsSync, mkdirSync } from 'fs' import axios from 'axios' -import { getSettings } from './settings-store' +import { ensureAgentDir, getSettings } from './settings-store' +import { getHfToken } from './hf-token' import { logger } from './logger' import { cleanPythonEnv, getVenvPythonExe } from './python-setup' @@ -52,10 +53,20 @@ export class PythonBridge { env: { ...cleanPythonEnv(), PYTHONUNBUFFERED: '1', + // mesh_ops uses Electron in Node mode for the existing meshoptimizer + // backend, so packaged builds do not depend on a system Node install. + MODLY_NODE_EXECUTABLE: process.execPath, + // We decode both pipes as UTF-8 below (Buffer.toString() default), and + // the API forwards extension output — tqdm bars included — through its + // own stderr. Without this the API would encode them with the Windows + // locale codec and the HUD log pane would show █ escapes. + PYTHONIOENCODING: 'utf-8', // No PYTHONPATH needed - the venv's Python has its own isolated site-packages MODELS_DIR: this.resolveModelsDir(), WORKSPACE_DIR: this.resolveWorkspaceDir(), EXTENSIONS_DIR: this.resolveExtensionsDir(), + // llm_server.py keeps everything of the local LLM here: engine, GGUF models, logs, config. + MODLY_LLM_DIR: this.resolveAgentDir(), SELECTED_MODEL_ID: process.env['SELECTED_MODEL_ID'] ?? '', HUGGING_FACE_HUB_TOKEN: this.resolveHfToken(), HF_TOKEN: this.resolveHfToken(), @@ -237,7 +248,11 @@ export class PythonBridge { return s.extensionsDir } + private resolveAgentDir(): string { + return ensureAgentDir(app.getPath('userData')) + } + private resolveHfToken(): string { - return getSettings(app.getPath('userData')).hfToken ?? '' + return getHfToken() // decrypted cache — settings.json holds the ciphertext } } diff --git a/electron/main/python-setup.ts b/electron/main/python-setup.ts index 3b85db1b..13e86a25 100644 --- a/electron/main/python-setup.ts +++ b/electron/main/python-setup.ts @@ -1,10 +1,11 @@ import { BrowserWindow, app } from 'electron' import { existsSync, readFileSync, writeFileSync } from 'fs' -import { cp, rm, mkdir } from 'fs/promises' +import { rm, mkdir } from 'fs/promises' import { join } from 'path' import { spawn, execSync } from 'child_process' import { createHash } from 'crypto' import { getSettings } from './settings-store' +import { copyRuntimeTree } from './copy-runtime' const SETUP_VERSION = 3 @@ -176,7 +177,7 @@ async function ensureStableEmbeddedPython(userData: string, win: BrowserWindow): await rm(stableDir, { recursive: true, force: true }) } await mkdir(stableDir, { recursive: true }) - await cp(getEmbeddedPythonDir(), stableDir, { recursive: true, preserveTimestamps: true }) + await copyRuntimeTree(getEmbeddedPythonDir(), stableDir) writeFileSync(versionFile, currentVersion, 'utf-8') console.log('[PythonSetup] Python runtime ready at:', stableDir) } diff --git a/electron/main/secure-store.ts b/electron/main/secure-store.ts new file mode 100644 index 00000000..7c70cd27 --- /dev/null +++ b/electron/main/secure-store.ts @@ -0,0 +1,78 @@ +import { safeStorage } from 'electron' +import { logger } from './logger' + +// Secrets that safeStorage couldn't encrypt (no OS keychain available, e.g. some +// Linux setups without a keyring) are stored with this prefix so decryptSecret +// can tell them apart from a real ciphertext and hand them back unchanged. +const PLAINTEXT_PREFIX = 'plain:' + +let warnedOnce = false + +// Electron 33 (pinned in package.json) only ships the synchronous safeStorage +// API — encryptStringAsync/decryptStringAsync/isAsyncEncryptionAvailable were +// added in a later Electron release. The sync calls are cheap (no disk I/O, +// just OS keychain/DPAPI crypto), so wrapping them in an async function here +// is only to keep the IPC handler signature uniform, not for a real await. + +export function encryptSecretSync(plainText: string): string { + if (!plainText) return plainText + + if (!safeStorage.isEncryptionAvailable()) { + if (!warnedOnce) { + warnedOnce = true + logger.warn('[secure-store] OS-level encryption unavailable — storing secrets in plain text') + } + return PLAINTEXT_PREFIX + plainText + } + + return safeStorage.encryptString(plainText).toString('hex') +} + +/** + * True when `stored` has the shape encryptSecretSync produces: a hex blob + * starting with Chromium OSCrypt's version tag ("v10", "v11", …). The tag + * matters — hex alone also matched plaintext API keys that happen to be hex + * (common for self-hosted endpoints), which then "failed to decrypt" and were + * wiped during the migration instead of encrypted. Used to tell "this was never encrypted" apart from "this IS one of + * our blobs but decryption just failed" — the two must not be confused: DPAPI + * keys are per-OS-user, so a restored backup or a different Windows account + * makes decryption fail on a perfectly valid ciphertext. Treating that as + * plaintext would hand the ciphertext out as a credential and then re-encrypt + * it, destroying the secret for good. + */ +function looksEncrypted(stored: string): boolean { + return stored.length >= 32 && /^76(3[0-9]){2}([0-9a-f]{2})+$/i.test(stored) +} + +/** + * Returns the plaintext, or `null` when `stored` is one of our blobs that could + * not be decrypted. A value that was never encrypted comes back unchanged, so + * callers can migrate it. + */ +export function decryptSecretSync(stored: string): string | null { + if (!stored) return stored + if (stored.startsWith(PLAINTEXT_PREFIX)) return stored.slice(PLAINTEXT_PREFIX.length) + + if (!looksEncrypted(stored)) return stored // legacy plaintext, saved before encryption existed + + try { + return safeStorage.decryptString(Buffer.from(stored, 'hex')) + } catch (err) { + logger.error( + '[secure-store] a stored secret could not be decrypted — it was encrypted by a ' + + `different OS user or machine. It must be re-entered. (${err})`, + ) + return null + } +} + +// Async wrappers kept for the IPC handlers, whose signatures are all async. +export async function encryptSecret(plainText: string): Promise { + return encryptSecretSync(plainText) +} + +/** `null` on an undecryptable blob — never the ciphertext, which callers would + * otherwise use as a credential and re-encrypt. */ +export async function decryptSecret(stored: string): Promise { + return decryptSecretSync(stored) +} diff --git a/electron/main/settings-store.test.mjs b/electron/main/settings-store.test.mjs new file mode 100644 index 00000000..06a78fc1 --- /dev/null +++ b/electron/main/settings-store.test.mjs @@ -0,0 +1,79 @@ +/** + * The agent's folder (local LLM engine, GGUF models, logs) sits beside the + * other data folders. Installs set up before it existed have models/, + * extensions/, … under the base they picked at first run, so an unsaved + * agentDir must land next to those, not in userData. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { existsSync, mkdtempSync, readFileSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' +import vm from 'node:vm' + +function loadModule() { + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/settings-store.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + const module = { exports: {} } + vm.runInNewContext(result.outputFiles[0].text, { module, exports: module.exports, require }) + return module.exports +} + +const { getSettings, setSettings, ensureAgentDir } = loadModule() + +test('a fresh install puts the agent folder in userData beside the others', () => { + const userData = mkdtempSync(join(tmpdir(), 'modly-settings-')) + const s = getSettings(userData) + assert.equal(s.agentDir, join(userData, 'agent')) + assert.equal(s.modelsDir, join(userData, 'models')) +}) + +test('an existing install gets the agent folder beside its chosen data folders', () => { + const userData = mkdtempSync(join(tmpdir(), 'modly-settings-')) + const base = join(userData, 'Documents', 'Modly') + writeFileSync(join(userData, 'settings.json'), JSON.stringify({ + modelsDir: join(base, 'models'), + workspaceDir: join(base, 'workspace'), + extensionsDir: join(base, 'extensions'), + })) + assert.equal(getSettings(userData).agentDir, join(base, 'agent')) +}) + +test('a saved agent folder is kept, and persists once any setting is written', () => { + const userData = mkdtempSync(join(tmpdir(), 'modly-settings-')) + const custom = join(userData, 'elsewhere', 'agent') + writeFileSync(join(userData, 'settings.json'), JSON.stringify({ agentDir: custom })) + assert.equal(getSettings(userData).agentDir, custom) + + setSettings(userData, { workflowsDir: join(userData, 'wf') }) + assert.equal(JSON.parse(readFileSync(join(userData, 'settings.json'), 'utf-8')).agentDir, custom) +}) + +test('an existing install gets its agent folder created and pinned at startup', () => { + const userData = mkdtempSync(join(tmpdir(), 'modly-settings-')) + const base = join(userData, 'Documents', 'Modly') + writeFileSync(join(userData, 'settings.json'), JSON.stringify({ modelsDir: join(base, 'models') })) + + assert.equal(ensureAgentDir(userData), join(base, 'agent')) + assert.equal(existsSync(join(base, 'agent')), true) + assert.equal(JSON.parse(readFileSync(join(userData, 'settings.json'), 'utf-8')).agentDir, join(base, 'agent')) + + // Pinned: moving the models folder afterwards leaves the agent folder where it is. + setSettings(userData, { modelsDir: join(userData, 'other-drive', 'models') }) + assert.equal(getSettings(userData).agentDir, join(base, 'agent')) +}) + +test('a fresh install gets its agent folder without settings.json being written', () => { + const userData = mkdtempSync(join(tmpdir(), 'modly-settings-')) + assert.equal(ensureAgentDir(userData), join(userData, 'agent')) + assert.equal(existsSync(join(userData, 'agent')), true) + assert.equal(existsSync(join(userData, 'settings.json')), false) +}) diff --git a/electron/main/settings-store.ts b/electron/main/settings-store.ts index c90599a5..a0279797 100644 --- a/electron/main/settings-store.ts +++ b/electron/main/settings-store.ts @@ -1,5 +1,5 @@ -import { join } from 'path' -import { readFileSync, writeFileSync, existsSync } from 'fs' +import { dirname, join } from 'path' +import { readFileSync, writeFileSync, existsSync, mkdirSync } from 'fs' export interface AppSettings { modelsDir: string @@ -7,6 +7,8 @@ export interface AppSettings { workflowsDir: string extensionsDir: string dependenciesDir: string + /** Local LLM engine, GGUF models, logs and config used by the agent. */ + agentDir: string hfToken?: string } @@ -15,7 +17,7 @@ function settingsPath(userData: string): string { } export function getSettings(userData: string): AppSettings { - const defaults: AppSettings = { + const defaults: Omit = { modelsDir: join(userData, 'models'), workspaceDir: join(userData, 'workspace'), workflowsDir: join(userData, 'workflows'), @@ -24,7 +26,7 @@ export function getSettings(userData: string): AppSettings { } const file = settingsPath(userData) - if (!existsSync(file)) return defaults + if (!existsSync(file)) return withAgentDir(defaults) try { const saved = JSON.parse(readFileSync(file, 'utf-8')) as Record @@ -33,12 +35,42 @@ export function getSettings(userData: string): AppSettings { saved['workspaceDir'] = saved['outputsDir'] delete saved['outputsDir'] } - return { ...defaults, ...saved } + return withAgentDir({ ...defaults, ...saved }) } catch { - return defaults + return withAgentDir(defaults) } } +/** + * Installs set up before agentDir existed have their data folders under the base + * chosen at first run (/models, /extensions, …). Defaulting agentDir + * to userData would put the agent's multi-GB engine and models apart from all of + * them, so it defaults to a sibling of modelsDir instead. + */ +function withAgentDir(settings: Omit & { agentDir?: string }): AppSettings { + return { ...settings, agentDir: settings.agentDir || join(dirname(settings.modelsDir), 'agent') } +} + +/** + * Create the agent folder, and pin it in settings.json for installs that predate + * it. Until pinned it is derived from modelsDir, so moving the models folder in + * Settings → Storage would silently take the agent folder (engine, GGUF models) + * somewhere else. A fresh install has no settings.json yet: setup writes it. + */ +export function ensureAgentDir(userData: string): string { + const { agentDir } = getSettings(userData) + const file = settingsPath(userData) + if (existsSync(file)) { + let saved: Record | null = null + try { + saved = JSON.parse(readFileSync(file, 'utf-8')) as Record + } catch { /* unreadable: leave it as is */ } + if (saved && !saved['agentDir']) setSettings(userData, { agentDir }) + } + mkdirSync(agentDir, { recursive: true }) + return agentDir +} + export function setSettings(userData: string, patch: Partial): AppSettings { const updated = { ...getSettings(userData), ...patch } writeFileSync(settingsPath(userData), JSON.stringify(updated, null, 2), 'utf-8') diff --git a/electron/main/setup-launcher.test.mjs b/electron/main/setup-launcher.test.mjs new file mode 100644 index 00000000..63d65ec2 --- /dev/null +++ b/electron/main/setup-launcher.test.mjs @@ -0,0 +1,338 @@ +/** + * Runs the setup launcher for real, with a stub standing in for pip, and checks + * what the extension's pip invocation was rewritten into. + * + * The commands exercised below are copied from the official extensions' + * setup.py, so a change that breaks AMD installs fails here rather than after a + * multi-gigabyte download on a user's machine. + */ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync, readFileSync, existsSync } from 'node:fs' +import { spawnSync } from 'node:child_process' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-launcher-test-')), 'setup-launcher.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('electron/main/setup-launcher.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { SETUP_LAUNCHER_SOURCE } = loadModule() + +function findPython() { + for (const candidate of ['python3', 'python']) { + const probe = spawnSync(candidate, ['--version'], { stdio: 'ignore' }) + if (probe.status === 0) return candidate + } + return null +} + +const PYTHON = findPython() + +// The launcher's _is_pip_command matches pip's executable spellings ahead of +// the subcommand, so the stub is named pip.py — matching how extensions invoke +// "/bin/pip". +const FAKE_PIP = ` +import json, sys +with open(sys.argv[1], "w") as handle: + json.dump(sys.argv[2:], handle) +` + +/** + * Runs a pip command through the launcher and returns the argv the stub + * actually received, i.e. the command after every rewrite. `invoke` picks how + * the stand-in setup.py issues the call — the shapes real extensions use. + */ +function runThroughLauncher(pipArgs, env = {}, invoke = 'run') { + const dir = mkdtempSync(join(tmpdir(), 'modly-launcher-run-')) + const fakePip = join(dir, 'pip.py') + const capture = join(dir, 'captured.json') + writeFileSync(fakePip, FAKE_PIP, 'utf8') + + const invocations = { + run: 'subprocess.run(CMD, check=True)', + run_kwarg: 'subprocess.run(args=CMD, check=True)', + popen: 'proc = subprocess.Popen(CMD)\nassert proc.wait() == 0', + check_call: 'subprocess.check_call(CMD)', + } + + // Stands in for the extension's setup.py: issues one pip call, exactly as the + // real ones do, and lets the launcher's patched subprocess handle it. + const setupPy = join(dir, 'setup.py') + writeFileSync(setupPy, [ + 'import subprocess, sys', + `PIP = ${JSON.stringify([fakePip, capture])}`, + `CMD = [sys.executable] + PIP + ${JSON.stringify(pipArgs)}`, + invocations[invoke], + ].join('\n'), 'utf8') + + const result = spawnSync(PYTHON, ['-c', SETUP_LAUNCHER_SOURCE, setupPy, '{}'], { + encoding: 'utf8', + env: { ...process.env, ...env }, + }) + assert.equal(result.status, 0, `launcher failed:\n${result.stderr}`) + assert.ok(existsSync(capture), `pip stub was never invoked:\n${result.stderr}`) + + // Drop the stub's own two arguments; keep the pip command itself. + return { argv: JSON.parse(readFileSync(capture, 'utf8')), stderr: result.stderr } +} + +const ROCM_ENV = { + MODLY_TORCH_FLAVOR: 'rocm', + MODLY_TORCH_INDEX_URL: 'https://download.pytorch.org/whl/rocm7.2', + MODLY_TORCH_SPECS: JSON.stringify(['torch', 'torchvision']), +} + +test('ROCm shim redirects a CUDA-pinned install (triposg / trellis2 shape)', { skip: !PYTHON }, () => { + const { argv, stderr } = runThroughLauncher( + ['install', 'torch==2.6.0', 'torchvision==0.21.0', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + ]) + assert.match(stderr, /Redirected PyTorch to ROCm wheels/) +}) + +test('ROCm shim rescues the CPU fallback hunyuan3d-mini forces on Windows', { skip: !PYTHON }, () => { + // That setup.py hardcodes the CPU index when torch_flavor is rocm on Windows; + // without this rewrite an AMD Windows user would silently get CPU-only torch. + const { argv } = runThroughLauncher( + ['install', 'torch==2.6.0', 'torchvision==0.21.0', '--index-url', 'https://download.pytorch.org/whl/cpu'], + { + ...ROCM_ENV, + MODLY_TORCH_INDEX_URL: 'https://repo.amd.com/rocm/whl-multi-arch/', + MODLY_TORCH_SPECS: JSON.stringify([ + 'torch[device-gfx1200]==2.11.0+rocm7.14.0', + 'torchvision[device-gfx1200]==0.26.0+rocm7.14.0', + ]), + }, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://repo.amd.com/rocm/whl-multi-arch/', + 'torch[device-gfx1200]==2.11.0+rocm7.14.0', + 'torchvision[device-gfx1200]==0.26.0+rocm7.14.0', + ]) +}) + +test('ROCm shim leaves an extension that already chose ROCm alone', { skip: !PYTHON }, () => { + // hunyuan3d-mini's own rocm branch on Linux. Its index is kept; only the + // requirements are normalised to what Modly resolved. + const { argv } = runThroughLauncher( + ['install', 'torch', 'torchvision', '--index-url', 'https://download.pytorch.org/whl/rocm7.2'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + 'torch', 'torchvision', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + ]) +}) + +test('ROCm shim replaces the pinned direct wheel URLs of the ARM64 path', { skip: !PYTHON }, () => { + const { argv } = runThroughLauncher( + [ + 'install', '--retries', '10', + '--extra-index-url', 'https://download.pytorch.org/whl/cu128', + 'https://download-r2.pytorch.org/whl/cu128/torch-2.7.0%2Bcu128-cp311-cp311-manylinux_2_28_aarch64.whl', + 'https://download-r2.pytorch.org/whl/cu128/torchvision-0.22.0-cp311-cp311-manylinux_2_28_aarch64.whl', + ], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', '--retries', '10', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + ]) +}) + +test('ROCm shim handles the --index-url=VALUE spelling', { skip: !PYTHON }, () => { + const { argv } = runThroughLauncher( + ['install', '--index-url=https://download.pytorch.org/whl/cu124', 'torch==2.6.0'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + ]) +}) + +test('ROCm shim leaves non-torch installs untouched', { skip: !PYTHON }, () => { + // The bulk of every setup.py: core deps, rembg, etc. Nothing here is ours to + // rewrite, and an extra --index-url would send them to the ROCm index. + const original = ['install', 'Pillow', 'numpy', 'trimesh', 'rembg', 'onnxruntime'] + const { argv } = runThroughLauncher(original, ROCM_ENV) + assert.deepEqual(argv, original) +}) + +test('ROCm shim keeps PyPI reachable when torch is mixed with other packages', { skip: !PYTHON }, () => { + // The ROCm index mirrors torch's dependency closure only: asking it for + // trimesh returns 403, so a mixed install would hard-fail without this. + const { argv } = runThroughLauncher( + ['install', 'torch==2.6.0', 'trimesh', 'diffusers', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + '--extra-index-url', 'https://pypi.org/simple', + 'torch', 'torchvision', + 'trimesh', 'diffusers', + ]) +}) + +test('ROCm shim does not add PyPI for a torch-only install', { skip: !PYTHON }, () => { + // Keeping PyPI out of a pure torch install avoids pip ever preferring a + // plain CUDA wheel over the ROCm one. + const { argv } = runThroughLauncher( + ['install', '--retries', '10', 'torch==2.6.0', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + ) + assert.ok(!argv.includes('--extra-index-url')) +}) + +test('shim is inert on a CUDA machine', { skip: !PYTHON }, () => { + const original = ['install', 'torch==2.6.0', '--index-url', 'https://download.pytorch.org/whl/cu124'] + const { argv } = runThroughLauncher(original, { + MODLY_TORCH_FLAVOR: 'cuda', + MODLY_TORCH_INDEX_URL: '', + MODLY_TORCH_SPECS: '[]', + }) + assert.deepEqual(argv, original) +}) + +test('ROCm shim covers subprocess.Popen and the args= keyword too', { skip: !PYTHON }, () => { + // Extension setup.py scripts stream pip output through Popen, or spell the + // command as run(args=[...]); both bypassed the patch when it only wrapped + // run/check_call/check_output positionally. + for (const invoke of ['popen', 'run_kwarg', 'check_call']) { + const { argv } = runThroughLauncher( + ['install', 'torch==2.6.0', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + invoke, + ) + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + ], `invocation shape "${invoke}" escaped the ROCm redirect`) + } +}) + +test('ROCm shim matches "python -u -m pip install" too', { skip: !PYTHON }, () => { + // hunyuan3d-style scripts call pip through the interpreter; "pip" is then the + // fourth token, which the old first-three-tokens check missed. + const dir = mkdtempSync(join(tmpdir(), 'modly-launcher-m-')) + const capture = join(dir, 'captured.json') + writeFileSync(join(dir, 'pip.py'), [ + 'import json, os, sys', + 'with open(os.environ["MODLY_TEST_CAPTURE"], "w") as handle:', + ' json.dump(sys.argv[1:], handle)', + ].join('\n'), 'utf8') + + const setupPy = join(dir, 'setup.py') + writeFileSync(setupPy, [ + 'import subprocess, sys', + 'subprocess.run([sys.executable, "-u", "-m", "pip", "install", "torch==2.6.0",', + ' "--index-url", "https://download.pytorch.org/whl/cu124"], check=True)', + ].join('\n'), 'utf8') + + const result = spawnSync(PYTHON, ['-c', SETUP_LAUNCHER_SOURCE, setupPy, '{}'], { + encoding: 'utf8', + env: { ...process.env, ...ROCM_ENV, MODLY_TEST_CAPTURE: capture, PYTHONPATH: dir }, + }) + assert.equal(result.status, 0, `launcher failed:\n${result.stderr}`) + + assert.deepEqual(JSON.parse(readFileSync(capture, 'utf8')), [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + ]) +}) + +test('ROCm shim demotes a foreign primary index instead of keeping it primary', { skip: !PYTHON }, () => { + // pip resolves a duplicated --index-url last-wins: a kept foreign primary + // index could shadow the injected ROCm one and silently hand torch back to + // CUDA/CPU wheels. Demoted to an extra index, its other packages stay + // reachable while ROCm stays primary. + const { argv } = runThroughLauncher( + ['install', 'torch==2.6.0', '--index-url', 'https://pypi.org/simple'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + 'torch', 'torchvision', + '--extra-index-url', 'https://pypi.org/simple', + ]) +}) + +test('ROCm shim keeps PyPI reachable when the extension chose ROCm itself', { skip: !PYTHON }, () => { + // The keeps-rocm-index branch needs the PyPI rescue just as much: the ROCm + // index answers 403 for trimesh/diffusers. + const { argv } = runThroughLauncher( + ['install', 'torch', 'trimesh', '--index-url', 'https://download.pytorch.org/whl/rocm7.2'], + ROCM_ENV, + ) + + assert.deepEqual(argv, [ + 'install', + '--extra-index-url', 'https://pypi.org/simple', + 'torch', 'torchvision', + 'trimesh', + '--index-url', 'https://download.pytorch.org/whl/rocm7.2', + ]) +}) + +test('ROCm shim recognises AMD\'s Windows index as already-ROCm', { skip: !PYTHON }, () => { + // repo.amd.com carries no "/whl/rocm" path segment; it must still count as a + // ROCm index or the shim would inject a second, conflicting --index-url. + const winEnv = { + ...ROCM_ENV, + MODLY_TORCH_INDEX_URL: 'https://repo.amd.com/rocm/whl-multi-arch/', + MODLY_TORCH_SPECS: JSON.stringify(['torch[device-gfx1200]==2.11.0+rocm7.14.0']), + } + const { argv } = runThroughLauncher( + ['install', 'torch==2.6.0', '--index-url', 'https://repo.amd.com/rocm/whl-multi-arch/'], + winEnv, + ) + + assert.deepEqual(argv, [ + 'install', + 'torch[device-gfx1200]==2.11.0+rocm7.14.0', + '--index-url', 'https://repo.amd.com/rocm/whl-multi-arch/', + ]) +}) + +test('--no-cache-dir is still stripped alongside the ROCm rewrite', { skip: !PYTHON }, () => { + const { argv } = runThroughLauncher( + ['install', '--no-cache-dir', 'torch==2.6.0', '--index-url', 'https://download.pytorch.org/whl/cu124'], + ROCM_ENV, + ) + + assert.ok(!argv.includes('--no-cache-dir')) + assert.ok(argv.includes('https://download.pytorch.org/whl/rocm7.2')) +}) diff --git a/electron/main/setup-launcher.ts b/electron/main/setup-launcher.ts new file mode 100644 index 00000000..a3150670 --- /dev/null +++ b/electron/main/setup-launcher.ts @@ -0,0 +1,258 @@ +/** + * Python launcher used to run an extension's setup.py. + * + * Extension setup scripts are third-party code we cannot edit, and they install + * PyTorch themselves from an index they hardcode. The launcher wraps them: it + * patches subprocess so every pip invocation passes through a few corrections + * before it runs — dropping CUDA-only indexes on macOS, keeping the shared wheel + * cache alive, and redirecting torch to ROCm wheels on AMD machines. + * + * Kept in its own module so setup-launcher.test.mjs can execute it for real + * against the command shapes the official extensions actually use. + */ + +export const SETUP_LAUNCHER_SOURCE = ` +import json +import os +import re +import runpy +import subprocess +import sys + +setup_py = sys.argv[1] +setup_args = sys.argv[2:] + +_TORCH_REQ_RE = re.compile(r"^(torch|torchvision|torchaudio)(\\[[^\\]]*\\])?\\s*([<>=!~].*)?$", re.I) + +def _is_cuda_torch_index(value): + return isinstance(value, str) and value.startswith("https://download.pytorch.org/whl/cu") + +def _is_torch_requirement(text): + # Matches "torch", "torch==2.6.0", "torch[device-gfx1200]==2.11.0+rocm7.14.0", + # and the pinned direct wheel URLs the ARM64 install path uses. + if _TORCH_REQ_RE.match(text): + return True + if text.startswith(("http://", "https://")): + return any(seg in text for seg in ("/torch-", "/torchvision-", "/torchaudio-")) + return False + +def _mentions_torch(command): + if not isinstance(command, (list, tuple)): + return False + return any(_is_torch_requirement(str(part)) for part in command) + +def _rewrite_command(command): + if sys.platform != "darwin" or not _mentions_torch(command): + return command + if not isinstance(command, (list, tuple)): + return command + + rewritten = [] + changed = False + i = 0 + while i < len(command): + part = command[i] + text = str(part) + if text in ("--index-url", "-i", "--extra-index-url") and i + 1 < len(command) and _is_cuda_torch_index(str(command[i + 1])): + changed = True + i += 2 + continue + if text.startswith("--index-url=") or text.startswith("--extra-index-url="): + value = text.split("=", 1)[1] + if _is_cuda_torch_index(value): + changed = True + i += 1 + continue + rewritten.append(part) + i += 1 + + if changed: + print("[Modly setup compat] Removed CUDA-only PyTorch index on macOS; pip will use macOS wheels.", file=sys.stderr) + return rewritten + return command + +# Matches the executable spellings pip arrives under: "pip", "pip3", "pip3.11", +# "pip.exe", or a stub script like "pip.py" — but not a requirement that merely +# starts with "pip" (pipdeptree) or a script like "pipeline.py". +_PIP_BASENAME_RE = re.compile(r"^pip[0-9.]*(\\.py|\\.exe)?$", re.I) + +def _is_pip_command(command): + # Scan everything ahead of the pip subcommand, so both "/bin/pip + # install" and "python -u -m pip install" match, while requirements after + # "install" are never mistaken for the executable. + if not isinstance(command, (list, tuple)): + return False + texts = [str(part) for part in command] + for i, text in enumerate(texts): + if text in ("install", "download", "wheel"): + break + if text == "-m" and i + 1 < len(texts) and texts[i + 1].lower() in ("pip", "pip._internal"): + return True + if _PIP_BASENAME_RE.match(os.path.basename(text)): + return True + return False + +def _strip_no_cache(command): + # Extension setup scripts often hardcode --no-cache-dir, which forces pip to + # re-download multi-GB wheels on every retry. Modly provides a shared cache + # via PIP_CACHE_DIR, so drop the flag and let pip use it. + if not _is_pip_command(command): + return command + if not any(str(part) == "--no-cache-dir" for part in command): + return command + print("[Modly setup compat] Removed --no-cache-dir so pip reuses the shared wheel cache.", file=sys.stderr) + return [part for part in command if str(part) != "--no-cache-dir"] + +# ROCm redirect. Most extension setup.py scripts predate AMD support and +# hardcode a CUDA index (hunyuan3d-mini even forces the CPU index on Windows), +# so on an AMD machine we swap the whole torch install for the ROCm one Modly +# resolved. An index the extension already pointed at ROCm is left alone. +_ROCM_INDEX = os.environ.get("MODLY_TORCH_INDEX_URL", "") +try: + _ROCM_SPECS = json.loads(os.environ.get("MODLY_TORCH_SPECS", "[]")) +except ValueError: + _ROCM_SPECS = [] + +def _is_pytorch_index(value): + return isinstance(value, str) and "download.pytorch.org/whl/" in value + +def _is_rocm_index(value): + # Covers the pytorch.org rocm indexes, AMD's Windows multi-arch index, and + # whatever MODLY_ROCM_INDEX was overridden to. + if not isinstance(value, str): + return False + if _ROCM_INDEX and value.rstrip("/") == _ROCM_INDEX.rstrip("/"): + return True + return "/whl/rocm" in value or "repo.amd.com/rocm" in value + +_PIP_VALUE_FLAGS = ( + "--index-url", "-i", "--extra-index-url", "--find-links", "-f", + "--retries", "--timeout", "--cache-dir", "--target", "-t", + "--requirement", "-r", "--constraint", "-c", "--progress-bar", + "--proxy", "--cert", "--client-cert", "--trusted-host", "--log", + "--no-binary", "--only-binary", "--prefix", "--root", "--upgrade-strategy", + "--python-version", "--platform", "--abi", "--implementation", +) + +def _has_non_torch_requirement(command): + # Only look past the subcommand, so the interpreter and pip executable + # paths ahead of it are never mistaken for requirements. + texts = [str(part) for part in command] + start = None + for index, text in enumerate(texts): + if text in ("install", "download", "wheel"): + start = index + 1 + break + if start is None: + return False + + skip_next = False + for text in texts[start:]: + if skip_next: + skip_next = False + continue + if text in _PIP_VALUE_FLAGS: + skip_next = True + continue + if text.startswith("-") or _is_torch_requirement(text): + continue + return True + return False + +def _rewrite_rocm(command): + if os.environ.get("MODLY_TORCH_FLAVOR") != "rocm" or not _ROCM_INDEX or not _ROCM_SPECS: + return command + if not isinstance(command, (list, tuple)): + return command + if not _is_pip_command(command) or not _mentions_torch(command): + return command + + rewritten = [] + insert_at = None + keeps_rocm_index = False + changed = False + i = 0 + while i < len(command): + text = str(command[i]) + if text in ("--index-url", "-i", "--extra-index-url") and i + 1 < len(command): + value = str(command[i + 1]) + if _is_rocm_index(value): + keeps_rocm_index = True + rewritten.extend(command[i:i + 2]) + elif _is_pytorch_index(value): + changed = True + elif text != "--extra-index-url": + # A foreign primary index (a PyPI mirror, pypi.nvidia.com…) + # must not stay primary: pip resolves a duplicated --index-url + # last-wins, so it could shadow the injected ROCm index and + # silently hand torch back to CUDA/CPU wheels. Demote it to an + # extra index so its other packages stay reachable. + changed = True + rewritten.extend(["--extra-index-url", command[i + 1]]) + else: + rewritten.extend(command[i:i + 2]) + i += 2 + continue + if "=" in text and text.split("=", 1)[0] in ("--index-url", "--extra-index-url"): + flag, value = text.split("=", 1) + if _is_rocm_index(value): + keeps_rocm_index = True + rewritten.append(command[i]) + elif _is_pytorch_index(value): + changed = True + elif flag != "--extra-index-url": + changed = True + rewritten.append("--extra-index-url=" + value) + else: + rewritten.append(command[i]) + i += 1 + continue + if _is_torch_requirement(text): + if insert_at is None: + insert_at = len(rewritten) + changed = True + i += 1 + continue + rewritten.append(command[i]) + i += 1 + + if not changed: + return command + if insert_at is None: + insert_at = len(rewritten) + index_args = [] if keeps_rocm_index else ["--index-url", _ROCM_INDEX] + if _has_non_torch_requirement(command) and not any("pypi.org/simple" in str(part) for part in rewritten): + # The ROCm index only mirrors torch's own dependency closure (numpy, + # pillow…), not application packages like trimesh or diffusers. A pip + # call that mixes both still needs PyPI reachable — also when the + # extension supplied the ROCm index itself. + index_args += ["--extra-index-url", "https://pypi.org/simple"] + injected = index_args + list(_ROCM_SPECS) + rewritten[insert_at:insert_at] = injected + print("[Modly setup compat] Redirected PyTorch to ROCm wheels: " + " ".join(injected), file=sys.stderr) + return rewritten + +def _transform_command(command): + return _strip_no_cache(_rewrite_rocm(_rewrite_command(command))) + +# Every subprocess entry point — run, call, check_call, check_output, and a +# direct subprocess.Popen(...) — constructs the module-global Popen, so patching +# that one class covers them all exactly once, whether the command arrives +# positionally or as the args= keyword. This also holds for code that did +# "from subprocess import run" before this launcher ran: run() still looks +# Popen up on the module at call time. +_OriginalPopen = subprocess.Popen + +class _PatchedPopen(_OriginalPopen): + def __init__(self, *args, **kwargs): + if args: + args = (_transform_command(args[0]),) + tuple(args[1:]) + elif "args" in kwargs: + kwargs["args"] = _transform_command(kwargs["args"]) + super().__init__(*args, **kwargs) + +subprocess.Popen = _PatchedPopen + +sys.argv = [setup_py] + setup_args +runpy.run_path(setup_py, run_name="__main__") +` diff --git a/electron/preload/artifact-registry-preload.test.ts b/electron/preload/artifact-registry-preload.test.ts index 511589e4..9dec46fd 100644 --- a/electron/preload/artifact-registry-preload.test.ts +++ b/electron/preload/artifact-registry-preload.test.ts @@ -13,7 +13,7 @@ test('preload exposes scoped workspace library list/read/open methods', async () send: () => undefined, on: () => undefined, removeAllListeners: () => undefined, - }, { setZoomFactor: () => undefined }) + }, { setZoomFactor: () => undefined }, { getPathForFile: () => '' }) await api.workspace.library.list() await api.workspace.library.read({ @@ -43,3 +43,27 @@ test('preload exposes scoped workspace library list/read/open methods', async () }, ]) }) + +test('preload resolves dropped file paths through Electron webUtils', () => { + const files: unknown[] = [] + const calls: string[] = [] + const api = createElectronApi({ + invoke: async (channel: string) => { + calls.push(channel) + return null + }, + send: () => undefined, + on: () => undefined, + removeAllListeners: () => undefined, + }, { setZoomFactor: () => undefined }, { + getPathForFile: (file) => { + files.push(file) + return 'C:\\images\\input.png' + }, + }) + const file = {} as Parameters[0] + + assert.equal(api.fs.getPathForFile(file), 'C:\\images\\input.png') + assert.deepEqual(files, [file]) + assert.deepEqual(calls, []) +}) diff --git a/electron/preload/electron-api.ts b/electron/preload/electron-api.ts index 89fa69ec..33944a57 100644 --- a/electron/preload/electron-api.ts +++ b/electron/preload/electron-api.ts @@ -17,7 +17,9 @@ export interface WebFrameLike { setZoomFactor(factor: number): void } -export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFrameLike) { +export type WebUtilsLike = Pick + +export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFrameLike, webUtils: WebUtilsLike) { return { // Window controls window: { @@ -43,6 +45,12 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra // Shell utilities shell: { openExternal: (url: string) => ipcRenderer.invoke('shell:openExternal', url) }, + // Slicer integration — open a model in OrcaSlicer via its deeplink + slicer: { + open: (url: string): Promise<{ success: boolean; error?: string }> => + ipcRenderer.invoke('slicer:open', url) as Promise<{ success: boolean; error?: string }>, + }, + // System info system: { memory: (): Promise<{ total: number; used: number; available: number }> => @@ -67,6 +75,8 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra // File system dialogs + local file reading fs: { + getPathForFile: (file: Parameters[0]): string => + webUtils.getPathForFile(file), selectImage: (): Promise => ipcRenderer.invoke('fs:selectImage') as Promise, selectMeshFile: (): Promise => @@ -93,6 +103,21 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra ipcRenderer.invoke('fs:readScreenshotDataUrl', filename) as Promise, }, + // Secure storage — OS-level encryption for secrets (API keys, …) + secureStore: { + encrypt: (plainText: string): Promise => ipcRenderer.invoke('secure:encrypt', plainText) as Promise, + // null = one of our blobs that couldn't be decrypted here (different OS + // user/machine). Never the ciphertext — see secure-store.ts. + decrypt: (stored: string): Promise => ipcRenderer.invoke('secure:decrypt', stored) as Promise, + }, + + // Agent — local LLM models + agent: { + // Opens a file picker and copies the chosen .gguf into the agent's models folder. + addModel: (): Promise<{ success: boolean; cancelled?: boolean; fileName?: string; error?: string }> => + ipcRenderer.invoke('agent:addModel') as Promise<{ success: boolean; cancelled?: boolean; fileName?: string; error?: string }>, + }, + // Settings settings: { get: (): Promise<{ modelsDir: string; workspaceDir: string; workflowsDir: string; extensionsDir: string; hfToken?: string }> => @@ -117,12 +142,19 @@ export function createElectronApi(ipcRenderer: IpcRendererLike, webFrame: WebFra model: { export: (args: { outputUrl: string; format: string }) => ipcRenderer.invoke('model:export', args), listDownloaded: () => ipcRenderer.invoke('model:listDownloaded'), - isDownloaded: (modelId: string, downloadCheck?: string) => ipcRenderer.invoke('model:isDownloaded', modelId, downloadCheck), - download: (repoId: string, modelId: string, skipPrefixes?: string[], includePrefixes?: string[]) => - ipcRenderer.invoke('model:download', { repoId, modelId, skipPrefixes, includePrefixes }), + isDownloaded: (modelId: string) => ipcRenderer.invoke('model:isDownloaded', modelId), + hasLocalData: (modelId: string) => ipcRenderer.invoke('model:hasLocalData', modelId), + sharedGroups: (extensionId: string) => ipcRenderer.invoke('model:sharedGroups', extensionId), + download: (modelId: string, variantId?: string) => (variantId === undefined + ? ipcRenderer.invoke('model:download', modelId) + : ipcRenderer.invoke('model:download', modelId, variantId)), pauseDownload: (modelId: string) => ipcRenderer.invoke('model:pauseDownload', modelId), cancelDownload: (modelId: string) => ipcRenderer.invoke('model:cancelDownload', modelId), delete: (modelId: string) => ipcRenderer.invoke('model:delete', modelId), + deleteSharedGroup: (extensionId: string, groupId: string) => ipcRenderer.invoke('model:deleteSharedGroup', extensionId, groupId), + deleteExtensionWeights: (extensionId: string) => ipcRenderer.invoke('model:deleteExtensionWeights', extensionId), + installedWeightVariants: (modelId: string) => ipcRenderer.invoke('model:installedWeightVariants', modelId), + deleteWeightVariant: (modelId: string, variantId: string) => ipcRenderer.invoke('model:deleteWeightVariant', modelId, variantId), unloadAll: () => ipcRenderer.invoke('model:unloadAll'), showInFolder: (modelId: string) => ipcRenderer.invoke('model:showInFolder', modelId), activeDownloads: (): Promise<{ modelId: string; percent: number; file?: string; fileIndex?: number; totalFiles?: number }[]> => @@ -155,6 +187,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/electron/preload/index.ts b/electron/preload/index.ts index ec34ebf2..03b80363 100644 --- a/electron/preload/index.ts +++ b/electron/preload/index.ts @@ -1,6 +1,6 @@ -import { contextBridge, ipcRenderer, webFrame } from 'electron' +import { contextBridge, ipcRenderer, webFrame, webUtils } from 'electron' import { createElectronApi } from './electron-api' // Expose a typed API to the renderer process via window.electron -contextBridge.exposeInMainWorld('electron', createElectronApi(ipcRenderer, webFrame)) +contextBridge.exposeInMainWorld('electron', createElectronApi(ipcRenderer, webFrame, webUtils)) diff --git a/package-lock.json b/package-lock.json index 5f7152ca..d7b522f6 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,12 +1,12 @@ { "name": "modly", - "version": "0.4.1", + "version": "0.4.2", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "modly", - "version": "0.4.1", + "version": "0.4.2", "dependencies": { "@electron-toolkit/utils": "^4.0.0", "@mkkellogg/gaussian-splats-3d": "^0.4.7", @@ -34,12 +34,13 @@ "@vitejs/plugin-react": "^4.3.4", "autoprefixer": "^10.4.20", "cross-env": "^10.1.0", - "electron": "^42.4.1", + "electron": "^42.11.1", "electron-builder": "^26.15.3", "electron-vite": "^5.0.0", "eslint": "^9.17.0", "eslint-plugin-react-hooks": "^7.1.1", "globals": "^17.6.0", + "jsdom": "^29.1.1", "postcss": "^8.4.49", "tailwindcss": "^3.4.17", "typescript": "^5.7.2", @@ -59,6 +60,57 @@ "url": "https://github.com/sponsors/sindresorhus" } }, + "node_modules/@asamuzakjp/css-color": { + "version": "5.1.11", + "resolved": "https://registry.npmjs.org/@asamuzakjp/css-color/-/css-color-5.1.11.tgz", + "integrity": "sha512-KVw6qIiCTUQhByfTd78h2yD1/00waTmm9uy/R7Ck/ctUyAPj+AEDLkQIdJW0T8+qGgj3j5bpNKK7Q3G+LedJWg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@asamuzakjp/generational-cache": "^1.0.1", + "@csstools/css-calc": "^3.2.0", + "@csstools/css-color-parser": "^4.1.0", + "@csstools/css-parser-algorithms": "^4.0.0", + "@csstools/css-tokenizer": "^4.0.0" + }, + "engines": { + "node": "^20.19.0 || ^22.12.0 || >=24.0.0" + } + }, + "node_modules/@asamuzakjp/dom-selector": { + "version": "7.1.1", + "resolved": "https://registry.npmjs.org/@asamuzakjp/dom-selector/-/dom-selector-7.1.1.tgz", + "integrity": "sha512-67RZDnYRc8H/8MLDgQCDE//zoqVFwajkepHZgmXrbwybzXOEwOWGPYGmALYl9J2DOLfFPPs6kKCqmbzV895hTQ==", + "dev": true, + "license": "MIT", + "dependencies": { + "@asamuzakjp/generational-cache": "^1.0.1", + "@asamuzakjp/nwsapi": "^2.3.9", + "bidi-js": "^1.0.3", + "css-tree": "^3.2.1", + "is-potential-custom-element-name": "^1.0.1" + }, + "engines": { + "node": "^20.19.0 || ^22.12.0 || >=24.0.0" + } + }, + "node_modules/@asamuzakjp/generational-cache": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/@asamuzakjp/generational-cache/-/generational-cache-1.0.1.tgz", + "integrity": "sha512-wajfB8KqzMCN2KGNFdLkReeHncd0AslUSrvHVvvYWuU8ghncRJoA50kT3zP9MVL0+9g4/67H+cdvBskj9THPzg==", + "dev": true, + "license": "MIT", + "engines": { + "node": "^20.19.0 || ^22.12.0 || >=24.0.0" + } + }, + "node_modules/@asamuzakjp/nwsapi": { + "version": "2.3.9", + "resolved": "https://registry.npmjs.org/@asamuzakjp/nwsapi/-/nwsapi-2.3.9.tgz", + "integrity": "sha512-n8GuYSrI9bF7FFZ/SjhwevlHc8xaVlb/7HmHelnc/PZXBD2ZR49NnN9sMMuDdEGPeeRQ5d0hqlSlEpgCX3Wl0Q==", + "dev": true, + "license": "MIT" + }, "node_modules/@babel/code-frame": { "version": "7.29.7", "resolved": "https://registry.npmjs.org/@babel/code-frame/-/code-frame-7.29.7.tgz", @@ -361,6 +413,159 @@ "node": ">=6.9.0" } }, + "node_modules/@bramus/specificity": { + "version": "2.4.2", + "resolved": "https://registry.npmjs.org/@bramus/specificity/-/specificity-2.4.2.tgz", + "integrity": "sha512-ctxtJ/eA+t+6q2++vj5j7FYX3nRu311q1wfYH3xjlLOsczhlhxAg2FWNUXhpGvAw3BWo1xBcvOV6/YLc2r5FJw==", + "dev": true, + "license": "MIT", + "dependencies": { + "css-tree": "^3.0.0" + }, + "bin": { + "specificity": "bin/cli.js" + } + }, + "node_modules/@csstools/color-helpers": { + "version": "6.1.2", + "resolved": "https://registry.npmjs.org/@csstools/color-helpers/-/color-helpers-6.1.2.tgz", + "integrity": "sha512-grhRy3OKmniaAEKXMjua5z/EODX0MSqBGjunw8+j/3HQjOnahs2AGhvEOIYVUWcU6ScApbhLhVrQTX8XqrMrow==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT-0", + "engines": { + "node": ">=20.19.0" + } + }, + "node_modules/@csstools/css-calc": { + "version": "3.4.3", + "resolved": "https://registry.npmjs.org/@csstools/css-calc/-/css-calc-3.4.3.tgz", + "integrity": "sha512-iex20d8CHVkyvg6B7UKV7uHnI2Bqo9g+EFfT9E0y+GvTvhZ/DwONJ+9aKb1dlqm0ZiGsL5RXjp0fCoJYnkeDjA==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT", + "engines": { + "node": ">=20.19.0" + }, + "peerDependencies": { + "@csstools/css-parser-algorithms": "^4.0.2", + "@csstools/css-tokenizer": "^4.0.2" + } + }, + "node_modules/@csstools/css-color-parser": { + "version": "4.2.6", + "resolved": "https://registry.npmjs.org/@csstools/css-color-parser/-/css-color-parser-4.2.6.tgz", + "integrity": "sha512-iiPQ3iRWwnJkeEn6RIu6SJPr7hYrLz6XZ9s/QZl+2/LI5KQVjpl2fdmDSZKuD4xP6GMmMPgHFFXg6k1Wkz0Trg==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT", + "dependencies": { + "@csstools/color-helpers": "^6.1.2", + "@csstools/css-calc": "^3.4.3" + }, + "engines": { + "node": ">=20.19.0" + }, + "peerDependencies": { + "@csstools/css-parser-algorithms": "^4.0.2", + "@csstools/css-tokenizer": "^4.0.2" + } + }, + "node_modules/@csstools/css-parser-algorithms": { + "version": "4.0.2", + "resolved": "https://registry.npmjs.org/@csstools/css-parser-algorithms/-/css-parser-algorithms-4.0.2.tgz", + "integrity": "sha512-40cSKyMvK+tq4qz6Awrlye2WGuOKt3FwPgtGg6KTfbHOWNw+Rk1rzbAtZnZ6IBhsY491HLRnDXwoyBAijmmILA==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT", + "engines": { + "node": ">=20.19.0" + }, + "peerDependencies": { + "@csstools/css-tokenizer": "^4.0.2" + } + }, + "node_modules/@csstools/css-syntax-patches-for-csstree": { + "version": "1.1.15", + "resolved": "https://registry.npmjs.org/@csstools/css-syntax-patches-for-csstree/-/css-syntax-patches-for-csstree-1.1.15.tgz", + "integrity": "sha512-J0u7HkVl2nzSlhsiTOp4AmwcUQ3D+mGEEKfBy/7To5/y7F2OHwyLrXfrhR0SMgr4p5Lo+eaMVSeai24zUcBIxA==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT-0", + "peerDependencies": { + "css-tree": "^3.2.1" + }, + "peerDependenciesMeta": { + "css-tree": { + "optional": true + } + } + }, + "node_modules/@csstools/css-tokenizer": { + "version": "4.0.2", + "resolved": "https://registry.npmjs.org/@csstools/css-tokenizer/-/css-tokenizer-4.0.2.tgz", + "integrity": "sha512-OoKoR0f76dCY666JlcbhmVTs2drYj1GUXZTYTcbUgJjh9Nv41aFfZ21bPQTERm5+L5cBDo466NltB2lplS5GBw==", + "dev": true, + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/csstools" + }, + { + "type": "opencollective", + "url": "https://opencollective.com/csstools" + } + ], + "license": "MIT", + "engines": { + "node": ">=20.19.0" + } + }, "node_modules/@electron-internal/extract-zip": { "version": "1.0.5", "resolved": "https://registry.npmjs.org/@electron-internal/extract-zip/-/extract-zip-1.0.5.tgz", @@ -1288,6 +1493,24 @@ "node": "^18.18.0 || ^20.9.0 || >=21.1.0" } }, + "node_modules/@exodus/bytes": { + "version": "1.16.0", + "resolved": "https://registry.npmjs.org/@exodus/bytes/-/bytes-1.16.0.tgz", + "integrity": "sha512-IcpW84uEn3N7ETtNZMlxKhfl6Pec8rUNGOTBtWbK1FKhJxIFAptZyVrvVRVBimAJxJCgc3PxepxkdWWG4DVzfA==", + "dev": true, + "license": "MIT", + "engines": { + "node": "^20.19.0 || ^22.12.0 || >=24.0.0" + }, + "peerDependencies": { + "@noble/hashes": "^1.8.0 || ^2.0.0" + }, + "peerDependenciesMeta": { + "@noble/hashes": { + "optional": true + } + } + }, "node_modules/@humanfs/core": { "version": "0.19.1", "resolved": "https://registry.npmjs.org/@humanfs/core/-/core-0.19.1.tgz", @@ -3787,6 +4010,20 @@ "node": ">= 8" } }, + "node_modules/css-tree": { + "version": "3.2.1", + "resolved": "https://registry.npmjs.org/css-tree/-/css-tree-3.2.1.tgz", + "integrity": "sha512-X7sjQzceUhu1u7Y/ylrRZFU2FS6LRiFVp6rKLPg23y3x3c3DOKAwuXGDp+PAGjh6CSnCjYeAul8pcT8bAl+lSA==", + "dev": true, + "license": "MIT", + "dependencies": { + "mdn-data": "2.27.1", + "source-map-js": "^1.2.1" + }, + "engines": { + "node": "^10 || ^12.20.0 || ^14.13.0 || >=15.0.0" + } + }, "node_modules/cssesc": { "version": "3.0.0", "resolved": "https://registry.npmjs.org/cssesc/-/cssesc-3.0.0.tgz", @@ -3909,6 +4146,20 @@ "node": ">=12" } }, + "node_modules/data-urls": { + "version": "7.0.0", + "resolved": "https://registry.npmjs.org/data-urls/-/data-urls-7.0.0.tgz", + "integrity": "sha512-23XHcCF+coGYevirZceTVD7NdJOqVn+49IHyxgszm+JIiHLoB2TkmPtsYkNWT1pvRSGkc35L6NHs0yHkN2SumA==", + "dev": true, + "license": "MIT", + "dependencies": { + "whatwg-mimetype": "^5.0.0", + "whatwg-url": "^16.0.0" + }, + "engines": { + "node": "^20.19.0 || ^22.12.0 || >=24.0.0" + } + }, "node_modules/debug": { "version": "4.4.3", "resolved": "https://registry.npmjs.org/debug/-/debug-4.4.3.tgz", @@ -3925,6 +4176,13 @@ } } }, + "node_modules/decimal.js": { + "version": "10.6.0", + "resolved": "https://registry.npmjs.org/decimal.js/-/decimal.js-10.6.0.tgz", + "integrity": "sha512-YpgQiITW3JXGntzdUmyUR1V812Hn8T1YVXhCu+wO3OpS4eU9l4YdD3qjyiKdV6mvV29zapkMeD390UVEf2lkUg==", + "dev": true, + "license": "MIT" + }, "node_modules/decompress-response": { "version": "6.0.0", "resolved": "https://registry.npmjs.org/decompress-response/-/decompress-response-6.0.0.tgz", @@ -4166,9 +4424,9 @@ } }, "node_modules/electron": { - "version": "42.9.3", - "resolved": "https://registry.npmjs.org/electron/-/electron-42.9.3.tgz", - "integrity": "sha512-REQUgPrCWOP0FajNcKCwwjkssirN+MVbemyd6bT+x51OVq58y6IqjtpA4vmm/KEzKXj2MIOebx/gF497CkRAPg==", + "version": "42.11.1", + "resolved": "https://registry.npmjs.org/electron/-/electron-42.11.1.tgz", + "integrity": "sha512-mRYYjGDRWCyU+h4FU/0ruqxL/rXc+h2VCLhTb19KHu18HCUFWQAeFS/mFwnEKT/7GQCwlIHOhHie5ICLPXW6aw==", "license": "MIT", "dependencies": { "@electron-internal/extract-zip": "^1.0.1", @@ -4396,6 +4654,19 @@ "once": "^1.4.0" } }, + "node_modules/entities": { + "version": "8.1.0", + "resolved": "https://registry.npmjs.org/entities/-/entities-8.1.0.tgz", + "integrity": "sha512-kxL7msIffSuh9aaFAMD7rxAIuTRMAHMeBtgHW2yUdWw732ZNh4MehkF2gdjvtdmikkaIP9bFDDJOPlsvm7avrA==", + "dev": true, + "license": "BSD-2-Clause", + "engines": { + "node": ">=20.19.0" + }, + "funding": { + "url": "https://github.com/fb55/entities?sponsor=1" + } + }, "node_modules/env-paths": { "version": "3.0.0", "resolved": "https://registry.npmjs.org/env-paths/-/env-paths-3.0.0.tgz", @@ -5352,6 +5623,19 @@ "dev": true, "license": "ISC" }, + "node_modules/html-encoding-sniffer": { + "version": "6.0.0", + "resolved": "https://registry.npmjs.org/html-encoding-sniffer/-/html-encoding-sniffer-6.0.0.tgz", + "integrity": "sha512-CV9TW3Y3f8/wT0BRFc1/KAVQ3TUHiXmaAb6VW9vtiMFf7SLoMd1PdAc4W3KFOFETBJUb90KatHqlsZMWV+R9Gg==", + "dev": true, + "license": "MIT", + "dependencies": { + "@exodus/bytes": "^1.6.0" + }, + "engines": { + "node": "^20.19.0 || ^22.12.0 || >=24.0.0" + } + }, "node_modules/http-cache-semantics": { "version": "4.2.0", "resolved": "https://registry.npmjs.org/http-cache-semantics/-/http-cache-semantics-4.2.0.tgz", @@ -5545,6 +5829,13 @@ "node": ">=0.12.0" } }, + "node_modules/is-potential-custom-element-name": { + "version": "1.0.1", + "resolved": "https://registry.npmjs.org/is-potential-custom-element-name/-/is-potential-custom-element-name-1.0.1.tgz", + "integrity": "sha512-bCYeRA2rVibKZd+s2625gGnGF/t7DSqDs4dP7CrLA1m7jKWz6pps0LpYLJN8Q64HtmPKJ1hrN3nzPNKFEKOUiQ==", + "dev": true, + "license": "MIT" + }, "node_modules/is-promise": { "version": "2.2.2", "resolved": "https://registry.npmjs.org/is-promise/-/is-promise-2.2.2.tgz", @@ -5648,6 +5939,57 @@ "js-yaml": "bin/js-yaml.js" } }, + "node_modules/jsdom": { + "version": "29.1.1", + "resolved": "https://registry.npmjs.org/jsdom/-/jsdom-29.1.1.tgz", + "integrity": "sha512-ECi4Fi2f7BdJtUKTflYRTiaMxIB0O6zfR1fX0GXpUrf6flp8QIYn1UT20YQqdSOfk2dfkCwS8LAFoJDEppNK5Q==", + "dev": true, + "license": "MIT", + "dependencies": { + "@asamuzakjp/css-color": "^5.1.11", + "@asamuzakjp/dom-selector": "^7.1.1", + "@bramus/specificity": "^2.4.2", + "@csstools/css-syntax-patches-for-csstree": "^1.1.3", + "@exodus/bytes": "^1.15.0", + "css-tree": "^3.2.1", + "data-urls": "^7.0.0", + "decimal.js": "^10.6.0", + "html-encoding-sniffer": "^6.0.0", + "is-potential-custom-element-name": "^1.0.1", + "lru-cache": "^11.3.5", + "parse5": "^8.0.1", + "saxes": "^6.0.0", + "symbol-tree": "^3.2.4", + "tough-cookie": "^6.0.1", + "undici": "^7.25.0", + "w3c-xmlserializer": "^5.0.0", + "webidl-conversions": "^8.0.1", + "whatwg-mimetype": "^5.0.0", + "whatwg-url": "^16.0.1", + "xml-name-validator": "^5.0.0" + }, + "engines": { + "node": "^20.19.0 || ^22.13.0 || >=24.0.0" + }, + "peerDependencies": { + "canvas": "^3.0.0" + }, + "peerDependenciesMeta": { + "canvas": { + "optional": true + } + } + }, + "node_modules/jsdom/node_modules/lru-cache": { + "version": "11.5.3", + "resolved": "https://registry.npmjs.org/lru-cache/-/lru-cache-11.5.3.tgz", + "integrity": "sha512-U4N8FgzmWxc8k1VH8Kr6lQg18U7Fjvby6wXHVRX/ZZ7IwWbRMgrRbP0Wrb5q5NVinryp4SQampHKdvtecItxUg==", + "dev": true, + "license": "BlueOak-1.0.0", + "engines": { + "node": "20 || >=22" + } + }, "node_modules/jsesc": { "version": "3.1.0", "resolved": "https://registry.npmjs.org/jsesc/-/jsesc-3.1.0.tgz", @@ -5877,6 +6219,13 @@ "node": ">= 0.4" } }, + "node_modules/mdn-data": { + "version": "2.27.1", + "resolved": "https://registry.npmjs.org/mdn-data/-/mdn-data-2.27.1.tgz", + "integrity": "sha512-9Yubnt3e8A0OKwxYSXyhLymGW4sCufcLG6VdiDdUGVkPhpqLxlvP5vl1983gQjJl3tqbrM731mjaZaP68AgosQ==", + "dev": true, + "license": "CC0-1.0" + }, "node_modules/merge2": { "version": "1.4.1", "resolved": "https://registry.npmjs.org/merge2/-/merge2-1.4.1.tgz", @@ -6336,6 +6685,19 @@ "node": ">=6" } }, + "node_modules/parse5": { + "version": "8.0.1", + "resolved": "https://registry.npmjs.org/parse5/-/parse5-8.0.1.tgz", + "integrity": "sha512-z1e/HMG90obSGeidlli3hj7cbocou0/wa5HacvI3ASx34PecNjNQeaHNo5WIZpWofN9kgkqV1q5YvXe3F0FoPw==", + "dev": true, + "license": "MIT", + "dependencies": { + "entities": "^8.0.0" + }, + "funding": { + "url": "https://github.com/inikulin/parse5?sponsor=1" + } + }, "node_modules/path-exists": { "version": "4.0.0", "resolved": "https://registry.npmjs.org/path-exists/-/path-exists-4.0.0.tgz", @@ -7197,6 +7559,19 @@ "node": ">=11.0.0" } }, + "node_modules/saxes": { + "version": "6.0.0", + "resolved": "https://registry.npmjs.org/saxes/-/saxes-6.0.0.tgz", + "integrity": "sha512-xAg7SOnEhrm5zI3puOOKyy1OMcMlIJZYNJY7xLBwSze0UjhPLnWfj2GF2EpT0jmzaJKIWKHLsaSSajf35bcYnA==", + "dev": true, + "license": "ISC", + "dependencies": { + "xmlchars": "^2.2.0" + }, + "engines": { + "node": ">=v12.22.7" + } + }, "node_modules/scheduler": { "version": "0.21.0", "resolved": "https://registry.npmjs.org/scheduler/-/scheduler-0.21.0.tgz", @@ -7486,6 +7861,13 @@ "react": ">=17.0" } }, + "node_modules/symbol-tree": { + "version": "3.2.4", + "resolved": "https://registry.npmjs.org/symbol-tree/-/symbol-tree-3.2.4.tgz", + "integrity": "sha512-9QNk5KwDF+Bvz+PyObkmSYjI5ksVUYtjW7AU22r2NKcfLJcXp96hkDWU3+XndOsUb+AQ9QhfzfCT2O+CNWT5Tw==", + "dev": true, + "license": "MIT" + }, "node_modules/tailwindcss": { "version": "3.4.19", "resolved": "https://registry.npmjs.org/tailwindcss/-/tailwindcss-3.4.19.tgz", @@ -7728,6 +8110,26 @@ "url": "https://github.com/sponsors/jonschlinkert" } }, + "node_modules/tldts": { + "version": "7.4.16", + "resolved": "https://registry.npmjs.org/tldts/-/tldts-7.4.16.tgz", + "integrity": "sha512-QwBER5KMR86IIjpIiO7H/Z3IMJPsZ1A6RKPAqzTTgOyUQUSt9FdnKcqhTaJmkY6HVrgouZHZR0ncK5QxvmnQeg==", + "dev": true, + "license": "MIT", + "dependencies": { + "tldts-core": "^7.4.16" + }, + "bin": { + "tldts": "bin/cli.js" + } + }, + "node_modules/tldts-core": { + "version": "7.4.16", + "resolved": "https://registry.npmjs.org/tldts-core/-/tldts-core-7.4.16.tgz", + "integrity": "sha512-MDolfaSJtlSK5Y0A1xl3277ekubZwobpBjugknDizI9O5Rm60a1m8k4ICK+MRsCDzPygT81mp3BBf5RKDlFRfA==", + "dev": true, + "license": "MIT" + }, "node_modules/tmp": { "version": "0.2.7", "resolved": "https://registry.npmjs.org/tmp/-/tmp-0.2.7.tgz", @@ -7760,6 +8162,32 @@ "node": ">=8.0" } }, + "node_modules/tough-cookie": { + "version": "6.0.2", + "resolved": "https://registry.npmjs.org/tough-cookie/-/tough-cookie-6.0.2.tgz", + "integrity": "sha512-exgYmnmL/sJpR3upZfXG5PoatXQii55xAiXGXzY+sROLZ/Y+SLcp9PgJNI9Vz37HpQ74WvDcLT8eqm+kV3FzrA==", + "dev": true, + "license": "BSD-3-Clause", + "dependencies": { + "tldts": "^7.0.5" + }, + "engines": { + "node": ">=16" + } + }, + "node_modules/tr46": { + "version": "6.0.0", + "resolved": "https://registry.npmjs.org/tr46/-/tr46-6.0.0.tgz", + "integrity": "sha512-bLVMLPtstlZ4iMQHpFHTR7GAGj2jxi8Dg0s2h2MafAE4uSWF98FC/3MomU51iQAMf8/qDUbKWf5GxuvvVcXEhw==", + "dev": true, + "license": "MIT", + "dependencies": { + "punycode": "^2.3.1" + }, + "engines": { + "node": ">=20" + } + }, "node_modules/troika-three-text": { "version": "0.52.4", "resolved": "https://registry.npmjs.org/troika-three-text/-/troika-three-text-0.52.4.tgz", @@ -7925,8 +8353,8 @@ "version": "7.29.0", "resolved": "https://registry.npmjs.org/undici/-/undici-7.29.0.tgz", "integrity": "sha512-IDxfleLmmbSskfWSUATiN1nfn2rDuvnMOqb5CWR92iIfojA0Ud+ulOAAEQ57LPr9rWmsreUyf5lwyao+7GNNVw==", + "devOptional": true, "license": "MIT", - "optional": true, "engines": { "node": ">=20.18.1" } @@ -8149,6 +8577,19 @@ "url": "https://github.com/sponsors/jonschlinkert" } }, + "node_modules/w3c-xmlserializer": { + "version": "5.0.0", + "resolved": "https://registry.npmjs.org/w3c-xmlserializer/-/w3c-xmlserializer-5.0.0.tgz", + "integrity": "sha512-o8qghlI8NZHU1lLPrpi2+Uq7abh4GGPpYANlalzWxyWteJOCsr/P+oPBA49TOLu5FTZO4d3F9MnWJfiMo4BkmA==", + "dev": true, + "license": "MIT", + "dependencies": { + "xml-name-validator": "^5.0.0" + }, + "engines": { + "node": ">=18" + } + }, "node_modules/webcrypto-core": { "version": "1.9.2", "resolved": "https://registry.npmjs.org/webcrypto-core/-/webcrypto-core-1.9.2.tgz", @@ -8173,6 +8614,41 @@ "resolved": "https://registry.npmjs.org/webgl-sdf-generator/-/webgl-sdf-generator-1.1.1.tgz", "integrity": "sha512-9Z0JcMTFxeE+b2x1LJTdnaT8rT8aEp7MVxkNwoycNmJWwPdzoXzMh0BjJSh/AEFP+KPYZUli814h8bJZFIZ2jA==" }, + "node_modules/webidl-conversions": { + "version": "8.0.1", + "resolved": "https://registry.npmjs.org/webidl-conversions/-/webidl-conversions-8.0.1.tgz", + "integrity": "sha512-BMhLD/Sw+GbJC21C/UgyaZX41nPt8bUTg+jWyDeg7e7YN4xOM05YPSIXceACnXVtqyEw/LMClUQMtMZ+PGGpqQ==", + "dev": true, + "license": "BSD-2-Clause", + "engines": { + "node": ">=20" + } + }, + "node_modules/whatwg-mimetype": { + "version": "5.0.0", + "resolved": "https://registry.npmjs.org/whatwg-mimetype/-/whatwg-mimetype-5.0.0.tgz", + "integrity": "sha512-sXcNcHOC51uPGF0P/D4NVtrkjSU2fNsm9iog4ZvZJsL3rjoDAzXZhkm2MWt1y+PUdggKAYVoMAIYcs78wJ51Cw==", + "dev": true, + "license": "MIT", + "engines": { + "node": ">=20" + } + }, + "node_modules/whatwg-url": { + "version": "16.0.1", + "resolved": "https://registry.npmjs.org/whatwg-url/-/whatwg-url-16.0.1.tgz", + "integrity": "sha512-1to4zXBxmXHV3IiSSEInrreIlu02vUOvrhxJJH5vcxYTBDAx51cqZiKdyTxlecdKNSjj8EcxGBxNf6Vg+945gw==", + "dev": true, + "license": "MIT", + "dependencies": { + "@exodus/bytes": "^1.11.0", + "tr46": "^6.0.0", + "webidl-conversions": "^8.0.1" + }, + "engines": { + "node": "^20.19.0 || ^22.12.0 || >=24.0.0" + } + }, "node_modules/which": { "version": "2.0.2", "resolved": "https://registry.npmjs.org/which/-/which-2.0.2.tgz", @@ -8221,6 +8697,16 @@ "dev": true, "license": "ISC" }, + "node_modules/xml-name-validator": { + "version": "5.0.0", + "resolved": "https://registry.npmjs.org/xml-name-validator/-/xml-name-validator-5.0.0.tgz", + "integrity": "sha512-EvGK8EJ3DhaHfbRlETOWAS5pO9MZITeauHKJyb8wyajUfQUenkIg2MvLDTZ4T/TgIcm3HU0TFBgWWboAZ30UHg==", + "dev": true, + "license": "Apache-2.0", + "engines": { + "node": ">=18" + } + }, "node_modules/xmlbuilder": { "version": "15.1.1", "resolved": "https://registry.npmjs.org/xmlbuilder/-/xmlbuilder-15.1.1.tgz", @@ -8231,6 +8717,13 @@ "node": ">=8.0" } }, + "node_modules/xmlchars": { + "version": "2.2.0", + "resolved": "https://registry.npmjs.org/xmlchars/-/xmlchars-2.2.0.tgz", + "integrity": "sha512-JZnDKK8B0RCDw84FNdDAIpZK+JuJw+s7Lz8nksI7SIuU3UXJJslUthsi+uWBUYOwPFwW7W7PRLRfUKpxjtjFCw==", + "dev": true, + "license": "MIT" + }, "node_modules/y18n": { "version": "5.0.8", "resolved": "https://registry.npmjs.org/y18n/-/y18n-5.0.8.tgz", diff --git a/package.json b/package.json index df135498..8c3a985b 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "modly", - "version": "0.4.2", + "version": "0.4.3", "description": "Local AI-powered 3D mesh generation from images", "main": "./out/main/index.js", "author": "Modly", @@ -13,7 +13,7 @@ "prepare-resources": "node scripts/download-python-embed.js", "test": "npm run test:py && npm run test:node", "test:py": "node scripts/run-pytests.mjs", - "test:node": "node --test --experimental-strip-types --experimental-loader ./scripts/node-ts-extensionless-loader.mjs src/shared/types/assetLibrary.test.ts src/areas/generate/assetLibraryProjection.test.ts src/areas/generate/assetLibraryService.test.ts src/areas/generate/assetLibraryUi.test.ts electron/main/artifact-registry-service.test.ts electron/main/extension-path-guard.test.ts electron/preload/artifact-registry-preload.test.ts && node --test electron/main/*.test.mjs src/**/*.test.mjs", + "test:node": "node --test --experimental-strip-types --experimental-loader ./scripts/node-ts-extensionless-loader.mjs src/shared/types/assetLibrary.test.ts src/areas/generate/assetLibraryProjection.test.ts src/areas/generate/assetLibraryService.test.ts src/areas/generate/assetLibraryUi.test.ts src/areas/generate/orcaSlicerLink.test.ts electron/main/artifact-registry-service.test.ts electron/main/extension-path-guard.test.ts electron/preload/artifact-registry-preload.test.ts && node --test electron/main/*.test.mjs src/**/*.test.mjs", "package": "cross-env CSC_IDENTITY_AUTO_DISCOVERY=false npm run build && npm run prepare-resources && electron-builder", "package:mac": "cross-env CSC_IDENTITY_AUTO_DISCOVERY=false npm run build && npm run prepare-resources && electron-builder --mac --arm64", "lint": "eslint ." @@ -45,12 +45,13 @@ "@vitejs/plugin-react": "^4.3.4", "autoprefixer": "^10.4.20", "cross-env": "^10.1.0", - "electron": "^42.4.1", + "electron": "^42.11.1", "electron-builder": "^26.15.3", "electron-vite": "^5.0.0", "eslint": "^9.17.0", "eslint-plugin-react-hooks": "^7.1.1", "globals": "^17.6.0", + "jsdom": "^29.1.1", "postcss": "^8.4.49", "tailwindcss": "^3.4.17", "typescript": "^5.7.2", diff --git a/scripts/react-test-env.mjs b/scripts/react-test-env.mjs new file mode 100644 index 00000000..d2bd8d2e --- /dev/null +++ b/scripts/react-test-env.mjs @@ -0,0 +1,103 @@ +/** + * Minimal React test environment for `node --test`. + * + * The repo tests everything except the UI layer, which is exactly where the + * bugs the user actually sees have come from: a hook returning a fresh callback + * on every render made a modal re-run its effect forever, flickering and + * hammering the API. That class of bug is invisible to a type-checker and to + * store-level tests — it only exists once a component renders more than once. + * + * Deliberately built on jsdom + react-dom directly rather than a testing + * library: what we need is "render, re-render, count the effects", not queries + * and user events. + */ +import { JSDOM } from 'jsdom' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdirSync, mkdtempSync, writeFileSync } from 'node:fs' +import { join, resolve } from 'node:path' + +/** Install a DOM into the globals React expects. Call once, at module scope. */ +export function setupDom() { + const dom = new JSDOM('', { url: 'http://localhost/' }) + const { window } = dom + + // Node 24 defines some of these (`navigator`) as getter-only on globalThis, + // so a plain assignment throws — go through defineProperty for all of them. + const define = (name, value) => + Object.defineProperty(globalThis, name, { value, writable: true, configurable: true }) + + define('window', window) + define('document', window.document) + define('navigator', window.navigator) + define('localStorage', window.localStorage) + define('requestAnimationFrame', (cb) => setTimeout(() => cb(Date.now()), 0)) + define('cancelAnimationFrame', (id) => clearTimeout(id)) + // React 18 refuses to run `act` without it, and warns on every update. + define('IS_REACT_ACT_ENVIRONMENT', true) + for (const name of ['HTMLElement', 'Element', 'Node', 'Event', 'MouseEvent', 'getComputedStyle']) { + define(name, window[name]) + } + return dom +} + +/** + * Bundle a source module and load it as CommonJS. + * + * `react` and `react-dom` stay external so the module under test and the test + * file share one React instance — two copies produce "invalid hook call", + * which reads as a bug in the code under test and is not one. + */ +export function loadModule(entryPath) { + // Inside the project, not the system temp dir: the bundle keeps `react` as a + // bare require, and that only resolves from a path under this node_modules. + const cacheRoot = resolve('node_modules/.cache/modly-react-tests') + mkdirSync(cacheRoot, { recursive: true }) + const outfile = join(mkdtempSync(join(cacheRoot, 'm-')), 'module.cjs') + const result = buildSync({ + entryPoints: [resolve(entryPath)], + bundle: true, + platform: 'node', + format: 'cjs', + jsx: 'automatic', + external: ['react', 'react-dom', 'react-dom/client', 'react/jsx-runtime'], + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return createRequire(import.meta.url)(outfile) +} + +/** + * Render `element`, and hand back a way to re-render it with the same root — + * which is the whole point: a hook that misbehaves does so on the SECOND render. + */ +export async function mount(element) { + const require = createRequire(import.meta.url) + const { createRoot } = require('react-dom/client') + const { act } = require('react-dom/test-utils') + const { cloneElement } = require('react') + + // Through globalThis: these are globals this module installed itself in + // setupDom(), not ambient browser ones. + const container = globalThis.document.createElement('div') + globalThis.document.body.appendChild(container) + const root = createRoot(container) + + await act(async () => { root.render(element) }) + + return { + container, + /** Re-render. The element is cloned by default: handed the very same + * element object, React bails out and nothing re-renders — which silently + * turns a re-render test into a no-op. */ + rerender: async (next) => { + await act(async () => { root.render(next ?? cloneElement(element)) }) + }, + /** Let effects, promises and state updates settle. */ + flush: async () => { await act(async () => { await Promise.resolve() }) }, + unmount: async () => { + await act(async () => { root.unmount() }) + container.remove() + }, + } +} diff --git a/src/areas/generate/GeneratePage.tsx b/src/areas/generate/GeneratePage.tsx index a6dbe4f3..76a13119 100644 --- a/src/areas/generate/GeneratePage.tsx +++ b/src/areas/generate/GeneratePage.tsx @@ -1,13 +1,14 @@ import { useState, useRef, useCallback, useEffect, useMemo } from 'react' import type { ReactNode } from 'react' import { useAppStore, DEFAULT_LIGHT_SETTINGS } from '@shared/stores/appStore' -import type { GenerationJob, LightSettings } from '@shared/stores/appStore' +import type { GenerationJob, LightSettings, PointLight } from '@shared/stores/appStore' import { useApi } from '@shared/hooks/useApi' import { ColorPicker } from '@shared/components/ui' import GenerationHUD from './components/GenerationHUD' import Viewer3D from './components/Viewer3D' import WorkflowPanel from './components/WorkflowPanel' import { getDefaultAssetLibraryService } from './assetLibraryService' +import { buildOrcaSlicerDeepLink, canOpenInOrcaSlicer } from './orcaSlicerLink' import { resolveAssetLibraryOpenTarget, type ProjectedAssetLibraryEntry } from './assetLibraryProjection' import { ASSET_LIBRARY_SORT_OPTIONS, @@ -27,6 +28,18 @@ const MIN_WIDTH = 220 const MAX_WIDTH = 520 const DEFAULT_WIDTH = 320 +const MAX_POINT_LIGHTS = 6 + +function createPointLight(): PointLight { + const angle = Math.random() * Math.PI * 2 + return { + id: crypto.randomUUID(), + position: [Math.cos(angle) * 1.5, 0.5, Math.sin(angle) * 1.5], + color: '#ffffff', + intensity: 1, + } +} + // --------------------------------------------------------------------------- // Export dropdown // --------------------------------------------------------------------------- @@ -41,9 +54,13 @@ const EXPORT_FORMATS = [ function ExportDropdown({ onExport, onClose, + onOpenInSlicer, + canOpenInSlicer, }: { onExport: (f: 'glb' | 'obj' | 'stl' | 'ply') => void onClose: () => void + onOpenInSlicer: () => void + canOpenInSlicer: boolean }) { return (
@@ -57,6 +74,23 @@ function ExportDropdown({ {desc} ))} + {canOpenInSlicer && ( + <> +
+ + + )}
) } @@ -179,10 +213,18 @@ function LightPopover({ settings, onChange, onClose, + pointLights, + onPointLightsChange, + selectedPointLightId, + onSelectPointLight, }: { settings: LightSettings onChange: (s: LightSettings) => void onClose: () => void + pointLights: PointLight[] + onPointLightsChange: (lights: PointLight[]) => void + selectedPointLightId: string | null + onSelectPointLight: (id: string | null) => void }) { function lightRow( label: string, @@ -251,6 +293,64 @@ function LightPopover({ {lightRow('Fill', 'fillColor', 'fillIntensity', 2)} {plainRow('Ambient', 'ambientIntensity', 1.5)} {plainRow('Environment', 'envIntensity', 2)} + + {/* ── Point lights ── */} +
+
+

Point lights

+ {pointLights.length < MAX_POINT_LIGHTS && ( + + )} +
+ + {pointLights.length === 0 && ( +

No point lights yet.

+ )} + + {pointLights.map((pl) => ( +
onSelectPointLight(pl.id)} + className={`flex flex-col gap-1.5 p-2 rounded-lg bg-zinc-800/40 border cursor-pointer transition-colors ${ + pl.id === selectedPointLightId + ? 'border-violet-500' + : 'border-zinc-700/40 hover:border-zinc-600' + }`} + > +
+ onPointLightsChange(pointLights.map((p) => p.id === pl.id ? { ...p, color: c } : p))} + /> + Point + {pl.intensity.toFixed(1)} + +
+ onPointLightsChange(pointLights.map((p) => p.id === pl.id ? { ...p, intensity: parseFloat(e.target.value) } : p))} + className="w-full h-1.5 accent-violet-500 cursor-pointer" /> +
+ ))} +
+
@@ -1073,6 +1206,10 @@ export default function GeneratePage(): JSX.Element { settings={lightSettings} onChange={setLightSettings} onClose={() => setOpenPanel(null)} + pointLights={pointLights} + onPointLightsChange={setPointLights} + selectedPointLightId={selectedPointLightId} + onSelectPointLight={handleSelectPointLight} /> )} @@ -1080,7 +1217,7 @@ export default function GeneratePage(): JSX.Element { {/* Tools bar — always visible; transform tools appear once a mesh is selected */}
- {hasModel && meshSelected && ( + {(meshSelected || selectedPointLightId) && ( <> - +
diff --git a/src/areas/generate/components/ChatPanel.tsx b/src/areas/generate/components/ChatPanel.tsx index 7618f75d..053e780a 100644 --- a/src/areas/generate/components/ChatPanel.tsx +++ b/src/areas/generate/components/ChatPanel.tsx @@ -1,6 +1,7 @@ import { useEffect, useMemo, useRef, useState } from 'react' import { useAppStore } from '@shared/stores/appStore' -import { useAgentStore } from '@shared/stores/agentStore' +import { useAgentStore, PROVIDERS, providerBaseUrl } from '@shared/stores/agentStore' +import { useLlmModels } from '@shared/stores/llmModelsStore' import { useWorkflowsStore } from '@shared/stores/workflowsStore' import { useExtensionsStore } from '@shared/stores/extensionsStore' import { useWorkflowRunStore } from '@areas/workflows/workflowRunStore' @@ -242,16 +243,22 @@ function WorkflowProgressCard({ name }: { name: string }): JSX.Element { // ─── Main component ────────────────────────────────────────────────────────── export default function ChatPanel(): JSX.Element { - const { ollamaUrl, defaultModel, defaultThinking } = useAgentStore() + const { provider, localModel, external, defaultThinking } = useAgentStore() + const { models: llmModels } = useLlmModels() + const isLocal = provider === 'local' + const externalConfig = external[provider] const [messages, setMessages] = useState([]) const [input, setInput] = useState('') const [isLoading, setIsLoading] = useState(false) const [error, setError] = useState(null) const [showAll, setShowAll] = useState(false) - const [model, setModel] = useState(defaultModel) + // null = follow the default from Settings, which hydrates asynchronously. + const [localPick, setLocalPick] = useState(null) + const model = isLocal ? (localPick ?? localModel) : (externalConfig?.model ?? '') + // code/cad models are node tools, not chat models. + const localModels = llmModels.filter((m) => m.downloaded && !(m.tags ?? []).some((t) => t === 'code' || t === 'cad')) const [showModelPicker, setShowModelPicker] = useState(false) - const [ollamaModels, setOllamaModels] = useState([]) const [pendingWorkflow, setPendingWorkflow] = useState<{ id: string; name: string } | null>(null) const [attachments, setAttachments] = useState([]) // data URLs const [isDragging, setIsDragging] = useState(false) @@ -349,7 +356,7 @@ export default function ChatPanel(): JSX.Element { content: m.content, } if (m.imageDataUrls?.length) { - entry.images = m.imageDataUrls.map((url) => url.split(',')[1]) + entry.images = m.imageDataUrls } return entry }) @@ -361,7 +368,15 @@ export default function ChatPanel(): JSX.Element { const res = await fetch(`${apiUrl}/agent/chat`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ messages: apiMessages, ollama_url: ollamaUrl, model, context, thinking: thinkingMode }), + body: JSON.stringify({ + messages: apiMessages, + model, + provider: isLocal + ? { type: 'local' } + : { type: 'external', base_url: providerBaseUrl(provider, externalConfig), api_key: externalConfig?.apiKey ?? '' }, + context, + thinking: thinkingMode, + }), }) if (!res.ok) throw new Error(`API error ${res.status}`) @@ -406,16 +421,6 @@ export default function ChatPanel(): JSX.Element { } } - async function fetchOllamaModels() { - try { - const res = await fetch(`${apiUrl}/agent/models?ollama_url=${encodeURIComponent(ollamaUrl)}`) - const data = await res.json() - setOllamaModels(data.models ?? []) - } catch { - setOllamaModels([]) - } - } - function handleFiles(files: File[]) { files.forEach((file) => { if (!file.type.startsWith('image/')) return @@ -646,10 +651,10 @@ export default function ChatPanel(): JSX.Element { {/* Model selector */}
+ )}
) diff --git a/src/areas/generate/orcaSlicerLink.test.ts b/src/areas/generate/orcaSlicerLink.test.ts new file mode 100644 index 00000000..514618eb --- /dev/null +++ b/src/areas/generate/orcaSlicerLink.test.ts @@ -0,0 +1,82 @@ +import assert from 'node:assert/strict' +import test from 'node:test' + +import { + SLICER_FORMAT, + buildOrcaSlicerDeepLink, + canOpenInOrcaSlicer, + encodeWorkspacePathToken, +} from './orcaSlicerLink.ts' + +test('builds an orcaslicer://open deeplink whose file= is a percent-encoded, query-less URL ending in model.stl', () => { + const link = buildOrcaSlicerDeepLink('http://localhost:8765', '/workspace/Workflows/checkpoints/hero.glb') + assert.ok(link.startsWith('orcaslicer://open?file=')) + const modelUrl = decodeURIComponent(link.slice('orcaslicer://open?file='.length)) + // OrcaSlicer derives the import format from the URL's final path segment, so + // it must end in the real extension and carry no query string. + assert.ok(!modelUrl.includes('?'), 'model URL must not contain a query string') + assert.ok(modelUrl.endsWith('/model.stl'), 'model URL must end in model.stl') + assert.equal( + modelUrl, + `http://localhost:8765/export/slicer/stl/${encodeWorkspacePathToken('Workflows/checkpoints/hero.glb')}/model.stl`, + ) +}) + +test('token round-trips a workspace path through url-safe base64 (matches the API decode)', () => { + const path = 'Workflows/checkpoints/hero model (v2).glb' + const token = encodeWorkspacePathToken(path) + assert.ok(!/[+/=]/.test(token), 'token must be url-safe with no padding') + // Decode the way the Python API does: restore padding, then urlsafe-decode. + const padded = token + '='.repeat((4 - (token.length % 4)) % 4) + const decoded = Buffer.from(padded.replace(/-/g, '+').replace(/_/g, '/'), 'base64').toString('utf-8') + assert.equal(decoded, path) +}) + +test('strips a trailing slash from the api origin', () => { + const link = buildOrcaSlicerDeepLink('http://localhost:8765/', '/workspace/a.glb') + const modelUrl = decodeURIComponent(link.slice('orcaslicer://open?file='.length)) + assert.equal(modelUrl, `http://localhost:8765/export/slicer/stl/${encodeWorkspacePathToken('a.glb')}/model.stl`) +}) + +test('canOpenInOrcaSlicer accepts sliceable workspace meshes and rejects splats and empty', () => { + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/hero.glb'), true) + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/part.stl'), true) + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/part.obj'), true) + // Gaussian splats are point clouds, not printable meshes. + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/scan.ply'), false) + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/scan.splat'), false) + // A format the slicer route cannot convert. + assert.equal(canOpenInOrcaSlicer('/workspace/Foo/rig.fbx'), false) + assert.equal(canOpenInOrcaSlicer(undefined), false) + assert.equal(canOpenInOrcaSlicer(''), false) +}) + +test('canOpenInOrcaSlicer accepts an imported mesh served from outside the workspace', () => { + // `Import` serves user-picked files through /optimize/serve-file, never from + // the workspace — excluding those made the action invisible for the most + // direct "I have a model, slice it" path. + assert.equal(canOpenInOrcaSlicer('/optimize/serve-file?path=%2Ftmp%2Fx.glb'), true) + assert.equal(canOpenInOrcaSlicer('/optimize/serve-file?path=C%3A%5CUsers%5CMe%5Cmage.glb'), true) + assert.equal(canOpenInOrcaSlicer('/optimize/serve-file?path=%2Ftmp%2Fscan.splat'), false) +}) + +test('deeplink for an imported mesh tokenises the decoded absolute path', () => { + const absolute = String.raw`C:\Users\Me\Downloads\mage_v2.glb` + const outputUrl = `/optimize/serve-file?path=${encodeURIComponent(absolute)}` + const modelUrl = decodeURIComponent( + buildOrcaSlicerDeepLink('http://localhost:8765', outputUrl).slice('orcaslicer://open?file='.length), + ) + assert.equal( + modelUrl, + `http://localhost:8765/export/slicer/stl/${encodeWorkspacePathToken(absolute)}/model.stl`, + ) + assert.ok(!modelUrl.includes('?'), 'model URL must not contain a query string') +}) + +test('buildOrcaSlicerDeepLink refuses an output it cannot slice', () => { + assert.throws(() => buildOrcaSlicerDeepLink('http://localhost:8765', '/workspace/Foo/scan.splat')) +}) + +test('SLICER_FORMAT is a format OrcaSlicer can import', () => { + assert.equal(SLICER_FORMAT, 'stl') +}) diff --git a/src/areas/generate/orcaSlicerLink.ts b/src/areas/generate/orcaSlicerLink.ts new file mode 100644 index 00000000..d5a76c9b --- /dev/null +++ b/src/areas/generate/orcaSlicerLink.ts @@ -0,0 +1,79 @@ +// Builds the OrcaSlicer deeplink for a generated mesh. +// +// OrcaSlicer registers the `orcaslicer://open?file=` scheme; its handler +// downloads the http(s) URL in `file=` and imports it, deriving the filename — +// and therefore the mesh format — from the URL's FINAL path segment. That means +// the served URL must be path-only and end in a real `model.` with NO +// query string, and the whole thing must be percent-encoded. OrcaSlicer cannot +// import GLB, so we point at the backend's slicer-export route which converts to +// STL on the fly. + +/** Format handed to OrcaSlicer. STL is universal and OrcaSlicer auto-repairs it. */ +export const SLICER_FORMAT = 'stl' + +/** Prefix of an imported mesh served from outside the workspace. */ +const SERVE_FILE_PREFIX = '/optimize/serve-file?path=' + +/** + * Source formats the slicer route can convert. Gaussian splats (`.splat`, and + * the `.ply` they are delivered in) are point clouds, not printable meshes. + */ +const SLICEABLE_SOURCE = /\.(glb|gltf|obj|stl)$/i + +/** URL-safe base64 (no padding) of a UTF-8 string — matches the API's token decode. */ +export function encodeWorkspacePathToken(workspacePath: string): string { + const bytes = new TextEncoder().encode(workspacePath) + let binary = '' + for (const b of bytes) binary += String.fromCharCode(b) + return btoa(binary).replace(/\+/g, '-').replace(/\//g, '_').replace(/=+$/, '') +} + +/** + * The source path the slicer route should convert, or `undefined` when the + * output cannot be sliced. + * + * Two output shapes reach the viewer: a workspace URL (generated or + * workflow-produced meshes) and a `serve-file` URL (meshes the user imported + * from elsewhere on disk). Both are sliceable — the API accepts the absolute + * path of an import because it recorded that the user picked it themselves. + */ +function sliceableSourcePath(outputUrl: string | undefined): string | undefined { + if (!outputUrl) return undefined + + if (outputUrl.startsWith('/workspace/')) { + const workspacePath = outputUrl.slice('/workspace/'.length) + return SLICEABLE_SOURCE.test(workspacePath) ? workspacePath : undefined + } + + if (outputUrl.startsWith(SERVE_FILE_PREFIX)) { + const absolutePath = decodeURIComponent(outputUrl.slice(SERVE_FILE_PREFIX.length)) + return SLICEABLE_SOURCE.test(absolutePath) ? absolutePath : undefined + } + + return undefined +} + +/** + * Whether a generation output can be opened in OrcaSlicer: it must be a mesh in + * a format the slicer route converts, served either from the workspace or as a + * user-selected import. + */ +export function canOpenInOrcaSlicer(outputUrl: string | undefined): boolean { + return sliceableSourcePath(outputUrl) !== undefined +} + +/** + * Build the `orcaslicer://open?file=...` deeplink for a generated mesh. + * + * @param apiUrl Modly backend origin, e.g. `http://localhost:8765` + * @param outputUrl workspace or serve-file URL of the mesh + * @throws if `outputUrl` is not sliceable — guard with {@link canOpenInOrcaSlicer} + */ +export function buildOrcaSlicerDeepLink(apiUrl: string, outputUrl: string): string { + const sourcePath = sliceableSourcePath(outputUrl) + if (!sourcePath) throw new Error(`Not sliceable: ${outputUrl}`) + const token = encodeWorkspacePathToken(sourcePath) + const base = apiUrl.replace(/\/+$/, '') + const modelUrl = `${base}/export/slicer/${SLICER_FORMAT}/${token}/model.${SLICER_FORMAT}` + return `orcaslicer://open?file=${encodeURIComponent(modelUrl)}` +} diff --git a/src/areas/models/ModelsPage.tsx b/src/areas/models/ModelsPage.tsx index 895d4e8f..7688474c 100644 --- a/src/areas/models/ModelsPage.tsx +++ b/src/areas/models/ModelsPage.tsx @@ -1,12 +1,13 @@ import { useEffect, useMemo, useRef, useState } from 'react' import { createPortal } from 'react-dom' import { useExtensionsStore } from '@shared/stores/extensionsStore' -import type { AnyExtension, ModelExtension } from '@shared/types/electron.d' -import { formatModelName } from './utils' +import { useNavStore } from '@shared/stores/navStore' +import type { AnyExtension, ModelExtension, SharedWeightGroupState } from '@shared/types/electron.d' +import { deleteModelsThenUninstallExtension, formatModelName, installModelAndRefresh, installModelQueue } from './utils' import { ExtensionCard } from './components/ExtensionCard' import type { ExtensionNode } from './components/ExtensionCard' import { ExtensionDrawer } from './components/ExtensionDrawer' -import { ICONS } from './components/extensionShared' +import { ICONS, nodeHasManagedWeights, type DownloadMap } from './components/extensionShared' // ─── Filters & sorts ────────────────────────────────────────────────────────── @@ -48,18 +49,10 @@ export default function ModelsPage(): JSX.Element { ) // Model weight state (needed for node install status + uninstall cleanup) - const [installedVariantIds, setInstalledVariantIds] = useState([]) - const [downloading, setDownloading] = useState>({}) + const [installedNodeIds, setInstalledNodeIds] = useState([]) + const [localDataIds, setLocalDataIds] = useState([]) + const [sharedGroupStates, setSharedGroupStates] = useState>({}) + const [downloading, setDownloading] = useState({}) // Uninstall modal state const [uninstallTarget, setUninstallTarget] = useState(null) @@ -74,25 +67,49 @@ export default function ModelsPage(): JSX.Element { const [selectedId, setSelectedId] = useState(null) const searchRef = useRef(null) + // Another page asked to open an extension (e.g. a weight variant picked but not installed) + const extensionToOpen = useNavStore((s) => s.extensionToOpen) + const clearExtensionToOpen = useNavStore((s) => s.clearExtensionToOpen) + useEffect(() => { + if (!extensionToOpen) return + setSelectedId(extensionToOpen) + clearExtensionToOpen() + }, [extensionToOpen, clearExtensionToOpen]) + // GitHub extension install form const [showGHForm, setShowGHForm] = useState(false) const [ghUrl, setGhUrl] = useState('') 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 = {} for (const ext of exts) { + sharedStates[ext.id] = await window.electron.model.sharedGroups(ext.id) for (const node of ext.nodes) { - if (!node.hfRepo) continue + if (!nodeHasManagedWeights(node)) continue const fullId = `${ext.id}/${node.id}` - const ok = await window.electron.model.isDownloaded(fullId, node.downloadCheck) + const [ok, hasLocalData] = await Promise.all([ + window.electron.model.isDownloaded(fullId), + window.electron.model.hasLocalData(fullId), + ]) if (ok) ids.push(fullId) + if (hasLocalData) localIds.push(fullId) } } - setInstalledVariantIds(ids) + if (revision !== installedRefreshRevision.current) return + setInstalledNodeIds(ids) + setLocalDataIds(localIds) + setSharedGroupStates(sharedStates) + await useExtensionsStore.getState().refreshInstalledWeightVariants() } useEffect(() => { @@ -108,7 +125,10 @@ export default function ModelsPage(): JSX.Element { } refreshInstalledIds(exts) }) - window.electron.model.onProgress(({ modelId: id, percent, file, fileIndex, totalFiles, status, bytesDownloaded, totalBytes, stalledSeconds, paused, cancelled }) => { + window.electron.model.onWeightsChanged(() => { + void refreshInstalledIds(useExtensionsStore.getState().modelExtensions) + }) + window.electron.model.onProgress(({ modelId: id, variantId, percent, file, fileIndex, totalFiles, status, bytesDownloaded, totalBytes, stalledSeconds, paused, cancelled }) => { if (cancelled) { setDownloading((prev) => { const n = { ...prev }; delete n[id]; return n }) return @@ -127,6 +147,7 @@ export default function ModelsPage(): JSX.Element { totalBytes: totalBytes ?? current?.totalBytes, stalledSeconds: stalledSeconds ?? current?.stalledSeconds, paused, + variantId: variantId ?? current?.variantId, }, } }) @@ -137,7 +158,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 }, []) @@ -160,24 +184,39 @@ export default function ModelsPage(): JSX.Element { // ── Node install / download controls ────────────────────────────────────── - function handleInstallNode(node: ExtensionNode, fullId: string) { - if (!node.hfRepo) return - setDownloading((prev) => ({ ...prev, [fullId]: { ...(prev[fullId] ?? { percent: 0 }), paused: false, status: 'Starting…' } })) - window.electron.model.download(node.hfRepo!, fullId, node.hfSkipPrefixes, node.hfIncludePrefixes).then((result: { success: boolean; paused?: boolean; cancelled?: boolean }) => { + async function handleInstallNode(node: ExtensionNode, fullId: string, variantId?: string) { + if (!nodeHasManagedWeights(node)) return { success: true } + setDownloading((prev) => ({ ...prev, [fullId]: { ...(prev[fullId] ?? { percent: 0 }), variantId, paused: false, status: 'Starting…' } })) + try { + const result = await installModelAndRefresh( + () => window.electron.model.download(fullId, variantId), + () => refreshInstalledIds(useExtensionsStore.getState().modelExtensions), + ) if (!result.success && !result.paused && !result.cancelled) { - setGhErr('Download failed') - setDownloading((prev) => { const n = { ...prev }; delete n[fullId]; return n }) + 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 } + } } - function handleInstallAll(ext: AnyExtension) { - if (ext.type !== 'model') return - for (const node of ext.nodes) { - if (!node.hfRepo) continue - const fullId = `${ext.id}/${node.id}` - if (installedVariantIds.includes(fullId) || downloading[fullId]) continue - handleInstallNode(node, fullId) + async function handleInstallAll(ext: AnyExtension) { + 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) } } @@ -188,7 +227,9 @@ export default function ModelsPage(): JSX.Element { async function handleCancelDownload(fullId: string) { setDownloading((prev) => { const n = { ...prev }; delete n[fullId]; return n }) - await window.electron.model.cancelDownload(fullId) + const result = await window.electron.model.cancelDownload(fullId) + if (!result.success) setGhErr(result.error ?? 'Could not cancel download') + await refreshInstalledIds(useExtensionsStore.getState().modelExtensions) } async function handleUninstallNode(fullId: string) { @@ -196,6 +237,18 @@ 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 + } + + async function handleDeleteWeightVariant(fullId: string, variantId: string) { + const result = await window.electron.model.deleteWeightVariant(fullId, variantId) + await refreshInstalledIds(useExtensionsStore.getState().modelExtensions) + return result + } + // ── GitHub extension install ─────────────────────────────────────────────── async function handleGHInstall() { @@ -228,8 +281,11 @@ export default function ModelsPage(): JSX.Element { function openUninstallModal(extId: string) { const ext = allExtensions.find((e) => e.id === extId) if (ext?.type === 'model') { - const installedModels = ext.nodes.filter((n) => installedVariantIds.includes(`${extId}/${n.id}`)) - setModelsToDelete(new Set(installedModels.map((n) => `${extId}/${n.id}`))) + const localModels = ext.nodes.filter((n) => localDataIds.includes(`${extId}/${n.id}`)) + 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()) } @@ -237,10 +293,14 @@ export default function ModelsPage(): JSX.Element { } async function handleUninstallExtension(extId: string) { - for (const modelId of modelsToDelete) { - await window.electron.model.delete(modelId) - } - const result = await uninstallExt(extId) + const result = await deleteModelsThenUninstallExtension( + extId, + modelsToDelete, + (modelId) => modelId === `${extId}/*` + ? window.electron.model.deleteExtensionWeights(extId) + : window.electron.model.delete(modelId), + uninstallExt, + ) if (!result.success) { // Keep the dialog open so the failure is visible (locked folder, etc.) setUninstallError(result.error ?? 'Could not delete the extension folder.') @@ -311,7 +371,7 @@ export default function ModelsPage(): JSX.Element { } const cardHandlers = { - installedIds: installedVariantIds, + installedIds: installedNodeIds, downloading, disabled: isBusy, onInstall: handleInstallNode, @@ -616,8 +676,10 @@ export default function ModelsPage(): JSX.Element { {selectedExt && ( openUninstallModal(extId)} onRepaired={() => reloadExtensions()} onSynced={() => reloadExtensions()} @@ -636,8 +700,12 @@ export default function ModelsPage(): JSX.Element { {uninstallTarget && (() => { const ext = allExtensions.find((e) => e.id === uninstallTarget) const installedModels = ext?.type === 'model' - ? ext.nodes.filter((n) => installedVariantIds.includes(`${uninstallTarget}/${n.id}`)) + ? ext.nodes.filter((n) => localDataIds.includes(`${uninstallTarget}/${n.id}`)) : [] + 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/ExtensionCard.tsx b/src/areas/models/components/ExtensionCard.tsx index 7b0da394..5ceb729a 100644 --- a/src/areas/models/components/ExtensionCard.tsx +++ b/src/areas/models/components/ExtensionCard.tsx @@ -125,7 +125,12 @@ export function ExtensionCard({ className="flex items-center justify-between gap-2.5 px-2.5 py-1.5 rounded-lg bg-white/[0.02] border border-zinc-800" >

- {node.name} + + {node.name} + {isModel && node.weightVariants && ( + · {node.weightVariants.options.length} variants + )} +
{isModel && ( diff --git a/src/areas/models/components/ExtensionDrawer.tsx b/src/areas/models/components/ExtensionDrawer.tsx index 54adab9c..949d7de3 100644 --- a/src/areas/models/components/ExtensionDrawer.tsx +++ b/src/areas/models/components/ExtensionDrawer.tsx @@ -1,11 +1,13 @@ import { useEffect, useState } from 'react' -import type { AnyExtension, ExtensionNode } from '@shared/types/electron.d' +import type { AnyExtension, ExtensionNode, SharedWeightGroupState } from '@shared/types/electron.d' +import { useExtensionsStore } from '@shared/stores/extensionsStore' import { useNavStore } from '@shared/stores/navStore' import { DownloadMap, ICONS, IOBadge, NodeInstallControl, + NodeUiState, TypePill, extInstallSummary, formatBytes, @@ -16,14 +18,18 @@ import { finishExtensionRepair, isExtensionRepairable } from '../utils' interface Props { ext: AnyExtension installedIds: string[] + localDataIds: string[] downloading: DownloadMap + sharedGroups: SharedWeightGroupState[] loadError?: string disabled?: boolean - onInstall: (node: ExtensionNode, fullId: string) => void + onInstall: (node: ExtensionNode, fullId: string, variantId?: string) => void onInstallAll: (ext: AnyExtension) => void onPauseDownload: (fullId: string) => void onCancelDownload: (fullId: string) => void onUninstallNode: (fullId: string) => void + onDeleteSharedGroup: (extensionId: string, groupId: string) => Promise<{ success: boolean; error?: string }> + onDeleteWeightVariant: (fullId: string, variantId: string) => Promise<{ success: boolean; error?: string }> onUninstall: (extId: string) => void onRepaired: () => void | Promise onSynced: () => void @@ -31,15 +37,17 @@ interface Props { } export function ExtensionDrawer({ - ext, installedIds, downloading, loadError, disabled, + ext, installedIds, localDataIds, downloading, sharedGroups, loadError, disabled, onInstall, onInstallAll, onPauseDownload, onCancelDownload, - onUninstallNode, onUninstall, onRepaired, onSynced, onClose, + onUninstallNode, onDeleteSharedGroup, onDeleteWeightVariant, onUninstall, onRepaired, onSynced, onClose, }: Props): JSX.Element { const navigate = useNavStore((s) => s.navigate) + const installedWeightVariants = useExtensionsStore((s) => s.installedWeightVariants) const [repairing, setRepairing] = useState(false) const [repairError, setRepairError] = useState(null) const [syncing, setSyncing] = useState(false) const [syncError, setSyncError] = useState(null) + const [variantError, setVariantError] = useState(null) const isModel = ext.type === 'model' // Built-ins are corrupted-flagged too (builtin-sync repairs them on restart), @@ -84,7 +92,24 @@ export function ExtensionDrawer({ } } - const error = syncError ?? repairError ?? loadError + 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.') + } + + + async function handleDeleteWeightVariant(fullId: string, variantId: string) { + setVariantError(null) + const result = await onDeleteWeightVariant(fullId, variantId) + if (!result.success) setVariantError(result.error ?? 'Could not remove these weights') + } + + const error = variantError ?? syncError ?? repairError ?? loadError return ( <> @@ -168,6 +193,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}`} @@ -177,12 +242,19 @@ export function ExtensionDrawer({ const fullId = `${ext.id}/${node.id}` const state = getNodeState(ext.id, node, installedIds, downloading) const dl = state.kind === 'downloading' ? state.dl : null + const variants = isModel ? node.weightVariants : undefined + // A download started without a variant (Install all) fetches the default one. + const dlVariantId = dl && variants ? (dl.variantId ?? variants.default) : undefined + const dlVariant = variants?.options.find((option) => option.id === dlVariantId) + const installedLabels = variants?.options + .filter((option) => installedWeightVariants[fullId]?.includes(option.id)) + .map((option) => option.label) ?? [] const sub = state.kind === 'ready' ? 'Available on the node graph' - : state.kind === 'installed' ? 'Installed' + : state.kind === 'installed' ? (installedLabels.length > 0 ? `${installedLabels.join(', ')} installed` : 'Installed') : state.kind === 'available' ? 'Not installed' : dl?.paused ? 'Download paused' - : `Downloading… ${dl?.percent ?? 0}%` + : `Downloading${dlVariant ? ` ${dlVariant.label}` : ''}… ${dl?.percent ?? 0}%` return (
@@ -195,19 +267,21 @@ export function ExtensionDrawer({ {isModel && (
- onInstall(node, fullId)} - onPause={() => onPauseDownload(fullId)} - onResume={() => onInstall(node, fullId)} - onCancel={() => onCancelDownload(fullId)} - /> - {state.kind === 'installed' && ( + {!variants && ( + onInstall(node, fullId)} + onPause={() => onPauseDownload(fullId)} + onResume={() => onInstall(node, fullId)} + onCancel={() => onCancelDownload(fullId)} + /> + )} + {localDataIds.includes(fullId) && state.kind !== 'downloading' && (
+ {/* Weight variants — each one installs and deletes on its own */} + {variants && ( +
+ {variants.options.map((option) => { + const installed = installedWeightVariants[fullId]?.includes(option.id) ?? false + const variantDl = dlVariantId === option.id ? dl : null + const variantState: NodeUiState = variantDl + ? { kind: 'downloading', dl: variantDl } + : installed ? { kind: 'installed' } : { kind: 'available' } + return ( +
+
+ {option.label} + {option.sizeGb !== undefined && ( + {option.sizeGb} GB + )} + {option.vramGb !== undefined && ( + ~{option.vramGb} GB VRAM + )} + {option.id === variants.default && ( + default + )} +
+
+ onInstall(node, fullId, option.id)} + onPause={() => onPauseDownload(fullId)} + onResume={() => onInstall(node, fullId, option.id)} + onCancel={() => onCancelDownload(fullId)} + /> + {installed && !dl && ( + + )} +
+
+ ) + })} +
+ )} + {/* Download detail */} {dl && (
diff --git a/src/areas/models/components/extensionShared.tsx b/src/areas/models/components/extensionShared.tsx index 0df825b5..dfd91b04 100644 --- a/src/areas/models/components/extensionShared.tsx +++ b/src/areas/models/components/extensionShared.tsx @@ -12,6 +12,7 @@ export interface DownloadInfo { totalBytes?: number stalledSeconds?: number paused?: boolean + variantId?: string // set when the download targets one weight variant of the node } export type DownloadMap = Record @@ -22,6 +23,10 @@ export type NodeUiState = | { kind: 'downloading'; dl: DownloadInfo } | { kind: 'installed' } +export function nodeHasManagedWeights(node: ExtensionNode): boolean { + return Boolean(node.hfRepo || node.hasModelSources || node.weightGroups?.length) +} + export function getNodeState( extId: string, node: ExtensionNode, @@ -29,7 +34,7 @@ export function getNodeState( downloading: DownloadMap, ): NodeUiState { const fullId = `${extId}/${node.id}` - if (!node.hfRepo) return { kind: 'ready' } + if (!nodeHasManagedWeights(node)) return { kind: 'ready' } const dl = downloading[fullId] if (dl) return { kind: 'downloading', dl } if (installedIds.includes(fullId)) return { kind: 'installed' } diff --git a/src/areas/models/utils.test.mjs b/src/areas/models/utils.test.mjs index 7acf39a1..531cfe92 100644 --- a/src/areas/models/utils.test.mjs +++ b/src/areas/models/utils.test.mjs @@ -21,6 +21,7 @@ function loadModule() { } const { + deleteModelsThenUninstallExtension, finishExtensionRepair, formatModelName, isExtensionRepairable, @@ -85,3 +86,67 @@ test('Repair completion refreshes extension state even when setup fails', async assert.equal(refreshCount, 1) assert.equal(error, 'setup failed') }) + +test('failed selected-weight deletion aborts extension uninstall and preserves its error', async () => { + const deleteCalls = [] + let uninstallCalls = 0 + + const result = await deleteModelsThenUninstallExtension( + 'pixal3d', + new Set([ + 'pixal3d/generate', + 'pixal3d/refine', + 'pixal3d/export', + ]), + async (modelId) => { + deleteCalls.push(modelId) + return modelId === 'pixal3d/refine' + ? { success: false, error: 'Model weights are locked.' } + : { success: true } + }, + async () => { + uninstallCalls += 1 + return { success: true } + }, + ) + + assert.deepEqual(deleteCalls, ['pixal3d/generate', 'pixal3d/refine']) + 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 07bdec05..33c3cf4d 100644 --- a/src/areas/models/utils.ts +++ b/src/areas/models/utils.ts @@ -21,3 +21,59 @@ export async function finishExtensionRepair( await refreshExtensions() return result.success ? null : (result.error ?? 'Repair failed') } + +interface ActionResult { + success: boolean + error?: string +} + +export async function deleteModelsThenUninstallExtension( + extensionId: string, + modelIds: Iterable, + deleteModel: (modelId: string) => Promise, + uninstallExtension: (extensionId: string) => Promise, +): Promise { + for (const modelId of modelIds) { + const result = await deleteModel(modelId) + if (!result.success) { + return { + success: false, + error: result.error ?? 'Could not delete selected model weights.', + } + } + } + + 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/areas/settings/components/AgentSection.tsx b/src/areas/settings/components/AgentSection.tsx index e8f2e1b4..9376ecb4 100644 --- a/src/areas/settings/components/AgentSection.tsx +++ b/src/areas/settings/components/AgentSection.tsx @@ -1,226 +1,478 @@ -import { useState, useEffect } from 'react' -import { useAgentStore, type ThinkingMode } from '@shared/stores/agentStore' +import { useState, useEffect, useCallback } from 'react' +import { + useAgentStore, PROVIDERS, + type ThinkingMode, type ProviderId, type ExternalConfig, +} from '@shared/stores/agentStore' import { useAppStore } from '@shared/stores/appStore' +import { consumeSse, type SseEvent } from '@shared/services/llmDownloads' +import { SseProgressBar } from '@shared/components/ui/SseProgressBar' +import { ModelLibraryModal } from '@shared/components/ui/ModelLibraryModal' +import type { LlmModel } from '@shared/stores/llmModelsStore' +import { formatBytes } from '@shared/utils/format' +import { Section, Card, SegmentedControl } from '@shared/ui' // ─── Helpers ────────────────────────────────────────────────────────────────── -function Field({ label, hint, children }: { label: string; hint?: string; children: React.ReactNode }): JSX.Element { +/** A labelled block inside a Card, for controls that need the full width. */ +function Block({ label, description, hint, action, children }: { + label?: string + description?: string + hint?: string + action?: React.ReactNode + children?: React.ReactNode +}): JSX.Element { return ( -
- +
+ {(label || action) && ( +
+
+ {label &&

{label}

} + {description &&

{description}

} +
+ {action &&
{action}
} +
+ )} {children} - {hint &&

{hint}

} + {hint &&

{hint}

}
) } -function Group({ title, children }: { title: string; children: React.ReactNode }): JSX.Element { +function Badge({ tone, children }: { tone: 'accent' | 'warn' | 'muted'; children: React.ReactNode }): JSX.Element { + const cls = tone === 'accent' + ? 'bg-accent/15 text-accent-light border-accent/30' + : tone === 'warn' + ? 'bg-amber-500/10 text-amber-400 border-amber-500/30' + : 'bg-zinc-800 text-zinc-400 border-zinc-700' + const dot = tone === 'accent' ? 'bg-accent-light' : tone === 'warn' ? 'bg-amber-400' : null return ( -
-

{title}

+ + {dot && } {children} -
+ + ) +} + +function CopyButton({ text }: { text: string }): JSX.Element { + const [copied, setCopied] = useState(false) + useEffect(() => { + if (!copied) return + const t = setTimeout(() => setCopied(false), 1500) + return () => clearTimeout(t) + }, [copied]) + return ( + + ) +} + +function StepTitle({ n, children }: { n: number; children: React.ReactNode }): JSX.Element { + return ( +

+ {n}{children} +

) } +const inputCls = 'w-full bg-zinc-800 border border-zinc-700 text-zinc-200 text-xs rounded-lg px-3 py-2 focus:outline-none focus:border-accent/50' +const secondaryBtnCls = 'px-2.5 py-1 rounded-md text-[11px] font-medium bg-zinc-800 hover:bg-zinc-700 text-zinc-300 transition-colors disabled:opacity-50' +const primaryBtnCls = 'px-3 py-1.5 rounded-lg bg-accent hover:bg-accent-dark text-white text-xs font-medium transition-colors' + +type McpClient = 'claude' | 'codex' | 'opencode' + +const MCP_CLIENTS: { value: McpClient; label: string; path: string; config: string }[] = [ + { + value: 'claude', + label: 'Claude Desktop', + path: '~/.config/claude/claude_desktop_config.json', + config: `{\n "mcpServers": {\n "modly": {\n "command": "modly-mcp"\n }\n }\n}`, + }, + { + value: 'codex', + label: 'Codex CLI', + path: '~/.codex/config.toml', + config: `[mcp_servers.modly]\ncommand = "modly-mcp"`, + }, + { + value: 'opencode', + label: 'OpenCode', + path: '~/.config/opencode/config.json', + config: `{\n "$schema": "https://opencode.ai/config.json",\n "mcp": {\n "modly": {\n "type": "local",\n "command": ["modly-mcp"]\n }\n }\n}`, + }, +] + +const MCP_INSTALL = 'npm install -g modly-cli-mcp' + +const THINKING_OPTIONS: { value: ThinkingMode; label: string; desc: string }[] = [ + { value: 'auto', label: 'Auto', desc: 'The model decides whether to think' }, + { value: 'on', label: 'Enabled', desc: 'Forces thinking on every response' }, + { value: 'off', label: 'Disabled', desc: 'Disables thinking (faster responses)' }, +] + // ─── Component ──────────────────────────────────────────────────────────────── export function AgentSection(): JSX.Element { - const { ollamaUrl, defaultModel, defaultThinking, setOllamaUrl, setDefaultModel, setDefaultThinking } = useAgentStore() + const { + provider, localModel, external, defaultThinking, + setProvider, setExternal, setDefaultThinking, + } = useAgentStore() const apiUrl = useAppStore((s) => s.apiUrl) - const [urlDraft, setUrlDraft] = useState(ollamaUrl) - const [modelDraft, setModelDraft] = useState(defaultModel) - const [models, setModels] = useState([]) - const [testing, setTesting] = useState(false) - const [testResult, setTestResult] = useState<'ok' | 'error' | null>(null) + // Local engine + const [engineInstalled, setEngineInstalled] = useState(null) + const [engineInstall, setEngineInstall] = useState(null) + const [engineError, setEngineError] = useState(null) + const [models, setModels] = useState([]) + const [showLibrary, setShowLibrary] = useState(false) + const [maxModels, setMaxModels] = useState('auto') + const [resolvedMax, setResolvedMax] = useState(null) + const [vramGb, setVramGb] = useState(null) - useEffect(() => { - setUrlDraft(ollamaUrl) - }, [ollamaUrl]) + // External provider drafts + const extCfg = external[provider] + const [keyDraft, setKeyDraft] = useState(extCfg?.apiKey ?? '') + const [extModelDraft, setExtModelDraft] = useState(extCfg?.model ?? '') + const [baseUrlDraft, setBaseUrlDraft] = useState(extCfg?.baseUrl ?? '') + const [extModels, setExtModels] = useState([]) + const [extTesting, setExtTesting] = useState(false) + const [extResult, setExtResult] = useState<'ok' | 'error' | null>(null) - useEffect(() => { - setModelDraft(defaultModel) - }, [defaultModel]) + const [mcpClient, setMcpClient] = useState('claude') - async function fetchModels(url: string) { + const refreshLocal = useCallback(async () => { try { - const res = await fetch(`${apiUrl}/agent/models?ollama_url=${encodeURIComponent(url)}`) - const data = await res.json() - setModels(data.models ?? []) + const [s, m, c] = await Promise.all([ + fetch(`${apiUrl}/llm/status`).then((r) => r.json()), + fetch(`${apiUrl}/llm/models?downloaded=true`).then((r) => r.json()), + fetch(`${apiUrl}/llm/config`).then((r) => r.json()), + ]) + setEngineInstalled(Boolean(s.binary_installed)) + setModels(m.models ?? []) + setMaxModels(String(c.max_models ?? 'auto')) + setResolvedMax(c.resolved_max_models ?? null) + setVramGb(c.vram_gb ?? null) } catch { - setModels([]) + setEngineInstalled(null) + } + }, [apiUrl]) + + async function installEngine() { + setEngineInstall({ percent: 0, status: 'Starting…' }) + setEngineError(null) + try { + await consumeSse(`${apiUrl}/llm/binary/install`, (e) => { + if (e.error) setEngineError(e.error) + else setEngineInstall(e) + }) + } catch (e) { + setEngineError(e instanceof Error ? e.message : String(e)) + } finally { + setEngineInstall(null) + void refreshLocal() } } - async function handleTestConnection() { - setTesting(true) - setTestResult(null) + async function changeMaxModels(value: string) { + setMaxModels(value) try { - const res = await fetch(`${apiUrl}/agent/models?ollama_url=${encodeURIComponent(urlDraft)}`) + const res = await fetch(`${apiUrl}/llm/config`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ max_models: value === 'auto' ? 'auto' : Number(value) }), + }) const data = await res.json() + setResolvedMax(data.resolved_max_models ?? null) + } catch { /* API unreachable — keep the optimistic value, refreshLocal will resync */ } + } + + useEffect(() => { void refreshLocal() }, [refreshLocal]) + + // Sync drafts when switching provider + useEffect(() => { + const cfg = useAgentStore.getState().external[provider] + setKeyDraft(cfg?.apiKey ?? '') + setExtModelDraft(cfg?.model ?? '') + setBaseUrlDraft(cfg?.baseUrl ?? '') + setExtModels([]) + setExtResult(null) + }, [provider]) + + function saveExternal() { + const cfg: ExternalConfig = { apiKey: keyDraft.trim(), model: extModelDraft.trim() } + if (provider === 'custom') cfg.baseUrl = baseUrlDraft.trim().replace(/\/$/, '') + setExternal(provider, cfg) + } + + async function handleTestExternal() { + setExtTesting(true) + setExtResult(null) + try { + const base = provider === 'custom' ? baseUrlDraft.trim().replace(/\/$/, '') : PROVIDERS[provider].baseUrl + // POST, not a query string: a GET would put the API key in the uvicorn + // access log, which ends up in runtime.log and in users' bug reports. + const res = await fetch(`${apiUrl}/agent/external/models`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ base_url: base, api_key: keyDraft.trim() }), + }) + const data: { models?: string[] } = await res.json() const found = (data.models ?? []).length > 0 - setTestResult(found ? 'ok' : 'error') - if (found) setModels(data.models) + setExtModels(data.models ?? []) + setExtResult(found ? 'ok' : 'error') } catch { - setTestResult('error') + setExtResult('error') } finally { - setTesting(false) + setExtTesting(false) } } - function handleSaveOllama() { - setOllamaUrl(urlDraft) - setDefaultModel(modelDraft) - fetchModels(urlDraft) - } + const selectedModel = models.find((m) => m.id === localModel) + const client = MCP_CLIENTS.find((c) => c.value === mcpClient) ?? MCP_CLIENTS[0] - const THINKING_OPTIONS: { value: ThinkingMode; label: string; desc: string }[] = [ - { value: 'auto', label: 'Auto', desc: 'The model decides whether to think' }, - { value: 'on', label: 'Enabled', desc: 'Forces thinking on every response' }, - { value: 'off', label: 'Disabled', desc: 'Disables thinking (faster responses)' }, - ] + return ( +
+
- const mcpConfigs = { - opencode: `{\n "$schema": "https://opencode.ai/config.json",\n "mcp": {\n "modly": {\n "type": "local",\n "command": ["modly-mcp"]\n }\n }\n}`, - codex: `[mcp_servers.modly]\ncommand = "modly-mcp"`, - claude: `{\n "mcpServers": {\n "modly": {\n "command": "modly-mcp"\n }\n }\n}`, - } + {/* ── Left column ── */} +
- return ( -
-
-

Agent

-

Configure the local LLM and Chat mode settings.

-
+ + + + + - {/* Ollama */} - - -
- { setUrlDraft(e.target.value); setTestResult(null) }} - className="flex-1 bg-zinc-900 border border-zinc-700/60 rounded-lg px-3 py-2 text-[12.5px] text-zinc-200 focus:outline-none focus:border-zinc-500" - placeholder="http://localhost:11434" - /> - -
- {testResult === 'ok' && ( -

Connection successful — {models.length} model{models.length > 1 ? 's' : ''} found

- )} - {testResult === 'error' && ( -

Could not reach Ollama at this address

- )} -
- - - {models.length > 0 ? ( - + {engineInstalled === false && ( + + {engineInstall ? ( + + ) : ( + + )} + {engineError &&

{engineError}

} +
+ )} + {engineInstalled === null && ( + + + + )} + + setShowLibrary(true)} className={secondaryBtnCls}>Browse…} + > + {selectedModel ? ( +
+ + + + +
+ {selectedModel.name} + + {[ + selectedModel.size_bytes ? formatBytes(selectedModel.size_bytes) : null, + selectedModel.quant, + selectedModel.vram_estimate_mb ? `~${(selectedModel.vram_estimate_mb / 1000).toFixed(1)} GB VRAM` : null, + ].filter(Boolean).join(' · ')} + +
+ In use +
+ ) : ( +

+ {models.length === 0 + ? 'No model yet — open Browse… to add or download one.' + : 'No model selected — open Browse… and select one.'} +

+ )} +
+ + + + + ) : ( - setModelDraft(e.target.value)} - className="bg-zinc-900 border border-zinc-700/60 rounded-lg px-3 py-2 text-[12.5px] text-zinc-200 focus:outline-none focus:border-zinc-500" - placeholder="gemma4:e4b" - /> + + {provider === 'custom' && ( + + { setBaseUrlDraft(e.target.value); setExtResult(null) }} + placeholder="http://192.168.1.20:8080/v1" + className={inputCls} + /> + + )} + + +
+ { setKeyDraft(e.target.value); setExtResult(null) }} + placeholder={PROVIDERS[provider].noKey ? '' : 'sk-…'} + className={inputCls} + /> + +
+ {extResult === 'ok' && ( +

Connected — {extModels.length} model{extModels.length > 1 ? 's' : ''} available

+ )} + {extResult === 'error' && ( +

+ {provider === 'ollama' + ? 'Could not list models — is Ollama running?' + : `Could not list models — check the key${provider === 'custom' ? ' and URL' : ''}`} +

+ )} +
+ + + {extModels.length > 0 ? ( + + ) : ( + setExtModelDraft(e.target.value)} + placeholder={provider === 'anthropic' ? 'claude-sonnet-5' : provider === 'openai' ? 'gpt-5.2' : provider === 'ollama' ? 'qwen2.5:3b' : 'model name'} + className={inputCls} + /> + )} + + + + + +
)} -
- - -
- - {/* Thinking */} - - -
- {THINKING_OPTIONS.map((opt) => ( -
-
- {([ - { label: 'Claude Desktop', key: 'claude' as const, hint: '~/.config/claude/claude_desktop_config.json' }, - { label: 'Codex CLI', key: 'codex' as const, hint: '~/.codex/config.toml' }, - { label: 'OpenCode', key: 'opencode' as const, hint: '~/.config/opencode/config.json' }, - ] as const).map(({ label, key, hint }) => ( -
-

{label} — {hint}

-
-
-                  {mcpConfigs[key]}
+        {/* ── Right column ── */}
+        
+ Control Modly from Claude Desktop, Codex or OpenCode. Community package by DrHepa.} + aside={Community} + > + + Install the package +
+ + ${MCP_INSTALL} + + +
+
+ + + Add it to your client +
+ ({ value, label }))} + ariaLabel="MCP client" + /> +
+
+
+

+ {client.label} + {client.path} +

+ +
+
+                  {client.config}
                 
-
-
- ))} + +
- -
+
+ + {showLibrary && ( + { setShowLibrary(false); void refreshLocal() }} /> + )} +
) } diff --git a/src/areas/workflows/WorkflowsPage.tsx b/src/areas/workflows/WorkflowsPage.tsx index c9d99c1c..8fc4cea6 100644 --- a/src/areas/workflows/WorkflowsPage.tsx +++ b/src/areas/workflows/WorkflowsPage.tsx @@ -27,7 +27,9 @@ import ImageNode from './nodes/ImageNode' import TextNode from './nodes/TextNode' import AddToSceneNode from './nodes/AddToSceneNode' import Load3DMeshNode from './nodes/Load3DMeshNode' +import LoadSceneNode from './nodes/LoadSceneNode' import PreviewImageNode from './nodes/PreviewImageNode' +import ImagePreviewNode from './nodes/ImagePreviewNode' import WaitNode from './nodes/WaitNode' import WhileNode from './nodes/WhileNode' import ForEachNode from './nodes/ForEachNode' @@ -37,7 +39,7 @@ import WorkflowEdge from './nodes/WorkflowEdge' const DRAG_KEY = 'modly/extension-id' const DRAG_NODE_KEY = 'modly/node-type' -const NODE_TYPES = { extensionNode: ExtensionNode, imageNode: ImageNode, textNode: TextNode, outputNode: AddToSceneNode, meshNode: Load3DMeshNode, previewNode: PreviewImageNode, waitNode: WaitNode, whileNode: WhileNode, forEachNode: ForEachNode } +const NODE_TYPES = { extensionNode: ExtensionNode, imageNode: ImageNode, textNode: TextNode, outputNode: AddToSceneNode, meshNode: Load3DMeshNode, sceneNode: LoadSceneNode, previewNode: PreviewImageNode, imagePreviewNode: ImagePreviewNode, waitNode: WaitNode, whileNode: WhileNode, forEachNode: ForEachNode } // Loop-container node types: resizable frames whose children form a loop body. // (For Each iterators are plain source nodes, not containers.) @@ -61,14 +63,15 @@ function findWhileContainerAt(nodes: Node[], pos: { x: number; y: number }): Nod // ─── IO badge ───────────────────────────────────────────────────────────────── -const IO_STYLES: Record<'image' | 'text' | 'mesh' | 'audio', string> = { +const IO_STYLES: Record<'image' | 'text' | 'mesh' | 'audio' | 'scene', string> = { audio: 'bg-emerald-500/15 text-emerald-400 border-emerald-500/25', image: 'bg-sky-500/15 text-sky-400 border-sky-500/25', mesh: 'bg-violet-500/15 text-violet-400 border-violet-500/25', text: 'bg-amber-500/15 text-amber-400 border-amber-500/25', + scene: 'bg-pink-500/15 text-pink-400 border-pink-500/25', } -function IoBadge({ type }: { type: 'image' | 'text' | 'mesh' | 'audio' }) { +function IoBadge({ type }: { type: 'image' | 'text' | 'mesh' | 'audio' | 'scene' }) { return ( {type} @@ -98,8 +101,10 @@ const PANEL_BUILTIN_NODES = [ { type: 'imageNode', label: 'Image', color: '#38bdf8', icon: <> }, { type: 'textNode', label: 'Text', color: '#fbbf24', icon: <> }, { type: 'meshNode', label: 'Load 3D Mesh', color: '#a78bfa', icon: <> }, + { type: 'sceneNode', label: 'Load Scene', color: '#f472b6', icon: <> }, { type: 'outputNode', label: 'Add to Scene', color: '#a78bfa', icon: <> }, { type: 'previewNode', label: 'Preview Views', color: '#38bdf8', icon: <> }, + { type: 'imagePreviewNode', label: 'Preview Image', color: '#38bdf8', icon: <> }, { type: 'waitNode', label: 'Wait', color: '#71717a', icon: <> }, { type: 'whileNode', label: 'While', color: '#f59e0b', icon: <> }, { type: 'forEachNode', label: 'For Each', color: '#38bdf8', icon: <> }, @@ -344,8 +349,10 @@ const BUILTIN_NODES = [ { type: 'imageNode', label: 'Image', color: '#38bdf8', description: 'Image input' }, { type: 'textNode', label: 'Text', color: '#fbbf24', description: 'Text input' }, { type: 'meshNode', label: 'Load 3D Mesh', color: '#a78bfa', description: 'Load a 3D mesh file or use current model' }, + { type: 'sceneNode', label: 'Load Scene', color: '#f472b6', description: 'Load and validate a workspace scene directory' }, { type: 'outputNode', label: 'Add to Scene', color: '#a78bfa', description: 'Output node — adds the mesh to the 3D scene' }, { type: 'previewNode', label: 'Preview Views', color: '#38bdf8', description: 'Displays multi-view image outputs in a 2×3 grid' }, + { type: 'imagePreviewNode', label: 'Preview Image', color: '#38bdf8', description: 'Displays a single image output in the workflow' }, { type: 'waitNode', label: 'Wait', color: '#71717a', description: 'Pauses the workflow until you click Continue' }, { type: 'whileNode', label: 'While', color: '#f59e0b', description: 'Container: wrap nodes to loop them N times or with Continue/Retry' }, { type: 'forEachNode', label: 'For Each', color: '#38bdf8', description: 'Iterates a folder (image / text / mesh) alphabetically, one item per run of the downstream nodes' }, @@ -700,7 +707,9 @@ function getNodeOutputType(node: Node | undefined, allExts: WorkflowExtension[]) if (!node) return undefined if (node.type === 'imageNode') return 'image' if (node.type === 'meshNode') return 'mesh' + if (node.type === 'sceneNode') return 'scene' if (node.type === 'textNode') return 'text' + if (node.type === 'imagePreviewNode') return 'image' return allExts.find((e) => e.id === (node.data as WFNodeData)?.extensionId)?.output } @@ -712,6 +721,7 @@ function getNodeInputType( if (!node) return undefined if (node.type === 'outputNode') return 'mesh' if (node.type === 'previewNode') return 'image' + if (node.type === 'imagePreviewNode') return 'image' const ext = allExts.find((e) => e.id === (node.data as WFNodeData)?.extensionId) if (ext?.inputs && ext.inputs.length > 1 && targetHandle) { const idx = parseInt(targetHandle.replace('input-', ''), 10) @@ -1369,6 +1379,7 @@ const MINI_NODE_TINTS: Record = { imageNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, textNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, meshNode: { fill: 'rgba(52,211,153,0.22)', stroke: '#34d399' }, + sceneNode: { fill: 'rgba(244,114,182,0.22)', stroke: '#f472b6' }, extensionNode: { fill: 'rgba(167,139,250,0.24)', stroke: '#a78bfa' }, outputNode: { fill: 'rgba(56,189,248,0.22)', stroke: '#38bdf8' }, previewNode: { fill: 'rgba(56,189,248,0.22)', stroke: '#38bdf8' }, diff --git a/src/areas/workflows/mockExtensions.ts b/src/areas/workflows/mockExtensions.ts index 2bbc8cc6..9bdb9a0a 100644 --- a/src/areas/workflows/mockExtensions.ts +++ b/src/areas/workflows/mockExtensions.ts @@ -1,6 +1,6 @@ import type { ModelExtension, ProcessExtension } from '@shared/stores/extensionsStore' export type { ParamSchema } from '@shared/types/electron.d' -import type { ParamSchema } from '@shared/types/electron.d' +import type { ParamSchema, WeightVariantsInfo } from '@shared/types/electron.d' export interface WorkflowExtension { id: string // "ext_id/node_id" @@ -10,13 +10,14 @@ export interface WorkflowExtension { nodeId: string // "node_id" name: string description: string - input: 'image' | 'text' | 'mesh' | 'audio' - inputs?: ('image' | 'text' | 'mesh' | 'audio')[] // multi-input; overrides input when set + input: 'image' | 'text' | 'mesh' | 'audio' | 'scene' + inputs?: ('image' | 'text' | 'mesh' | 'audio' | 'scene')[] // multi-input; overrides input when set inputLabels?: string[] // display labels per input slot - output: 'image' | 'text' | 'mesh' | 'audio' + output: 'image' | 'text' | 'mesh' | 'audio' | 'scene' params: ParamSchema[] builtin: boolean type: 'model' | 'process' + weightVariants?: WeightVariantsInfo } function applyParamDefaults( @@ -77,6 +78,7 @@ export function buildAllWorkflowExtensions( params: applyParamDefaults(node.paramsSchema as ParamSchema[], node.paramDefaults), builtin: ext.builtin, type: 'model', + weightVariants: node.weightVariants, }) } } diff --git a/src/areas/workflows/nodeBehaviors.ts b/src/areas/workflows/nodeBehaviors.ts index fc75eea0..4846eeef 100644 --- a/src/areas/workflows/nodeBehaviors.ts +++ b/src/areas/workflows/nodeBehaviors.ts @@ -23,9 +23,10 @@ export interface NodeBehavior { } const BEHAVIORS: Record = { - waitNode: { passthrough: true, branchStarter: true }, - outputNode: { sceneOutput: true, branchConsumer: true }, - extensionNode: { branchConsumer: true }, + waitNode: { passthrough: true, branchStarter: true }, + outputNode: { sceneOutput: true, branchConsumer: true }, + extensionNode: { branchConsumer: true }, + imagePreviewNode: { passthrough: true }, } export const isPassthrough = (type: string | undefined): boolean => !!type && !!BEHAVIORS[type]?.passthrough diff --git a/src/areas/workflows/nodes/ExtensionNode.tsx b/src/areas/workflows/nodes/ExtensionNode.tsx index abe2f9f6..17e5f3fa 100644 --- a/src/areas/workflows/nodes/ExtensionNode.tsx +++ b/src/areas/workflows/nodes/ExtensionNode.tsx @@ -1,11 +1,13 @@ import { useCallback, useEffect, useRef, useLayoutEffect, useState } from 'react' import { Handle, Position, useReactFlow } from '@xyflow/react' import { useExtensionsStore } from '@shared/stores/extensionsStore' +import { useNavStore } from '@shared/stores/navStore' import { buildAllWorkflowExtensions } from '../mockExtensions' import type { ParamSchema } from '../mockExtensions' import type { WFNodeData } from '@shared/types/electron.d' import { PICKER_LABELS, openParamPicker, resolvePickerIntent } from '@shared/utils/paramPicker' -import { PickerIcon } from '@shared/components/ui' +import { FloatInput, IntInput, PickerIcon } from '@shared/components/ui' +import { isMissingWeightVariant, withWeightVariantAvailability } from '@shared/utils/weightVariants' import { useWorkflowRunStore } from '../workflowRunStore' import BaseNode from './BaseNode' @@ -16,6 +18,7 @@ const HANDLE_COLOR: Record = { image: '#38bdf8', mesh: '#a78bfa', text: '#fbbf24', + scene: '#f472b6', } const TAG_CLS: Record = { @@ -23,6 +26,7 @@ const TAG_CLS: Record = { image: 'border-sky-500/30 bg-sky-500/10 text-sky-400', mesh: 'border-violet-500/30 bg-violet-500/10 text-violet-400', text: 'border-amber-500/30 bg-amber-500/10 text-amber-400', + scene: 'border-pink-500/30 bg-pink-500/10 text-pink-400', } // ─── Param control ──────────────────────────────────────────────────────────── @@ -32,55 +36,6 @@ const TAG_CLS: Record = { // node instead. const inputCls = 'nodrag w-full bg-zinc-800 border border-zinc-700 rounded-lg px-2 py-1 text-[11px] text-zinc-200 focus:outline-none focus:border-accent/60' -function IntInput({ value, onChange, className }: { value: number; onChange: (v: number) => void; className: string }) { - const [text, setText] = useState(String(value)) - const prevValue = useRef(value) - if (prevValue.current !== value && parseInt(text, 10) !== value) { - prevValue.current = value - setText(String(value)) - } - return ( - { - const raw = e.target.value - if (raw !== '' && raw !== '-' && !/^-?\d+$/.test(raw)) return - setText(raw) - const n = parseInt(raw, 10) - if (!isNaN(n)) { prevValue.current = n; onChange(n) } - }} - className={className} - /> - ) -} - -function FloatInput({ value, onChange, className }: { value: number; onChange: (v: number) => void; className: string }) { - const [text, setText] = useState(String(value)) - // Sync when external value changes (e.g. reset) - const prevValue = useRef(value) - if (prevValue.current !== value && parseFloat(text.replace(',', '.')) !== value) { - prevValue.current = value - setText(String(value)) - } - return ( - { - const raw = e.target.value.replace(',', '.') - if (raw !== '' && raw !== '-' && raw !== '.' && !/^-?\d*\.?\d*$/.test(raw)) return - setText(e.target.value) - const num = parseFloat(raw) - if (!isNaN(num)) { prevValue.current = num; onChange(num) } - }} - className={className} - /> - ) -} - /** Dropdown of the files inside the folder held by another param (dir_from). */ function FileSelectControl({ param, value, dirValue, onChange }: { param: ParamSchema @@ -151,7 +106,8 @@ function ParamControl({ param, value, onChange, resolvedParams }: { ) } if (param.type === 'float') { - return onChange(v)} className={inputCls} /> + return onChange(v)} className={inputCls} + min={param.min} max={param.max} step={param.step} label={param.label} /> } // int return onChange(v)} className={inputCls} /> @@ -167,8 +123,10 @@ export default function ExtensionNode({ id, data, selected }: { id: string; data const [handleTops, setHandleTops] = useState([]) const { modelExtensions, processExtensions } = useExtensionsStore() + const installedVariants = useExtensionsStore((s) => (data.extensionId ? s.installedWeightVariants[data.extensionId] : undefined)) const allExtensions = buildAllWorkflowExtensions(modelExtensions, processExtensions) const ext = allExtensions.find((e) => e.id === data.extensionId) + const openExtension = useNavStore((s) => s.openExtension) const inputs = ext?.inputs // defined → multi-input mode const isMulti = inputs && inputs.length > 1 @@ -313,7 +271,21 @@ export default function ExtensionNode({ id, data, selected }: { id: string; data
- patchParam(param.id, v)} resolvedParams={resolvedParams} /> + patchParam(param.id, v)} + resolvedParams={resolvedParams} + /> + {/* Offer the install instead of leaving the graph on selection. */} + {ext && isMissingWeightVariant(param.id, val, ext.weightVariants, installedVariants) && ( + + )}
) diff --git a/src/areas/workflows/nodes/ImageNode.tsx b/src/areas/workflows/nodes/ImageNode.tsx index 8c072fe2..7ff44892 100644 --- a/src/areas/workflows/nodes/ImageNode.tsx +++ b/src/areas/workflows/nodes/ImageNode.tsx @@ -2,16 +2,10 @@ import { useCallback, useLayoutEffect, useRef, useState } from 'react' import { Handle, Position, useReactFlow } from '@xyflow/react' import type { WFNodeData } from '@shared/types/electron.d' import BaseNode from './BaseNode' +import { mimeFromPath } from './imageUtils' const OUTPUT_COLOR = '#38bdf8' -function mimeFromPath(p: string): string { - const ext = p.split('.').pop()?.toLowerCase() ?? '' - if (ext === 'jpg' || ext === 'jpeg') return 'image/jpeg' - if (ext === 'webp') return 'image/webp' - return 'image/png' -} - export default function ImageNode({ id, data, selected }: { id: string; data: WFNodeData; selected?: boolean }) { const { updateNodeData } = useReactFlow() const ioRowRef = useRef(null) diff --git a/src/areas/workflows/nodes/ImagePreviewNode.tsx b/src/areas/workflows/nodes/ImagePreviewNode.tsx new file mode 100644 index 00000000..05cb7987 --- /dev/null +++ b/src/areas/workflows/nodes/ImagePreviewNode.tsx @@ -0,0 +1,117 @@ +import { useEffect, useLayoutEffect, useRef, useState } from 'react' +import { Handle, Position, useEdges, useNodes } from '@xyflow/react' +import { useWorkflowRunStore, fromWorkspaceUrl } from '../workflowRunStore' +import { resolveDataSource } from '../nodeBehaviors' +import BaseNode from './BaseNode' +import { mimeFromPath } from './imageUtils' +import type { WFNode } from '@shared/types/electron.d' + +const IO_COLOR = '#38bdf8' + +/** + * Single-image passthrough preview: an image in, the same image forwarded out + * unchanged. Distinct from PreviewImageNode ("Preview Views"), which is a + * terminal node for multi-view strip outputs (e.g. MV-Adapter) and has no + * output handle. + * + * Reads the source file straight off disk and renders it as a `data:` URL + * (rather than an API-served URL) so the preview works without depending on + * the backend being up or CSP allowances for a remote origin. + */ +export default function ImagePreviewNode({ id, selected }: { id: string; selected?: boolean }) { + const nodeImageOutputs = useWorkflowRunStore((s) => s.nodeImageOutputs) + const edges = useEdges() + const nodes = useNodes() + const ioRowRef = useRef(null) + const [handleTop, setHandleTop] = useState('50%') + const [dataUrl, setDataUrl] = useState(undefined) + + useLayoutEffect(() => { + if (ioRowRef.current) { + const center = ioRowRef.current.offsetTop + ioRowRef.current.offsetHeight / 2 + setHandleTop(`${center}px`) + } + }, []) + + const incomingEdge = edges.find((e) => e.target === id) + const nodeMap = new Map(nodes.map((n) => [n.id, n as unknown as WFNode])) + const realSourceId = incomingEdge + ? resolveDataSource(incomingEdge.source, edges, nodeMap) + : undefined + const sourceNode = realSourceId ? nodes.find((n) => n.id === realSourceId) : undefined + const workspaceUrl = realSourceId ? nodeImageOutputs[realSourceId] : undefined + const imageFilePath = sourceNode?.type === 'imageNode' + ? (sourceNode.data as { params?: { filePath?: string } })?.params?.filePath + : undefined + + useEffect(() => { + let cancelled = false + if (!workspaceUrl && !imageFilePath) { + setDataUrl(undefined) + return + } + ;(async () => { + try { + let absPath: string + if (workspaceUrl) { + const settings = await window.electron.settings.get() + absPath = fromWorkspaceUrl(workspaceUrl, settings.workspaceDir) + } else { + absPath = imageFilePath! + } + const base64 = await window.electron.fs.readFileBase64(absPath) + if (!cancelled) setDataUrl(`data:${mimeFromPath(absPath)};base64,${base64}`) + } catch { + if (!cancelled) setDataUrl(undefined) + } + })() + return () => { cancelled = true } + }, [workspaceUrl, imageFilePath]) + + return ( + + + + + + } + subheader={ +
+ image + → + image +
+ } + handles={ + <> + + + + } + > +
+ {dataUrl ? ( + preview + ) : ( +

+ Connect an image and run to preview. +

+ )} +
+
+ ) +} diff --git a/src/areas/workflows/nodes/LoadSceneNode.tsx b/src/areas/workflows/nodes/LoadSceneNode.tsx new file mode 100644 index 00000000..40f3af1c --- /dev/null +++ b/src/areas/workflows/nodes/LoadSceneNode.tsx @@ -0,0 +1,123 @@ +import { useCallback, useLayoutEffect, useRef, useState } from 'react' +import { Handle, Position, useReactFlow } from '@xyflow/react' +import type { ReactFlowInstance } from '@xyflow/react' +import type { WFNode, WFNodeData } from '@shared/types/electron.d' + +import BaseNode from './BaseNode' +import { + applySceneValidationResult, + invalidateValidatedScenePath, + resolveSceneSourceManifest, +} from '../workflowSceneSource' + +const OUTPUT_COLOR = '#f472b6' + +async function validateAndPersistScenePath(args: { + id: string + nextPath: string + updateNodeData: ReactFlowInstance['updateNodeData'] +}): Promise { + args.updateNodeData(args.id, (node) => ({ + params: invalidateValidatedScenePath(node.data.params, args.nextPath), + })) + + const settings = await window.electron.settings.get() + const resolution = await resolveSceneSourceManifest({ + scenePath: args.nextPath, + workspaceDir: settings.workspaceDir, + readFileBase64: window.electron.fs.readFileBase64, + }) + + args.updateNodeData(args.id, (node) => { + const params = applySceneValidationResult(node.data.params, args.nextPath, resolution) + return params ? { params } : {} + }) +} + +export default function LoadSceneNode({ id, data, selected }: { id: string; data: WFNodeData; selected?: boolean }) { + const { updateNodeData } = useReactFlow() + const ioRowRef = useRef(null) + const [handleTop, setHandleTop] = useState('50%') + + useLayoutEffect(() => { + if (ioRowRef.current) { + const center = ioRowRef.current.offsetTop + ioRowRef.current.offsetHeight / 2 + setHandleTop(`${center}px`) + } + }, []) + + const scenePath = typeof data.params.path === 'string' ? data.params.path : '' + const manifestPath = typeof data.params.manifestPath === 'string' ? data.params.manifestPath : undefined + const sceneRoot = typeof data.params.sceneRoot === 'string' ? data.params.sceneRoot : undefined + const error = typeof data.params.error === 'string' ? data.params.error : undefined + + const browseDirectory = useCallback(async () => { + const path = await window.electron.fs.selectDirectory() + if (!path) return + await validateAndPersistScenePath({ id, nextPath: path, updateNodeData }) + }, [id, updateNodeData]) + + const validatePath = useCallback(async () => { + if (!scenePath.trim()) return + await validateAndPersistScenePath({ id, nextPath: scenePath, updateNodeData }) + }, [id, scenePath, updateNodeData]) + + return ( + + + + + + + } + subheader={ +
+ scene +
+ } + handles={ + + } + > +
+ updateNodeData(id, { params: invalidateValidatedScenePath(data.params, event.target.value) })} + className="nodrag w-full rounded-lg border border-zinc-700 bg-zinc-800 px-2.5 py-2 text-[10px] text-zinc-200 placeholder-zinc-600 focus:outline-none focus:border-pink-500/40" + /> +
+ + +
+ {manifestPath ? ( +
+
Manifest: {manifestPath}
+ {sceneRoot &&
sceneRoot: {sceneRoot}
} +
+ ) : ( +
+ Loads an existing workspace scene manifest for downstream scene nodes. +
+ )} + {error &&
{error}
} +
+
+ ) +} diff --git a/src/areas/workflows/nodes/imageUtils.ts b/src/areas/workflows/nodes/imageUtils.ts new file mode 100644 index 00000000..44dd5327 --- /dev/null +++ b/src/areas/workflows/nodes/imageUtils.ts @@ -0,0 +1,8 @@ +// Shared helpers for node components that read image files from disk. + +export function mimeFromPath(p: string): string { + const ext = p.split('.').pop()?.toLowerCase() ?? '' + if (ext === 'jpg' || ext === 'jpeg') return 'image/jpeg' + if (ext === 'webp') return 'image/webp' + return 'image/png' +} diff --git a/src/areas/workflows/nodes/mesh-optimizer/manifest.json b/src/areas/workflows/nodes/mesh-optimizer/manifest.json index 1b13004e..e2480dda 100644 --- a/src/areas/workflows/nodes/mesh-optimizer/manifest.json +++ b/src/areas/workflows/nodes/mesh-optimizer/manifest.json @@ -2,7 +2,7 @@ "id": "mesh-optimizer", "name": "Mesh Optimizer", "type": "process", - "entry": "processor.js", + "entry": "processor.py", "version": "1.0.0", "author": "Modly", "description": "Reduces mesh triangle count using quadric simplification (meshoptimizer).", diff --git a/src/areas/workflows/nodes/mesh-optimizer/processor.py b/src/areas/workflows/nodes/mesh-optimizer/processor.py new file mode 100644 index 00000000..c2842873 --- /dev/null +++ b/src/areas/workflows/nodes/mesh-optimizer/processor.py @@ -0,0 +1,22 @@ +"""Thin workflow adapter for the shared meshoptimizer operation.""" + +import os +import sys +from pathlib import Path + + +api_dir = os.environ.get("MODLY_API_DIR") +if not api_dir: + for parent in Path(__file__).resolve().parents: + candidate = parent / "api" + if candidate.is_dir(): + api_dir = str(candidate) + break +if api_dir and api_dir not in sys.path: + sys.path.insert(0, api_dir) + +from services.mesh_ops.processor import run_processor + + +if __name__ == "__main__": + run_processor("decimate", "mesh-optimizer") diff --git a/src/areas/workflows/nodes/mesh-optimizer/processor.ts b/src/areas/workflows/nodes/mesh-optimizer/processor.ts deleted file mode 100644 index f6f16d4e..00000000 --- a/src/areas/workflows/nodes/mesh-optimizer/processor.ts +++ /dev/null @@ -1,88 +0,0 @@ -import path = require('path') - -interface ProcessInput { filePath?: string; text?: string } -interface ProcessResult { filePath?: string; text?: string } -interface ProcessContext { - workspaceDir: string - tempDir: string - log: (msg: string) => void - progress: (pct: number, label: string) => void -} - -const processor = async ( - input: ProcessInput, - params: Record, - context: ProcessContext, -): Promise => { - if (!input.filePath) throw new Error('mesh-optimizer: input.filePath is required') - - const targetFaces = Math.max(100, Math.round(Number(params['target_faces'] ?? 10000))) - context.log(`Target: ${targetFaces} triangles — input: ${input.filePath}`) - - // Lazy requires — resolved from the extension's own node_modules - const { NodeIO } = require('@gltf-transform/core') - const { ALL_EXTENSIONS } = require('@gltf-transform/extensions') - const { simplify, weld } = require('@gltf-transform/functions') - const { MeshoptSimplifier } = require('meshoptimizer') - - // MeshoptSimplifier loads a WASM binary asynchronously - await MeshoptSimplifier.ready - - context.progress(10, 'Loading mesh…') - const io = new NodeIO().registerExtensions(ALL_EXTENSIONS) - const doc = await io.read(input.filePath) - - // Count current triangles across all primitives - let currentFaces = 0 - for (const mesh of doc.getRoot().listMeshes()) { - for (const prim of mesh.listPrimitives()) { - const indices = prim.getIndices() - if (indices) { - currentFaces += Math.round(indices.getCount() / 3) - } else { - const pos = prim.getAttribute('POSITION') - if (pos) currentFaces += Math.round(pos.getCount() / 3) - } - } - } - context.log(`Current triangles: ${currentFaces}`) - - if (currentFaces <= targetFaces) { - context.log('Already within target — skipping simplification') - context.progress(100, 'Done') - return { filePath: input.filePath } - } - - const ratio = Math.min(1, targetFaces / currentFaces) - context.log(`Simplification ratio: ${ratio.toFixed(4)} (~${Math.round(currentFaces * ratio)} triangles)`) - - // error tolerance scales with aggressiveness: tighter simplification needs more room - const error = Math.max(0.001, 1 - ratio) - - // Skip weld on large meshes — deduplication is O(N²) and stalls for millions of faces - if (currentFaces < 500_000) { - context.progress(25, 'Welding vertices…') - await doc.transform(weld()) - } else { - context.log(`Skipping weld (${currentFaces} faces > 500k threshold)`) - } - - context.progress(55, 'Simplifying mesh…') - await doc.transform( - simplify({ simplifier: MeshoptSimplifier, ratio, error, lockBorder: false }), - ) - - context.progress(85, 'Writing output…') - // Save to workspaceDir/Workflows/ so the result lands in the workspace - const outDir = path.join(context.workspaceDir, 'Workflows') - require('fs').mkdirSync(outDir, { recursive: true }) - const outPath = path.join(outDir, `mesh-optimizer-${Date.now()}.glb`) - await io.write(outPath, doc) - - context.progress(100, 'Done') - context.log(`Output: ${outPath}`) - - return { filePath: outPath } -} - -export = processor diff --git a/src/areas/workflows/nodes/mesh-repair/processor.py b/src/areas/workflows/nodes/mesh-repair/processor.py index b09979e4..6ff43d02 100644 --- a/src/areas/workflows/nodes/mesh-repair/processor.py +++ b/src/areas/workflows/nodes/mesh-repair/processor.py @@ -1,153 +1,22 @@ -""" -Mesh Repair — built-in process extension. +"""Thin workflow adapter for the shared repair mesh operation.""" -Fixes common topology issues in AI-generated meshes: - - Duplicate vertices and faces - - Non-manifold edges - - Degenerate (zero-area) faces - - Simple boundary holes - -Note: structural holes from FlexiCubes/TRELLIS voxel extraction cannot be -reliably closed in post-processing. Increase the generator's remesh resolution -to reduce them at the source. - -Protocol: reads one JSON line from stdin, writes JSON lines to stdout. - stdin : { input, params, workspaceDir, tempDir } - stdout: { type: "progress"|"log"|"done"|"error", ... } -""" -import json import os -import shutil import sys -import tempfile from pathlib import Path -def emit(obj: dict) -> None: - print(json.dumps(obj), flush=True) - - -def progress(pct: int, label: str) -> None: - emit({"type": "progress", "percent": pct, "label": label}) - - -def log(msg: str) -> None: - emit({"type": "log", "message": msg}) - - -def done(file_path: str) -> None: - emit({"type": "done", "result": {"filePath": file_path}}) - - -def error(msg: str) -> None: - emit({"type": "error", "message": msg}) - - -def main() -> None: - raw = sys.stdin.readline() - data = json.loads(raw) - - input_data = data.get("input", {}) - params = data.get("params", {}) - workspace_dir = data.get("workspaceDir", "") - - input_path = input_data.get("filePath") - if not input_path or not Path(input_path).is_file(): - error(f"mesh-repair: input file not found: {input_path}") - return - - do_remove_dupes = bool(params.get("remove_duplicates", True)) - do_fix_non_manifold = bool(params.get("fix_non_manifold", True)) - do_remove_degen = bool(params.get("remove_degenerate", True)) - do_fill_holes = bool(params.get("fill_holes", True)) - max_hole_size = int(params.get("max_hole_size", 2000)) - - out_dir = Path(workspace_dir) / "Workflows" - out_dir.mkdir(parents=True, exist_ok=True) - from time import time - out_path = str(out_dir / f"mesh-repair-{int(time() * 1000)}.glb") - - try: - import pymeshlab - except ImportError: - error("mesh-repair: pymeshlab is not available on this system") - return - - import trimesh - - progress(10, "Loading mesh…") - loaded = trimesh.load(input_path) - if isinstance(loaded, trimesh.Scene): - geoms = list(loaded.geometry.values()) - geom = trimesh.util.concatenate(geoms) if len(geoms) > 1 else geoms[0] - else: - geom = loaded - - tmp_dir = tempfile.mkdtemp() - try: - ply_in = os.path.join(tmp_dir, "input.ply") - ply_out = os.path.join(tmp_dir, "output.ply") - geom.export(ply_in) - - ms = pymeshlab.MeshSet() - ms.load_new_mesh(ply_in) - - log(f"Input: {ms.current_mesh().vertex_number()} verts, {ms.current_mesh().face_number()} faces") - - if do_remove_dupes: - progress(20, "Removing duplicates…") - ms.meshing_remove_duplicate_vertices() - ms.meshing_remove_duplicate_faces() - - if do_remove_degen: - progress(40, "Removing degenerate faces…") - ms.meshing_remove_null_faces() - ms.meshing_remove_folded_faces() - - if do_fix_non_manifold: - progress(60, "Fixing non-manifold edges…") - # method=0 removes offending faces (low memory); method=1 detaches (OOMs on dense meshes) - try: - ms.meshing_repair_non_manifold_edges(method=0) - except Exception as e: - log(f"Non-manifold edge repair skipped: {e}") - try: - ms.meshing_repair_non_manifold_vertices() - except Exception as e: - log(f"Non-manifold vertex repair skipped: {e}") - - if do_fill_holes: - progress(75, "Filling holes…") - try: - ms.meshing_close_holes( - maxholesize=max_hole_size, - newfaceselected=False, - selfintersection=False, - ) - except Exception as e: - log(f"Hole fill skipped (mesh may still be non-manifold): {e}") - - after = ms.current_mesh().face_number() - log(f"Output: {ms.current_mesh().vertex_number()} verts, {after} faces") - - progress(85, "Exporting…") - ms.save_current_mesh(ply_out) - _loaded = trimesh.load(ply_out, process=False) - if isinstance(_loaded, trimesh.Scene): - _geoms = list(_loaded.geometry.values()) - _loaded = _geoms[0] if len(_geoms) == 1 else trimesh.util.concatenate(_geoms) - result = trimesh.Trimesh(vertices=_loaded.vertices, faces=_loaded.faces, process=False) - finally: - shutil.rmtree(tmp_dir, ignore_errors=True) +api_dir = os.environ.get("MODLY_API_DIR") +if not api_dir: + for parent in Path(__file__).resolve().parents: + candidate = parent / "api" + if candidate.is_dir(): + api_dir = str(candidate) + break +if api_dir and api_dir not in sys.path: + sys.path.insert(0, api_dir) - result.export(out_path) - progress(100, "Done") - done(out_path) +from services.mesh_ops.processor import run_processor if __name__ == "__main__": - try: - main() - except Exception as exc: - import traceback - error(f"{exc}\n{traceback.format_exc()}") + run_processor("repair", "mesh-repair") diff --git a/src/areas/workflows/nodes/mesh-smoother/processor.py b/src/areas/workflows/nodes/mesh-smoother/processor.py index 3ae78b61..b59dcb38 100644 --- a/src/areas/workflows/nodes/mesh-smoother/processor.py +++ b/src/areas/workflows/nodes/mesh-smoother/processor.py @@ -1,124 +1,22 @@ -""" -Mesh Smoother — built-in process extension. +"""Thin workflow adapter for the shared smooth mesh operation.""" -Reduces sharp artifacts (zipper triangles, sawtooth edges) produced by -AI mesh generators via Taubin or Laplacian smoothing. - -Protocol: reads one JSON line from stdin, writes JSON lines to stdout. - stdin : { input, params, workspaceDir, tempDir } - stdout: { type: "progress"|"log"|"done"|"error", ... } -""" -import json import os -import shutil import sys -import tempfile from pathlib import Path -def emit(obj: dict) -> None: - print(json.dumps(obj), flush=True) - - -def progress(pct: int, label: str) -> None: - emit({"type": "progress", "percent": pct, "label": label}) - - -def log(msg: str) -> None: - emit({"type": "log", "message": msg}) - - -def done(file_path: str) -> None: - emit({"type": "done", "result": {"filePath": file_path}}) - - -def error(msg: str) -> None: - emit({"type": "error", "message": msg}) - - -def main() -> None: - raw = sys.stdin.readline() - data = json.loads(raw) - - input_data = data.get("input", {}) - params = data.get("params", {}) - workspace_dir = data.get("workspaceDir", "") - - input_path = input_data.get("filePath") - if not input_path or not Path(input_path).is_file(): - error(f"mesh-smoother: input file not found: {input_path}") - return - - iterations = int(params.get("iterations", 5)) - lambda_ = float(params.get("lambda_", 0.5)) - mode = str(params.get("mode", "taubin")) - - out_dir = Path(workspace_dir) / "Workflows" - out_dir.mkdir(parents=True, exist_ok=True) - from time import time - out_path = str(out_dir / f"mesh-smoother-{int(time() * 1000)}.glb") - - log(f"Mode: {mode}, iterations: {iterations}, strength: {lambda_}") - - try: - import pymeshlab - except ImportError: - error("mesh-smoother: pymeshlab is not available on this system") - return - - import trimesh - - progress(10, "Loading mesh…") - loaded = trimesh.load(input_path) - if isinstance(loaded, trimesh.Scene): - geoms = list(loaded.geometry.values()) - geom = trimesh.util.concatenate(geoms) if len(geoms) > 1 else geoms[0] - else: - geom = loaded - - tmp_dir = tempfile.mkdtemp() - try: - ply_in = os.path.join(tmp_dir, "input.ply") - ply_out = os.path.join(tmp_dir, "output.ply") - geom.export(ply_in) - - ms = pymeshlab.MeshSet() - ms.load_new_mesh(ply_in) - - progress(30, f"Smoothing ({mode})…") - - if mode == "taubin": - ms.apply_coord_taubin_smoothing( - lambda_=lambda_, - mu=-lambda_ - 0.01, - stepsmoothnum=iterations, - ) - else: - ms.apply_coord_laplacian_smoothing( - stepsmoothnum=iterations, - cotangentweight=False, - ) - - progress(80, "Exporting…") - ms.save_current_mesh(ply_out) - # Load raw geometry only — avoids scipy dependency triggered by face→vertex color conversion - _loaded = trimesh.load(ply_out, process=False) - if isinstance(_loaded, trimesh.Scene): - _geoms = list(_loaded.geometry.values()) - _loaded = _geoms[0] if len(_geoms) == 1 else trimesh.util.concatenate(_geoms) - result = trimesh.Trimesh(vertices=_loaded.vertices, faces=_loaded.faces, process=False) - finally: - shutil.rmtree(tmp_dir, ignore_errors=True) +api_dir = os.environ.get("MODLY_API_DIR") +if not api_dir: + for parent in Path(__file__).resolve().parents: + candidate = parent / "api" + if candidate.is_dir(): + api_dir = str(candidate) + break +if api_dir and api_dir not in sys.path: + sys.path.insert(0, api_dir) - result.export(out_path) - log(f"Output: {out_path} ({len(result.faces)} faces)") - progress(100, "Done") - done(out_path) +from services.mesh_ops.processor import run_processor if __name__ == "__main__": - try: - main() - except Exception as exc: - import traceback - error(f"{exc}\n{traceback.format_exc()}") + run_processor("smooth", "mesh-smoother") diff --git a/src/areas/workflows/preflight.test.mjs b/src/areas/workflows/preflight.test.mjs index 75b09d71..3b4df8aa 100644 --- a/src/areas/workflows/preflight.test.mjs +++ b/src/areas/workflows/preflight.test.mjs @@ -131,3 +131,38 @@ test('multi-input extension requires every declared input type', () => { assert.ok(!issues.some((i) => i.key === 'proc:missing:image')) assert.ok(issues.some((i) => i.key === 'proc:missing:text')) }) + +test('scene input requires a validated Load Scene source and rejects image wiring', () => { + const { validateWorkflowPreflight } = loadModule() + const model = { id: 'model', type: 'extensionNode', position: { x: 0, y: 0 }, data: { extensionId: 'pack/process-node' } } + const scene = { id: 'scene', type: 'sceneNode', position: { x: 0, y: 0 }, data: { params: { manifestPath: 'Workflows/room/scene-manifest.json' } } } + for (const output of ['scene', 'mesh']) { + const extension = ext({ input: 'scene', output, type: 'model' }) + assert.deepEqual(validateWorkflowPreflight(wf([scene, model], [{ id: 'scene-edge', source: 'scene', target: 'model' }]), [extension]), []) + } + + const extension = ext({ input: 'scene', output: 'scene', type: 'model' }) + const issues = validateWorkflowPreflight(wf([imageNode(), model], [{ id: 'image-edge', source: 'img', target: 'model' }]), [extension]) + assert.ok(issues.some((issue) => issue.key === 'model:missing:scene')) + assert.ok(issues.some((issue) => issue.key === 'model:type:image-edge')) +}) + +test('Load Scene must be validated before a workflow can run', () => { + const { validateWorkflowPreflight } = loadModule() + const scene = { id: 'scene', type: 'sceneNode', position: { x: 0, y: 0 }, data: { params: { path: 'Workflows/room' } } } + const issues = validateWorkflowPreflight(wf([scene], []), []) + assert.equal(issues[0].key, 'scene:scene-invalid') +}) + +test('renderer fails closed for unsupported process and mixed scene node shapes', () => { + const { validateWorkflowPreflight } = loadModule() + const scene = { id: 'scene', type: 'sceneNode', position: { x: 0, y: 0 }, data: { params: { manifestPath: 'Workflows/room/scene-manifest.json' } } } + const target = { id: 'target', type: 'extensionNode', position: { x: 0, y: 0 }, data: { extensionId: 'pack/process-node' } } + for (const extension of [ + ext({ input: 'scene', output: 'mesh', type: 'process' }), + ext({ input: 'scene', inputs: ['scene', 'text'], output: 'mesh', type: 'model' }), + ]) { + const issues = validateWorkflowPreflight(wf([scene, target], [{ id: 'e', source: 'scene', target: 'target' }]), [extension]) + assert.ok(issues.some((issue) => issue.key === 'target:unsupported-scene-shape')) + } +}) diff --git a/src/areas/workflows/preflight.ts b/src/areas/workflows/preflight.ts index b33c1851..068a9277 100644 --- a/src/areas/workflows/preflight.ts +++ b/src/areas/workflows/preflight.ts @@ -1,8 +1,9 @@ import type { Workflow, WFNode } from '@shared/types/electron.d' import { getWorkflowExtension, type WorkflowExtension } from './mockExtensions' +import { hasUnsupportedSceneShape } from './sceneShape' import { isPassthrough, isBranchConsumer, resolveDataSource, nearestUpstreamWaits } from './nodeBehaviors' -type DataType = 'image' | 'text' | 'mesh' | 'audio' +type DataType = 'image' | 'text' | 'mesh' | 'audio' | 'scene' export interface WorkflowPreflightIssue { key: string @@ -14,8 +15,10 @@ function nodeLabel(node: WFNode, allExtensions: WorkflowExtension[]): string { if (node.type === 'imageNode') return 'Image' if (node.type === 'textNode') return 'Text' if (node.type === 'meshNode') return 'Load 3D Mesh' + if (node.type === 'sceneNode') return 'Load Scene' if (node.type === 'outputNode') return 'Add to Scene' if (node.type === 'previewNode') return 'Preview Views' + if (node.type === 'imagePreviewNode') return 'Preview Image' if (node.type === 'forEachNode') { const mode = (node.data.params?.mode as string) ?? 'image' return mode === 'text' ? 'For Each Text' : mode === 'mesh' ? 'For Each Mesh' : 'For Each Image' @@ -27,6 +30,7 @@ function nodeLabel(node: WFNode, allExtensions: WorkflowExtension[]): string { } function formatType(type: DataType): string { + if (type === 'scene') return 'scene' if (type === 'mesh') return 'mesh' if (type === 'image') return 'image' if (type === 'audio') return 'audio' @@ -43,7 +47,9 @@ function getNodeOutputType(node: WFNode, allExtensions: WorkflowExtension[]): Da if (node.type === 'imageNode') return 'image' if (node.type === 'textNode') return 'text' if (node.type === 'meshNode' || node.type === 'outputNode') return 'mesh' + if (node.type === 'sceneNode') return 'scene' if (node.type === 'previewNode') return 'image' + if (node.type === 'imagePreviewNode') return 'image' if (node.type === 'forEachNode') { const mode = (node.data.params?.mode as DataType | undefined) ?? 'image' return mode === 'text' || mode === 'mesh' ? mode : 'image' @@ -94,6 +100,13 @@ export function validateWorkflowPreflight( }) } + if (node.type === 'sceneNode' && !((node.data.params?.manifestPath as string | undefined)?.trim())) { + pushIssue(issues, { + key: `${node.id}:scene-invalid`, nodeId: node.id, + message: 'Load Scene needs a validated scene directory.', + }) + } + // A node fed by two different Wait branches can't be scheduled into a single // branch — it would run before either branch produces its mesh. if ( @@ -119,6 +132,15 @@ export function validateWorkflowPreflight( continue } + if (hasUnsupportedSceneShape(ext)) { + pushIssue(issues, { + key: `${node.id}:unsupported-scene-shape`, + nodeId: node.id, + message: `${ext.name} uses an unsupported scene input or output declaration.`, + }) + continue + } + const incomingEdges = workflow.edges.filter((edge) => edge.target === node.id) const requiredTypes = [...new Set((ext.inputs ?? [ext.input]) as DataType[])] diff --git a/src/areas/workflows/sceneShape.ts b/src/areas/workflows/sceneShape.ts new file mode 100644 index 00000000..5c6814bf --- /dev/null +++ b/src/areas/workflows/sceneShape.ts @@ -0,0 +1,11 @@ +import type { WorkflowExtension } from './mockExtensions' + +/** + * Scene is model-only and must be the node's single `input` (never inside + * `inputs`). Mirrors assertSupportedSceneNodeShape on the install side. + */ +export function hasUnsupportedSceneShape(ext: Pick): boolean { + const usesSceneInput = ext.input === 'scene' || ext.inputs?.includes('scene') === true + if (ext.type === 'process') return usesSceneInput || ext.output === 'scene' + return usesSceneInput && (ext.inputs !== undefined || ext.input !== 'scene') +} diff --git a/src/areas/workflows/slotInputs.test.mjs b/src/areas/workflows/slotInputs.test.mjs new file mode 100644 index 00000000..16adecd2 --- /dev/null +++ b/src/areas/workflows/slotInputs.test.mjs @@ -0,0 +1,54 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +// slotInputs.ts has no runtime imports, so esbuild bundles it standalone. +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-slotinputs-test-')), 'slotInputs.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('src/areas/workflows/slotInputs.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { assignSlotFilePaths } = loadModule() + +test('a mesh slot becomes the mesh path', () => { + const r = assignSlotFilePaths(['mesh', 'image'], ['a.glb', 'b.png']) + assert.equal(r.nodeInputMeshPath, 'a.glb') + assert.equal(r.nodeInputPath, 'b.png') + assert.deepEqual(r.extraImagePaths, []) +}) + +test('the first image slot is the primary path and later images are extras', () => { + const r = assignSlotFilePaths(['image', 'image', 'image'], ['1.png', '2.png', '3.png']) + assert.equal(r.nodeInputPath, '1.png') + assert.deepEqual(r.extraImagePaths, ['2.png', '3.png']) + assert.equal(r.nodeInputMeshPath, undefined) +}) + +test('an audio slot becomes the primary path', () => { + // Regression: audio slots were resolved and then dropped, so a multi-input + // node such as [audio, text] failed with "needs an incoming audio connection". + const r = assignSlotFilePaths(['audio', 'text'], ['song.wav', undefined]) + assert.equal(r.nodeInputPath, 'song.wav') + assert.equal(r.nodeInputMeshPath, undefined) + assert.deepEqual(r.extraImagePaths, []) +}) + +test('text slots never claim a file path, and empty slots are skipped', () => { + const r = assignSlotFilePaths(['text', 'image'], ['leaked.png', undefined]) + assert.equal(r.nodeInputPath, undefined) + assert.equal(r.nodeInputMeshPath, undefined) + assert.deepEqual(r.extraImagePaths, []) +}) diff --git a/src/areas/workflows/slotInputs.ts b/src/areas/workflows/slotInputs.ts new file mode 100644 index 00000000..c07705e9 --- /dev/null +++ b/src/areas/workflows/slotInputs.ts @@ -0,0 +1,36 @@ +// Multi-input slot → file-path assignment for extension nodes. +// +// `inputTypes` is the node's declared `inputs` (one per handle); `inputPaths` is +// the file path resolved for each slot, indexed the same way (undefined where the +// slot carries text or nothing). Pure so it can be tested without the store. + +export type SlotInputType = 'image' | 'text' | 'mesh' | 'audio' + +export interface SlotFilePaths { + /** Primary file: what the extension receives as `filePath` when no mesh is present. */ + nodeInputPath?: string + /** Mesh slot, if any. Takes precedence as `filePath`; the primary file then rides in params. */ + nodeInputMeshPath?: string + /** Every image beyond the first resolved image slot. */ + extraImagePaths: string[] +} + +export function assignSlotFilePaths( + inputTypes: readonly SlotInputType[], + inputPaths: readonly (string | undefined)[], +): SlotFilePaths { + const out: SlotFilePaths = { extraImagePaths: [] } + for (let i = 0; i < inputTypes.length; i++) { + const fp = inputPaths[i] + if (!fp) continue + if (inputTypes[i] === 'mesh') { + out.nodeInputMeshPath = fp + } else if (inputTypes[i] === 'image') { + if (!out.nodeInputPath) out.nodeInputPath = fp + else out.extraImagePaths.push(fp) + } else if (inputTypes[i] === 'audio') { + if (!out.nodeInputPath) out.nodeInputPath = fp + } + } + return out +} diff --git a/src/areas/workflows/workflowRunStore.ts b/src/areas/workflows/workflowRunStore.ts index 4ae35b3a..1c64c0c5 100644 --- a/src/areas/workflows/workflowRunStore.ts +++ b/src/areas/workflows/workflowRunStore.ts @@ -2,10 +2,13 @@ import { create } from 'zustand' import axios, { AxiosInstance } from 'axios' import { useAppStore } from '@shared/stores/appStore' import { getWorkflowExtension } from './mockExtensions' -import { showCompletionNotification } from '@shared/utils/notification' +import { hasUnsupportedSceneShape } from './sceneShape' +import { showCompletionNotification, showErrorNotification } from '@shared/utils/notification' import type { WorkflowExtension } from './mockExtensions' import type { Workflow, WFNode, WFEdge } from '@shared/types/electron.d' import { isBranchStarter, isSceneOutput, resolveDataSource, reachesSceneOutput, nearestUpstreamWaits } from './nodeBehaviors' +import { assignSlotFilePaths } from './slotInputs' +import type { SlotInputType } from './slotInputs' // ─── Types ──────────────────────────────────────────────────────────────────── @@ -118,6 +121,14 @@ function toWorkspaceUrl(filePath: string, workspaceDir: string): string | undefi return `/workspace/${norm.slice(workspaceDir.length).replace(/^\//, '')}` } +// Inverse of toWorkspaceUrl — used by node components that read a `/workspace/...` +// output URL back off disk (e.g. to build a data: URL for a preview). +export function fromWorkspaceUrl(url: string, workspaceDir: string): string { + const wsDir = workspaceDir.replace(/\\/g, '/').replace(/\/+$/, '') + const rel = url.replace(/^\/workspace\//, '') + return `${wsDir}/${rel}` +} + interface RunContext { workflow: Workflow allExtensions: WorkflowExtension[] @@ -297,6 +308,9 @@ async function executeExtensionNode( selectedImagePath, selectedImageData } = ctx const ext = getWorkflowExtension(node.data.extensionId ?? '', allExtensions) + if (ext && hasUnsupportedSceneShape(ext)) { + throw new Error(`${ext.name} uses an unsupported scene input or output declaration`) + } // Freshest params at the moment the node starts (so loop iterations / Retry pick // up edits made while paused, not the values captured at run start). const liveParams = _liveParams.current.get(node.id) ?? node.data.params ?? {} @@ -309,6 +323,7 @@ async function executeExtensionNode( let nodeInputPath: string | undefined let nodeInputText: string | undefined let nodeInputMeshPath: string | undefined + let nodeInputScenePath: string | undefined // Per-slot texts for multi-text-input nodes (e.g. positive/negative prompts). // Indexed by target handle: input-0 → texts[0], input-1 → texts[1]. const nodeInputTexts: (string | undefined)[] = [] @@ -319,7 +334,10 @@ async function executeExtensionNode( const incomingEdges = workflow.edges.filter((e) => e.target === node.id) if (ext?.inputs && ext.inputs.length > 1) { - const inputTypes = ext.inputs + // Scene is intentionally a single-input-only model contract. The guard + // above rejects it before this multi-slot path, and filtering it here also + // narrows the remaining declarations to assignSlotFilePaths' exact ABI. + const inputTypes: SlotInputType[] = ext.inputs.filter((input) => input !== 'scene') // Resolved by target handle first, then typed by that slot's declared input -- // not by the arrival order of `incomingEdges`, which does not match slot order. const inputPaths = new Array(inputTypes.length).fill(undefined) @@ -336,36 +354,33 @@ async function executeExtensionNode( } } - for (let i = 0; i < inputTypes.length; i++) { - const fp = inputPaths[i] - if (!fp) continue - if (inputTypes[i] === 'mesh') { - nodeInputMeshPath = fp - } else if (inputTypes[i] === 'image') { - if (!nodeInputPath) nodeInputPath = fp - else extraImagePaths.push(fp) - } - } + const slots = assignSlotFilePaths(inputTypes, inputPaths) + nodeInputPath = slots.nodeInputPath + nodeInputMeshPath = slots.nodeInputMeshPath + extraImagePaths.push(...slots.extraImagePaths) } else { for (const edge of incomingEdges) { const src = resolveSource(edge.source) if (src?.filePath !== undefined) nodeInputPath = src.filePath if (src?.text !== undefined && src.text.trim().length > 0) nodeInputText = src.text + if (src?.outputType === 'scene') nodeInputScenePath = src.filePath } } const isModelNode = ext?.type === 'model' if (isModelNode) { + const isSceneInput = ext?.inputs ? ext.inputs.includes('scene') : ext?.input === 'scene' const isTextInput = ext?.inputs ? ext.inputs.every((i) => i === 'text') : ext?.input === 'text' - const activeImagePath = isTextInput ? undefined : (nodeInputPath ?? selectedImagePath) - if (!isTextInput && !selectedImageData && (!activeImagePath || activeImagePath.trim().length === 0)) { + if (isSceneInput && !nodeInputScenePath) throw new Error(`${ext?.name ?? 'Model'} needs an incoming scene connection`) + const activeImagePath = (isTextInput || isSceneInput) ? undefined : (nodeInputPath ?? selectedImagePath) + if (!isTextInput && !isSceneInput && !selectedImageData && (!activeImagePath || activeImagePath.trim().length === 0)) { throw new Error('No input image selected for model node') } let blob: Blob let fname: string - if (isTextInput || (selectedImageData && nodeInputPath === undefined)) { + if (isTextInput || isSceneInput || (selectedImageData && nodeInputPath === undefined)) { const base64 = selectedImageData && nodeInputPath === undefined ? selectedImageData : 'iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==' // 1x1 transparent PNG @@ -398,21 +413,30 @@ async function executeExtensionNode( ) const effectiveParams = { ...schemaDefaults, ...liveParams } - const fd = new FormData() - fd.append('image', blob, fname) - fd.append('model_id', node.data.extensionId ?? '') - fd.append('collection', 'Workflows') - fd.append('remesh', 'none') - fd.append('enable_texture', 'false') - fd.append('texture_resolution', '1024') - fd.append('params', JSON.stringify({ ...effectiveParams, ...extraParams })) - setRunState((s) => ({ ...s, blockProgress: 5, blockStep: 'Submitting to model…' })) - - const { data } = await client.post<{ job_id: string }>( - '/generate/from-image', fd, - { headers: { 'Content-Type': 'multipart/form-data' } }, - ) + let submission: { data: { job_id: string } } + if (isSceneInput) { + const normalized = nodeInputScenePath!.replace(/\\/g, '/') + const inputPath = normalized.startsWith(`${workspaceDir}/`) + ? normalized.slice(workspaceDir.length + 1) + : normalized.replace(/^\/workspace\//, '') + submission = await client.post('/generate/from-artifact', { + input_kind: 'scene', input_path: inputPath, + model_id: node.data.extensionId ?? '', collection: 'Workflows', + params: { ...effectiveParams, ...extraParams }, + }) + } else { + const fd = new FormData() + fd.append('image', blob, fname) + fd.append('model_id', node.data.extensionId ?? '') + fd.append('collection', 'Workflows') + fd.append('remesh', 'none') + fd.append('enable_texture', 'false') + fd.append('texture_resolution', '1024') + fd.append('params', JSON.stringify({ ...effectiveParams, ...extraParams })) + submission = await client.post('/generate/from-image', fd, { headers: { 'Content-Type': 'multipart/form-data' } }) + } + const { data } = submission _activeJobId.current = data.job_id while (true) { @@ -591,7 +615,7 @@ export const useWorkflowRunStore = create((set, get) => { if (!outputUrl) { for (const [, o] of ctx.nodeOutputs) { if (o.filePath) { - if (o.outputType === 'audio') { + if (o.outputType === 'audio' || o.outputType === 'scene') { outputPath = o.filePath continue } @@ -799,6 +823,13 @@ export const useWorkflowRunStore = create((set, get) => { if (fp) nodeOutputs.set(node.id, { filePath: fp, outputType: 'mesh' }) } } + if (node.type === 'sceneNode') { + const manifestPath = node.data.params?.manifestPath as string | undefined + if (manifestPath) nodeOutputs.set(node.id, { + filePath: `${workspaceDir}/${manifestPath.replace(/^\/+/, '')}`, + outputType: 'scene', + }) + } } const ctx: RunContext = { @@ -942,6 +973,7 @@ export const useWorkflowRunStore = create((set, get) => { if (!_cancel.current) { set((s) => ({ runState: { ...s.runState, status: 'error', error: String(err) }, activeNodeId: null })) useAppStore.getState().updateCurrentJob({ status: 'error', error: String(err) }) + void showErrorNotification(String(err), 'Workflow run failed') } } }, diff --git a/src/areas/workflows/workflowSceneRun.test.mjs b/src/areas/workflows/workflowSceneRun.test.mjs new file mode 100644 index 00000000..cb43b9d5 --- /dev/null +++ b/src/areas/workflows/workflowSceneRun.test.mjs @@ -0,0 +1,60 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { build } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +const dir = mkdtempSync(join(tmpdir(), 'modly-scene-run-')) +const stub = (name, source) => { const path = join(dir, name); writeFileSync(path, source); return path } +const appStoreStub = stub('app.ts', ` +export const appState: any = { apiUrl: 'http://modly.test', currentJob: null, + setCurrentJob(value: any) { this.currentJob = value }, + updateCurrentJob(value: any) { this.currentJob = { ...(this.currentJob ?? {}), ...value } } } +export const useAppStore: any = (selector: any) => selector(appState) +useAppStore.getState = () => appState +`) +const axiosStub = stub('axios.ts', `const axios: any = { create: () => (globalThis as any).__sceneClient }; export default axios; export type AxiosInstance = any`) +const extStub = stub('ext.ts', `export const getWorkflowExtension = (id: string, all: any[]) => all.find((value) => value.id === id); export type WorkflowExtension = any`) +const notifyStub = stub('notify.ts', `export const showCompletionNotification = async () => {}; export const showErrorNotification = async () => {}`) +const aliases = new Map([ + ['axios', axiosStub], ['@shared/stores/appStore', appStoreStub], + ['./mockExtensions', extStub], ['@shared/utils/notification', notifyStub], +]) +const outfile = join(dir, 'store.cjs') +writeFileSync(outfile, (await build({ + entryPoints: [resolve('src/areas/workflows/workflowRunStore.ts')], bundle: true, + platform: 'node', format: 'cjs', write: false, + plugins: [{ name: 'aliases', setup(build) { build.onResolve({ filter: /.*/ }, (args) => aliases.has(args.path) ? { path: aliases.get(args.path) } : null) } }], +})).outputFiles[0].text) +const { useWorkflowRunStore } = createRequire(import.meta.url)(outfile) + +test('scene model uses typed artifact route and preserves scene output', async () => { + const posts = [] + globalThis.window = { electron: { + settings: { get: async () => ({ workspaceDir: '/workspace' }) }, + fs: { deleteDirectory: async () => ({ success: true }), listFiles: async () => [], readFileBase64: async () => { throw new Error('scene must not be read as image bytes') } }, + } } + globalThis.__sceneClient = { + post: async (url, body) => { posts.push({ url, body }); return { data: { job_id: 'scene-job' } } }, + get: async () => ({ data: { status: 'done', progress: 100, output_url: '/workspace/Workflows/result/scene-manifest.json' } }), + } + const workflow = { + id: 'wf', name: 'Scene', description: '', createdAt: '', updatedAt: '', + nodes: [ + { id: 'source', type: 'sceneNode', position: { x: 0, y: 0 }, data: { enabled: true, params: { manifestPath: 'Workflows/input/scene-manifest.json' } } }, + { id: 'model', type: 'extensionNode', position: { x: 1, y: 0 }, data: { enabled: true, extensionId: 'pixal/world', params: {} } }, + ], + edges: [{ id: 'e', source: 'source', target: 'model' }], + } + const extension = { id: 'pixal/world', name: 'World', type: 'model', input: 'scene', output: 'scene', params: [] } + await useWorkflowRunStore.getState().run(workflow, [extension]) + assert.equal(posts[0].url, '/generate/from-artifact') + assert.deepEqual(posts[0].body, { + input_kind: 'scene', input_path: 'Workflows/input/scene-manifest.json', + model_id: 'pixal/world', collection: 'Workflows', params: {}, + }) + assert.equal(useWorkflowRunStore.getState().runState.outputPath, '/workspace/Workflows/result/scene-manifest.json') + assert.equal(useWorkflowRunStore.getState().runState.outputUrl, undefined) +}) diff --git a/src/areas/workflows/workflowSceneSource.test.mjs b/src/areas/workflows/workflowSceneSource.test.mjs new file mode 100644 index 00000000..778e13d3 --- /dev/null +++ b/src/areas/workflows/workflowSceneSource.test.mjs @@ -0,0 +1,70 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +const outfile = join(mkdtempSync(join(tmpdir(), 'modly-scene-source-')), 'scene.cjs') +writeFileSync(outfile, buildSync({ entryPoints: [resolve('src/areas/workflows/workflowSceneSource.ts')], bundle: true, platform: 'node', format: 'cjs', write: false }).outputFiles[0].text) +const { + applySceneValidationResult, + invalidateValidatedScenePath, + resolveSceneSourceManifest, +} = createRequire(import.meta.url)(outfile) +const encoded = Buffer.from(JSON.stringify({ schema: 'modly.scene-manifest.v1', sceneRoot: '.', assets: [] })).toString('base64') + +test('Load Scene resolves directory and manifest without image bytes', async () => { + for (const scenePath of ['Workflows/room', 'Workflows/room/scene-manifest.json']) { + const result = await resolveSceneSourceManifest({ scenePath, workspaceDir: '/workspace', readFileBase64: async () => encoded }) + assert.equal(result.ok, true) + assert.equal(result.manifestWorkspacePath, 'Workflows/room/scene-manifest.json') + } +}) + +test('Load Scene refuses unsafe paths before reading', async () => { + let reads = 0 + for (const scenePath of ['../outside', '/etc/passwd', 'C:/outside', 'Workflows/%2e%2e', 'Workflows/room/']) { + const result = await resolveSceneSourceManifest({ scenePath, workspaceDir: '/workspace', readFileBase64: async () => { reads++; return encoded } }) + assert.equal(result.ok, false, scenePath) + } + assert.equal(reads, 0) +}) + +test('editing a validated Load Scene path clears every derived scene reference', () => { + const params = invalidateValidatedScenePath({ + path: 'Workflows/old', + manifestPath: 'Workflows/old/scene-manifest.json', + sceneRoot: '.', + sourceKind: 'directory', + }, 'Workflows/new') + + assert.equal(params.path, 'Workflows/new') + assert.equal(params.manifestPath, undefined) + assert.equal(params.sceneRoot, undefined) + assert.equal(params.sourceKind, undefined) +}) + +test('an async validation result cannot overwrite a subsequently edited path', async () => { + const oldResolution = await resolveSceneSourceManifest({ + scenePath: 'Workflows/old', + workspaceDir: '/workspace', + readFileBase64: async () => encoded, + }) + assert.equal(oldResolution.ok, true) + + const currentParams = invalidateValidatedScenePath({ + path: 'Workflows/old', + manifestPath: 'Workflows/old/scene-manifest.json', + sceneRoot: '.', + sourceKind: 'directory', + }, 'Workflows/new') + + assert.equal( + applySceneValidationResult(currentParams, 'Workflows/old', oldResolution), + undefined, + ) + assert.equal(currentParams.path, 'Workflows/new') + assert.equal(currentParams.manifestPath, undefined) +}) diff --git a/src/areas/workflows/workflowSceneSource.ts b/src/areas/workflows/workflowSceneSource.ts new file mode 100644 index 00000000..efe2b3ed --- /dev/null +++ b/src/areas/workflows/workflowSceneSource.ts @@ -0,0 +1,246 @@ +import type { SceneArtifactManifestInitialView, SceneArtifactManifestPreview, SceneArtifactManifestV1 } from '../../shared/types/artifacts' + +export const SCENE_MANIFEST_FILE_NAME = 'scene-manifest.json' + +export type SceneSourceKind = 'manifest' | 'directory' + +export type ResolveSceneSourceSuccess = { + ok: true + sourceKind: SceneSourceKind + inputWorkspacePath: string + manifestWorkspacePath: string + manifestAbsolutePath: string + sceneRoot: string + manifest: SceneArtifactManifestV1 +} + +export type ResolveSceneSourceFailure = { + ok: false + error: string +} + +export type ResolveSceneSourceResult = ResolveSceneSourceSuccess | ResolveSceneSourceFailure + +export function applySceneValidationResult( + currentParams: Record, + validatedPath: string, + resolution: ResolveSceneSourceResult, +): Record | undefined { + if (currentParams.path !== validatedPath) return undefined + + if (!resolution.ok) { + return { + ...currentParams, + manifestPath: undefined, + sceneRoot: undefined, + sourceKind: undefined, + error: resolution.error, + } + } + + return { + ...currentParams, + path: resolution.inputWorkspacePath, + manifestPath: resolution.manifestWorkspacePath, + sceneRoot: resolution.sceneRoot, + sourceKind: resolution.sourceKind, + error: undefined, + } +} + +export function invalidateValidatedScenePath( + params: Record, + nextPath: string, +): Record { + return { + ...params, + path: nextPath, + manifestPath: undefined, + sceneRoot: undefined, + sourceKind: undefined, + error: undefined, + } +} + +type ResolveSceneSourceArgs = { + scenePath: string + workspaceDir: string + readFileBase64: (filePath: string) => Promise +} + +function isAbsolutePath(value: string): boolean { + return value.startsWith('/') || /^[A-Za-z]:\//.test(value) +} + +function trimTrailingSlashes(value: string): string { + return value.replace(/\/+$/, '') +} + +function isSafeRelativePath(value: unknown, allowDot = false): value is string { + if (typeof value !== 'string' || !value || value !== value.trim() || value.includes('\u0000')) return false + const normalized = value.replace(/\\/g, '/') + if (allowDot && normalized === '.') return true + if (isAbsolutePath(normalized) || /^[A-Za-z][A-Za-z0-9+.-]*:/.test(normalized) + || /%(?:25|2e|2f|5c|00)/i.test(normalized) || /%(?![0-9a-f]{2})/i.test(normalized)) return false + return normalized.split('/').every((segment) => segment.length > 0 && segment !== '.' && segment !== '..') +} + +function normalizeWorkspaceRelativePath(value: string | undefined, workspaceDir: string): string | undefined { + const normalizedValue = value?.replace(/\\/g, '/') + if (!normalizedValue) return undefined + + const normalizedWorkspace = trimTrailingSlashes(workspaceDir.replace(/\\/g, '/')) + let relativePath: string | undefined + + if (normalizedValue.startsWith('/workspace/')) { + relativePath = normalizedValue.slice('/workspace/'.length) + } else if (normalizedValue === normalizedWorkspace) { + return undefined + } else if (normalizedValue.startsWith(`${normalizedWorkspace}/`)) { + relativePath = normalizedValue.slice(normalizedWorkspace.length + 1) + } else if (!isAbsolutePath(normalizedValue)) { + relativePath = normalizedValue + } + + if (!relativePath) return undefined + return isSafeRelativePath(relativePath) ? relativePath : undefined +} + +function resolveSceneSourceKind(inputWorkspacePath: string): SceneSourceKind | undefined { + if (inputWorkspacePath === SCENE_MANIFEST_FILE_NAME || inputWorkspacePath.endsWith(`/${SCENE_MANIFEST_FILE_NAME}`)) { + return 'manifest' + } + if (inputWorkspacePath.toLowerCase().endsWith('.json')) return undefined + return 'directory' +} + +function decodeBase64Utf8(base64: string): string { + const bytes = Uint8Array.from(atob(base64), (char) => char.charCodeAt(0)) + return new TextDecoder().decode(bytes) +} + +function isPlainObject(value: unknown): value is Record { + return typeof value === 'object' && value !== null && !Array.isArray(value) +} + +function isSafePreviewReference(value: unknown): value is string { + return isSafeRelativePath(value) +} + +function isSceneManifestPreview(value: unknown): value is SceneArtifactManifestPreview { + return isPlainObject(value) + && (!('image' in value) || isSafePreviewReference(value.image)) + && (!('video' in value) || isSafePreviewReference(value.video)) +} + +function isFiniteTriple(value: unknown): value is [number, number, number] { + return Array.isArray(value) && value.length === 3 + && value.every((component) => typeof component === 'number' && Number.isFinite(component)) +} + +function isSceneManifestInitialView(value: unknown): value is SceneArtifactManifestInitialView { + if (!isPlainObject(value)) return false + const { position, target, up } = value + if (!isFiniteTriple(position) || !isFiniteTriple(target)) return false + if (position.every((component, index) => component === target[index])) return false + return up === undefined || (isFiniteTriple(up) && up.some((component) => component !== 0)) +} + +function normalizeSceneRoot(sceneRoot: unknown): string | undefined { + return isSafeRelativePath(sceneRoot, true) ? sceneRoot.replace(/\\/g, '/') : undefined +} + +function validateSceneManifest(manifest: unknown): { ok: true; manifest: SceneArtifactManifestV1; sceneRoot: string } | { ok: false; error: string } { + if (!isPlainObject(manifest)) { + return { ok: false, error: 'Scene manifest must be a JSON object.' } + } + if (manifest.schema !== 'modly.scene-manifest.v1') { + return { ok: false, error: 'Scene manifest schema must be modly.scene-manifest.v1.' } + } + + const rawSceneRoot = manifest.sceneRoot + const sceneRoot = normalizeSceneRoot(rawSceneRoot) + if (typeof rawSceneRoot !== 'string' || !sceneRoot) { + return { ok: false, error: 'Scene manifest sceneRoot must be a safe relative path.' } + } + if (!Array.isArray(manifest.assets)) { + return { ok: false, error: 'Scene manifest assets must be an array.' } + } + if (manifest.assets.some((asset) => isPlainObject(asset) + && (('workspacePath' in asset && !isSafeRelativePath(asset.workspacePath)) + || ('path' in asset && !isSafeRelativePath(asset.path))))) { + return { ok: false, error: 'Scene manifest asset paths must be safe relative file references.' } + } + + const { preview, initialView, ...metadata } = manifest + if (preview !== undefined && !isPlainObject(preview)) { + return { ok: false, error: 'Scene manifest preview must be a JSON object.' } + } + if (preview !== undefined && !isSceneManifestPreview(preview)) { + return { ok: false, error: 'Scene manifest preview image/video must be safe relative file references.' } + } + if (initialView !== undefined && !isPlainObject(initialView)) { + return { ok: false, error: 'Scene manifest initialView must be a JSON object.' } + } + if (initialView !== undefined && !isSceneManifestInitialView(initialView)) { + return { ok: false, error: 'Scene manifest initialView requires distinct finite numeric position/target triples and optional non-zero finite up.' } + } + + return { + ok: true, + sceneRoot, + manifest: { + ...metadata, + schema: 'modly.scene-manifest.v1', + sceneRoot: rawSceneRoot, + assets: manifest.assets, + ...(preview !== undefined ? { preview } : {}), + ...(initialView !== undefined ? { initialView } : {}), + }, + } +} + +export async function resolveSceneSourceManifest(args: ResolveSceneSourceArgs): Promise { + const inputWorkspacePath = normalizeWorkspaceRelativePath(args.scenePath, args.workspaceDir) + if (!inputWorkspacePath) { + return { ok: false, error: 'Load Scene requires a safe workspace-relative scene path.' } + } + + const sourceKind = resolveSceneSourceKind(inputWorkspacePath) + if (!sourceKind) { + return { ok: false, error: `Load Scene accepts ${SCENE_MANIFEST_FILE_NAME} or a scene directory.` } + } + + const manifestWorkspacePath = sourceKind === 'manifest' + ? inputWorkspacePath + : `${inputWorkspacePath}/${SCENE_MANIFEST_FILE_NAME}` + const normalizedWorkspace = trimTrailingSlashes(args.workspaceDir.replace(/\\/g, '/')) + const manifestAbsolutePath = `${normalizedWorkspace}/${manifestWorkspacePath}` + + let manifestRaw: string + try { + manifestRaw = decodeBase64Utf8(await args.readFileBase64(manifestAbsolutePath)) + } catch (error) { + return { ok: false, error: `Unable to read scene manifest: ${String(error)}` } + } + + let parsed: unknown + try { + parsed = JSON.parse(manifestRaw) + } catch (error) { + return { ok: false, error: `Scene manifest is not valid JSON: ${String(error)}` } + } + + const validation = validateSceneManifest(parsed) + if (!validation.ok) return validation + + return { + ok: true, + sourceKind, + inputWorkspacePath, + manifestWorkspacePath, + manifestAbsolutePath, + sceneRoot: validation.sceneRoot, + manifest: validation.manifest, + } +} diff --git a/src/shared/components/ui/ModelLibraryModal.tsx b/src/shared/components/ui/ModelLibraryModal.tsx new file mode 100644 index 00000000..28d68999 --- /dev/null +++ b/src/shared/components/ui/ModelLibraryModal.tsx @@ -0,0 +1,429 @@ +import { useCallback, useEffect, useRef, useState } from 'react' +import { createPortal } from 'react-dom' +import { useAppStore } from '@shared/stores/appStore' +import { useAgentStore } from '@shared/stores/agentStore' +import { useLlmModels, type LlmModel } from '@shared/stores/llmModelsStore' +import { useLlmDownloadsStore } from '@shared/services/llmDownloads' +import { formatBytes as fmtBytes } from '@shared/utils/format' +import { SegmentedControl } from '@shared/ui' +import { SseProgressBar } from './SseProgressBar' +import { vramFit } from './vramFit' +import { agentGrade } from './agentGrade' + +// ─── Types ──────────────────────────────────────────────────────────────────── + +// Single source of truth lives in the shared catalog store; re-exported here so +// existing importers (AgentSection) keep resolving `LlmModel` from this module. +export type { LlmModel } + +interface LlmStatus { + vram_gb: number | null +} + +type Filter = 'all' | 'fits' | 'vision' | 'cad' + +// ─── Helpers ────────────────────────────────────────────────────────────────── + +export function formatBytes(n?: number): string { + return n ? fmtBytes(n) : '—' +} + +type Tone = 'accent' | 'warn' | 'muted' | 'outline' + +const TONES: Record = { + accent: 'bg-accent/15 text-accent-light border-accent/30', + warn: 'bg-amber-500/10 text-amber-400 border-amber-500/30', + muted: 'bg-zinc-800 text-zinc-400 border-zinc-700', + outline: 'bg-transparent text-accent-light border-accent/40', +} + +function Tag({ tone, colors, icon, title, children }: { + tone?: Tone + /** Explicit color classes, for verdicts that bring their own (vramFit). */ + colors?: string + icon?: React.ReactNode + title?: string + children: React.ReactNode +}): JSX.Element { + return ( + + {icon} + {children} + + ) +} + +const ICONS = { + check: , + warn: , + star: , + eye: , + cube: , + file: , +} + +const outlineBtnCls = 'inline-flex items-center gap-1.5 px-3.5 py-1.5 rounded-lg border border-accent/40 text-accent-light text-xs font-medium hover:bg-accent/10 hover:border-accent/60 transition-colors disabled:opacity-50' +const ghostBtnCls = 'px-3.5 py-1.5 rounded-lg border border-zinc-700 text-zinc-300 text-xs font-medium hover:text-white hover:border-zinc-500 transition-colors' + +// ─── Component ──────────────────────────────────────────────────────────────── + +export function ModelLibraryModal({ onClose }: { onClose: () => void }): JSX.Element { + const apiUrl = useAppStore((s) => s.apiUrl) + const localModel = useAgentStore((s) => s.localModel) + const setLocalModel = useAgentStore((s) => s.setLocalModel) + + const [status, setStatus] = useState(null) + const [adding, setAdding] = useState(false) + const [error, setError] = useState(null) + const [query, setQuery] = useState('') + const [filter, setFilter] = useState('all') + + // The model list comes from the shared catalog store, so a download or delete + // here immediately updates every other picker (chat, extension params, + // chat) and vice-versa — no independent per-surface fetch. + const { models, refresh: refreshModels } = useLlmModels() + + // Downloads live in a module-level store (src/shared/services/llmDownloads.ts) + // so they keep running — and stay visible on reopen — after this modal closes. + const downloads = useLlmDownloadsStore((s) => s.downloads) + const downloadError = useLlmDownloadsStore((s) => s.error) + const startDownload = useLlmDownloadsStore((s) => s.start) + const pauseDownload = useLlmDownloadsStore((s) => s.pause) + const cancelDownload = useLlmDownloadsStore((s) => s.cancel) + const dismissDownloadError = useLlmDownloadsStore((s) => s.dismissError) + + // Aborts every in-flight fetch (status/model list + any SSE stream) when the + // modal unmounts, so closing it mid-download doesn't leak a fetch or call + // setState on a component that's gone. + const aliveRef = useRef(true) + const abortControllersRef = useRef(new Set()) + const dialogRef = useRef(null) + + function withAbort(run: (signal: AbortSignal) => Promise): Promise { + const controller = new AbortController() + abortControllersRef.current.add(controller) + return run(controller.signal).finally(() => { abortControllersRef.current.delete(controller) }) + } + + useEffect(() => { + aliveRef.current = true + const controllers = abortControllersRef.current + return () => { + aliveRef.current = false + for (const controller of controllers) controller.abort() + controllers.clear() + } + }, []) + + // Escape closes the modal; Tab is trapped inside it while it's open. + useEffect(() => { + function onKeyDown(e: KeyboardEvent) { + if (e.key === 'Escape') { onClose(); return } + if (e.key !== 'Tab' || !dialogRef.current) return + const focusable = dialogRef.current.querySelectorAll( + 'button, [href], input, select, textarea, [tabindex]:not([tabindex="-1"])', + ) + if (focusable.length === 0) return + const first = focusable[0] + const last = focusable[focusable.length - 1] + if (e.shiftKey && document.activeElement === first) { e.preventDefault(); last.focus() } + else if (!e.shiftKey && document.activeElement === last) { e.preventDefault(); first.focus() } + } + document.addEventListener('keydown', onKeyDown) + dialogRef.current?.focus() + return () => document.removeEventListener('keydown', onKeyDown) + }, [onClose]) + + const refresh = useCallback(async () => { + void refreshModels() // shared model catalog (propagates to every picker) + try { + const s = await withAbort((signal) => + fetch(`${apiUrl}/llm/status`, { signal }).then((r) => r.json()), + ) + if (!aliveRef.current) return + setStatus(s) + } catch { + if (!aliveRef.current) return + setStatus(null) + } + }, [apiUrl, refreshModels]) + + // Once, on open (and if the API URL changes) — `refresh` is stable, see + // useLlmModels. It used to be rebuilt on every render, so this effect re-ran + // on every render and each pass forced another /llm/models + /llm/status. + useEffect(() => { void refresh() }, [refresh]) + + async function handleAdd() { + setAdding(true) + setError(null) + try { + const res = await window.electron.agent.addModel() + if (!aliveRef.current) return + if (res.error) setError(res.error) + if (res.success) void refresh() + } finally { + if (aliveRef.current) setAdding(false) + } + } + + async function handleDelete(id: string) { + await fetch(`${apiUrl}/llm/models/${encodeURIComponent(id)}`, { method: 'DELETE' }).catch(() => {}) + void refresh() + } + + // Downloads run in the shared store, possibly finishing while this modal is + // closed — refresh the model list whenever one drops out (done/error/cancelled) + // while we're mounted, so "downloaded" flips without waiting for a remount. + const prevDownloadIdsRef = useRef>(new Set()) + useEffect(() => { + const prev = prevDownloadIdsRef.current + const current = new Set(Object.keys(downloads).filter((id) => downloads[id] !== undefined)) + let finished = false + for (const id of prev) if (!current.has(id)) finished = true + prevDownloadIdsRef.current = current + if (finished) void refresh() + }, [downloads, refresh]) + + // VRAM is only known on NVIDIA (nvidia-smi); without it there is no fit verdict + // to show or filter on. + const vramGb = status?.vram_gb ?? null + + const filterOptions: { value: Filter; label: string }[] = [ + { value: 'all', label: 'All' }, + ...(vramGb ? [{ value: 'fits' as const, label: 'Fits my GPU' }] : []), + { value: 'vision', label: 'Vision' }, + { value: 'cad', label: 'CAD' }, + ] + + const q = query.trim().toLowerCase() + const matches = (m: LlmModel): boolean => { + const tags = m.tags ?? [] + if (filter === 'vision' && !tags.includes('vision')) return false + if (filter === 'cad' && !tags.includes('cad')) return false + if (filter === 'fits') { + const fit = vramFit(m.vram_estimate_mb, vramGb) + if (!fit || fit.label === "Won't fit") return false + } + return !q || m.name.toLowerCase().includes(q) || (m.description ?? '').toLowerCase().includes(q) + } + + const installed = models.filter((m) => m.downloaded && matches(m)) + const suggested = models.filter((m) => !m.downloaded && matches(m)) + const filtering = filter !== 'all' || q !== '' + + function renderCard(m: LlmModel): JSX.Element { + const dl = downloads[m.id] + const tags = m.tags ?? [] + const inUse = m.downloaded && localModel === m.id + // Code/CAD models are tools for workflow nodes, not chat models: the chat's + // own picker leaves them out, so they cannot become the agent's model here. + const nodeOnly = tags.some((t) => t === 'code' || t === 'cad') + const fit = vramFit(m.vram_estimate_mb, vramGb) + // Size and VRAM say nothing about how well a model drives the agent — a 4B + // outscores a 20B here. The tooltip keeps a measured rate and an estimate + // visibly apart. + const grade = agentGrade(m) + + return ( +
+
+

{m.name}

+ {inUse && ( + }>In use + )} +
+ +
+ {tags.includes('default') && Recommended} + {fit && {fit.label}} + {tags.includes('vision') && Vision} + {tags.includes('cad') && CAD} + {m.source === 'custom' && Custom} + {grade && {grade.label}} +
+ + {m.description && ( +

{m.description}

+ )} + +
+ {[m.size_bytes ? formatBytes(m.size_bytes) : null, m.quant].filter(Boolean).join(' · ')} + {m.vram_estimate_mb ? ( + + ~{(m.vram_estimate_mb / 1024).toFixed(1)}{vramGb ? ` / ${vramGb}` : ''} GB VRAM + + ) : null} +
+ + {dl && } + +
+ {m.downloaded ? ( + <> + {inUse ? ( + Used by the chat agent + ) : nodeOnly ? ( + For workflow nodes, not the chat + ) : ( + + )} + + + ) : dl ? ( + <> + {dl.paused ? ( + + ) : ( + + )} + + + ) : ( + + )} +
+
+ ) + } + + return createPortal( +
{ if (e.target === e.currentTarget) onClose() }} + > +
+ +
+ + {/* Header */} +
+
+
+

Models

+

+ Local models shared by the whole app — chat agent and extensions. +

+
+
+ {vramGb != null && ( + // self-stretch: as tall as the Add button beside it, whatever its padding. + + + + + {vramGb} GB VRAM + + )} + + +
+
+ +
+
+ + + + setQuery(e.target.value)} + placeholder="Search models" + className="w-full bg-zinc-800 border border-zinc-700 text-zinc-200 text-xs rounded-lg pl-8 pr-3 py-2 placeholder:text-zinc-500 focus:outline-none focus:border-accent/50" + /> +
+ +
+ + {status === null && ( +

+ Cannot reach the Modly API. + {/* The library does not poll, so a backend that was still starting + * up needs a way back in short of reopening the modal. */} + +

+ )} + {(error || downloadError) && ( +

+ {error || downloadError} + +

+ )} +
+ + {/* Installed first, then the catalog's suggestions not yet downloaded */} +
+ {([ + { + title: 'Installed', + items: installed, + empty: filtering ? 'No installed model matches.' : 'No model installed yet — add a .gguf file or download a suggestion below.', + }, + { + title: 'Suggested', + items: suggested, + empty: filtering ? 'No suggested model matches.' : 'Every suggested model is installed.', + }, + ]).map((section) => ( +
+

+ {section.title} + {section.items.length} +

+ {section.items.length > 0 ? ( +
+ {section.items.map(renderCard)} +
+ ) : ( +

{section.empty}

+ )} +
+ ))} +
+
+
, + document.body, + ) +} diff --git a/src/shared/components/ui/NumberInput.tsx b/src/shared/components/ui/NumberInput.tsx new file mode 100644 index 00000000..3630c639 --- /dev/null +++ b/src/shared/components/ui/NumberInput.tsx @@ -0,0 +1,99 @@ +import { useRef, useState } from 'react' + +/** + * Text-based numeric inputs for extension params. They keep the raw text the + * user is typing (so "-", "0." or "1," don't get swallowed) and only emit once + * it parses. External value changes (e.g. reset) re-sync the text. + */ + +export function IntInput({ value, onChange, className }: { value: number; onChange: (v: number) => void; className: string }) { + const [text, setText] = useState(String(value)) + const prevValue = useRef(value) + if (prevValue.current !== value && parseInt(text, 10) !== value) { + prevValue.current = value + setText(String(value)) + } + return ( + { + const raw = e.target.value + if (raw !== '' && raw !== '-' && !/^-?\d+$/.test(raw)) return + setText(raw) + const n = parseInt(raw, 10) + if (!isNaN(n)) { prevValue.current = n; onChange(n) } + }} + className={className} + /> + ) +} + +/** + * Float input. When the param declares both `min` and `max`, a range slider is + * shown next to the text field; typed values are clamped to those bounds. + */ +export function FloatInput({ value, onChange, className, min, max, step, label }: { + value: number + onChange: (v: number) => void + className: string + min?: number + max?: number + step?: number + label: string +}) { + const [text, setText] = useState(String(value)) + const prevValue = useRef(value) + if (prevValue.current !== value && parseFloat(text.replace(',', '.')) !== value) { + prevValue.current = value + setText(String(value)) + } + + const clamp = (n: number) => Math.min(max ?? Infinity, Math.max(min ?? -Infinity, n)) + const emit = (n: number) => { + const clamped = clamp(n) + prevValue.current = clamped + onChange(clamped) + } + + const textInput = ( + { + const raw = e.target.value.replace(',', '.') + if (raw !== '' && raw !== '-' && raw !== '.' && !/^-?\d*\.?\d*$/.test(raw)) return + setText(e.target.value) + const num = parseFloat(raw) + if (!isNaN(num)) emit(num) + }} + // Show the clamped value once the user is done typing an out-of-range one. + onBlur={() => { if (parseFloat(text.replace(',', '.')) !== value) setText(String(value)) }} + className={className} + /> + ) + + if (min === undefined || max === undefined || max <= min) return textInput + + return ( +
+ 0 ? step : (max - min) / 100} + value={Number.isFinite(value) ? clamp(value) : min} + onChange={(e) => { + const num = e.currentTarget.valueAsNumber + if (Number.isFinite(num)) { setText(String(num)); emit(num) } + }} + aria-label={label} + // nodrag: inside a React Flow node, dragging the thumb would move the node. + className="nodrag min-w-0 flex-1 accent-accent cursor-pointer" + /> +
{textInput}
+
+ ) +} diff --git a/src/shared/components/ui/SseProgressBar.tsx b/src/shared/components/ui/SseProgressBar.tsx new file mode 100644 index 00000000..60f8596b --- /dev/null +++ b/src/shared/components/ui/SseProgressBar.tsx @@ -0,0 +1,21 @@ +import type { SseEvent } from '@shared/services/llmDownloads' +import { formatBytes } from '@shared/utils/format' + +/** Progress of an SSE-driven download or install (engine, GGUF models). */ +export function SseProgressBar({ event }: { event: SseEvent }): JSX.Element { + return ( +
+
+ {event.status} + + {event.totalBytes + ? `${formatBytes(event.bytesDownloaded ?? 0)} / ${formatBytes(event.totalBytes)}` + : `${event.percent ?? 0}%`} + +
+
+
+
+
+ ) +} diff --git a/src/shared/components/ui/agentGrade.test.mjs b/src/shared/components/ui/agentGrade.test.mjs new file mode 100644 index 00000000..3ee1aee4 --- /dev/null +++ b/src/shared/components/ui/agentGrade.test.mjs @@ -0,0 +1,53 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +// Same bundling trick as vramFit.test.mjs — a pure helper, no React deps. +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-agentgrade-test-')), 'agentGrade.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('src/shared/components/ui/agentGrade.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { agentGrade } = loadModule() + +test('a model without a tier gets no badge at all', () => { + assert.equal(agentGrade(undefined), null) + assert.equal(agentGrade({}), null) + assert.equal(agentGrade({ agent_tier: 'brilliant' }), null) // unknown tier, not a badge +}) + +test('a measured model shows its own score', () => { + const g = agentGrade({ agent_tier: 'excellent', agent_score: 0.98, agent_source: 'measured' }) + assert.equal(g.label, 'Agent: excellent (98%)') + assert.match(g.title, /Modly's tool-calling suite/) +}) + +test('an unmeasured model never borrows the credibility of a measurement', () => { + const g = agentGrade({ agent_tier: 'solid', agent_score: 0.9, agent_source: 'estimate' }) + assert.equal(g.label, 'Agent: solid') // no percentage + assert.match(g.title, /not measured in Modly/) +}) + +test('the note is carried into the tooltip', () => { + const g = agentGrade({ agent_tier: 'limited', agent_source: 'estimate', agent_note: 'older generation' }) + assert.match(g.title, /older generation/) + assert.equal(g.label, 'Agent: limited') +}) + +test('each tier gets its own colour', () => { + const classes = ['excellent', 'solid', 'limited'].map((t) => agentGrade({ agent_tier: t }).className) + assert.equal(new Set(classes).size, 3) +}) diff --git a/src/shared/components/ui/agentGrade.ts b/src/shared/components/ui/agentGrade.ts new file mode 100644 index 00000000..874b103d --- /dev/null +++ b/src/shared/components/ui/agentGrade.ts @@ -0,0 +1,44 @@ +/** + * How well a model drives the agent — the one thing the model list never said. + * + * Size and VRAM are already shown, and both are poor proxies: a 4B tops the + * tool-calling tests while a 20B sits below it. `agent_tier` comes from the + * catalog; when Modly's own eval suite has been run against the model, the + * measured pass rate is shown with it, and everything else is flagged as an + * estimate so the two are never confused. + */ + +export type AgentTier = 'excellent' | 'solid' | 'limited' + +export interface AgentGradeInput { + agent_tier?: string | null + agent_score?: number | null // 0..1, Modly's eval suite + agent_note?: string | null + agent_source?: string | null // 'measured' | 'estimate' +} + +export interface AgentGrade { + label: string + className: string + title: string +} + +const TIERS: Record = { + excellent: { label: 'Agent: excellent', className: 'border-emerald-500/30 bg-emerald-500/10 text-emerald-400' }, + solid: { label: 'Agent: solid', className: 'border-amber-500/30 bg-amber-500/10 text-amber-400' }, + limited: { label: 'Agent: limited', className: 'border-zinc-600/40 bg-zinc-600/10 text-zinc-400' }, +} + +export function agentGrade(m: AgentGradeInput | null | undefined): AgentGrade | null { + const tier = m?.agent_tier as AgentTier | undefined + if (!tier || !(tier in TIERS)) return null + const { label, className } = TIERS[tier] + + const measured = m?.agent_source === 'measured' && typeof m?.agent_score === 'number' + const score = measured ? `${Math.round((m!.agent_score as number) * 100)}% on Modly's tool-calling suite` : null + const title = [score ?? 'Estimated from published benchmarks — not measured in Modly', m?.agent_note] + .filter(Boolean) + .join(' · ') + + return { label: measured ? `${label} (${Math.round((m!.agent_score as number) * 100)}%)` : label, className, title } +} diff --git a/src/shared/components/ui/index.ts b/src/shared/components/ui/index.ts index ef9d77f7..7c3a1147 100644 --- a/src/shared/components/ui/index.ts +++ b/src/shared/components/ui/index.ts @@ -4,3 +4,4 @@ export { ConfirmModal } from './ConfirmModal' export { ColorPicker } from './ColorPicker' export { PickerIcon } from './PickerIcon' export { Toast } from './Toast' +export { IntInput, FloatInput } from './NumberInput' diff --git a/src/shared/components/ui/vramFit.test.mjs b/src/shared/components/ui/vramFit.test.mjs new file mode 100644 index 00000000..551fd69d --- /dev/null +++ b/src/shared/components/ui/vramFit.test.mjs @@ -0,0 +1,48 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +// Bundle the pure vramFit helper (no React deps) into CJS — same approach as +// autoWire.test.mjs / preflight.test.mjs. +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-vramfit-test-')), 'vramFit.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('src/shared/components/ui/vramFit.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { vramFit } = loadModule() + +test('returns null when the estimate or VRAM is unknown', () => { + assert.equal(vramFit(undefined, 12), null) + assert.equal(vramFit(4200, undefined), null) + assert.equal(vramFit(4200, null), null) + assert.equal(vramFit(0, 12), null) +}) + +test('"Fits" when the model sits at or under 90% of VRAM', () => { + // 12 GB → 12288 MB; 90% = 11059 MB + assert.equal(vramFit(4200, 12).label, 'Fits') + assert.equal(vramFit(11059, 12).label, 'Fits') +}) + +test('"Tight" between 90% and 100% of VRAM', () => { + assert.equal(vramFit(11060, 12).label, 'Tight') + assert.equal(vramFit(12288, 12).label, 'Tight') +}) + +test('"Won\'t fit" above 100% of VRAM', () => { + assert.equal(vramFit(13500, 12).label, "Won't fit") + assert.equal(vramFit(12289, 12).label, "Won't fit") +}) diff --git a/src/shared/components/ui/vramFit.ts b/src/shared/components/ui/vramFit.ts new file mode 100644 index 00000000..1f5c9bc0 --- /dev/null +++ b/src/shared/components/ui/vramFit.ts @@ -0,0 +1,20 @@ +/** + * "Does this model fit in VRAM?" verdict, computed from the model's estimated + * footprint against the detected VRAM. Only meaningful on GPUs we can measure + * (NVIDIA via nvidia-smi); returns null otherwise so no misleading badge shows + * on Vulkan/CPU where `vramGb` is unknown. + */ +export interface VramFit { + label: string + className: string +} + +export function vramFit(estimateMb?: number, vramGb?: number | null): VramFit | null { + if (!estimateMb || !vramGb) return null + const vramMb = vramGb * 1024 + if (estimateMb <= vramMb * 0.9) + return { label: 'Fits', className: 'border-emerald-500/30 bg-emerald-500/10 text-emerald-400' } + if (estimateMb <= vramMb) + return { label: 'Tight', className: 'border-amber-500/30 bg-amber-500/10 text-amber-400' } + return { label: "Won't fit", className: 'border-red-500/30 bg-red-500/10 text-red-400' } +} diff --git a/src/shared/hooks/useGeneration.ts b/src/shared/hooks/useGeneration.ts index 36bdbcb9..6f3d3a37 100644 --- a/src/shared/hooks/useGeneration.ts +++ b/src/shared/hooks/useGeneration.ts @@ -1,7 +1,7 @@ import { useCallback, useRef } from 'react' import { useAppStore } from '@shared/stores/appStore' import { useApi } from './useApi' -import { showCompletionNotification } from '@shared/utils/notification' +import { showCompletionNotification, showErrorNotification } from '@shared/utils/notification' export function useGeneration() { const { currentJob, setCurrentJob, updateCurrentJob, generationOptions, selectedImageData, pushMeshUrl, clearMeshHistory } = useAppStore() @@ -53,6 +53,7 @@ export function useGeneration() { status: 'error', error: errorMessage }) + void showErrorNotification(errorMessage, 'Generation failed') } }, // eslint-disable-next-line react-hooks/exhaustive-deps -- useApi re-creates its fns each render, so this re-memoizes anyway (values stay fresh) @@ -85,6 +86,7 @@ export function useGeneration() { if (result.status === 'error') { updateCurrentJob({ status: 'error', error: result.error }) + void showErrorNotification(result.error ?? 'Unknown error', 'Generation failed') break } diff --git a/src/shared/services/llmDownloads.ts b/src/shared/services/llmDownloads.ts new file mode 100644 index 00000000..18de4412 --- /dev/null +++ b/src/shared/services/llmDownloads.ts @@ -0,0 +1,104 @@ +import { create } from 'zustand' +import { useAppStore } from '@shared/stores/appStore' + +// GGUF model downloads must survive the Model Library modal being closed and +// reopened — the backend already keeps downloading in the background once +// started (see api/routers/llm.py), it just needs a watcher that isn't tied to +// a component's lifecycle. This module-level store owns that watch, the same +// way agentChat.ts owns the chat's SSE stream outside React. + +export interface SseEvent { + percent?: number + status?: string + bytesDownloaded?: number + totalBytes?: number + error?: string + cancelled?: boolean + paused?: boolean +} + +/** Reads one SSE `data: {...}` frame at a time and calls `onEvent` for each. */ +export async function consumeSse(url: string, onEvent: (data: SseEvent) => void, signal?: AbortSignal): Promise { + const res = await fetch(url, { signal }) + if (!res.ok || !res.body) throw new Error(`HTTP ${res.status}`) + const reader = res.body.getReader() + const decoder = new TextDecoder() + let buf = '' + outer: for (;;) { + const { done, value } = await reader.read() + if (done) break + buf += decoder.decode(value, { stream: true }) + const parts = buf.split('\n\n') + buf = parts.pop() ?? '' + for (const part of parts) { + const line = part.split('\n').find((l) => l.startsWith('data: ')) + if (!line) continue + let data: SseEvent + try { data = JSON.parse(line.slice(6)) } catch { continue /* malformed frame */ } + onEvent(data) + // The stream stays open after an error frame — stop reading instead of + // spinning on a download the server has already given up on. + if (data.error) { void reader.cancel().catch(() => {}); break outer } + } + } +} + +interface LlmDownloadsStore { + downloads: Record + error: string | null + start: (modelId: string) => void + pause: (modelId: string) => Promise + cancel: (modelId: string) => Promise + dismissError: () => void +} + +// Model ids this renderer is currently watching. The backend is reconnect-safe +// (a GET while a download is already in flight attaches instead of restarting +// it), this just avoids opening a second redundant connection from this same +// store if start() is called twice in a row (e.g. StrictMode double-invoke). +const _watching = new Set() + +export const useLlmDownloadsStore = create((set) => ({ + downloads: {}, + error: null, + + start(modelId) { + if (_watching.has(modelId)) return + _watching.add(modelId) + const apiUrl = useAppStore.getState().apiUrl + + // Resuming from a paused entry: reset it to a live "connecting" state so the + // Resume button flips back to Pause immediately, before the first SSE frame. + set((s) => ({ downloads: { ...s.downloads, [modelId]: { percent: s.downloads[modelId]?.percent ?? 0, status: 'Starting…' } } })) + + let last: SseEvent | undefined + void consumeSse(`${apiUrl}/llm/download?model_id=${encodeURIComponent(modelId)}`, (e) => { + if (e.error) { set({ error: e.error }); return } + last = e + set((s) => ({ downloads: { ...s.downloads, [modelId]: e } })) + }) + .catch((e) => set({ error: e instanceof Error ? e.message : String(e) })) + .finally(() => { + _watching.delete(modelId) + // Keep the paused entry so a Resume button stays visible; the .part file + // is preserved server-side and start() resumes it. Any other terminal + // (done/cancelled/error) drops the row. + set((s) => ({ downloads: { ...s.downloads, [modelId]: last?.paused ? last : undefined } })) + }) + }, + + async pause(modelId) { + const apiUrl = useAppStore.getState().apiUrl + await fetch(`${apiUrl}/llm/download/pause?model_id=${encodeURIComponent(modelId)}`, { method: 'POST' }).catch(() => {}) + }, + + async cancel(modelId) { + const apiUrl = useAppStore.getState().apiUrl + await fetch(`${apiUrl}/llm/download/cancel?model_id=${encodeURIComponent(modelId)}`, { method: 'POST' }).catch(() => {}) + // A paused download has no live watcher to hit its .finally cleanup — drop + // the row here so Cancel visibly clears it. + set((s) => ({ downloads: { ...s.downloads, [modelId]: undefined } })) + }, + + dismissError: () => set({ error: null }), +})) diff --git a/src/shared/stores/agentStore.ts b/src/shared/stores/agentStore.ts index f6395c09..23b2cfc0 100644 --- a/src/shared/stores/agentStore.ts +++ b/src/shared/stores/agentStore.ts @@ -1,29 +1,155 @@ import { create } from 'zustand' -import { persist } from 'zustand/middleware' +import { persist, createJSONStorage } from 'zustand/middleware' export type ThinkingMode = 'auto' | 'on' | 'off' +export type ProviderId = 'local' | 'ollama' | 'openai' | 'anthropic' | 'mistral' | 'groq' | 'openrouter' | 'custom' + +export interface ExternalConfig { + apiKey: string + model: string + baseUrl?: string // only used by 'custom' +} + +export const PROVIDERS: Record = { + local: { label: 'Local (llama.cpp)', baseUrl: '' }, + // Ollama serves an OpenAI-compatible API — reuses models already pulled with it. + ollama: { label: 'Ollama', baseUrl: 'http://localhost:11434/v1', noKey: true }, + openai: { label: 'ChatGPT (OpenAI)', baseUrl: 'https://api.openai.com/v1' }, + anthropic: { label: 'Claude (Anthropic)', baseUrl: 'https://api.anthropic.com/v1' }, + mistral: { label: 'Mistral', baseUrl: 'https://api.mistral.ai/v1' }, + groq: { label: 'Groq', baseUrl: 'https://api.groq.com/openai/v1' }, + openrouter: { label: 'OpenRouter', baseUrl: 'https://openrouter.ai/api/v1' }, + custom: { label: 'Custom endpoint', baseUrl: '' }, +} + +export const DEFAULT_LOCAL_MODEL = 'qwen3-4b' + +// Catalog ids removed in the Qwen3/gpt-oss refresh → their closest replacement +const RETIRED_LOCAL_MODELS: Record = { + 'qwen2.5-3b': 'qwen3-4b', + 'qwen2.5-7b': 'qwen3-4b', + 'qwen2.5-14b': 'qwen3-14b', + 'llama-3.1-8b': 'qwen3-4b', + 'deepseek-r1-distill-qwen-7b': 'qwen3-14b', +} + +/** Resolve the base URL for a provider (custom uses its own field). */ +export function providerBaseUrl(provider: ProviderId, external: ExternalConfig | undefined): string { + if (provider === 'custom') return external?.baseUrl ?? '' + return PROVIDERS[provider].baseUrl +} + interface AgentSettings { - ollamaUrl: string - defaultModel: string - defaultThinking: ThinkingMode + provider: ProviderId + localModel: string // catalog id or custom: + external: Partial> + defaultThinking: ThinkingMode + + setProvider: (provider: ProviderId) => void + setLocalModel: (model: string) => void + setExternal: (provider: ProviderId, cfg: ExternalConfig) => void + setDefaultThinking: (mode: ThinkingMode) => void +} + +// ─── Secure persistence ──────────────────────────────────────────────────────── +// External provider API keys are the one sensitive field in this store. They're +// encrypted at rest via Electron's safeStorage (OS keychain/DPAPI/libsecret) — +// everything else (provider, localModel, defaultThinking) stays plain, it isn't +// a secret. The ciphertext itself still lives in localStorage as a hex string; +// only the plaintext key never touches disk unencrypted. + +type PersistedExternal = Partial> + +function hasSecureStore(): boolean { + return typeof window !== 'undefined' && !!window.electron?.secureStore +} + +async function transformApiKeys( + external: PersistedExternal | undefined, + transform: (key: string) => Promise, +): Promise { + if (!external || !hasSecureStore()) return external + const entries = await Promise.all( + Object.entries(external).map(async ([provider, cfg]) => { + if (!cfg?.apiKey) return [provider, cfg] as const + const key = await transform(cfg.apiKey) + // null = stored blob we can't decrypt here (different OS user/machine). + // Drop it: sending a ciphertext as an Authorization header just 401s, and + // keeping it in state would re-encrypt it on the next write. + return [provider, { ...cfg, apiKey: key ?? '' }] as const + }), + ) + return Object.fromEntries(entries) as PersistedExternal +} - setOllamaUrl: (url: string) => void - setDefaultModel: (model: string) => void - setDefaultThinking: (mode: ThinkingMode) => void +const secureAgentStorage = { + getItem: async (name: string): Promise => { + const raw = localStorage.getItem(name) + if (!raw) return raw + try { + const envelope = JSON.parse(raw) + if (envelope?.state?.external) { + envelope.state.external = await transformApiKeys(envelope.state.external, (k) => window.electron.secureStore.decrypt(k)) + } + return JSON.stringify(envelope) + } catch { + return raw + } + }, + setItem: async (name: string, value: string): Promise => { + try { + const envelope = JSON.parse(value) + if (envelope?.state?.external) { + envelope.state.external = await transformApiKeys(envelope.state.external, (k) => window.electron.secureStore.encrypt(k)) + } + localStorage.setItem(name, JSON.stringify(envelope)) + } catch { + localStorage.setItem(name, value) + } + }, + removeItem: async (name: string): Promise => { localStorage.removeItem(name) }, } export const useAgentStore = create()( persist( (set) => ({ - ollamaUrl: 'http://localhost:11434', - defaultModel: 'gemma4:e4b', + provider: 'local', + localModel: DEFAULT_LOCAL_MODEL, + external: {}, defaultThinking: 'auto', - setOllamaUrl: (url) => set({ ollamaUrl: url }), - setDefaultModel: (model) => set({ defaultModel: model }), - setDefaultThinking: (mode) => set({ defaultThinking: mode }), + setProvider: (provider) => set({ provider }), + setLocalModel: (model) => set({ localModel: model }), + setExternal: (provider, cfg) => set((s) => ({ external: { ...s.external, [provider]: cfg } })), + setDefaultThinking: (mode) => set({ defaultThinking: mode }), }), - { name: 'modly-agent-settings' }, + { + name: 'modly-agent-settings', + version: 3, + storage: createJSONStorage(() => secureAgentStorage), + // v0 stored { ollamaUrl, defaultModel } — drop them, keep only thinking. + // v2 retired the Qwen2.5/Llama3.1 catalog ids. + // v3 is a shape-less bump: the version change alone forces one re-persist + // through secureAgentStorage.setItem, which encrypts any legacy + // plaintext API key saved before safeStorage existed. Self-limiting — + // once stored as v3 it never re-runs, so it costs one write, not one per boot. + migrate: (persisted: unknown, version) => { + if (version === 0 && persisted && typeof persisted === 'object') { + const old = persisted as { defaultThinking?: ThinkingMode } + return { + provider: 'local' as ProviderId, + localModel: DEFAULT_LOCAL_MODEL, + external: {}, + defaultThinking: old.defaultThinking ?? 'auto', + } + } + const state = persisted as AgentSettings + if (version < 2 && state?.localModel && RETIRED_LOCAL_MODELS[state.localModel]) { + return { ...state, localModel: RETIRED_LOCAL_MODELS[state.localModel] } + } + return state + }, + }, ), ) diff --git a/src/shared/stores/appStore.ts b/src/shared/stores/appStore.ts index d06a5211..bb5ddd0e 100644 --- a/src/shared/stores/appStore.ts +++ b/src/shared/stores/appStore.ts @@ -46,6 +46,13 @@ export interface LightSettings { envIntensity: number } +export interface PointLight { + id: string + position: [number, number, number] + color: string + intensity: number +} + export interface AppToast { id: number message: string @@ -147,6 +154,8 @@ interface AppState { // 3D viewer lighting lightSettings: LightSettings setLightSettings: (settings: LightSettings) => void + pointLights: PointLight[] + setPointLights: (lights: PointLight[]) => void // Actions initApp: () => Promise @@ -249,6 +258,8 @@ export const useAppStore = create()( lightSettings: DEFAULT_LIGHT_SETTINGS, setLightSettings: (settings) => set({ lightSettings: settings }), + pointLights: [], + setPointLights: (lights) => set({ pointLights: lights }), currentJob: null, selectedImagePath: null, diff --git a/src/shared/stores/extensionsStore.ts b/src/shared/stores/extensionsStore.ts index 91657a0a..ada611f2 100644 --- a/src/shared/stores/extensionsStore.ts +++ b/src/shared/stores/extensionsStore.ts @@ -24,8 +24,11 @@ interface ExtensionsStore { installProgress: InstallProgress | null installError: string | null loadErrors: Record + /** Installed weight variant ids, keyed by "ext_id/node_id", for nodes that declare variants */ + installedWeightVariants: Record loadExtensions: () => Promise + refreshInstalledWeightVariants: () => Promise installFromGitHub: (url: string) => Promise<{ success: boolean; error?: string }> installFromLocal: () => Promise<{ success: boolean; error?: string; cancelled?: boolean; needsRepair?: boolean }> uninstall: (extensionId: string) => Promise<{ success: boolean; error?: string }> @@ -54,6 +57,7 @@ export const useExtensionsStore = create((set, get) => ({ installProgress: null, installError: null, loadErrors: {}, + installedWeightVariants: {}, // ── Load list ────────────────────────────────────────────────────────────── @@ -66,11 +70,30 @@ export const useExtensionsStore = create((set, get) => ({ ...extensions, loading: false, }) + await get().refreshInstalledWeightVariants() } catch { set({ loading: false }) } }, + async refreshInstalledWeightVariants() { + const entries = await Promise.all( + get().modelExtensions.flatMap((ext) => ext.nodes + .filter((node) => node.weightVariants) + .map(async (node) => { + const fullId = `${ext.id}/${node.id}` + return [fullId, await window.electron.model.installedWeightVariants(fullId)] as const + })), + ) + // A node whose state could not be read stays absent: undefined reads as "unknown", + // which the UI keeps neutral, while [] would claim no variant is installed. + set({ + installedWeightVariants: Object.fromEntries( + entries.filter((entry): entry is readonly [string, string[]] => entry[1] !== null), + ), + }) + }, + // ── Install from GitHub ──────────────────────────────────────────────────── async installFromGitHub(url: string) { diff --git a/src/shared/stores/llmModelsStore.react.test.mjs b/src/shared/stores/llmModelsStore.react.test.mjs new file mode 100644 index 00000000..b8c305f9 --- /dev/null +++ b/src/shared/stores/llmModelsStore.react.test.mjs @@ -0,0 +1,82 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { createRequire } from 'node:module' +import { setupDom, loadModule, mount } from '../../../scripts/react-test-env.mjs' + +setupDom() +const React = createRequire(import.meta.url)('react') + +/** A fresh module per test: the catalog store caches across calls by design, so + * a shared instance would make each test depend on the ones before it. */ +const freshHook = () => loadModule('src/shared/stores/llmModelsStore.ts').useLlmModels + +/** Counts every request the hook makes, and answers with a fixed catalog. */ +function countingFetch() { + const calls = [] + globalThis.fetch = (url) => { + calls.push(String(url)) + return Promise.resolve({ json: () => Promise.resolve({ models: [{ id: 'qwen3-4b', downloaded: true }] }) }) + } + return calls +} + +test('the hook fetches the catalog once, however many times it renders', async () => { + const useLlmModels = freshHook() + const calls = countingFetch() + const Probe = () => { useLlmModels(); return null } + + const view = await mount(React.createElement(Probe)) + await view.rerender() + await view.rerender() + await view.flush() + + assert.equal(calls.length, 1) + await view.unmount() +}) + +test('refresh and the model list keep their identity across renders', async () => { + // The regression: `refresh` was rebuilt on every render, so any caller + // holding it in a dependency array re-ran its effect forever. + const useLlmModels = freshHook() + countingFetch() + const seen = [] + const Probe = () => { seen.push(useLlmModels()); return null } + + const view = await mount(React.createElement(Probe)) + await view.flush() + await view.rerender() + + const last = seen[seen.length - 1] + const previous = seen[seen.length - 2] + assert.equal(last.refresh, previous.refresh) + assert.equal(last.models, previous.models) + await view.unmount() +}) + +test('a consumer that refreshes from an effect settles instead of looping', async () => { + // Exactly the shape of ModelLibraryModal: an effect keyed on `refresh` that + // stores something in state. With an unstable `refresh` this rendered — and + // fetched — without end; the modal flickered for as long as it was open. + const useLlmModels = freshHook() + const calls = countingFetch() + let renders = 0 + + const Modal = () => { + const { refresh } = useLlmModels() + const [, setStatus] = React.useState(null) + renders++ + // Fails fast and says why: an unstable `refresh` makes this loop forever, + // and a test that hangs for a minute before dying explains nothing. + if (renders > 50) throw new Error('render loop — the effect keeps re-running') + React.useEffect(() => { void refresh(); setStatus({ checkedAt: renders }) }, [refresh]) + return null + } + + const view = await mount(React.createElement(Modal)) + await view.flush() + await view.flush() + + assert.equal(calls.length, 2) // the hook's own load, plus one forced refresh + assert.ok(renders < 10, `expected a handful of renders, got ${renders}`) + await view.unmount() +}) diff --git a/src/shared/stores/llmModelsStore.test.mjs b/src/shared/stores/llmModelsStore.test.mjs new file mode 100644 index 00000000..372d42ba --- /dev/null +++ b/src/shared/stores/llmModelsStore.test.mjs @@ -0,0 +1,106 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +// Bundle the store with its real zustand dependency — same approach as +// workflowsStore.test.mjs. Only fetchModels() is exercised, so the React hook +// at the bottom of the module is never rendered. +function loadStore() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-llmstore-test-')), 'llmModelsStore.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('src/shared/stores/llmModelsStore.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile).useLlmModelsStore +} + +/** A /llm/models endpoint whose answers are released one at a time. */ +function deferredFetch() { + const calls = [] + globalThis.fetch = () => { + let release + const body = new Promise((r) => { release = r }) + calls.push({ release }) + return Promise.resolve({ json: () => body }) + } + return calls +} + +const model = (id, downloaded) => ({ + id, name: id, hf_filename: `${id}.gguf`, downloaded, source: 'catalog', +}) + +test('concurrent plain fetches share one request', async () => { + const useStore = loadStore() + const calls = deferredFetch() + + const a = useStore.getState().fetchModels('http://api') + const b = useStore.getState().fetchModels('http://api') + assert.equal(calls.length, 1) + + calls[0].release({ models: [model('qwen', false)] }) + await Promise.all([a, b]) + assert.equal(useStore.getState().models[0].downloaded, false) +}) + +test('a forced refresh re-fetches instead of joining the pending request', async () => { + const useStore = loadStore() + const calls = deferredFetch() + + // Mount-time fetch, still in flight when the download finishes. + const initial = useStore.getState().fetchModels('http://api') + const refreshed = useStore.getState().fetchModels('http://api', { force: true }) + + calls[0].release({ models: [model('qwen', false)] }) + await initial + assert.equal(calls.length, 2, 'force must issue its own request') + + calls[1].release({ models: [model('qwen', true)] }) + await refreshed + // Joining the pre-download request left `downloaded: false`, and the preflight + // kept refusing to run the workflow. + assert.equal(useStore.getState().models[0].downloaded, true) +}) + +test('a forced refresh still settles when the pending request fails', async () => { + const useStore = loadStore() + const calls = [] + globalThis.fetch = () => { + let settle + const p = new Promise((resolve, reject) => { settle = { resolve, reject } }) + calls.push(settle) + return p + } + + const initial = useStore.getState().fetchModels('http://api') + const refreshed = useStore.getState().fetchModels('http://api', { force: true }) + + calls[0].reject(new Error('offline')) + await initial + calls[1].resolve({ json: async () => ({ models: [model('qwen', true)] }) }) + await refreshed + + assert.equal(useStore.getState().models[0].downloaded, true) + assert.equal(useStore.getState().error, null) +}) + +test('a cached catalog is not re-fetched without force', async () => { + const useStore = loadStore() + const calls = deferredFetch() + + const first = useStore.getState().fetchModels('http://api') + calls[0].release({ models: [model('qwen', true)] }) + await first + + await useStore.getState().fetchModels('http://api') + assert.equal(calls.length, 1) +}) diff --git a/src/shared/stores/llmModelsStore.ts b/src/shared/stores/llmModelsStore.ts new file mode 100644 index 00000000..89da3ed1 --- /dev/null +++ b/src/shared/stores/llmModelsStore.ts @@ -0,0 +1,112 @@ +import { create } from 'zustand' +import { useCallback, useEffect, useMemo } from 'react' +import { useAppStore } from './appStore' + +export interface LlmModel { + id: string + name: string + description?: string + hf_filename: string + size_bytes?: number + quant?: string + vram_estimate_mb?: number + downloaded: boolean + source: 'catalog' | 'custom' + tags?: string[] + /** How well this model drives the agent — see components/ui/agentGrade.ts. + * `agent_score` is only shown when `agent_source` is 'measured', i.e. the + * model was actually run against Modly's own eval suite. */ + agent_tier?: 'excellent' | 'solid' | 'limited' + agent_score?: number + agent_source?: 'measured' | 'estimate' + agent_note?: string +} + +interface LlmModelsStore { + models: LlmModel[] + loading: boolean + error: string | null + fetchedApiUrl: string | null + fetchModels: (apiUrl: string, opts?: { force?: boolean }) => Promise +} + +// Module-level so concurrent callers (multiple nodes/components mounting at +// once) share one in-flight request instead of firing N identical fetches. +let inFlight: Promise | null = null +// Identifies the request currently owning `inFlight`, so a superseded one does +// not clear a newer forced refresh on its way out. +let inFlightId = 0 + +export const useLlmModelsStore = create((set, get) => ({ + models: [], + loading: false, + error: null, + fetchedApiUrl: null, + + async fetchModels(apiUrl, opts) { + const state = get() + if (!opts?.force && state.fetchedApiUrl === apiUrl && state.models.length > 0) return + // A forced refresh follows a mutation (download finished, model deleted), so + // it must not settle for a request issued BEFORE it: joining the in-flight + // one kept the pre-mutation `downloaded` flags, and the preflight went on + // reporting "…isn't downloaded" for a model that had just landed. + const pending = inFlight + if (pending && !opts?.force) return pending + + set({ loading: true, error: null }) + const id = ++inFlightId + const request = (async () => { + if (pending) { + await pending.catch(() => {}) + set({ loading: true, error: null }) // the request we waited on may have failed + } + try { + const res = await fetch(`${apiUrl}/llm/models`) + const data: { models?: LlmModel[] } = await res.json() + set({ models: data.models ?? [], fetchedApiUrl: apiUrl, loading: false, error: null }) + } catch (e) { + set({ error: e instanceof Error ? e.message : String(e), loading: false }) + } finally { + if (inFlightId === id) inFlight = null + } + })() + inFlight = request + return request + }, +})) + +/** + * Shared local-LLM catalog: fetched once per apiUrl and cached across every + * consumer (LLM node, extension param pickers, chat model picker, Settings…) + * instead of each component firing its own `/llm/models` request. + * + * `tag` mirrors the backend's own filter (`GET /llm/models?tag=`): custom + * GGUFs are always kept since their capabilities aren't known ahead of time. + */ +export function useLlmModels(tag?: string): { + models: LlmModel[] + loading: boolean + error: string | null + refresh: () => Promise +} { + const apiUrl = useAppStore((s) => s.apiUrl) + const models = useLlmModelsStore((s) => s.models) + const loading = useLlmModelsStore((s) => s.loading) + const error = useLlmModelsStore((s) => s.error) + const fetchModels = useLlmModelsStore((s) => s.fetchModels) + + useEffect(() => { void fetchModels(apiUrl) }, [apiUrl, fetchModels]) + + // Both memoised because callers put them in dependency arrays. A `refresh` + // rebuilt on every render made ModelLibraryModal's `useEffect(…, [refresh])` + // re-run on every render: each pass forced a /llm/models + /llm/status fetch, + // whose setState triggered the next one. The modal flickered and hammered the + // API for as long as it stayed open. + const filtered = useMemo( + () => (tag ? models.filter((m) => m.source === 'custom' || (m.tags ?? []).includes(tag)) : models), + [models, tag], + ) + const refresh = useCallback(() => fetchModels(apiUrl, { force: true }), [apiUrl, fetchModels]) + + return { models: filtered, loading, error, refresh } +} diff --git a/src/shared/stores/navStore.ts b/src/shared/stores/navStore.ts index c6854424..62afa2f5 100644 --- a/src/shared/stores/navStore.ts +++ b/src/shared/stores/navStore.ts @@ -4,10 +4,16 @@ export type Page = 'generate' | 'workflows' | 'models' | 'settings' interface NavState { currentPage: Page + extensionToOpen: string | null navigate: (page: Page) => void + openExtension: (extensionId: string) => void + clearExtensionToOpen: () => void } export const useNavStore = create((set) => ({ currentPage: 'generate', - navigate: (page) => set({ currentPage: page }) + extensionToOpen: null, + navigate: (page) => set({ currentPage: page }), + openExtension: (extensionId) => set({ currentPage: 'models', extensionToOpen: extensionId }), + clearExtensionToOpen: () => set({ extensionToOpen: null }), })) diff --git a/src/shared/stores/workflowsStore.ts b/src/shared/stores/workflowsStore.ts index b8753375..f6fb829c 100644 --- a/src/shared/stores/workflowsStore.ts +++ b/src/shared/stores/workflowsStore.ts @@ -98,7 +98,7 @@ interface LegacyWorkflow { // Source-only nodes have no target handle; sink-only nodes have no source handle. // An edge into/out of the wrong side can't resolve a handle and makes React Flow // warn ("Couldn't create edge for target handle id: null") on every render. -export const NODE_TYPES_WITHOUT_TARGET = new Set(['imageNode', 'textNode', 'meshNode', 'inputNode', 'forEachNode']) +export const NODE_TYPES_WITHOUT_TARGET = new Set(['imageNode', 'textNode', 'meshNode', 'sceneNode', 'inputNode', 'forEachNode']) export const NODE_TYPES_WITHOUT_SOURCE = new Set(['outputNode', 'previewNode']) function sanitizeEdges(nodes: WFNode[], edges: WFEdge[]): WFEdge[] { diff --git a/src/shared/types/artifacts.ts b/src/shared/types/artifacts.ts index a8dfc6eb..57c6b8d4 100644 --- a/src/shared/types/artifacts.ts +++ b/src/shared/types/artifacts.ts @@ -4,3 +4,17 @@ export interface ArtifactProvenance { source?: string [key: string]: unknown } + +export interface SceneArtifactManifestPreview { image?: string; video?: string } +export interface SceneArtifactManifestInitialView { + position: [number, number, number] + target: [number, number, number] + up?: [number, number, number] +} +export interface SceneArtifactManifestV1 { + schema: 'modly.scene-manifest.v1' + sceneRoot: string + assets: unknown[] + preview?: SceneArtifactManifestPreview + initialView?: SceneArtifactManifestInitialView +} diff --git a/src/shared/types/electron.d.ts b/src/shared/types/electron.d.ts index 1a4d6fde..c171b664 100644 --- a/src/shared/types/electron.d.ts +++ b/src/shared/types/electron.d.ts @@ -14,16 +14,30 @@ import type { export interface ExtensionNode { id: string name: string - input: 'image' | 'text' | 'mesh' | 'audio' - inputs?: ('image' | 'text' | 'mesh' | 'audio')[] // multi-input nodes; overrides input when set + input: 'image' | 'text' | 'mesh' | 'audio' | 'scene' + inputs?: ('image' | 'text' | 'mesh' | 'audio' | 'scene')[] // multi-input nodes; overrides input when set inputLabels?: string[] // display labels per input slot (e.g. positive/negative) - output: 'image' | 'text' | 'mesh' | 'audio' + output: 'image' | 'text' | 'mesh' | 'audio' | 'scene' paramsSchema: ParamSchema[] paramDefaults?: Record hfRepo?: string downloadCheck?: string hfSkipPrefixes?: string[] hfIncludePrefixes?: string[] + hasModelSources?: boolean + weightGroups?: string[] + weightVariants?: WeightVariantsInfo +} + +export interface WeightVariantsInfo { + param: string // params_schema id whose value selects the variant + default: string + options: { id: string; label: string; sizeGb?: number; vramGb?: number }[] +} + +export interface SharedWeightGroup { + id: string + dependentNodeIds: string[] } export interface ModelExtension { @@ -38,6 +52,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 */ @@ -154,6 +169,9 @@ declare global { shell: { openExternal: (url: string) => Promise } + slicer: { + open: (url: string) => Promise<{ success: boolean; error?: string }> + } system: { memory: () => Promise<{ total: number; used: number; available: number }> } @@ -180,6 +198,7 @@ declare global { offLog: () => void } fs: { + getPathForFile: (file: File) => string selectImage: () => Promise selectMeshFile: () => Promise saveModel: (defaultName: string) => Promise @@ -197,6 +216,15 @@ declare global { get: () => Promise<{ modelsDir: string; workspaceDir: string; workflowsDir: string; extensionsDir: string; hfToken?: string }> set: (patch: { modelsDir?: string; workspaceDir?: string; workflowsDir?: string; extensionsDir?: string; hfToken?: string }) => Promise<{ modelsDir: string; workspaceDir: string; workflowsDir: string; extensionsDir: string; hfToken?: string }> } + /** decrypt returns null when the stored blob can't be decrypted here. */ + secureStore: { + encrypt: (plainText: string) => Promise + decrypt: (stored: string) => Promise + } + agent: { + /** Opens a file picker and copies the chosen .gguf into the agent's models folder. */ + addModel: () => Promise<{ success: boolean; cancelled?: boolean; fileName?: string; error?: string }> + } cache: { clear: () => Promise<{ success: boolean; error?: string }> } @@ -206,16 +234,24 @@ declare global { model: { export: (args: { outputUrl: string; format: string }) => Promise<{ success: boolean; error?: string }> listDownloaded: () => Promise<{ id: string; name: string; size_gb: number }[]> - activeDownloads: () => Promise<{ modelId: string; percent: number; file?: string; fileIndex?: number; totalFiles?: number }[]> - isDownloaded: (modelId: string, downloadCheck?: string) => Promise - download: (repoId: string, modelId: string, skipPrefixes?: string[], includePrefixes?: string[]) => Promise<{ success: boolean; error?: string }> + activeDownloads: () => Promise<{ modelId: string; variantId?: string; percent: number; file?: string; fileIndex?: number; totalFiles?: number }[]> + isDownloaded: (modelId: string) => Promise + hasLocalData: (modelId: string) => Promise + sharedGroups: (extensionId: string) => Promise + download: (modelId: string, variantId?: string) => Promise<{ success: boolean; error?: string; paused?: boolean; cancelled?: boolean }> pauseDownload: (modelId: string) => Promise<{ success: boolean; error?: string }> cancelDownload: (modelId: string) => Promise<{ success: boolean; error?: string }> delete: (modelId: string) => Promise<{ success: boolean; error?: string }> + deleteSharedGroup: (extensionId: string, groupId: string) => Promise<{ success: boolean; error?: string }> + deleteExtensionWeights: (extensionId: string) => Promise<{ success: boolean; error?: string }> + /** Installed variant ids, or null when the node's install state could not be read */ + installedWeightVariants: (modelId: string) => Promise + deleteWeightVariant: (modelId: string, variantId: string) => Promise<{ success: boolean; error?: string }> unloadAll: () => Promise<{ success: boolean; error?: string }> showInFolder: (modelId: string) => Promise onProgress: (cb: (data: { modelId: string + variantId?: string percent: number file?: string fileIndex?: number @@ -228,6 +264,8 @@ declare global { cancelled?: boolean }) => void) => void offProgress: () => void + onWeightsChanged: (cb: () => void) => void + offWeightsChanged: () => void } app: { info: () => Promise<{ @@ -320,3 +358,11 @@ declare global { } } } + +export interface SharedWeightGroupState { + id: string + targetId: string + dependentModelIds: string[] + downloaded: boolean + hasLocalData: boolean +} diff --git a/src/shared/ui/index.tsx b/src/shared/ui/index.tsx index a9609995..a2c6c181 100644 --- a/src/shared/ui/index.tsx +++ b/src/shared/ui/index.tsx @@ -12,13 +12,22 @@ export function Section({ title, subtitle, children }: { title: string; subtitle ) } -export function Card({ title, description, children }: { title?: string; description?: string; children: React.ReactNode }): JSX.Element { +export function Card({ title, description, aside, children }: { + title?: string + description?: React.ReactNode + /** Rendered at the right of the header (status badge, …). */ + aside?: React.ReactNode + children: React.ReactNode +}): JSX.Element { return (
- {(title || description) && ( -
- {title &&

{title}

} - {description &&

{description}

} + {(title || description || aside) && ( +
+
+ {title &&

{title}

} + {description &&

{description}

} +
+ {aside &&
{aside}
}
)}
diff --git a/src/shared/utils/notification.ts b/src/shared/utils/notification.ts index 864cd161..e5179974 100644 --- a/src/shared/utils/notification.ts +++ b/src/shared/utils/notification.ts @@ -7,7 +7,7 @@ * Skipped when the app window already has focus: the user is looking right at * it, so a toast on top would just be noise. */ -export async function showCompletionNotification(body: string, title = 'Modly'): Promise { +async function notifyIfUnfocused(title: string, body: string): Promise { if (typeof document !== 'undefined' && document.hasFocus()) return try { await window.electron.notifications.show(title, body) @@ -15,3 +15,11 @@ export async function showCompletionNotification(body: string, title = 'Modly'): // Notifications not available (e.g. unsupported platform) } } + +export function showCompletionNotification(body: string, title = 'Modly'): Promise { + return notifyIfUnfocused(title, body) +} + +export function showErrorNotification(body: string, title = 'Modly'): Promise { + return notifyIfUnfocused(title, body) +} diff --git a/src/shared/utils/weightVariants.test.mjs b/src/shared/utils/weightVariants.test.mjs new file mode 100644 index 00000000..1659391b --- /dev/null +++ b/src/shared/utils/weightVariants.test.mjs @@ -0,0 +1,63 @@ +import test from 'node:test' +import assert from 'node:assert/strict' +import { buildSync } from 'esbuild' +import { createRequire } from 'node:module' +import { mkdtempSync, writeFileSync } from 'node:fs' +import { tmpdir } from 'node:os' +import { join, resolve } from 'node:path' + +function loadModule() { + const outfile = join(mkdtempSync(join(tmpdir(), 'modly-weight-variants-test-')), 'weightVariants.cjs') + const require = createRequire(import.meta.url) + const result = buildSync({ + entryPoints: [resolve('src/shared/utils/weightVariants.ts')], + bundle: true, + platform: 'node', + format: 'cjs', + write: false, + }) + writeFileSync(outfile, result.outputFiles[0].text, 'utf8') + return require(outfile) +} + +const { withWeightVariantAvailability, isMissingWeightVariant } = loadModule() + +const variants = { + param: 'gguf_quant', + default: 'Q5_K_M', + options: [{ id: 'Q4_K_M', label: 'Q4_K_M' }, { id: 'Q5_K_M', label: 'Q5_K_M' }], +} + +const quantParam = { + id: 'gguf_quant', + label: 'Quantization', + type: 'select', + default: 'Q5_K_M', + options: [ + { value: 'Q4_K_M', label: 'Q4_K_M' }, + { value: 'Q5_K_M', label: 'Q5_K_M' }, + { value: 'auto', label: 'Auto' }, + ], +} + +test('labels declared variants that are not installed and keeps their values', () => { + const marked = withWeightVariantAvailability(quantParam, variants, ['Q5_K_M']) + assert.deepEqual(marked.options.map((option) => option.label), ['Q4_K_M (not installed)', 'Q5_K_M', 'Auto']) + assert.deepEqual(marked.options.map((option) => option.value), ['Q4_K_M', 'Q5_K_M', 'auto']) +}) + +test('returns the param untouched when availability is unknown or the param selects nothing', () => { + assert.equal(withWeightVariantAvailability(quantParam, variants, undefined), quantParam) + assert.equal(withWeightVariantAvailability(quantParam, undefined, []), quantParam) + const steps = { id: 'steps', label: 'Steps', type: 'int', default: 25 } + assert.equal(withWeightVariantAvailability(steps, variants, []), steps) +}) + +test('flags a selected variant only when availability is known and it is not installed', () => { + assert.equal(isMissingWeightVariant('gguf_quant', 'Q4_K_M', variants, ['Q5_K_M']), true) + assert.equal(isMissingWeightVariant('gguf_quant', 'Q5_K_M', variants, ['Q5_K_M']), false) + assert.equal(isMissingWeightVariant('gguf_quant', 'auto', variants, []), false) + assert.equal(isMissingWeightVariant('steps', 'Q4_K_M', variants, []), false) + assert.equal(isMissingWeightVariant('gguf_quant', 'Q4_K_M', variants, undefined), false) + assert.equal(isMissingWeightVariant('gguf_quant', 'Q4_K_M', undefined, []), false) +}) diff --git a/src/shared/utils/weightVariants.ts b/src/shared/utils/weightVariants.ts new file mode 100644 index 00000000..252fd6d5 --- /dev/null +++ b/src/shared/utils/weightVariants.ts @@ -0,0 +1,35 @@ +import type { ParamSchema, WeightVariantsInfo } from '@shared/types/electron.d' + +/** + * Suffix the options of a node's variant-selecting param that are declared + * weight variants but not installed. Other params, and values that are not + * variant ids, are returned untouched. `installed` undefined means not known yet. + */ +export function withWeightVariantAvailability( + param: ParamSchema, + variants: WeightVariantsInfo | undefined, + installed: string[] | undefined, +): ParamSchema { + if (!variants || !installed || param.id !== variants.param || !param.options) return param + const variantIds = new Set(variants.options.map((option) => option.id)) + return { + ...param, + options: param.options.map((option) => { + const value = String(option.value) + if (!variantIds.has(value) || installed.includes(value)) return option + return { ...option, label: `${option.label ?? value} (not installed)` } + }), + } +} + +/** True when `value` selects a declared weight variant that is known not to be installed. */ +export function isMissingWeightVariant( + paramId: string, + value: unknown, + variants: WeightVariantsInfo | undefined, + installed: string[] | undefined, +): boolean { + if (!variants || !installed || paramId !== variants.param) return false + const id = String(value) + return variants.options.some((option) => option.id === id) && !installed.includes(id) +} diff --git a/tsconfig.builtins.json b/tsconfig.builtins.json index 5a7a957a..15c7e29d 100644 --- a/tsconfig.builtins.json +++ b/tsconfig.builtins.json @@ -12,7 +12,6 @@ "esModuleInterop": true }, "include": [ - "src/areas/workflows/nodes/mesh-exporter/**/*.ts", - "src/areas/workflows/nodes/mesh-optimizer/**/*.ts" + "src/areas/workflows/nodes/mesh-exporter/**/*.ts" ] }