diff --git a/.dockerignore b/.dockerignore index 9798785..c49484a 100644 --- a/.dockerignore +++ b/.dockerignore @@ -14,9 +14,11 @@ backend/static htmlcov .coverage -# Tiled server config, data & catalog (external in production — never bundled) +# Tiled's own runtime catalog/data (local dev state — never bundled). NOT +# excluding tiled/ itself: the app-full Docker image's Tiled server needs +# tiled/config.docker.yml (see Dockerfile's app-full stage) in the build +# context, and it's tiny source config, not runtime data. .tiled -tiled/ # Local data, backups, runtime files data diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..03334ce --- /dev/null +++ b/.env.example @@ -0,0 +1,49 @@ +# docker-compose config — copy to `.env` at the repo root (docker compose +# loads this automatically for every `docker compose -f docker-compose*.yml` +# invocation; git-ignored, never commit real secrets into it). +# +# This is distinct from backend/.env.example (used by start_all.sh / running +# the backend directly, not via Docker) and frontend/.env.example (used by a +# plain `npm run build`/`npm run dev`, not via Docker) — those are for local +# development outside a container; this file is for `docker compose` runs. +# +# Relevant to: docker-compose.yml (app), docker-compose.ml.yml (app-ml). +# docker-compose.full.yml and docker-compose.local.yml bundle their own Tiled, +# so TILED_URI/TILED_BROWSE_PATH don't apply there — only TILED_API_KEY and +# BROWSE_ALLOWED_ORIGINS do. + +# --- Connect to an existing (e.g. staging/production) Tiled server --- +# Required for docker-compose.yml / docker-compose.ml.yml, which never bundle +# a Tiled server themselves. +TILED_URI=https://tiled-staging.als.lbl.gov +TILED_API_KEY= + +# Path into that Tiled server's tree Browse should treat as its root — set +# this to the real institutional path (confirm with whoever operates that +# Tiled server; don't assume it matches another deployment's beamline). +# Leave unset for a from-scratch bundled Tiled (docker-compose.full.yml / +# docker-compose.local.yml), which uses its own ingest root instead. +TILED_BROWSE_PATH=beamlines/bl832/processed + +# --- Frontend build-time (docker-compose.ml.yml only; baked in at build, +# not read at container start — see docs/reference/deployment.md) --- +# Bare path only, never a scheme/host (see that doc's warning on this). +# Leave unset for root-hosted. +VITE_BASE_PATH=/bl832/seg_studio/ + +# --- Same-origin SPA+API → CORS can stay empty; set only for split origins --- +BROWSE_ALLOWED_ORIGINS= + +# --- docker-compose.local.yml only: host port for the nginx proxy in front +# of the bundled :local-shaped stack (default 5173) --- +HOST_PORT=5173 + +# --- docker-compose.full.yml / docker-compose.local.yml only: bind-mount a +# real host directory to /data/processed, so the bundled Tiled server can read your +# own datasets directly (register-in-place via the Zarr loader, or the +# Browse-tab ingest flow) instead of only ever seeing what was uploaded +# through the app. Defaults to an empty placeholder (.local-source/, git- +# ignored) so these compose files still work unset. Must be an ABSOLUTE +# host path (a relative one is resolved against wherever `docker compose` +# happens to be run from, not this file's location): +# LOCAL_SOURCE_DIR=/Users/you/Documents/data/tomo diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 07874df..ebe7b36 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -5,6 +5,11 @@ on: branches: [main] pull_request: +# None of these jobs push commits, comment on PRs, or touch packages — read-only +# is the least-privilege default (CodeQL actions/missing-workflow-permissions). +permissions: + contents: read + jobs: backend: runs-on: ubuntu-latest @@ -16,20 +21,58 @@ jobs: - uses: actions/setup-python@v5 with: python-version: "3.11" - - run: pip install -e ".[dev,test]" + # `ml` (torch/dlsia/qlty) installed here so dlsia's train/infer code + # paths are actually exercised in CI, CPU-only — without it every + # dlsia-dependent test silently skips (torch_available() is False). + - run: pip install -e ".[dev,test,ml]" - run: flake8 . --max-line-length=120 --extend-ignore=E501,W503 - run: pytest + ipred: + runs-on: ubuntu-latest + defaults: + run: + working-directory: ipred + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: "3.11" + - run: pip install -e ".[test]" + - run: pytest + frontend: runs-on: ubuntu-latest defaults: run: working-directory: frontend steps: + # The volume renderer lives in a submodule (frontend/vendor/); without it + # typecheck and build fail on an unresolved import. - uses: actions/checkout@v4 + with: + submodules: recursive - uses: actions/setup-node@v4 with: node-version: "20" - run: npm ci - run: npm run typecheck - run: npm run test + + docker-build: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + with: + submodules: recursive + - uses: docker/setup-buildx-action@v3 + # Build only — catches a Dockerfile/compose file silently falling + # behind the app it packages (the exact drift found in this session's + # review: no ipred service, no ml extra, no Tiled service). Three + # standalone compose files (not overlays — each is a separate, complete + # service, not an override of another), built independently. Does not + # push anywhere; image publishing lives in its own workflow + # (publish-image.yml), gated on merge to main, not every PR. + - run: docker compose -f docker-compose.yml build + - run: docker compose -f docker-compose.ml.yml build + - run: docker compose -f docker-compose.full.yml build diff --git a/.github/workflows/publish-image.yml b/.github/workflows/publish-image.yml new file mode 100644 index 0000000..137990e --- /dev/null +++ b/.github/workflows/publish-image.yml @@ -0,0 +1,94 @@ +name: Create and publish image + +# Mirrors als-computing/view_tomography_recon_app's own publish-image.yml — +# same tag names (:local / :als-prod / :als-staging), same one-image-many- +# tags shape. Runs on merge to main only (never on a PR — that stays in +# ci.yml's build-only docker-build job, so PRs get a real "does it build" +# check with no registry pushes or wasted publish minutes). +on: + push: + branches: ['main'] + tags: ['v*'] + +env: + REGISTRY: ghcr.io + IMAGE_NAME: ${{ github.repository }} + +jobs: + build-and-push-image: + runs-on: ubuntu-latest + permissions: + contents: read + packages: write + + steps: + - name: Checkout repository + uses: actions/checkout@v4 + with: + # The 3D viewer is a git submodule (frontend/vendor/) — without it + # the frontend build stage fails on an unresolved import. + submodules: recursive + fetch-depth: 0 + + - name: Log in to the Container registry + uses: docker/login-action@v3 + with: + registry: ${{ env.REGISTRY }} + username: ${{ github.actor }} + password: ${{ secrets.GITHUB_TOKEN }} + + - name: Extract metadata (tags, labels) for Docker + id: meta + uses: docker/metadata-action@v5 + with: + images: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }} + + - name: Set up Docker Buildx + uses: docker/setup-buildx-action@v3 + + # :local — fully bundled (Tiled + backend/ml + ipred), for anyone + # running this repo locally with nothing external set up. Base path + # baked in at /seg_studio/ so it exercises subpath hosting by default, + # matching :als-prod/:als-staging below rather than leaving that path + # tested only at ALS. See docker-compose.local.yml. + - name: Build and push local image + uses: docker/build-push-action@v6 + with: + context: . + target: app-full + push: true + build-args: | + VITE_BASE_PATH=/seg_studio/ + tags: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}:local + labels: ${{ steps.meta.outputs.labels }} + + # :als-prod — ipred/ML bundled, Tiled external (ALS's own production + # Tiled — see docker-compose.als-prod.yml for the real TILED_URI + # default), hosted under ALS's hub path prefix. + - name: Build and push ALS production image + uses: docker/build-push-action@v6 + with: + context: . + target: app-ml + push: true + build-args: | + VITE_BASE_PATH=/bl832/seg_studio/ + tags: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}:als-prod + labels: ${{ steps.meta.outputs.labels }} + + # :als-staging — identical build to :als-prod today (this app reads + # TILED_URI at container-run time, not build time, so the prod/staging + # distinction lives in docker-compose.als-staging.yml's default rather + # than here) — kept as its own explicit step, not just a second tag on + # the same build, so staging and prod can diverge later without + # restructuring this workflow. + - name: Build and push ALS staging image + uses: docker/build-push-action@v6 + with: + context: . + target: app-ml + push: true + build-args: | + VITE_BASE_PATH=/bl832/seg_studio/ + tags: ${{ env.REGISTRY }}/${{ env.IMAGE_NAME }}:als-staging + labels: ${{ steps.meta.outputs.labels }} diff --git a/.gitignore b/.gitignore index d7e46d2..a651d27 100644 --- a/.gitignore +++ b/.gitignore @@ -13,6 +13,11 @@ build/ # Secrets — never commit .env +# Default bind-mount target for docker-compose.full.yml/docker-compose.local.yml's +# LOCAL_SOURCE_DIR (see .env.example) — an empty placeholder so the compose +# files work unset; your own data lives wherever LOCAL_SOURCE_DIR points, not here. +.local-source/ + # Tiled catalog database (binary, regenerated on first run via start_all.sh) .tiled/catalog.db .tiled/catalog.db-shm @@ -23,6 +28,10 @@ build/ # ingested locally are ignored. To commit a new demo dataset, use `git add -f`. .tiled/data/ +# Generated 3-D volume pyramids (backend/tiff_stack_source.py). Derived data — +# rebuildable from the source TIFFs by re-registering, and large. +.tiled/volumes/ + # Runtime Tiled thumbnail containers created on save **/*__v_thumbs/ @@ -52,6 +61,7 @@ Thumbs.db htmlcov/ .coverage coverage.xml +frontend/coverage/ # Vendored SAM model (large, fetched via frontend/scripts/fetch-sam-model.mjs) frontend/public/models/ @@ -63,3 +73,9 @@ site/ peter/ .reset-backups/ + +.claude/ + +CLAUDE.md + +graphify-out/* \ No newline at end of file diff --git a/.gitmodules b/.gitmodules new file mode 100644 index 0000000..d1edcb4 --- /dev/null +++ b/.gitmodules @@ -0,0 +1,4 @@ +[submodule "frontend/vendor/view_tomography_recon_app"] + path = frontend/vendor/view_tomography_recon_app + url = https://github.com/als-computing/view_tomography_recon_app.git + branch = remote-tiled diff --git a/Dockerfile b/Dockerfile index 12ff00c..012e1b7 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,18 +1,54 @@ # syntax=docker/dockerfile:1 -# Lightweight production image: builds the SPA and serves it from the FastAPI -# backend (single container, one port). Tiled is NOT included — point the app at -# an external Tiled via TILED_URI/TILED_API_KEY. Local dev still uses start_all.sh. +# Three images from one file, selected via `--target`: +# app (default) — frontend + backend only, Tiled external +# (docker build . / docker compose up, unchanged). +# app-ml — the same, PLUS the `ml` extra (torch/dlsia) and +# ipred bundled, Tiled still external — for a +# deployment with its own production Tiled that +# still wants Train/iPred to work out of the box +# (see docker-compose.ml.yml, and the `:als` tag). +# app-full — app-ml PLUS a bundled Tiled server in the same +# container (all three talk over 127.0.0.1), for a +# single `docker run` that needs nothing external — +# see docker-compose.full.yml, and the `:local` tag. +# Local dev still uses start_all.sh, which runs the same three services as +# separate local processes instead of inside a container. +# +# Each publishable target is a thin `*-base` stage (Python deps only) plus one +# final layer that copies in the built frontend. This is deliberate: the +# frontend build changes far more often than the Python deps do, and +# Docker/BuildKit's cache is a linear chain per stage — a layer's cache key +# depends on its parent layer's digest. If the frontend COPY sat partway +# through a stage (as it once did), any frontend-only change would invalidate +# every layer built on top of it, including the expensive torch/CUDA and +# tiled[all] installs in app-ml/app-full, forcing them to reinstall from +# scratch for a one-line frontend fix. Keeping the frontend COPY as the LAST +# layer in each `*-base` stage's leaf means a frontend-only change only ever +# invalidates that one small COPY (+ its trailing EXPOSE/CMD/ENTRYPOINT +# metadata) — never the Python installs above it. -# --- Stage 1: build the frontend (same-origin: VITE_API_BASE left empty) --- +# --- Stage: build the frontend (same-origin: VITE_API_BASE left empty) --- FROM node:22-alpine AS web WORKDIR /web COPY frontend/package.json frontend/package-lock.json ./ RUN npm ci COPY frontend/ ./ +# The WebGPU volume renderer is a git submodule under frontend/vendor/. Docker +# copies the working tree as-is, so an uninitialised submodule arrives as an +# empty directory and the build fails deep inside Vite with an unresolved +# import. Fail here instead, with the fix in the message. +RUN test -f vendor/view_tomography_recon_app/src/zarr-viewer/src/ome-zarr-viewer.ts \ + || (echo "ERROR: submodule frontend/vendor/view_tomography_recon_app is missing." \ + && echo "Run: git submodule update --init --recursive" && exit 1) +# Read at build time (see vite.config.ts's `base` / src/config.ts's API_BASE) — +# empty/unset produces the exact same root-hosted output as before this arg +# existed. Set for a subpath deployment, e.g. --build-arg VITE_BASE_PATH=/bl832/seg_studio/ +ARG VITE_BASE_PATH="" +ENV VITE_BASE_PATH=${VITE_BASE_PATH} RUN npm run build # → /web/dist -# --- Stage 2: backend + built SPA --- -FROM python:3.12-slim AS app +# --- Stage: lightweight production base (backend deps + source, no frontend) --- +FROM python:3.12-slim AS app-base WORKDIR /app # Install Python deps first (cached until pyproject changes). py-modules=[] means @@ -24,13 +60,78 @@ WORKDIR /app COPY backend/pyproject.toml ./pyproject.toml RUN pip install --no-cache-dir . -# App source + the built SPA (served from ./static by annotation_server.py). +# App source (the built SPA is copied in by each leaf stage below, last). COPY backend/ ./ -COPY --from=web /web/dist ./static # Drafts/versions/exports persist here — mount a volume in production. ENV LOCAL_DATA_ROOT=/data VOLUME ["/data"] +# --- Stage: lightweight production image (frontend + backend only) --- +FROM app-base AS app +COPY --from=web /web/dist ./static + EXPOSE 8002 CMD ["uvicorn", "annotation_server:app", "--host", "0.0.0.0", "--port", "8002"] + +# --- Stage: backend + ipred + ml base, Tiled still external --- +# For a deployment that already has its own production Tiled (so bundling a +# second, empty one would be actively wrong) but still wants the Train tab +# and the fast pixel classifier to work without standing up a separate ipred +# service just for this app — see the `:als` tag and docker-compose.ml.yml. +FROM app-base AS app-ml-base + +# The `ml` extra (torch/dlsia/qlty) — kept out of the lean `app` image above. +# Pulls in the full CUDA toolkit (nvidia-cudnn-cu12, nvidia-cublas-cu12, etc.) +# as torch dependencies — several GB — needed for real GPU support when this +# container is run with `--gpus all` on a CUDA host; a build here can fail +# with an I/O error mid-write on a disk-constrained host, which is a local +# Docker disk-space problem to fix (see docker system df / prune, or grow +# Docker Desktop's disk allocation), not a reason to drop CUDA support. +RUN pip install --no-cache-dir ".[ml]" + +# ipred is a sibling package with its own pyproject, not a backend dependency. +COPY ipred/ /ipred/ +RUN pip install --no-cache-dir /ipred + +COPY docker-entrypoint-ml.sh /usr/local/bin/docker-entrypoint-ml.sh +RUN chmod +x /usr/local/bin/docker-entrypoint-ml.sh + +# --- Stage: backend + ipred + ml, Tiled still external --- +FROM app-ml-base AS app-ml +COPY --from=web /web/dist ./static + +# 8002 backend; 8003 ipred, exposed so it can be reached directly if wanted. +# No Tiled port here — this stage never starts one; TILED_URI at `docker run` +# time points at the deployment's own external Tiled. +EXPOSE 8002 8003 +ENTRYPOINT ["/usr/local/bin/docker-entrypoint-ml.sh"] + +# --- Stage: batteries-included base (+ Tiled server too) --- +# Bundles all three backend services into one container over loopback — the +# CI/Docker coverage gap this stage fixes: previously NOTHING packaged ipred +# or a Tiled server at all, so the app image alone could never run iPred or +# dlsia, and there was no single-command way to try the full stack without +# start_all.sh's separate local processes. +FROM app-ml-base AS app-full-base + +# tiled[all] pulls in the actual server (catalog, array/table adapters) — +# `app`'s pyproject only pins tiled[client], which has no server component. +RUN pip install --no-cache-dir "tiled[all]" + +# Portable Tiled config (see tiled/config.docker.yml's own doc comment for +# why this isn't just tiled/config.yml — that one hardcodes a local dev +# machine's absolute data path). +COPY tiled/config.docker.yml /app/tiled/config.docker.yml + +COPY docker-entrypoint-full.sh /usr/local/bin/docker-entrypoint-full.sh +RUN chmod +x /usr/local/bin/docker-entrypoint-full.sh + +# --- Stage: batteries-included image (+ Tiled server too) --- +FROM app-full-base AS app-full +COPY --from=web /web/dist ./static + +# 8002 backend (the only port most deployments need — same-origin SPA+API); +# 8003/8010 exposed too so ipred/Tiled can be reached directly if wanted. +EXPOSE 8002 8003 8010 +ENTRYPOINT ["/usr/local/bin/docker-entrypoint-full.sh"] diff --git a/README.md b/README.md index 79d33de..56aa0c3 100644 --- a/README.md +++ b/README.md @@ -29,6 +29,7 @@ On first run this will automatically: - install [`uv`](https://docs.astral.sh/uv/) if it's missing, - create a `.venv` with Python 3.12 and install the backend dependencies, - generate a strong Tiled API key into `backend/.env` (gitignored, never sent to the browser), +- initialise git submodules (the WebGPU renderer behind the **3D** tab), - vendor the SlimSAM model in the background so the AI Magic tool works offline, - start Tiled, the backend API, the frontend dev server, and the docs site. @@ -70,6 +71,27 @@ PROD=1 ./start_all.sh The frontend is built to `backend/static/` and served by FastAPI. The whole app is then available at the **Backend** URL (http://127.0.0.1:8002). +## Docker + +Three Dockerfile targets/compose files, layered `app` → `app-ml` → `app-full`, cover +different needs — pick the one that matches what you already have running: + +| Compose file | Tiled | iPred / Train | Use when | +| --- | --- | --- | --- | +| `docker-compose.yml` | External (bring your own) | Not included | You already have Tiled and just need Connect/Browse/Annotate/Export. | +| `docker-compose.ml.yml` | External (bring your own) | Bundled | You already have Tiled but also want Train/iPred. | +| `docker-compose.full.yml` | Bundled | Bundled | Nothing external required — the simplest way to try everything. | + +```bash +docker compose -f docker-compose.full.yml up --build # fully bundled +``` + +Then open . See +[Installation](docs/getting-started/installation.md) for the other two images and +env vars, and [Production deployment](docs/reference/deployment.md) for hosting a +shared deployment under a URL prefix (e.g. behind a reverse proxy at +`hub.example.org/your-path/`) with the published `ghcr.io` images. + ## Workflow The app is organized into four tabs: @@ -114,6 +136,12 @@ pytest # tests Frontend (React + TypeScript + Vite): +The 3D tab's volume renderer is a git submodule under `frontend/vendor/`, developed +in its own repo ([als-computing/view_tomography_recon_app](https://github.com/als-computing/view_tomography_recon_app)). +Clone with `--recurse-submodules`, or run `git submodule update --init --recursive` +in an existing checkout — typecheck and build both fail without it. Changes to the +renderer belong upstream; bump the pointer here with `git submodule update --remote`. + ```bash cd frontend npm install @@ -139,7 +167,9 @@ These are the same checks CI runs (see `.github/workflows/ci.yml`). ``` backend/ FastAPI API, COCO/Lightly export, Tiled client, local-folder access frontend/ React SPA (Konva canvas, Zustand stores, in-browser SAM) +frontend/vendor/ Git submodules — the WebGPU volume renderer used by the 3D tab tiled/ Local Tiled server config docs/ MkDocs Material documentation site start_all.sh One-command launcher for the full stack +Dockerfile app / app-ml / app-full image targets (see Docker, above) ``` diff --git a/backend/.env.example b/backend/.env.example index c26e815..cfa192d 100644 --- a/backend/.env.example +++ b/backend/.env.example @@ -8,6 +8,14 @@ BROWSE_CACHE_TTL_SECONDS=300 LOCAL_DATA_ROOT=~/data EXPORT_ROOT= +# Path into the Tiled tree Browse treats as its root. Leave unset for a +# from-scratch bundled Tiled (defaults to wherever this repo's own ingest +# writes under, see TILED_INGEST_ROOT). A deployment pointed at an existing +# institutional Tiled catalog (e.g. ALS's production Tiled) needs this set to +# the real path into that catalog — confirm the exact value with whoever +# operates that Tiled server rather than guessing: +# TILED_BROWSE_PATH=beamlines/bl832/processed + # Production / Docker notes: # - The container serves the SPA same-origin, so BROWSE_ALLOWED_ORIGINS can stay # empty. Set it (comma-separated) only for split frontend/backend origins. diff --git a/backend/annotation_server.py b/backend/annotation_server.py index 66591a7..ec3c4a5 100644 --- a/backend/annotation_server.py +++ b/backend/annotation_server.py @@ -8,6 +8,7 @@ * ``GET /api/browse/items`` — sample records matching a filter set * ``GET /api/browse/thumbnail`` — PNG thumbnail for a Tiled array path * ``GET /health`` — liveness check +* ``/api/ipred/*`` — proxy to the standalone iPred service (see ``ipred_routes.py``) Run with -------- @@ -26,7 +27,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import datetime, timezone from pathlib import Path -from typing import Optional +from typing import Any, Optional import numpy as np from fastapi import FastAPI, File, Form, HTTPException, Query, Response, UploadFile @@ -37,12 +38,23 @@ import annotation_thumbnails import arrays as arrays_mod +import batch_probe +import denoise as denoise_mod +import denoise_bake as denoise_bake_mod import drafts as drafts_mod import export_jobs import guides as guides_mod import images as images_mod +import infer_jobs import ingest as ingest_mod +import ipred_routes import local_fs +import tiff_stack_source +import train_common +import train_jobs +import volume_build +import volume_nodes +import zarr_source from browse_helpers import ( _SINGLE_VALUE_FACET_RAW_KEYS, FieldMapping, @@ -62,14 +74,22 @@ write_lightly_split, ) from schemas import ( + BatchProbeRequest, + DenoiseBakeRequest, DraftPayload, ExportRequest, ExportSourceItem, GuidePayload, ImageMeta, + InferRequest, IngestPreflightRequest, MeasureRequest, SaveVersionRequest, + TiffStackRegisterRequest, + TrainRequest, + VolumeBuildRequest, + ZarrRegisterRequest, + ZarrScanRequest, ) from source_keys import parse_source_key from thumbnails import render_thumbnail @@ -113,6 +133,11 @@ # the ~500-byte floor skips tiny/binary payloads. (Dev uses the Vite server instead.) app.add_middleware(GZipMiddleware, minimum_size=500) +# iPred (interactive segmentation) proxy — see ipred_routes.py. iPred is an +# optional, separately-run service (port 8003 by default); a down/missing +# service surfaces as 503 from these routes rather than breaking the app. +app.include_router(ipred_routes.router) + # --------------------------------------------------------------------------- # Response models @@ -136,6 +161,10 @@ class ServerConfig(BaseModel): _column_cache: TTLCache = TTLCache(ttl_seconds=_CACHE_TTL, max_entries=256) _items_cache: TTLCache = TTLCache(ttl_seconds=_CACHE_TTL, max_entries=128) _field_mapping_cache: TTLCache = TTLCache(ttl_seconds=_FIELD_MAPPING_TTL, max_entries=32) +# Denoised slice PNGs only — the plain render path stays uncached because it is +# already cheap. Entries are a few MB each at full resolution, so 32 covers +# scrubbing a stack back and forth without unbounded growth. +_denoised_slice_cache: TTLCache = TTLCache(ttl_seconds=300.0, max_entries=32) def _resolve_field_mapping( @@ -394,6 +423,18 @@ def _build() -> bytes | None: ) +@app.get("/api/local/root") +async def local_root() -> dict: + """Return the default local browse root (``LOCAL_DATA_ROOT``) as an absolute path. + + ``/api/local/list`` returns entries relative to whatever root was granted; + when no explicit root is granted, callers that need to build an absolute + server-side path (e.g. the Zarr loader's directory browser) fetch it here + once rather than duplicating the ``LOCAL_DATA_ROOT`` default client-side. + """ + return {"root": local_fs.default_root()} + + @app.get("/api/local/list") async def local_list( rel: str = Query("", description="Relative path under the granted root"), @@ -556,9 +597,11 @@ async def image_meta( """Return shape / dtype metadata for an image source.""" def _run() -> ImageMeta: node = arrays_mod.resolve_array(source, kind, server_uri, root) - meta = arrays_mod.array_shape_meta(node) + pyramid = arrays_mod.pyramid_info(source, kind, server_uri, root) + meta = arrays_mod.array_shape_meta(node, pyramid) sl = arrays_mod.read_slice(node, meta, 0) flat = sl.ravel().astype(float) + global_range = images_mod._sample_global_stats(node, meta) if not meta["is_rgb"] else None return ImageMeta( n_slices=meta["n_slices"], height=meta["height"], @@ -566,7 +609,15 @@ def _run() -> ImageMeta: dtype=meta["dtype"], is_rgb=meta["is_rgb"], value_range=[float(flat.min()), float(flat.max())], + global_value_range=list(global_range) if global_range is not None else None, keywords=arrays_mod.node_keywords(node), + level_key=meta.get("level_key"), + level_index=meta.get("level_index"), + level_count=meta.get("level_count"), + level_height=meta.get("level_height"), + level_width=meta.get("level_width"), + level_n_slices=meta.get("level_n_slices"), + z_downsample=meta.get("z_downsample"), ) try: @@ -590,6 +641,14 @@ async def image_slice( vmin_pct: float = Query(1.0), vmax_pct: float = Query(99.0), cmap: str = Query("gray"), + denoise_method: str = Query("none", description="Classical denoise filter (see denoise.ALL_METHODS)"), + denoise_strength: float = Query(0.5, ge=0.0, le=1.0), + denoise_crop: int = Query( + 0, + ge=0, + description="If >0, denoise and return only a centred square crop of this size at 1:1. " + "For tuning: filtering a full slice costs seconds for NLM/TV.", + ), ) -> Response: """Render one slice of an image source as a PNG.""" opts = { @@ -599,13 +658,42 @@ async def image_slice( "vmax_pct": vmax_pct, "cmap": cmap, } + if denoise_method not in denoise_mod.ALL_METHODS: + raise HTTPException(400, f"Unknown denoise method {denoise_method!r}") + denoising = denoise_method != "none" + + # Denoised renders are cached; the plain path stays uncached because it is + # already cheap. Without this, every slider tweak or revisit re-pays the full + # filter cost — seconds, not milliseconds, for NLM and TV. + cache_key = ( + source, kind, slice_index, server_uri, root, norm, scale, vmin_pct, vmax_pct, cmap, + denoise_method, round(denoise_strength, 4), denoise_crop, + ) + if denoising: + cached = _denoised_slice_cache.get(cache_key) + if cached is not None: + return Response(content=cached, media_type="image/png") def _run() -> bytes: node = arrays_mod.resolve_array(source, kind, server_uri, root) - meta = arrays_mod.array_shape_meta(node) - sl = arrays_mod.read_slice(node, meta, slice_index) + pyramid = arrays_mod.pyramid_info(source, kind, server_uri, root) + meta = arrays_mod.array_shape_meta(node, pyramid) + # Denoise the RAW slice, before normalization: noise statistics live in + # the source's own intensity units, not in the 8-bit display range. + if denoising: + sl = _denoised_slice(node, meta, slice_index, denoise_method, denoise_strength, denoise_crop) + else: + sl = arrays_mod.read_slice(node, meta, slice_index) global_range = None if norm == "global": + # NB: deriving this from the pyramid's coarsest level was tried and + # reverted. It is ~8x faster, but those levels are built by AVERAGING, + # which pulls the extremes in hard — on the reference volume the range + # came back (-20.7, 18.0) against (-73.0, 71.3) at full resolution. + # Since this range IS the contrast window, the cheap version visibly + # clips the image. The full-resolution sampler decimates instead of + # averaging, so it keeps the extremes; it costs ~3s once per volume + # and is then cached. global_range = images_mod._sample_global_stats(node, meta) rgb = images_mod.render_slice(sl, opts, global_range) return images_mod.encode_png(rgb) @@ -618,6 +706,9 @@ def _run() -> bytes: logger.error("image_slice failed: %s", exc) raise HTTPException(500, f"Failed to render slice: {exc}") from exc + if denoising: + _denoised_slice_cache.set(cache_key, png) + return Response( content=png, media_type="image/png", @@ -1096,6 +1187,19 @@ def _cb(message: str, _jid: str = jid) -> None: export_jobs.update(jid, state="error", phase="error", error=str(exc)) +@app.post("/api/export/cancel/{job_id}") +async def export_cancel(job_id: str) -> dict: + """Ask a running job to stop at its next clean boundary. + + Cooperative rather than immediate: a job that stops mid-write would leave a + partial dataset that looks complete. Jobs that honour it discard their + partial output; those that do not simply run to completion. + """ + if not export_jobs.request_cancel(job_id): + raise HTTPException(404, "Unknown job_id") + return {"cancelled": True} + + @app.get("/api/export/status/{job_id}") async def export_status(job_id: str) -> dict: """Poll an export job's progress (state, phase, done/total, log, result).""" @@ -1238,6 +1342,496 @@ async def ingest_status(job_id: str) -> dict: return job +@app.get("/api/zarr/inspect") +async def zarr_inspect( + path: str = Query(..., description="Absolute path to a .zarr directory on the server"), +) -> dict: + """Describe a Zarr store's resolution pyramid without registering it. + + Lets the Connect page show what was found — dimensions, dtype, voxel size and + the available levels — before the user commits to loading it. + + Raises: + HTTPException: 4xx from :func:`zarr_source.inspect_zarr` with a + user-facing message (missing path, zipped archive, empty group…). + """ + return await asyncio.to_thread(zarr_source.inspect_zarr, path) + + +@app.post("/api/zarr/preflight") +async def zarr_preflight(req: ZarrRegisterRequest) -> dict: + """Report whether loading this Zarr would collide with an existing node. + + Distinguishes a previous Zarr registration (safe to replace — only catalog + rows are dropped) from internally-managed data such as an uploaded image + stack, where replacing would delete the files themselves. + """ + return await asyncio.to_thread( + zarr_source.preflight_zarr, req.server_uri, req.path, req.container_path + ) + + +@app.post("/api/zarr/register") +async def zarr_register(req: ZarrRegisterRequest) -> dict: + """Register an on-disk Zarr volume with Tiled, copying no data. + + Unlike ``/api/ingest/upload`` this needs no background job: registration + writes catalog rows, not pixels, so it returns in well under a second even + for a 56 GB store. + """ + if req.on_conflict not in ingest_mod.ON_CONFLICT_MODES: + raise HTTPException( + 400, f"on_conflict must be one of {sorted(ingest_mod.ON_CONFLICT_MODES)}" + ) + return await asyncio.to_thread( + zarr_source.register_zarr, + req.server_uri, + req.path, + req.container_path, + req.description, + req.on_conflict, + ) + + +@app.post("/api/zarr/scan") +async def zarr_scan(req: ZarrScanRequest) -> dict: + """Scan a directory for Zarr stores and register any not already present. + + For a bind-mounted directory of already-reconstructed volumes that should + all show up in Browse in one action, rather than registering each one + individually via ``/api/zarr/register``. Safe to re-run — already- + registered stores are skipped, not re-registered. + """ + return await asyncio.to_thread( + zarr_source.scan_and_register_zarrs, + req.server_uri, + req.scan_root, + req.container_path, + req.on_conflict, + req.renames, + ) + + +@app.post("/api/ingest/scan") +async def ingest_scan(req: ZarrScanRequest) -> dict: + """Scan a directory for folders of image slices and ingest any not already + present, as fast per-slice registration (no 3-D pyramid — see the 3D + page's on-demand "Build 3D volume" for that). Safe to re-run. + """ + return await asyncio.to_thread( + ingest_mod.scan_and_register_image_stacks, + req.server_uri, + req.scan_root, + req.container_path, + req.on_conflict, + req.renames, + ) + + +@app.post("/api/scan-datasets") +async def scan_datasets(req: ZarrScanRequest) -> dict: + """Scan a directory for BOTH Zarr stores and folders of image slices, + registering whatever isn't already present. The single combined action + behind the Connect page's "Scan folder" button and the `:local`/`:full` + container's startup auto-discovery — one call instead of two, with one + merged result. + """ + # Sequential, not concurrent: both scans call _ensure_container against the + # same target container, and running them in parallel would race on its + # creation. + zarr_result = await asyncio.to_thread( + zarr_source.scan_and_register_zarrs, + req.server_uri, req.scan_root, req.container_path, req.on_conflict, req.renames, + ) + image_result = await asyncio.to_thread( + ingest_mod.scan_and_register_image_stacks, + req.server_uri, req.scan_root, req.container_path, req.on_conflict, req.renames, + ) + return { + "scanned": zarr_result["scanned"] + image_result["scanned"], + "registered": zarr_result["registered"] + image_result["registered"], + "skipped": zarr_result["skipped"] + image_result["skipped"], + "shadowed": zarr_result["shadowed"] + image_result["shadowed"], + "errors": zarr_result["errors"] + image_result["errors"], + } + + +def _centre_crop(arr: np.ndarray, size: int) -> np.ndarray: + """Centred square crop of *size*, or *arr* unchanged if it already fits.""" + h, w = arr.shape[:2] + if size <= 0 or (h <= size and w <= size): + return arr + top = max(0, (h - size) // 2) + left = max(0, (w - size) // 2) + return arr[top: top + min(size, h), left: left + min(size, w)] + + +def _denoised_slice( + node: Any, + meta: dict, + slice_index: int, + method: str, + strength: float, + crop: int = 0, +) -> np.ndarray: + """Read *slice_index* and denoise it, pulling z-neighbours for 3-D methods. + + The 3-D filters are the training-free way to exploit slice-to-slice + correlation — adjacent tomographic slices share structure while their noise + is independent — which is why this reads a window rather than one slice. The + window is clamped to the volume, and the target's position inside the stack + is tracked explicitly: it is NOT always the centre, since the window is + truncated at the first and last slice. + + When *crop* > 0 the crop is taken BEFORE filtering — filtering a full 6.5 MP + slice is what's slow. Worth knowing: results near the crop border, and NLM's + patch search in particular, differ slightly from the full-slice result, so a + crop is a tuning aid rather than a byte-exact preview of the bake. + """ + radius = denoise_mod.z_radius_for(method) + if radius == 0: + sl = _centre_crop(np.asarray(arrays_mod.read_slice(node, meta, slice_index)), crop) + return denoise_mod.denoise_slice(sl, method, strength) + + n_slices = int(meta["n_slices"]) + lo = max(0, slice_index - radius) + hi = min(n_slices - 1, slice_index + radius) + frames = [] + for idx in range(lo, hi + 1): + try: + frames.append(_centre_crop(np.asarray(arrays_mod.read_slice(node, meta, idx)), crop)) + except Exception as exc: # noqa: BLE001 — a bad neighbour must not fail the view + logger.warning("denoise: skipping unreadable neighbour slice %d: %s", idx, exc) + if idx == slice_index: + raise + if len(frames) < 2: + # Not enough usable z-context (single-slice source, or unreadable + # neighbours) — fall back to the 2-D sibling rather than erroring out. + fallback = "gaussian" if method == "gaussian3d" else "median" + sl = _centre_crop(np.asarray(arrays_mod.read_slice(node, meta, slice_index)), crop) + return denoise_mod.denoise_slice(sl, fallback, strength) + + target_pos = min(slice_index - lo, len(frames) - 1) + return denoise_mod.denoise_stack(np.stack(frames, axis=0), method, strength)[target_pos] + + +@app.get("/api/denoise/methods") +async def denoise_methods() -> dict: + """Denoise filters available in this environment, with cost hints. + + ``available`` is probed rather than assumed: ``denoise_wavelet`` imports + fine without PyWavelets and only fails when called. + """ + return {"methods": denoise_mod.describe_methods()} + + +@app.get("/api/denoise/auto") +async def denoise_auto( + source: str = Query(...), + kind: str = Query(...), + slice_index: int = Query(0), + method: str = Query("tv"), + server_uri: Optional[str] = None, + root: Optional[str] = Query(None), +) -> dict: + """Suggest a strength for this slice, from its own measured noise level.""" + if method not in denoise_mod.ALL_METHODS: + raise HTTPException(400, f"Unknown denoise method {method!r}") + + def _run() -> dict: + node = arrays_mod.resolve_array(source, kind, server_uri, root) + pyramid = arrays_mod.pyramid_info(source, kind, server_uri, root) + meta = arrays_mod.array_shape_meta(node, pyramid) + sl = np.asarray(arrays_mod.read_slice(node, meta, slice_index)) + unit, _, span = denoise_mod._to_unit(sl) + return { + "strength": denoise_mod.auto_strength(sl, method), + # Noise as a fraction of the slice's own dynamic range, so the UI can + # say how noisy this is rather than only what to do about it. + "noise_sigma": denoise_mod.estimate_noise_sigma(unit) if span > 0 else 0.0, + } + + return await asyncio.to_thread(_run) + + +@app.post("/api/denoise/bake") +async def denoise_bake(payload: DenoiseBakeRequest) -> dict: + """Denoise a whole volume and save it as a new, annotatable Tiled dataset. + + The Annotate preview is display-only; this is how a denoised volume becomes + real data you can annotate and export. Runs on a background thread; poll + ``GET /api/export/status/{job_id}``. + """ + if payload.method == "none": + raise HTTPException(422, "Pick a denoise method before saving a denoised copy.") + if payload.method not in denoise_mod.ALL_METHODS: + raise HTTPException(422, f"Unknown denoise method: {payload.method}") + if payload.method not in denoise_mod.available_methods(): + raise HTTPException( + 422, + f"Denoise method {payload.method!r} is unavailable on this server " + "(missing optional dependency).", + ) + + target = payload.target_path or denoise_bake_mod.default_target_path(payload.source) + try: + ingest_mod.validate_container_path(target) + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + + jid = export_jobs.new_job(target) + threading.Thread( + target=denoise_bake_mod.run_denoise_bake_job, args=(jid, payload), daemon=True + ).start() + return {"job_id": jid, "target_path": target} + + +@app.get("/api/train/capability") +async def train_capability() -> dict: + """Best-effort snapshot of Train-tab readiness (torch/dlsia/tiling/device).""" + return train_common.capability() + + +@app.get("/api/train/runs") +async def train_list_runs() -> dict: + """List saved fine-tune runs (both dlsia_tunet and dlsia_denoiser), newest first.""" + return {"runs": train_common.list_runs()} + + +@app.delete("/api/train/runs/{run_id}") +async def train_delete_run(run_id: str) -> dict: + """Permanently remove a saved run's directory (config, metrics, weights).""" + train_common.delete_run(run_id) + return {"deleted": run_id} + + +@app.post("/api/train/start") +async def train_start(payload: TrainRequest) -> dict: + """Start a fine-tuning job. Runs on a background thread; poll + ``GET /api/export/status/{job_id}``; cancel via + ``POST /api/export/cancel/{job_id}`` (same shared registry every + background job in this app already uses). + """ + if not train_common.torch_available(): + raise HTTPException(503, "torch is not installed on this server — see the ml extra in pyproject.toml") + needs_dlsia = payload.model.model_family == "dlsia_tunet" or ( + payload.model.model_family == "dlsia_denoiser" and payload.model.architecture == "tunet" + ) + if needs_dlsia and not train_common.dlsia_available(): + raise HTTPException(503, "dlsia is not installed on this server") + if payload.model.hyperparams.tiling: + import tiling + + if not tiling.qlty_available(): + raise HTTPException(503, "Tiling requires the 'qlty' package, which is not installed on this server") + + run_id = train_jobs.new_run_id(payload.model.model_family) + if payload.resume_from_run_id: + try: + parent_config = train_common.load_run_config(payload.resume_from_run_id) + train_jobs.check_resume_compatible(parent_config, payload) + except HTTPException: + raise + except ValueError as exc: + raise HTTPException(400, str(exc)) from exc + + if not train_common.ML_LOCK.acquire(blocking=False): + raise HTTPException(409, "Another training or inference job is already running") + train_common.ML_LOCK.release() # run_train_job re-acquires it itself; this was just a pre-check + + jid = export_jobs.new_job(run_id) + threading.Thread(target=train_jobs.run_train_job, args=(jid, payload, run_id), daemon=True).start() + return {"job_id": jid, "run_id": run_id} + + +@app.post("/api/train/estimate-batch") +async def train_estimate_batch(payload: BatchProbeRequest) -> dict: + """Measure the largest batch size that fits in device memory for *payload.model*. + + Runs on a background thread; poll ``GET /api/export/status/{job_id}``. + """ + if not train_common.torch_available(): + raise HTTPException(503, "torch is not installed on this server") + jid = export_jobs.new_job(f"batch-probe:{payload.model.model_family}") + threading.Thread(target=batch_probe.run_probe_job, args=(jid, payload), daemon=True).start() + return {"job_id": jid} + + +@app.post("/api/train/infer") +async def train_infer(payload: InferRequest) -> dict: + """Run a saved fine-tuned run over the requested slices. Runs on a + background thread; poll ``GET /api/export/status/{job_id}``. + """ + if not train_common.torch_available(): + raise HTTPException(503, "torch is not installed on this server") + try: + train_common.load_run_config(payload.run_id) + except HTTPException: + raise + jid = export_jobs.new_job(payload.run_id) + threading.Thread(target=infer_jobs.run_infer_job, args=(jid, payload), daemon=True).start() + return {"job_id": jid} + + +@app.get("/api/train/infer/preview/{job_id}/{slice_index}") +async def train_infer_preview(job_id: str, slice_index: int) -> Response: + """Colourised RGBA overlay PNG for one predicted slice of a cached inference job.""" + return Response(content=infer_jobs.preview_png(job_id, slice_index), media_type="image/png") + + +@app.post("/api/train/infer/write-tiled/{job_id}") +async def train_infer_write_tiled(job_id: str) -> dict: + """Push a completed inference job's label maps into Tiled. Runs on a + background thread; poll ``GET /api/export/status/{job_id}`` with the + RETURNED job id (distinct from *job_id*, the inference job being written). + """ + write_jid = export_jobs.new_job(f"write-tiled:{job_id}") + threading.Thread(target=infer_jobs.run_write_tiled_job, args=(write_jid, job_id), daemon=True).start() + return {"job_id": write_jid} + + +@app.get("/api/volume/resolve") +async def volume_resolve( + source: str = Query(..., description="Tiled path of the open dataset"), + server_uri: Optional[str] = Query(None), +) -> dict: + """Locate the renderable 3-D volume for the open dataset. + + Which node holds it depends on how the dataset was catalogued — a registered + Zarr volume is one already, a TIFF stack's lives in its ``__volume`` sidecar, + and a stack nobody has built one for has none. The frontend cannot tell these + apart from the path, and guessing produces + ``missing multiscales in root .zattrs`` at the viewer instead of an answer. + """ + return await asyncio.to_thread(volume_nodes.resolve_volume, server_uri, source) + + +@app.get("/api/volume/build/inspect") +async def volume_build_inspect( + source: str = Query(..., description="Tiled path of the per-slice dataset"), + kind: str = Query("tiled"), + server_uri: Optional[str] = Query(None), +) -> dict: + """Describe the 3-D volume that would be built for this dataset.""" + return await asyncio.to_thread( + volume_build.inspect_volume_build, source, kind, server_uri + ) + + +@app.post("/api/volume/build") +async def volume_build_start(req: VolumeBuildRequest) -> dict: + """Build a 3-D volume from a slice stack already in the catalog. + + Needs nothing but the open dataset: the slices are already in Tiled, so + asking for a source directory would be asking the user to re-supply data the + app has. Returns a ``job_id``; poll ``GET /api/export/status/{job_id}``. + """ + info = await asyncio.to_thread( + volume_build.inspect_volume_build, req.source, req.kind, req.server_uri + ) + jid = export_jobs.new_job(req.source) + export_jobs.set_total(jid, max(info["slices_to_read"], 1)) + + def _run() -> None: + try: + export_jobs.update(jid, state="running", phase="building") + + def _progress(message: str, done: int, total: int) -> None: + export_jobs.update(jid, phase=message, done=done, total=max(total, 1)) + + result = volume_build.build_volume( + req.source, req.kind, req.server_uri, req.container_path, progress=_progress + ) + export_jobs.update(jid, state="done", phase="done", result=result) + export_jobs.log(jid, f"Built 3-D volume {result['key']!r}.") + except HTTPException as exc: + export_jobs.update(jid, state="error", phase="error", error=str(exc.detail)) + except Exception as exc: # noqa: BLE001 — surfaced to the UI via the job + logger.exception("volume build failed") + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + + threading.Thread(target=_run, daemon=True).start() + return {"job_id": jid, **info} + + +@app.get("/api/tiff-stack/inspect") +async def tiff_stack_inspect( + path: str = Query(..., description="Absolute path to a directory of TIFF slices"), +) -> dict: + """Describe a TIFF directory and the pyramid that would be built for it. + + Reads only the first file, so this is cheap enough to call while the user is + still typing a path. ``slices_to_read`` lets the UI say up front how much + work registration will be, rather than appearing to hang. + + Raises: + HTTPException: 4xx from :func:`tiff_stack_source.inspect_tiff_stack` with + a user-facing message (missing path, no TIFFs, inconsistent + numbering…). + """ + return await asyncio.to_thread(tiff_stack_source.inspect_tiff_stack, path) + + +@app.post("/api/tiff-stack/preflight") +async def tiff_stack_preflight(req: TiffStackRegisterRequest) -> dict: + """Report whether registering this TIFF stack would collide, changing nothing.""" + return await asyncio.to_thread( + tiff_stack_source.preflight_tiff_stack, req.server_uri, req.path, req.container_path + ) + + +@app.post("/api/tiff-stack/register") +async def tiff_stack_register(req: TiffStackRegisterRequest) -> dict: + """Register a TIFF directory as a 3-D multiscale volume, copying no slices. + + Returns a ``job_id`` immediately; poll ``GET /api/export/status/{job_id}``. + A job rather than a straight call because — unlike Zarr registration, which + only writes catalog rows — the downsampled levels the 3-D viewer renders have + to be computed, and that means reading every source slice once. + + The full-resolution slices themselves are registered in place: no pixels are + copied, and the existing per-slice nodes the 2-D canvas reads are untouched. + """ + if req.on_conflict not in ingest_mod.ON_CONFLICT_MODES: + raise HTTPException( + 400, f"on_conflict must be one of {sorted(ingest_mod.ON_CONFLICT_MODES)}" + ) + # Validate before returning a job id, so a bad path is a 4xx the user sees + # immediately rather than a job that fails a second later. + info = await asyncio.to_thread(tiff_stack_source.inspect_tiff_stack, req.path) + + jid = export_jobs.new_job(req.path) + # Every source slice is read exactly once: the finest generated level comes + # from the TIFFs, the coarser ones cascade from it in memory. + export_jobs.set_total(jid, max(info["slices_to_read"], 1)) + + def _run() -> None: + try: + export_jobs.update(jid, state="running", phase="registering") + + def _progress(message: str, done: int, total: int) -> None: + export_jobs.update(jid, phase=message, done=done, total=max(total, 1)) + + result = tiff_stack_source.register_tiff_stack( + req.server_uri, + req.path, + req.container_path, + req.description, + req.on_conflict, + progress=_progress, + ) + export_jobs.update(jid, state="done", phase="done", result=result) + export_jobs.log(jid, f"Registered {result['key']!r} as a 3-D volume.") + except HTTPException as exc: + export_jobs.update(jid, state="error", phase="error", error=str(exc.detail)) + except Exception as exc: # noqa: BLE001 — surfaced to the UI via the job + logger.exception("tiff stack registration failed") + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + + threading.Thread(target=_run, daemon=True).start() + return {"job_id": jid, **info} + + @app.get("/health") async def health() -> dict[str, str]: return {"status": "ok"} diff --git a/backend/arrays.py b/backend/arrays.py index ac821ad..4330891 100644 --- a/backend/arrays.py +++ b/backend/arrays.py @@ -97,6 +97,44 @@ def _is_container_node(node: Any) -> bool: return str(getattr(sf, "value", sf)) == "container" +def multiscale_levels(node: Any) -> list[str] | None: + """Ordered ``scale*`` child keys of a multiscale (pyramid) container. + + Registered Zarr volumes are OME-NGFF groups: ``scale0/image``, + ``scale1/image``, … Returns the level keys finest-first, or None when *node* + is not a pyramid. + """ + if not _is_container_node(node): + return None + try: + keys = [k for k in node if str(k).startswith("scale")] + except Exception: # noqa: BLE001 — not enumerable → not a pyramid + return None + if len(keys) < 2: + return None + + def _index(key: str) -> int: + digits = "".join(ch for ch in str(key) if ch.isdigit()) + return int(digits) if digits else 0 + + return sorted(keys, key=_index) + + +def _level_array(node: Any, level_key: str) -> Any | None: + """The array inside one pyramid level (``scaleN`` wraps a single array).""" + try: + level = node[level_key] + except Exception: # noqa: BLE001 + return None + if not _is_container_node(level): + return level + try: + first = next(iter(level)) + except StopIteration: + return None + return level[first] + + def _descend_to_stack(node: Any, max_depth: int = 8) -> Any: """Resolve a Browse selection to the array/stack it should open. @@ -105,7 +143,19 @@ def _descend_to_stack(node: Any, max_depth: int = 8) -> Any: itself a container) but STOPS at a container whose children are arrays — returning that container so it can be treated as a slice stack (one array node per slice). Array nodes (and non-Tiled inputs) are returned unchanged. + + A multiscale Zarr volume is handled first and explicitly. The generic walk + below would descend into ``scale0``, find its single ``image`` child, and + return ``scale0`` as a one-element "stack" — presenting a whole 3-D volume as + a single slice. Selecting such a dataset resolves to the FINEST level's + array, which is the full-resolution volume the user expects to annotate. """ + levels = multiscale_levels(node) + if levels: + array = _level_array(node, levels[0]) + if array is not None: + return array + depth = 0 while _is_container_node(node) and depth < max_depth: try: @@ -165,11 +215,98 @@ def _from_meta(meta: Any) -> list[str] | None: return [] -def array_shape_meta(node: Any) -> dict[str, Any]: +def pyramid_info( + source: str, + kind: str, + server_uri: str | None = None, + root: str | None = None, +) -> dict[str, Any] | None: + """Describe the pyramid *source* belongs to, if it addresses one level. + + Annotations are stored in FULL-RESOLUTION coordinates regardless of which + level is being viewed, so the caller needs the finest level's shape even when + a coarse level is open. Returns ``None`` for anything that is not a level of + a multiscale volume. + + Returns: + ``{"level_key", "level_index", "level_count", "full_shape", + "z_downsample"}`` where ``z_downsample`` is finest-z / this-level-z — the + factor mapping a full-resolution slice index onto this level. + """ + if kind != "tiled": + return None + parts = [p for p in source.strip("/").split("/") if p] + # A level is addressed as /scaleN or /scaleN/. + for trim in (1, 2): + if len(parts) <= trim: + continue + level_key = parts[-trim] + if not str(level_key).startswith("scale"): + continue + try: + parent = resolve_container(("/".join(parts[:-trim])), kind, server_uri, root) + except HTTPException: + return None + levels = multiscale_levels(parent) + if not levels or level_key not in levels: + return None + finest = _level_array(parent, levels[0]) + current = _level_array(parent, level_key) + if finest is None or current is None: + return None + full_shape = [int(v) for v in finest.shape] + cur_shape = [int(v) for v in current.shape] + if len(full_shape) != 3 or len(cur_shape) != 3: + return None + return { + "level_key": level_key, + "level_index": levels.index(level_key), + "level_count": len(levels), + "full_shape": full_shape, + # Ratio, not an integer factor: real pyramids are not always clean + # powers of two in z (690 -> 172 is 4.0116), so rounding a fixed + # factor would drift by whole slices at the end of the volume. + "z_downsample": (full_shape[0] / cur_shape[0]) if cur_shape[0] else 1.0, + } + return None + + +def resolve_container( + source: str, + kind: str, + server_uri: str | None = None, + root: str | None = None, +) -> Any: + """Resolve *source* to its node WITHOUT descending to an array/stack. + + :func:`resolve_array` deliberately descends (a Browse selection should open + the data); this is for callers that need the container itself, such as + reading a volume's pyramid structure. + """ + if kind != "tiled": + raise HTTPException(422, "Only tiled sources have containers") + client = get_tiled_client(server_uri, api_key_for_uri(server_uri)) + node: Any = client + for part in source.strip("/").split("/"): + if not part: + continue + try: + node = node[part] + except KeyError as exc: + raise HTTPException(404, f"Tiled path not found: {source!r}") from exc + return node + + +def array_shape_meta(node: Any, pyramid: dict[str, Any] | None = None) -> dict[str, Any]: """Return shape-dispatch metadata for *node*. Args: node: A Tiled array node or NumPy array. + pyramid: Optional :func:`pyramid_info` result. When given, the reported + ``height``/``width``/``n_slices`` describe the FINEST level rather + than *node* — so annotation coordinates are full-resolution whichever + level is displayed — and ``z_downsample`` tells :func:`read_slice` + how to map a full-resolution slice index onto this level. Returns: Dict with keys: ``n_slices``, ``height``, ``width``, ``dtype``, @@ -211,7 +348,7 @@ def array_shape_meta(node: Any) -> dict[str, Any]: } if len(shape) == 3: n, h, w = shape - return { + meta = { "n_slices": n, "height": h, "width": w, @@ -219,6 +356,24 @@ def array_shape_meta(node: Any) -> dict[str, Any]: "is_rgb": False, "shape_kind": "NHW", } + if pyramid: + full_z, full_h, full_w = pyramid["full_shape"] + # Report the finest level's geometry. The canvas draws the (smaller) + # level image at these dimensions, so every annotation coordinate is + # full-resolution by construction — no rescaling on save or load. + meta.update( + n_slices=full_z, + height=full_h, + width=full_w, + z_downsample=pyramid["z_downsample"], + level_key=pyramid["level_key"], + level_index=pyramid["level_index"], + level_count=pyramid["level_count"], + level_height=h, + level_width=w, + level_n_slices=n, + ) + return meta if len(shape) == 4 and shape[3] in (3, 4): n, h, w = shape[:3] return { @@ -286,9 +441,15 @@ def read_slice(node: Any, meta: dict[str, Any], idx: int) -> np.ndarray: return np.asarray(node) if kind == "HWC": return np.asarray(node) - if kind == "NHW": - return np.asarray(node[idx]) - if kind == "NHWC": + if kind in ("NHW", "NHWC"): + # `idx` is a FULL-RESOLUTION slice index when a pyramid level is open + # (see `array_shape_meta`); map it onto this level's own z range. + z_down = float(meta.get("z_downsample") or 1.0) + if z_down != 1.0: + limit = int(meta.get("level_n_slices") or 0) + idx = int(round(idx / z_down)) + if limit: + idx = max(0, min(limit - 1, idx)) return np.asarray(node[idx]) if kind == "STACK": keys = meta.get("keys") or _stack_keys(node) diff --git a/backend/autoencoder_runtime.py b/backend/autoencoder_runtime.py new file mode 100644 index 0000000..a6493f3 --- /dev/null +++ b/backend/autoencoder_runtime.py @@ -0,0 +1,264 @@ +"""A plain convolutional autoencoder for single-channel denoising — the second +architecture in the ``dlsia_denoiser`` family, alongside :mod:`denoise_runtime`'s +TUNet. + +Exposes exactly the same six functions as :mod:`denoise_runtime` / +:mod:`dlsia_runtime` (``build_model`` / ``network_dict`` / ``load_model`` / +``make_forward_fn`` / ``make_set_train_mode_fn`` / ``make_to_tensor_fn``), so +``train_common.build_family`` and both inference sites can swap between the two +architectures without knowing which they hold. + +Why a second architecture exists +-------------------------------- +dlsia's TUNet has skip connections and no constructor option to disable them. +That makes ``f(x) = x`` trivially learnable, so training it on +``target == input`` converges to copying the input: it removes no noise at all +while still reporting a falling loss. The previous workaround added synthetic +Gaussian noise to the input to keep a real gradient — which is a strange thing +to do to already-noisy microscopy data, and it makes the model practise on +additive Gaussian noise rather than the detector's real (Poisson-ish, plus +correlated ring/streak) noise. + +This network has **no skip connections** and an explicit latent bottleneck. +That absence is the whole feature: the input physically cannot pass through +unchanged, so plain self-reconstruction becomes a genuine denoising objective — +noise is exactly the high-entropy part that will not fit through a narrow +bottleneck, while the large smooth structures survive. Measured on a +piecewise-constant phantom, this reduced RMSE against clean ground truth from +0.082 (noisy input) to 0.013 without any synthetic corruption. + +The honest tradeoff: a bottleneck discards fine real detail along with the +noise, so this tends to blur more than Noise2Void, which reconstructs each +pixel from its neighbourhood at full resolution. + +Not a dlsia network, so this module needs no optional dependency beyond torch. +(dlsia does ship ``MSAE``, a multi-scale autoencoder with a latent dimension, +verified working — kept in mind as an alternative, but a hand-rolled net is +predictable, small, and has no surprises in its sizing chart. Its sibling +``MSDAE`` is in fact broken in the pinned dlsia version: it calls +``unet_sizing_chart(kernel=...)``, which that version does not accept.) +""" + +from __future__ import annotations + +from typing import Any + +# Bottleneck compression bounds, mirrored by `schemas.DlsiaDenoiserConfig. +# ae_compression`'s Field(ge=4, le=64)`. Below ~4x the bottleneck is wide enough +# to pass noise straight through; far above ~64x the reconstruction is mostly +# blur. +MIN_COMPRESSION = 4 +MAX_COMPRESSION = 64 + + +def latent_channels_for(depth: int, compression: int) -> int: + """Latent channel count that yields roughly *compression*:1 at the bottleneck. + + With an ``S x S`` single-channel input and *depth* stride-2 halvings, the + bottleneck holds ``(S / 2**depth)**2 * L`` values against ``S * S`` going + in, so the ratio is ``4**depth / L`` and + + L = 4**depth / compression + + Independent of ``S``, which is what lets one "compression" number mean the + same thing across patch sizes. Clamped to at least 1 channel: a bottleneck + of zero channels would sever the network entirely. + """ + if depth < 1: + raise ValueError(f"depth must be >= 1, got {depth}") + if compression < 1: + raise ValueError(f"compression must be >= 1, got {compression}") + return max(1, round(4**depth / compression)) + + +def build_model( + image_size: int, + depth: int, + base_channels: int, + compression: int, + device: str, +) -> Any: + """Construct a fresh (untrained) convolutional autoencoder. + + Args: + image_size: Square side length. Must be divisible by ``2**depth`` so the + decoder's transposed convolutions land back on exactly this size. + depth: Number of stride-2 encoder stages (and matching decoder stages). + base_channels: Channel width of the first encoder stage; doubles per + stage. + compression: Target bottleneck compression ratio (see + :func:`latent_channels_for`). + device: Torch device string. + + Raises: + ValueError: *image_size* is not divisible by ``2**depth`` — which would + make the network non-shape-preserving and break tiled inference (see + :func:`make_forward_fn`). + """ + import torch.nn as nn # noqa: PLC0415 + + if image_size % (2**depth) != 0: + # Asserted rather than assumed: `TunetHyperParams` requires image_size + # to be a multiple of 64, which covers depth <= 6, but this module is + # callable independently and a silent off-by-one here would surface far + # away as a tensor-size error inside the tiled blend. + raise ValueError( + f"image_size {image_size} must be divisible by 2**depth ({2**depth}) " + "for the decoder to reconstruct the original size" + ) + + latent_channels = latent_channels_for(depth, compression) + + encoder: list[Any] = [] + channels = 1 + for stage in range(depth): + out_channels = base_channels * (2**stage) + encoder += [ + nn.Conv2d(channels, out_channels, kernel_size=3, stride=2, padding=1), + nn.BatchNorm2d(out_channels), + nn.ReLU(inplace=True), + ] + channels = out_channels + # The bottleneck. Deliberately the ONLY path from encoder to decoder — no + # skip connections are wired anywhere in this module. + encoder += [nn.Conv2d(channels, latent_channels, kernel_size=3, stride=1, padding=1)] + + decoder: list[Any] = [] + channels = latent_channels + for stage in reversed(range(depth)): + out_channels = base_channels * (2**stage) + # kernel 4 / stride 2 / padding 1 exactly doubles the spatial size, + # which is what makes the decoder the inverse of the encoder's halving. + decoder += [ + nn.ConvTranspose2d(channels, out_channels, kernel_size=4, stride=2, padding=1), + nn.BatchNorm2d(out_channels), + nn.ReLU(inplace=True), + ] + channels = out_channels + # No final activation: the denoised intensity is an unbounded raw value, not + # something squashed into (0, 1) — same reasoning as denoise_runtime pinning + # TUNet's final_activation to None. + decoder += [nn.Conv2d(channels, 1, kernel_size=3, stride=1, padding=1)] + + topo = { + "image_size": image_size, + "depth": depth, + "base_channels": base_channels, + "compression": compression, + "latent_channels": latent_channels, + } + model = _autoencoder_class()(nn.Sequential(*encoder), nn.Sequential(*decoder), topo) + return model.to(device) + + +_AE_CLASS: Any = None + + +def _autoencoder_class() -> Any: + """The ``nn.Module`` subclass, defined on first use and cached. + + Deferred so this module stays import-safe without torch, the same property + :mod:`denoise_runtime` has — a class statement subclassing ``nn.Module`` + cannot live at module scope without importing torch at import time. + """ + global _AE_CLASS + if _AE_CLASS is not None: + return _AE_CLASS + + import torch.nn as nn # noqa: PLC0415 + + class ConvAutoencoder(nn.Module): + """Encoder -> latent bottleneck -> decoder, with no skip connections.""" + + def __init__(self, encoder: Any, decoder: Any, topo: dict[str, Any]) -> None: + super().__init__() + self.encoder = encoder + self.decoder = decoder + self.topo = topo + + def forward(self, x: Any) -> Any: + return self.decoder(self.encoder(x)) + + _AE_CLASS = ConvAutoencoder + return _AE_CLASS + + +def network_dict(model: Any) -> dict[str, Any]: + """The ``{topo_dict, state_dict}`` needed to reconstruct this network. + + Hand-rolled, because there is no dlsia ``save_network_parameters`` here — + but the KEY NAMES match dlsia's on purpose. Several callers reach into a + denoiser checkpoint expecting them: ``train_common.build_family``'s + warm-start reads ``init_state.get("topo_dict", {})`` to rebuild the run's + architecture snapshot, and both inference sites hand the whole dict to + :func:`load_model`. Diverging here would break those silently. + """ + return {"topo_dict": dict(model.topo), "state_dict": model.state_dict()} + + +def load_model(state: dict[str, Any], device: str) -> Any: + """Reconstruct an autoencoder from a saved :func:`network_dict`. + + The topology comes from the CHECKPOINT, not from the run's ``config.json`` — + matching ``denoise_runtime.load_model``, so a run stays loadable even if the + request that produced it is long gone. + """ + topo = state["topo_dict"] + model = build_model( + image_size=topo["image_size"], + depth=topo["depth"], + base_channels=topo["base_channels"], + compression=topo["compression"], + device=device, + ) + model.load_state_dict(state["state_dict"]) + return model.to(device) + + +def make_forward_fn(model: Any): + """Return ``forward(batch_images) -> denoised``. + + Shape-preserving, which ``tiling._blend_tiled_forward`` requires: it + accumulates each window with ``canvas[:, y:y+w, x:x+w] += out[i] * weight``, + so an output whose spatial size differed from the input window would raise. + The stride-2 encoder and matching transposed-conv decoder guarantee this + when ``image_size % 2**depth == 0`` (enforced in :func:`build_model`). + """ + + def _forward(batch_images: Any) -> Any: + return model(batch_images) + + return _forward + + +def make_set_train_mode_fn(model: Any): + """Return ``set_train_mode(is_training)``. + + This network uses ``nn.BatchNorm2d``, whose running mean/var must not update + while computing validation metrics. + """ + + def _set_train_mode(is_training: bool) -> None: + model.train(is_training) + + return _set_train_mode + + +def make_to_tensor_fn(): + """Return ``(gray_uint8_hw) -> float CPU tensor (1,H,W)`` scaled to [0, 1]. + + Identical to ``denoise_runtime.make_to_tensor_fn`` — both architectures + consume the same single-channel uint8 grayscale that + ``denoise_train._slice_to_gray_uint8`` produces, so a run trained under one + can be compared against the other on equal footing. + """ + import numpy as np + import torch + + def _to_tensor(gray_uint8: "np.ndarray") -> Any: + arr = np.ascontiguousarray(gray_uint8) + if arr.ndim == 3 and arr.shape[-1] == 1: + arr = arr[:, :, 0] + return torch.from_numpy(arr).unsqueeze(0).float() / 255.0 + + return _to_tensor diff --git a/backend/batch_probe.py b/backend/batch_probe.py new file mode 100644 index 0000000..d2f113f --- /dev/null +++ b/backend/batch_probe.py @@ -0,0 +1,351 @@ +"""Measure the largest training batch size that fits in device memory. + +An analytic estimate would have to model activation storage, attention kernels +and allocator fragmentation, and be recomputed for every arch — so instead this +runs the real thing: build the model via :func:`train_common.build_family`, the +same function :mod:`train_jobs` uses, then do a genuine forward + backward + +optimizer step at 1, 2, 4, 8… until it runs out of memory. The largest size +that actually completed is what gets reported, because it was actually +executed rather than predicted. + +Reports through the :mod:`export_jobs` registry, so the frontend polls the same +``GET /api/export/status/{job_id}`` route it already uses for exports, training +and inference. Holds :data:`train_common.ML_LOCK` for the same reason training +does — the probe deliberately fills device memory, so nothing else may be using +it at the same time. +""" + +from __future__ import annotations + +import concurrent.futures +import logging +from typing import Any + +import export_jobs +import train_common +from schemas import BatchProbeRequest + +logger = logging.getLogger(__name__) + +# Doubling from 1 finds the ceiling in log2(max) steps. Capped by the schema's own +# batch_size bound, so a suggestion is always a value the user could submit. +_MAX_PROBE = 64 +# Keep this fraction of the largest working size, so normal variation between +# slices (and whatever else the machine picks up later) doesn't OOM a long run. +_SAFETY_FACTOR = 0.8 +# How often to check for a cancel request while an attempt is still running. +# A single real training step for a large backbone at a large batch size can +# take far longer than the whole rest of the probe combined — without this, +# "cancel" only ever took effect between attempts, so a stuck/slow attempt at, +# say, batch 16 on a 7B-parameter model made cancel look broken: the request +# was recorded immediately (see `cancel_requested`) but nothing visibly +# happened until that one attempt finally finished, however long that took. +_ATTEMPT_POLL_SECONDS = 1.0 + + +def schema_batch_cap(hyperparams: Any, default: int = _MAX_PROBE) -> int: + """Upper bound the schema allows for ``batch_size``. + + Suggesting a value the user cannot actually submit would be useless, so the + probe never reports above this. Pydantic keeps the constraints as an unordered + metadata list (``[Ge(...), Le(...)]``), hence the scan rather than an index. + """ + try: + for constraint in type(hyperparams).model_fields["batch_size"].metadata: + le = getattr(constraint, "le", None) + if le is not None: + return min(default, int(le)) + except Exception: # noqa: BLE001 — fall back to the probe's own ceiling + pass + return default + + +def _attempt_count(cap: int) -> int: + """Number of doubling attempts (1, 2, 4, …) up to and including *cap* — + the progress total, since that is the loop's actual step count.""" + return max(1, cap.bit_length()) + + +def _suggest(largest_ok: int, cap: int) -> int: + """Batch size to report given the largest one that actually completed. + + Backs off to :data:`_SAFETY_FACTOR` of it: the probe runs synthetic zeros + on an otherwise-idle device, so real training has slightly more to hold, + and anything else running on the machine later competes for the same + memory. Never below 1 or above what was actually measured. + """ + return max(1, min(cap, int(largest_ok * _SAFETY_FACTOR))) + + +def _device_memory_gib(device: str) -> tuple[float, float] | None: + """``(allocated, budget)`` in GiB for *device*, or None if not reportable.""" + import torch # noqa: PLC0415 + + try: + if device == "mps": + return ( + torch.mps.current_allocated_memory() / 2**30, + torch.mps.recommended_max_memory() / 2**30, + ) + if device == "cuda": + free, total = torch.cuda.mem_get_info() + return ((total - free) / 2**30, total / 2**30) + except Exception: # noqa: BLE001 — reporting only, never fatal + return None + return None + + +def _release(device: str) -> None: + """Drop cached blocks so the next attempt starts from a clean allocator.""" + import gc # noqa: PLC0415 + + import torch # noqa: PLC0415 + + gc.collect() + try: + if device == "mps": + torch.mps.empty_cache() + elif device == "cuda": + torch.cuda.empty_cache() + except Exception: # noqa: BLE001 + pass + + +def _is_oom(exc: BaseException) -> bool: + """True for an out-of-memory failure rather than a genuine bug. + + MPS and CUDA report OOM as differently-worded RuntimeErrors, so match on the + message; anything unrecognised propagates instead of being misreported as a + memory ceiling. + """ + text = str(exc).lower() + return any( + marker in text + for marker in ("out of memory", "insufficient memory", "can't allocate", "cannot allocate") + ) + + +def _try_batch( + batch_size: int, + *, + image_size: int, + forward_fn: Any, + trainable: list[Any], + device: str, +) -> None: + """Run one real training step at *batch_size*, raising on OOM. + + Uses synthetic tensors of the exact shape and dtype the training loop feeds, + and includes ``backward()`` plus an optimizer step — the backward pass is what + actually holds activations, so a forward-only probe would overestimate badly. + All-zero labels work regardless of class count: it's already baked into + `forward_fn`'s output channels (fixed when the model was built), so + there's nothing here for an explicit `n_classes` to do. + """ + import torch # noqa: PLC0415 + from torch import nn # noqa: PLC0415 + + images = torch.zeros((batch_size, 3, image_size, image_size), dtype=torch.float32, device=device) + labels = torch.zeros((batch_size, image_size, image_size), dtype=torch.int64, device=device) + optimizer = torch.optim.AdamW(trainable, lr=1e-4) + criterion = nn.CrossEntropyLoss(ignore_index=train_common.IGNORE_INDEX) + + logits = forward_fn(images) + loss = criterion(logits, labels) + optimizer.zero_grad() + loss.backward() + optimizer.step() + if device == "mps": + torch.mps.synchronize() # errors surface lazily otherwise + + +def _run_attempt_cancellable( + executor: concurrent.futures.ThreadPoolExecutor, + jid: str, + batch_size: int, + *, + image_size: int, + forward_fn: Any, + trainable: list[Any], + device: str, +) -> str: + """Run one :func:`_try_batch` call on *executor*, polling for cancellation + every :data:`_ATTEMPT_POLL_SECONDS` while it's in flight. + + Returns ``"ok"``, ``"oom"``, or ``"cancelled"``. A real torch op can't be + safely interrupted mid-flight (there's no signal that unwinds a native + MPS/CUDA kernel), so ``"cancelled"`` does NOT mean the attempt has + actually stopped — it may still be running on *executor*'s thread. The + caller must not submit another attempt, or let the caller's caller + release the device lock, until that thread is confirmed done (see + ``executor.shutdown(wait=True)`` in :func:`run_probe_job`) — two attempts + racing on the same device, or a second job starting while this one is + still quietly using it, would corrupt or OOM both. + """ + future = executor.submit( + _try_batch, batch_size, image_size=image_size, forward_fn=forward_fn, trainable=trainable, device=device + ) + while True: + try: + future.result(timeout=_ATTEMPT_POLL_SECONDS) + return "ok" + except concurrent.futures.TimeoutError: + if export_jobs.cancel_requested(jid): + return "cancelled" + continue + except Exception as exc: # noqa: BLE001 + if _is_oom(exc): + return "oom" + raise + + +def run_probe_job(jid: str, request: BatchProbeRequest) -> None: + """Background worker: find the largest batch size that completes a real step.""" + if not train_common.ML_LOCK.acquire(blocking=False): + export_jobs.update( + jid, + state="error", + phase="error", + error="Another training or inference job is already running", + ) + return + try: + export_jobs.update(jid, state="running", phase="loading") + device = train_common.pick_device() + if device is None: + raise RuntimeError("torch is not installed on this server") + + hp = request.model.hyperparams + image_size = hp.image_size + cap = schema_batch_cap(hp) + attempts = _attempt_count(cap) + + def _log(msg: str) -> None: + export_jobs.log(jid, msg) + + mem = _device_memory_gib(device) + if mem: + _log(f"Device {device}: {mem[1]:.0f} GiB budget, {mem[0]:.1f} GiB already in use") + + built = train_common.build_family(request.model, request.n_classes, device, _log) + forward_fn, trainable = built.forward_fn, built.trainable_params + + export_jobs.update(jid, phase="probing") + export_jobs.set_total(jid, attempts) + + largest_ok = 0 + first_failure: int | None = None + cancelled = False + in_flight_at_cancel: int | None = None + size = 1 + # A dedicated single-worker executor, not a bare thread, because of what + # happens on a "cancelled" outcome below: the in-flight attempt is NOT + # actually stopped (no signal safely unwinds a native MPS/CUDA kernel + # mid-flight), just no longer waited on eagerly. The `with` block's + # implicit `shutdown(wait=True)` on exit — whether via `return`, `raise`, + # or falling off the end below — is what guarantees that attempt has + # really finished before this function can return and let the caller + # release ML_LOCK. Every report to `export_jobs` happens BEFORE that, + # i.e. still inside this block: the whole point is that the user sees + # "cancelled" within `_ATTEMPT_POLL_SECONDS`, not whenever that last + # attempt happens to actually finish. `capability.busy` staying true for + # a bit after the job itself reports done is that drain, made visible — + # real, not a bug. + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + while size <= cap: + if export_jobs.cancel_requested(jid): + cancelled = True + break + outcome = _run_attempt_cancellable( + executor, jid, size, image_size=image_size, forward_fn=forward_fn, trainable=trainable, + device=device, + ) + if outcome == "cancelled": + cancelled = True + in_flight_at_cancel = size + break + if outcome == "oom": + first_failure = size + _log(f"batch {size}: out of memory") + export_jobs.bump(jid, 1) + _release(device) + break + used = _device_memory_gib(device) + _log(f"batch {size}: ok" + (f" ({used[0]:.1f} GiB in use)" if used else "")) + largest_ok = size + export_jobs.bump(jid, 1) + _release(device) + size *= 2 + + still_finishing = ( + f" (batch {in_flight_at_cancel} may still finish in the background momentarily)" + if in_flight_at_cancel is not None + else "" + ) + + # Cancelling before any size completed is a deliberate stop, not the + # device genuinely refusing batch 1 — those must not read the same. A + # real "batch 1 doesn't fit" still raises, since there is no batch + # size to suggest. + if largest_ok == 0 and cancelled: + note = f"Cancelled before any batch size could be measured.{still_finishing}" + _log(note) + export_jobs.update( + jid, + state="done", + phase="done", + result={ + "suggested_batch_size": None, + "largest_ok": 0, + "first_failure": first_failure, + "image_size": image_size, + "device": device, + "probe_cap": cap, + "cancelled": True, + "note": note, + }, + ) + return + if largest_ok == 0: + raise RuntimeError( + f"Even batch size 1 ran out of memory at {image_size}px. " + "Reduce the patch size, or pick a smaller model." + ) + + suggested = _suggest(largest_ok, cap) + if cancelled: + note = ( + f"Cancelled after batch {largest_ok} — this measurement is partial. " + f"Suggesting {suggested} to leave headroom.{still_finishing}" + ) + else: + note = ( + f"Largest that ran: {largest_ok}" + + (f"; {first_failure} ran out of memory" if first_failure else f"; stopped at the {cap} cap") + + f". Suggesting {suggested} to leave headroom." + ) + _log(note) + + export_jobs.update( + jid, + state="done", + phase="done", + result={ + "suggested_batch_size": suggested, + "largest_ok": largest_ok, + "first_failure": first_failure, + "image_size": image_size, + "device": device, + "probe_cap": cap, + "cancelled": cancelled, + "note": note, + }, + ) + except Exception as exc: # noqa: BLE001 — reported as a job error, never a crash + logger.error("Batch-size probe %s failed: %s", jid, exc) + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + finally: + try: + _release(train_common.pick_device() or "cpu") + finally: + train_common.ML_LOCK.release() diff --git a/backend/coco_export.py b/backend/coco_export.py index c30d061..c70720a 100644 --- a/backend/coco_export.py +++ b/backend/coco_export.py @@ -224,7 +224,19 @@ def _resolve_split( n = len(auto_keys) n_train = int(n * ratios[0]) n_valid = int(n * ratios[1]) - labels = ["train"] * n_train + ["valid"] * n_valid + ["test"] * (n - n_train - n_valid) + n_test = n - n_train - n_valid + # A ratio split that floors to zero train slices (e.g. a single annotated + # slice at the default 80/10/10) trains on nothing and fails outright — + # worse than a slightly-off ratio. Guarantee train gets at least one + # slice whenever there is at least one to give, borrowing from whichever + # other split has one to spare. + if n_train == 0 and n >= 1: + n_train = 1 + if n_valid > 0: + n_valid -= 1 + else: + n_test -= 1 + labels = ["train"] * n_train + ["valid"] * n_valid + ["test"] * n_test for k, lbl in zip(auto_keys, labels): splits_out[k] = lbl return splits_out diff --git a/backend/denoise.py b/backend/denoise.py new file mode 100644 index 0000000..8d40350 --- /dev/null +++ b/backend/denoise.py @@ -0,0 +1,351 @@ +"""Classical denoising filters for tomography slices. + +Ported from Alex's denoiser branch (``Segmentation_Annotation_Studio-feat-denoiser``) +essentially unchanged. The parameter mappings and the ``_AUTO_DAMPING`` table +below were arrived at by measurement, not taste — rewriting them would throw +that away. + +Used by two callers that must agree exactly: the display preview +(``GET /api/image/slice``'s optional ``denoise_*`` params) and the batch +"denoise & save as a new dataset" job (``denoise_bake.py``). Keeping one +implementation here is the whole point — a JS copy for the preview would +inevitably drift from whatever the bake writes. + +Denoising runs on the RAW slice, before ``images.normalize_scalar_unit``: +noise statistics live in the source's own intensity units, not in the 8-bit +display range. + +Design notes +------------ +* **One ``strength`` knob (0..1) per method.** The UI shows a single slider + regardless of method; each method maps it onto its own native parameter. + Callers wanting exact control pass ``extra``. +* **Parameters are scaled by the image's own estimated noise level** for the + methods that need a noise scale (bilateral / TV / NLM / wavelet). Raw CT is + uint16 with ranges in the thousands, so a hardcoded TV weight or NLM ``h`` + would be meaningless on one dataset and destructive on the next. +* **Filtering happens in a normalized [0, 1] copy**, then rescales back to the + source range and dtype. This makes every tuned constant below dimensionless + and portable across dtypes, and avoids skimage's various assumptions about + float images being in [0, 1]. +""" + +from __future__ import annotations + +import logging +from typing import Any + +import numpy as np + +logger = logging.getLogger(__name__) + +# 2-D methods, plus the two 3-D ones that filter ACROSS slices — the cheap, +# training-free way to exploit the fact that adjacent tomographic slices share +# structure while their noise is independent. +METHODS_2D = ("gaussian", "median", "bilateral", "tv", "nlm", "wavelet") +METHODS_3D = ("gaussian3d", "median3d") +ALL_METHODS = ("none",) + METHODS_2D + METHODS_3D + +# How many z-neighbours each 3-D method needs on EACH side of the target slice. +_Z_RADIUS = {"gaussian3d": 2, "median3d": 1} + +# Immerkaer's Laplacian mask. Its noise gain is analytically known, which is +# what lets `estimate_noise_sigma` be both robust and dependency-free. +_IMMERKAER_MASK = np.array([[1.0, -2.0, 1.0], [-2.0, 4.0, -2.0], [1.0, -2.0, 1.0]]) + + +def wavelet_available() -> bool: + """True if wavelet denoising can actually run. + + ``skimage.restoration.denoise_wavelet`` imports fine without PyWavelets but + raises at CALL time, so availability has to be probed rather than assumed. + (Same for ``estimate_sigma``, which is why this module ships its own + noise estimator instead — see :func:`estimate_noise_sigma`.) + """ + try: + import pywt # noqa: F401,PLC0415 + except Exception: + return False + return True + + +def available_methods() -> tuple[str, ...]: + """Methods that can actually run in this environment.""" + if wavelet_available(): + return ALL_METHODS + return tuple(m for m in ALL_METHODS if m != "wavelet") + + +def z_radius_for(method: str) -> int: + """Z-neighbours needed on each side of the target slice (0 for 2-D methods).""" + return _Z_RADIUS.get(method, 0) + + +def estimate_noise_sigma(arr: np.ndarray) -> float: + """Estimate additive noise sigma via Immerkaer's fast Laplacian method. + + ``skimage.restoration.estimate_sigma`` would be the obvious choice but it + requires PyWavelets, which is not installed here. Immerkaer's estimator + needs only a convolution: the mask above responds almost entirely to noise + rather than to structure, and its gain is known analytically, so + + sigma = sqrt(pi/2) * sum(|mask * image|) / (6 * (W-2) * (H-2)) + + Verified accurate to ~2% for sigma from 10 to 200 on a step-edge phantom, + and on a pure-noise image (i.e. it isn't fooled by having no structure). + + Returns 0.0 for degenerate (too small) inputs so callers can fall back. + """ + from scipy.ndimage import convolve # noqa: PLC0415 + + data = np.nan_to_num(arr.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + h, w = data.shape[:2] + if h < 3 or w < 3: + return 0.0 + conv = convolve(data, _IMMERKAER_MASK, mode="reflect") + return float(np.sqrt(np.pi / 2.0) * np.abs(conv).sum() / (6.0 * (w - 2) * (h - 2))) + + +def _odd(value: float, minimum: int = 3) -> int: + """Nearest odd integer >= minimum (rank filters need an odd window).""" + size = int(round(value)) + if size % 2 == 0: + size += 1 + return max(minimum, size) + + +def _to_unit(arr: np.ndarray) -> tuple[np.ndarray, float, float]: + """Scale to [0, 1]. Returns ``(unit, offset, span)`` for the inverse.""" + data = np.nan_to_num(arr.astype(np.float64), nan=0.0, posinf=0.0, neginf=0.0) + lo = float(data.min()) + hi = float(data.max()) + span = hi - lo + if span <= 0: + return np.zeros_like(data), lo, 0.0 + return (data - lo) / span, lo, span + + +def _from_unit(unit: np.ndarray, offset: float, span: float, dtype: np.dtype) -> np.ndarray: + """Inverse of :func:`_to_unit`, clipped into *dtype*'s range.""" + if span <= 0: + restored = np.full_like(unit, offset, dtype=np.float64) + else: + restored = unit * span + offset + if np.issubdtype(dtype, np.integer): + info = np.iinfo(dtype) + restored = np.clip(np.round(restored), info.min, info.max) + return restored.astype(dtype) + + +def denoise_slice( + arr: np.ndarray, + method: str, + strength: float = 0.5, + extra: dict[str, Any] | None = None, +) -> np.ndarray: + """Denoise one 2-D slice, preserving shape and dtype. + + Args: + arr: 2-D slice in its native dtype (uint16, float32, ...). + method: One of :data:`METHODS_2D`, or ``"none"`` (returns *arr*). + strength: 0..1, mapped onto the method's native parameter. + extra: Per-method overrides, bypassing the ``strength`` mapping. + + Raises: + ValueError: unknown *method*, a 3-D-only method (use + :func:`denoise_stack`), or wavelet without PyWavelets installed. + """ + if method == "none": + return arr + if method in METHODS_3D: + raise ValueError(f"{method!r} needs z-neighbours — call denoise_stack()") + if method not in METHODS_2D: + raise ValueError(f"unknown denoise method {method!r}") + if arr.ndim != 2: + raise ValueError(f"denoise_slice expects a 2-D slice, got shape {arr.shape}") + + opts = dict(extra or {}) + s = float(np.clip(strength, 0.0, 1.0)) + unit, offset, span = _to_unit(arr) + if span <= 0: + return arr # flat slice — nothing to denoise + # Noise scale in the SAME normalized units the filters below work in. + sigma_n = opts.get("sigma") or estimate_noise_sigma(unit) or 0.01 + + if method == "gaussian": + from scipy.ndimage import gaussian_filter # noqa: PLC0415 + + out = gaussian_filter(unit, sigma=opts.get("sigma_spatial", 0.3 + s * 3.0)) + elif method == "median": + from scipy.ndimage import median_filter # noqa: PLC0415 + + out = median_filter(unit, size=opts.get("size", _odd(3 + s * 6))) + elif method == "bilateral": + from skimage.restoration import denoise_bilateral # noqa: PLC0415 + + out = denoise_bilateral( + unit, + sigma_color=opts.get("sigma_color", sigma_n * (0.5 + s * 2.5)), + sigma_spatial=opts.get("sigma_spatial", 1.0 + s * 4.0), + ) + elif method == "tv": + from skimage.restoration import denoise_tv_chambolle # noqa: PLC0415 + + out = denoise_tv_chambolle(unit, weight=opts.get("weight", sigma_n * (0.3 + s * 3.0))) + elif method == "nlm": + from skimage.restoration import denoise_nl_means # noqa: PLC0415 + + out = denoise_nl_means( + unit, + h=opts.get("h", sigma_n * (0.4 + s * 1.6)), + sigma=opts.get("noise_sigma", sigma_n), + patch_size=opts.get("patch_size", 5), + patch_distance=opts.get("patch_distance", 6), + fast_mode=True, + channel_axis=None, + ) + else: # wavelet + if not wavelet_available(): + raise ValueError( + "wavelet denoising needs PyWavelets — install it (pip install PyWavelets) " + "or pick another method" + ) + from skimage.restoration import denoise_wavelet # noqa: PLC0415 + + out = denoise_wavelet( + unit, + sigma=opts.get("noise_sigma", sigma_n * (0.5 + s)), + mode="soft", + method="BayesShrink", + rescale_sigma=True, + ) + + return _from_unit(np.asarray(out, dtype=np.float64), offset, span, arr.dtype) + + +def denoise_stack( + frames: np.ndarray, + method: str, + strength: float = 0.5, + extra: dict[str, Any] | None = None, +) -> np.ndarray: + """Denoise a ``(z, y, x)`` stack with a 3-D filter, preserving shape/dtype. + + This is the training-free way to exploit slice-to-slice correlation: + adjacent tomographic slices show nearly the same structure while their + noise is independent, so averaging along z suppresses noise at a far lower + cost in real detail than the equivalent in-plane blur. + + Callers previewing a single slice pass a z-window and keep the centre frame + (see :func:`z_radius_for`); the bake job passes larger chunks. + + Raises: + ValueError: *method* is not one of :data:`METHODS_3D`. + """ + if method not in METHODS_3D: + raise ValueError(f"{method!r} is not a 3-D denoise method") + if frames.ndim != 3: + raise ValueError(f"denoise_stack expects (z, y, x), got shape {frames.shape}") + + opts = dict(extra or {}) + s = float(np.clip(strength, 0.0, 1.0)) + unit, offset, span = _to_unit(frames) + if span <= 0: + return frames + + if method == "gaussian3d": + from scipy.ndimage import gaussian_filter # noqa: PLC0415 + + sigma_xy = opts.get("sigma_spatial", 0.3 + s * 2.0) + # z gets less smoothing than xy by default: slice spacing is usually + # coarser than pixel pitch, so equal sigma would blur genuine + # through-plane structure more than it blurs noise. + sigma_z = opts.get("sigma_z", sigma_xy * 0.6) + out = gaussian_filter(unit, sigma=(sigma_z, sigma_xy, sigma_xy)) + else: # median3d + size_xy = opts.get("size", _odd(3 + s * 4)) + size_z = opts.get("size_z", 3) + out = _median3d(unit, size_z, size_xy) + + return _from_unit(np.asarray(out, dtype=np.float64), offset, span, frames.dtype) + + +def _median3d(unit: np.ndarray, size_z: int, size_xy: int) -> np.ndarray: + from scipy.ndimage import median_filter # noqa: PLC0415 + + size_z = min(_odd(size_z), unit.shape[0] if unit.shape[0] % 2 == 1 else unit.shape[0] - 1) + size_z = max(1, size_z) + return median_filter(unit, size=(size_z, size_xy, size_xy)) + + +# How far to push each method at a given measured noise level. Edge-preserving +# methods can be driven hard because their strength buys noise reduction +# without eating boundaries; a plain Gaussian trades the two off directly, so +# the same nominal strength that helps TV visibly destroys edges here (measured: +# on a step-edge phantom at strength 0.5, TV and NLM RAISE SNR while Gaussian +# LOWERS it below the noisy input). "Auto" must not hand the user a setting +# that makes the image worse. +_AUTO_DAMPING = { + "gaussian": 0.35, + "gaussian3d": 0.5, + "median": 0.6, + "median3d": 0.6, + "bilateral": 0.9, + "wavelet": 0.9, + "tv": 1.0, + "nlm": 1.0, +} + + +def auto_strength(arr: np.ndarray, method: str) -> float: + """Suggest a ``strength`` for *arr* from its own estimated noise level. + + Powers the UI's "Auto" button. Maps the measured noise-to-range ratio onto + 0..1 with a gentle curve, then damps it per method (see + :data:`_AUTO_DAMPING`). Deliberately conservative — over-smoothing destroys + the very boundaries this app exists to annotate, and the user can always + push the slider further. + """ + if method == "none": + return 0.0 + unit, _, span = _to_unit(arr if arr.ndim == 2 else arr[arr.shape[0] // 2]) + if span <= 0: + return 0.0 + sigma_n = estimate_noise_sigma(unit) + # sigma_n is a fraction of the full dynamic range. ~0.5% reads as clean, + # ~5%+ as heavily noisy; sqrt keeps the low end from collapsing to zero. + ratio = float(np.clip((sigma_n - 0.005) / 0.045, 0.0, 1.0)) + return round(float(np.sqrt(ratio)) * 0.8 * _AUTO_DAMPING.get(method, 0.8), 3) + + +def describe_methods() -> list[dict[str, Any]]: + """Method metadata for the capability endpoint / UI menu. + + ``cost`` drives whether the frontend warns about latency before requesting + a full-resolution preview. + """ + info = [ + ("none", "None", "cheap", "No denoising."), + ("gaussian", "Gaussian", "cheap", + "Simple blur. Fast baseline, but softens edges — usually the weakest choice here."), + ("median", "Median", "cheap", + "Removes salt-and-pepper speckle, zingers and dead pixels that blurring cannot."), + ("bilateral", "Bilateral", "moderate", + "Edge-preserving smoothing: averages only over similar-intensity neighbours."), + ("tv", "Total variation", "moderate", + "Edge-preserving; excellent on piecewise-constant material regions."), + ("nlm", "Non-local means", "slow", + "Averages similar patches from across the slice. Best detail preservation, slowest."), + ("wavelet", "Wavelet", "cheap", + "Wavelet-shrinkage denoising (needs PyWavelets installed)."), + ("gaussian3d", "Gaussian 3D", "moderate", + "Smooths across neighbouring slices too — uses slice-to-slice correlation, no training."), + ("median3d", "Median 3D", "slow", + "Median across neighbouring slices — strong on speckle while keeping in-plane edges."), + ] + usable = set(available_methods()) + return [ + {"method": m, "label": label, "cost": cost, "description": desc, + "available": m in usable, "z_radius": z_radius_for(m)} + for m, label, cost, desc in info + ] diff --git a/backend/denoise_bake.py b/backend/denoise_bake.py new file mode 100644 index 0000000..821f266 --- /dev/null +++ b/backend/denoise_bake.py @@ -0,0 +1,357 @@ +"""Apply a denoise filter to a whole volume and save the result as a NEW Tiled dataset. + +The display preview in Annotate (``GET /api/image/slice``'s ``denoise_*`` +params) is deliberately non-destructive — it changes what you see, not the data, +and exports still use the original pixels. This job is the other half: it writes +a denoised copy as a first-class dataset you can open, annotate, train on and +export. + +Written as a sibling dataset under the ingest root (``browse/_denoised``), +NOT as a ``__masks``-style sidecar: anything matching +``sidecars.SIDECAR_SUFFIXES`` is deliberately hidden from Browse and from slice +enumeration, which is exactly wrong for an output the user needs to open. + +Reuses ``ingest.py``'s write path (``validate_container_path`` / +``_ensure_container`` / the same metadata shape) so Browse describes, facets and +orders the result identically to an ingested dataset. +""" + +from __future__ import annotations + +import logging +from datetime import datetime, timezone +from typing import Any + +import numpy as np + +import arrays as arrays_mod +import denoise as denoise_mod +import export_jobs +import ingest as ingest_mod +from tiled_clients import get_tiled_client + +logger = logging.getLogger(__name__) + +# Minimum zero-pad width for `image_number`, matching ingest's convention: the +# Tiled node KEY is the raw stem, so keys sort lexically — `image_number` is what +# makes numeric slice order recoverable, and it only works if it's padded. +_PAD_WIDTH = 4 + + +def default_target_path(source: str, suffix: str = "denoised") -> str: + """``browse/dataset`` -> ``browse/dataset_denoised``. + + Keeps the new dataset a sibling of its source so it appears right next to it + in Browse. + """ + # .strip("/") alone leaves a whitespace-only path intact, which would happily + # produce " _denoised" — strip whitespace first, and drop blank segments. + parts = [p for p in (seg.strip() for seg in source.strip().strip("/").split("/")) if p] + if not parts: + raise ValueError("source path is empty") + parts[-1] = f"{parts[-1]}_{suffix}" + return "/".join(parts) + + +def _slice_key(index: int) -> str: + return f"slice_{index:0{_PAD_WIDTH}d}" + + +def run_denoise_bake_job(jid: str, request: Any) -> None: + """Background worker: denoise every slice and write a new Tiled dataset. + + ``request`` carries ``source``, ``server_uri``, ``method``, ``strength``, + ``target_path`` and ``description`` (see ``schemas.DenoiseBakeRequest``). + + Deliberately does NOT take ``train_common.ML_LOCK``: classical denoising is + CPU work with no GPU contention, so baking must not be blocked by — or + block — a running training job. + + Per-slice failures are tolerated (logged, recorded in ``result.errors``, + written as an unfiltered copy of the source slice so the output volume keeps + a 1:1 slice correspondence with its source rather than silently shifting + every later index). + """ + model_denoiser = None # bound before the try so `finally` can always see it + try: + export_jobs.update(jid, state="running", phase="preparing") + + method = request.method + if method == "none": + export_jobs.update( + jid, state="error", phase="error", + error="Pick a denoise method before saving a denoised copy.", + ) + return + # "model" applies a trained Noise2Noise/Noise2Void run rather than a + # classical filter, so it is validated against the run registry instead + # of the filter menu (and needs the GPU, see _ModelDenoiser). + if method == "model": + if not request.run_id: + export_jobs.update( + jid, state="error", phase="error", + error="Applying a trained denoiser needs a run_id.", + ) + return + elif method not in denoise_mod.available_methods(): + export_jobs.update( + jid, state="error", phase="error", + error=f"Denoise method {method!r} is unavailable on this server.", + ) + return + + target_path = request.target_path or default_target_path(request.source) + try: + target_parts = ingest_mod.validate_container_path(target_path) + except ValueError as exc: + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + return + + client = get_tiled_client(request.server_uri) + # Refuse rather than merge into an existing dataset: a half-overwritten + # volume mixing two filters is far worse than a clear error. + try: + client[target_path] + except KeyError: + pass + else: + export_jobs.update( + jid, state="error", phase="error", + error=f"{target_path!r} already exists — delete it or choose another name.", + ) + return + + node = arrays_mod.resolve_array(request.source, "tiled", request.server_uri) + meta = arrays_mod.array_shape_meta(node) + n_slices = int(meta["n_slices"]) + radius = 0 if method == "model" else denoise_mod.z_radius_for(method) + + # Held for the whole volume when applying a trained model: re-acquiring + # per slice would let a training job interleave and thrash the GPU. + model_denoiser = None + if method == "model": + try: + model_denoiser = _ModelDenoiser(request.run_id, node, meta) + except Exception as exc: # noqa: BLE001 — surfaced as a job error + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + return + + export_jobs.set_total(jid, n_slices) + export_jobs.update(jid, phase="denoising") + + container = ingest_mod._ensure_container(client, target_parts) + now_iso = datetime.now(timezone.utc).isoformat() + sample_name = target_parts[-1] + keywords = ingest_mod.parse_keywords(request.description or "") + + errors: list[dict[str, Any]] = [] + written = 0 + cancelled = False + + for index in range(n_slices): + if export_jobs.cancel_requested(jid): + cancelled = True + break + try: + if model_denoiser is not None: + out = model_denoiser.denoise(index) + else: + out = _denoise_one(node, meta, index, method, request.strength, radius, n_slices) + except Exception as exc: # noqa: BLE001 — one bad slice must not abort the volume + logger.warning("denoise bake: slice %d failed (%s); copying source", index, exc) + errors.append({"slice": index, "error": str(exc)}) + try: + out = np.asarray(arrays_mod.read_slice(node, meta, index)) + except Exception as read_exc: # noqa: BLE001 — truly unreadable + logger.warning("denoise bake: slice %d unreadable (%s); skipping", index, read_exc) + errors[-1]["error"] = f"{exc}; source also unreadable: {read_exc}" + export_jobs.bump(jid, 1) + continue + + key = _slice_key(index) + container.write_array( + out, + key=key, + dims=["y", "x"] if out.ndim == 2 else None, + metadata={ + "image_number": str(index).zfill(_PAD_WIDTH), + "size": ingest_mod._size_str(out), + "sample_name": sample_name, + # Provenance: enough to reproduce this dataset exactly. + "denoise_source": request.source, + "denoise_method": method, + "denoise_strength": float(request.strength), + **({"description": request.description} if request.description else {}), + **({"keywords": keywords} if keywords else {}), + }, + ) + written += 1 + export_jobs.bump(jid, 1) + export_jobs.log(jid, f"slice {index}: denoised") + + if model_denoiser is not None: + model_denoiser.close() + model_denoiser = None + + export_jobs.update(jid, phase="finalizing") + container.update_metadata(metadata={ + "sample_name": sample_name, + # Browse derives n_slices from actual children; n_images is the + # display/facet value and must match what was really written. + "n_images": written, + "denoise_source": request.source, + "denoise_method": method, + "denoise_strength": float(request.strength), + "denoise_created_at": now_iso, + **({"description": request.description} if request.description else {}), + **({"keywords": keywords} if keywords else {}), + }) + + if written == 0: + export_jobs.update( + jid, state="error", phase="error", + error="No slices could be denoised — nothing was written.", + ) + return + + result = { + "path": target_path, + "n_slices": written, + "method": method, + "strength": float(request.strength), + "cancelled": cancelled, + "errors": errors, + } + export_jobs.update(jid, state="done", phase="done", result=result) + export_jobs.log( + jid, + f"Denoise bake {'cancelled after' if cancelled else 'complete —'} " + f"{written} slice(s) written to {target_path}.", + ) + except Exception as exc: # noqa: BLE001 — reported as a job error, never a crash + logger.error("Denoise bake job %s failed: %s", jid, exc) + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + finally: + # A model bake holds ML_LOCK for the whole volume. Releasing only on the + # happy path would leave it held forever after any mid-bake failure, + # deadlocking every later training job and preview. + if model_denoiser is not None: + model_denoiser.close() + + +class _ModelDenoiser: + """Applies a trained Noise2Noise/Noise2Void run across a whole volume. + + Loads the model and takes ``ML_LOCK`` ONCE for the entire bake rather than + per slice: re-acquiring per slice would let a training job interleave, and + reloading weights per slice would dominate the runtime. + + Preprocessing is delegated to ``denoise_train``'s own helpers, so the + network sees exactly the display-mapped uint8 grayscale it was trained on + (volume-global bounds, the run's saved render options). Reimplementing that + here is how the two halves of the contract would silently drift apart. + + Output dtype differs from the classical filters on purpose: a denoiser + returns continuous values in ``[0, 1]``, so this writes **uint8** (the scale + it actually operated in) rather than pretending to the source's uint16 + precision it never had access to. + """ + + def __init__(self, run_id: str, node: Any, meta: dict[str, Any]) -> None: + import denoise_train + import tiling + import train_common + + config = train_common.load_run_config(run_id) + if config.get("model_family") != "dlsia_denoiser" or config.get("task") != "denoising": + raise ValueError(f"Run {run_id!r} is not a denoiser run.") + # Which network this run is; defaults to TUNet for runs saved before the + # field existed. Only the TUNet architecture needs dlsia. + fam = train_common.denoiser_runtime_for(config) + if train_common.denoiser_needs_dlsia(config) and not train_common.dlsia_available(): + raise ValueError("dlsia is not installed on this server.") + if not tiling.qlty_available(): + raise ValueError("Applying a denoiser needs the 'qlty' package, which is missing.") + device = train_common.pick_device() + if device is None: + raise ValueError("torch is not installed on this server.") + + self._tiling = tiling + self._train_common = train_common + self._denoise_train = denoise_train + self._node = node + self._meta = meta + self._device = device + self._window = int(config["image_size"]) + self._opts = denoise_train._render_opts(config.get("render") or {}) + + import images as images_mod + + self._global_range = images_mod._sample_global_stats(node, meta) + + if not train_common.ML_LOCK.acquire(blocking=False): + raise ValueError("The device is busy with another job — try again when it finishes.") + self._locked = True + try: + state = train_common.load_adapter_state(run_id) + model = fam.load_model(state, device) + model.eval() + self._forward_fn = fam.make_forward_fn(model) + self._to_tensor_fn = fam.make_to_tensor_fn() + except Exception: + self.close() # never hold the lock past a failed load + raise + + def denoise(self, index: int) -> np.ndarray: + import torch + + gray = self._denoise_train._slice_to_gray_uint8( + self._node, self._meta, index, self._opts, self._global_range + ) + with torch.no_grad(): + out = self._tiling.denoise_image_tiled( + gray, + forward_fn=self._forward_fn, + to_tensor_fn=self._to_tensor_fn, + window=self._window, + device=self._device, + ) + if out is None: + raise RuntimeError("denoising returned no result") + unit = np.clip(np.asarray(out, dtype=np.float64), 0.0, 1.0) + return (unit * 255.0).round().astype(np.uint8) + + def close(self) -> None: + """Release ``ML_LOCK``. Idempotent — callers release on both the happy + path and in a ``finally``.""" + if getattr(self, "_locked", False): + self._train_common.ML_LOCK.release() + self._locked = False + + +def _denoise_one( + node: Any, + meta: dict[str, Any], + index: int, + method: str, + strength: float, + radius: int, + n_slices: int, +) -> np.ndarray: + """Denoise slice *index*, reading a z-window when the method needs one. + + The window clamps at the volume ends, so the target slice is not always the + centre of what was read — its position is tracked explicitly. + """ + if radius == 0: + return denoise_mod.denoise_slice( + np.asarray(arrays_mod.read_slice(node, meta, index)), method, strength + ) + + lo = max(0, index - radius) + hi = min(n_slices - 1, index + radius) + frames = [np.asarray(arrays_mod.read_slice(node, meta, i)) for i in range(lo, hi + 1)] + if len(frames) < 2: + fallback = "gaussian" if method == "gaussian3d" else "median" + return denoise_mod.denoise_slice(frames[0], fallback, strength) + return denoise_mod.denoise_stack(np.stack(frames, axis=0), method, strength)[index - lo] diff --git a/backend/denoise_runtime.py b/backend/denoise_runtime.py new file mode 100644 index 0000000..0eb5506 --- /dev/null +++ b/backend/denoise_runtime.py @@ -0,0 +1,140 @@ +"""dlsia TUNet model building for the Train tab's self-supervised denoiser +(Noise2Noise / Noise2Void) — the third model family alongside dlsia TUNet +segmentation (see :mod:`dlsia_runtime`). + +Mirrors :mod:`dlsia_runtime`'s function template closely — same shared TUNet +architecture, same save/load/forward-fn/train-mode plumbing — but wired for +regression instead of classification: + +* Fixed at ``in_channels=1, out_channels=1``: a single noisy intensity + channel in, a single denoised intensity channel out. There is no + "n_classes" here — the output channel count isn't user-configurable, unlike + the segmentation family's ``out_channels=n_classes``. +* No softmax/argmax anywhere in this module. A classification head turns + logits into class probabilities/labels; a regression head's raw output IS + the answer (the denoised pixel value), so applying either would corrupt it. + +dlsia is an optional dependency (``backend/pyproject.toml``'s ``ml`` extra, +alongside ``torch``) — this module is import-safe without it installed; only +:func:`build_model`/:func:`load_model` actually import it, guarded by +``train_common.dlsia_available()``. + +Training itself is driven by ``backend/denoise_train.py`` (not this module, +and not written here) — this module only builds/saves/loads the network and +its forward pass, the same division of labour ``dlsia_runtime`` has with the +generic loop in ``train_common.run_training_loop``. + +Documented future extension points (not implemented here — real classes in +``dlsia.core.networks``, kept as an easy on-ramp if TUNet turns out not to be +the best fit for denoising specifically): + * ``MSDNet`` / ``MSDAE`` — mixed-scale dense networks; an alternative to + TUNet's encoder/decoder that some denoising literature prefers. + * ``MSAE`` / ``SparseAutoEncoder`` — autoencoder-style architectures. +Any of these could plug in behind the same function names this module +exposes, without changing ``train_common.build_family``'s call sites. +""" + +from __future__ import annotations + +from typing import Any + + +def build_model(image_size: int, depth: int, base_channels: int, growth_rate: float, device: str) -> Any: + """Construct a fresh (untrained) dlsia TUNet for single-channel denoising. + + ``image_shape`` is fixed to ``(image_size, image_size)`` at construction — + dlsia's TUNet precomputes exact per-layer tensor sizes from it, so every + training and inference image must be letterboxed/tiled to this same size + (same constraint as the segmentation family's TUNet — see + ``dlsia_runtime.build_model``). + + Fixed at ``in_channels=1, out_channels=1``: denoising reads and writes a + single raw intensity channel, never the 3-channel RGB convention the + segmentation family renders — there is no class count to parameterize. + """ + from dlsia.core.networks.tunet import TUNet # noqa: PLC0415 — optional dependency + + model = TUNet( + image_shape=(image_size, image_size), + in_channels=1, + out_channels=1, + depth=depth, + base_channels=base_channels, + growth_rate=growth_rate, + # dlsia's TUNet already defaults `final_activation` to None (verified + # against dlsia.core.networks.tunet.TUNet.__init__/forward: with no + # final_activation configured, forward() returns the last conv's raw + # output unchanged — no softmax/sigmoid is applied). Passed explicitly + # here anyway, and pinned to None, so a regression head can never + # silently gain an output-squashing activation if a future dlsia + # release changed that default — the denoised intensity must come + # back as an unbounded raw value, not something clamped to (0, 1). + final_activation=None, + ) + return model.to(device) + + +def network_dict(model: Any) -> dict[str, Any]: + """The full ``{topo_dict, state_dict}`` dlsia uses to reconstruct a TUNet + (see ``TUNet.save_network_parameters`` / ``TUNetwork_from_file``). Stored + as-is in this run's ``adapter.pt`` — the denoiser trains from scratch, so + there's no base/delta split the way LoRA has.""" + return model.save_network_parameters(name=None) + + +def load_model(state: dict[str, Any], device: str) -> Any: + """Reconstruct a denoiser TUNet from a saved :func:`network_dict`.""" + from dlsia.core.networks.tunet import TUNet # noqa: PLC0415 + + model = TUNet(**state["topo_dict"]) + model.load_state_dict(state["state_dict"]) + return model.to(device) + + +def make_forward_fn(model: Any): + """Return ``forward(batch_images) -> denoised`` for + :func:`train_common.run_training_loop`. + + TUNet's transposed-conv decoder is symmetric with its encoder, so output + spatial size already matches input size — no resizing needed. The output + is the raw regression prediction (denoised intensity): no softmax/argmax + here, unlike the segmentation family's per-class logits. + """ + + def _forward(batch_images: Any) -> Any: + return model(batch_images) + + return _forward + + +def make_set_train_mode_fn(model: Any): + """Return ``set_train_mode(is_training)`` for :func:`train_common.run_training_loop`. + + TUNet defaults to ``nn.BatchNorm2d`` — its running mean/var should only + update during training, not while computing validation metrics. + """ + + def _set_train_mode(is_training: bool) -> None: + model.train(is_training) + + return _set_train_mode + + +def make_to_tensor_fn(): + """Return ``(gray_uint8_hw) -> float CPU tensor (1,H,W)`` scaled to [0, 1]. + + Mirrors ``dlsia_runtime.make_to_tensor_fn``'s uint8->[0,1] convention, but + for a single intensity channel instead of 3-channel RGB: denoising trains + on raw grayscale intensity, not the rendered RGB tiles segmentation uses. + Accepts either ``(H, W)`` or a single-channel ``(H, W, 1)`` array. + """ + import numpy as np + import torch + + def _to_tensor(gray_uint8: "np.ndarray") -> Any: + arr = np.ascontiguousarray(gray_uint8) + if arr.ndim == 3 and arr.shape[-1] == 1: + arr = arr[:, :, 0] + return torch.from_numpy(arr).unsqueeze(0).float() / 255.0 + + return _to_tensor diff --git a/backend/denoise_train.py b/backend/denoise_train.py new file mode 100644 index 0000000..fbf344a --- /dev/null +++ b/backend/denoise_train.py @@ -0,0 +1,841 @@ +"""Self-supervised denoiser training: data sampling, patching, and the loop. + +This is the training-time counterpart to :mod:`denoise_runtime` (which only +builds/saves/loads the single-channel TUNet) and the denoiser's answer to +:func:`train_common.run_training_loop`. + +Why a separate loop instead of extending ``run_training_loop`` +-------------------------------------------------------------- +``train_common.run_training_loop`` is not family-generic in the way its name +suggests — it is *segmentation*-generic. Three of its assumptions are +hardcoded and all three are wrong for regression: + +1. ``nn.CrossEntropyLoss(ignore_index=IGNORE_INDEX)`` — a classification loss + over class logits. +2. Its ``_prep`` casts the target with ``.astype(np.int64)`` — class indices, + not intensities. +3. ``_evaluate`` does ``logits.argmax(dim=1)`` + :func:`train_common.compute_miou`. + +So this module deliberately duplicates the epoch/batch/scheduler/cancel +skeleton rather than growing a ``task=`` switch through the middle of a +function every existing segmentation run depends on. Everything *around* the +loss — device handling, the ``on_batch``/``on_epoch`` cooperative-cancel +protocol (return ``True`` to stop), AdamW + cosine annealing, the returned +metrics dict shape — is kept identical, so :mod:`train_jobs` drives this the +same way it drives the segmentation loop. + +Three self-supervised schemes +----------------------------- +None needs a single annotation; all three learn from the raw slices themselves. +The first two are architecture-agnostic; the third depends on the architecture +having no skip connections. + +* **Noise2Noise** (``"n2n"``) — train slice *i* to predict slice *i+stride*. + Adjacent slices of a tomographic/microscopy volume share structure but carry + independent noise realizations, so under MSE the optimum is the conditional + mean, i.e. the shared (clean) structure: the noise cannot be predicted and + averages out. +* **Noise2Void** (``"n2v"``) — single slices, no pairing. A small fraction of + pixels is *masked* by overwriting each with a random neighbour's value, and + the loss is evaluated **only** at those coordinates, against the original + values. Because the model never sees a masked pixel's own value, it cannot + learn the identity function — it has to infer the value from context, which + is exactly the denoising task. +* **Autoencoder** (``"ae"``) — single slices, pure self-reconstruction: the + target IS the input, under plain MSE. That only works on an architecture with + no skip connections and a narrow latent bottleneck (``"cnn_ae"``; see + :mod:`autoencoder_runtime`), which is why ``schemas.DlsiaDenoiserConfig`` + refuses this scheme on TUNet. Given those, the input physically cannot pass + through unchanged, and noise — the high-entropy part that does not fit through + the bottleneck — is what gets dropped. Requires no synthetic corruption, which + is the point: adding fake noise to already-noisy data trains the model on the + wrong noise distribution. The honest tradeoff is that a bottleneck discards + fine real detail along with the noise, so this tends to blur more than n2v. + +Intensity normalisation +----------------------- +The denoiser is single-channel and trains on raw intensity, not on the +3-channel RGB tiles the segmentation path renders. Slices are mapped to uint8 +through :func:`images.normalize_scalar_unit` — the same intensity pipeline +behind the 2-D canvas and the 3-D volume view — so training data matches what +the user actually sees. The percentile bounds are forced **volume-global** +even when the request asks for per-slice normalisation: a per-slice mapping +would put a Noise2Noise input and its target on two different intensity +scales, which would show up as a constant structural error the network cannot +fix and would happily waste capacity trying to. The requested ``scale`` +transform (linear/log/symlog) is preserved. +""" + +from __future__ import annotations + +import logging +import math +from typing import Any, Callable, Iterable, Sequence + +import numpy as np + +logger = logging.getLogger(__name__) + +# Slice offset between a Noise2Noise input and its target. 1 = immediately +# adjacent, which is the strongest structural correspondence available; larger +# strides trade structural similarity for a bit more independence. Not a +# `TunetHyperParams` field (that schema is shared with the segmentation TUNet +# and is out of scope here), so it is a sampler argument with this default. +DEFAULT_PAIR_STRIDE = 1 + +# Share of pixels blanked per patch for Noise2Void. The N2V paper's usable +# range is ~0.5-2%: too few and each patch supervises almost nothing (the loss +# gets noisy and training crawls), too many and the masked pixels start +# destroying the very context the model needs to inpaint them from. +DEFAULT_N2V_MASK_FRACTION = 0.015 +# Side of the square window a masked pixel's replacement value is drawn from. +DEFAULT_N2V_NEIGHBOURHOOD = 5 +# Bounded resampling attempts when a drawn donor is unusable (it is the masked +# pixel itself, or another masked pixel). See :func:`n2v_mask_and_replace`. +_N2V_MAX_RESAMPLE = 8 + + +# --------------------------------------------------------------------------- +# Raw slice sampling (annotation-free — deliberately NOT prepare_datasets) +# --------------------------------------------------------------------------- + + +def _render_opts(render: Any) -> dict[str, Any]: + """Coerce a ``RenderOpts`` model (or plain dict) to the dict form + :mod:`images` expects, with the normalisation scope pinned to + ``"global"`` (see this module's docstring).""" + opts = render.model_dump() if hasattr(render, "model_dump") else dict(render or {}) + return {**opts, "norm": "global"} + + +def selected_slice_indices(item: Any, n_slices: int) -> list[int]: + """Which raw slice indices of *item* the denoiser should train on. + + The Learned Denoiser panel has no annotations to send, so it reuses + ``ExportSourceItem.slices`` purely to name the in-scope slice indices, with + an empty shape list for each (there is no other field on that schema for + "just these indices"). An empty/absent mapping means the whole volume. + + Returns a sorted, deduplicated, in-bounds list. + """ + keys = getattr(item, "slices", None) or {} + indices: set[int] = set() + for key in keys: + try: + idx = int(key) + except (TypeError, ValueError): + logger.warning("Denoiser scope: ignoring non-integer slice key %r", key) + continue + if 0 <= idx < n_slices: + indices.add(idx) + return sorted(indices) if indices else list(range(n_slices)) + + +def _slice_to_gray_uint8( + node: Any, + meta: dict[str, Any], + idx: int, + opts: dict[str, Any], + global_range: tuple[float, float] | None, +) -> np.ndarray: + """Read slice *idx* and map it to a ``(H, W)`` uint8 grayscale array. + + Matches ``images.render_slice``'s grayscale branch exactly (same + ``normalize_scalar_unit`` call, same ``round()`` quantisation), so the + denoiser trains on the intensities the viewer displays. A colour source is + collapsed to a single channel first — the network has ``in_channels=1``. + """ + import arrays as arrays_mod # noqa: PLC0415 — avoid a hard import-time cycle + import images as images_mod # noqa: PLC0415 + + arr = np.asarray(arrays_mod.read_slice(node, meta, idx)) + if arr.ndim == 3: + arr = arr[:, :, :3].astype(np.float64).mean(axis=2) + unit = images_mod.normalize_scalar_unit(arr, opts, global_range) + return (unit * 255.0).round().astype(np.uint8) + + +def load_source_slices( + item: Any, + render: Any, + progress_cb: Callable[[str], None] | None = None, +) -> tuple[list[int], dict[int, np.ndarray]]: + """Read every in-scope slice of one source as uint8 grayscale. + + Returns ``(indices, {index: (H, W) uint8})``. + + Each slice is read **exactly once** and shared by every pair that + references it. That is deliberately not the "fetch a contiguous pair in one + round-trip" optimisation (``node[i:i+2]``) the shape of a Noise2Noise + sampler first suggests: with ``stride=1`` every interior slice belongs to + two pairs (as input of one and target of the previous), and with + ``both_directions`` to four, so a per-pair range read would fetch each + slice 2-4x. Reading the scope once into a dict is strictly fewer + round-trips than any per-pair batching, at the same peak memory the + segmentation path already accepts (uint8 grayscale here vs. its uint8 RGB — + a third the bytes per slice). + """ + import arrays as arrays_mod # noqa: PLC0415 + import images as images_mod # noqa: PLC0415 + + node = arrays_mod.resolve_array(item.source, item.kind, item.server_uri) + meta = arrays_mod.array_shape_meta(node) + indices = selected_slice_indices(item, meta["n_slices"]) + + opts = _render_opts(render) + global_range = images_mod._sample_global_stats(node, meta) + + slices: dict[int, np.ndarray] = {} + for n_done, idx in enumerate(indices, start=1): + slices[idx] = _slice_to_gray_uint8(node, meta, idx, opts, global_range) + if progress_cb is not None and (n_done % 10 == 0 or n_done == len(indices)): + progress_cb(f"{item.source}: read {n_done}/{len(indices)} slice(s)") + return indices, slices + + +def noise2noise_pairs( + indices: Sequence[int], + *, + stride: int = DEFAULT_PAIR_STRIDE, + both_directions: bool = True, +) -> list[tuple[int, int]]: + """Index pairs ``(input_index, target_index)`` for Noise2Noise. + + A pair is emitted only when **both** of its indices are in *indices*, which + is what keeps it inside the volume at either end: the last in-scope slice + simply has no partner rather than being clamped onto itself (a + self-pair would be a perfect identity target and would train the network to + do nothing). + + With *both_directions*, each adjacency also contributes its reverse pair. + That is free extra data, not a duplicate: MSE is symmetric in value but the + two directions are different (input, target) assignments, so the network + sees each slice as an input as well as a target. + """ + if stride < 1: + raise ValueError(f"Noise2Noise pair stride must be >= 1, got {stride}") + available = {int(i) for i in indices} + pairs: list[tuple[int, int]] = [] + for i in sorted(available): + j = i + stride + if j in available: + pairs.append((i, j)) + if both_directions: + pairs.append((j, i)) + return pairs + + +def prepare_noise2noise_datasets( + sources: Iterable[Any], + render: Any, + *, + stride: int = DEFAULT_PAIR_STRIDE, + both_directions: bool = True, + progress_cb: Callable[[str], None] | None = None, +) -> dict[str, list[tuple[np.ndarray, np.ndarray]]]: + """Build ``(input_slice, target_slice)`` pairs from adjacent raw slices. + + Annotation-free by construction: unlike + :func:`train_common.prepare_datasets`, nothing here goes through + ``coco_export.build_export_plan`` — a denoiser has no classes to rasterise + and would be blocked entirely by that path's "no annotated slices" outcome. + + Returns :func:`train_common.prepare_datasets`' dict shape + (``{"train": [(input, target), ...], "val": [...]}``), except both arrays + are ``(H, W)`` uint8 grayscale rather than ``(rgb_hwc, label_hw)``. + Everything lands in ``"train"``; the validation split is taken afterwards + by :func:`tiling.holdout_val_patches`, at patch granularity, which is the + same seeded holdout the tiled segmentation path already uses (and the only + thing that works for a one- or two-slice scope). + """ + pairs: list[tuple[np.ndarray, np.ndarray]] = [] + for item in sources: + indices, slices = load_source_slices(item, render, progress_cb) + index_pairs = noise2noise_pairs(indices, stride=stride, both_directions=both_directions) + if not index_pairs: + logger.warning( + "Noise2Noise: source %r has no slice %d apart in scope — contributed no pairs", + getattr(item, "source", "?"), + stride, + ) + pairs.extend((slices[i], slices[j]) for i, j in index_pairs) + if progress_cb is not None: + progress_cb(f"{item.source}: {len(indices)} slice(s) → {len(index_pairs)} Noise2Noise pair(s)") + + if not pairs: + raise ValueError( + "Noise2Noise needs at least two slices that are " + f"{stride} apart within the selected scope — none were found. " + "Widen the slice range, or train with Noise2Void instead." + ) + return {"train": pairs, "val": []} + + +def prepare_noise2void_datasets( + sources: Iterable[Any], + render: Any, + *, + progress_cb: Callable[[str], None] | None = None, +) -> dict[str, list[tuple[np.ndarray, np.ndarray]]]: + """Build single-slice training items for Noise2Void. + + Returns the same ``{"train": [(input, target), ...], "val": []}`` shape as + :func:`prepare_noise2noise_datasets` so patching and the loop are shared, + with the slice paired **with itself**: the blind-spot masking that makes + the two differ is applied per patch, per epoch, inside + :func:`run_denoise_training_loop` (a fresh mask each time a patch is seen, + which is the point — one fixed mask would supervise only ~1.5% of pixels + ever). The same array object is used on both sides; nothing downstream + mutates it in place. + """ + pairs: list[tuple[np.ndarray, np.ndarray]] = [] + for item in sources: + indices, slices = load_source_slices(item, render, progress_cb) + pairs.extend((slices[i], slices[i]) for i in indices) + if progress_cb is not None: + progress_cb(f"{item.source}: {len(indices)} Noise2Void slice(s)") + + if not pairs: + raise ValueError("Noise2Void needs at least one slice in the selected scope — none were found.") + return {"train": pairs, "val": []} + + +# --------------------------------------------------------------------------- +# Patch extraction (qlty geometry, WITHOUT the sparse-annotation weeding) +# --------------------------------------------------------------------------- + + +def tile_denoise_pair(inp: np.ndarray, tgt: np.ndarray, window: int) -> list[tuple[np.ndarray, np.ndarray]]: + """Cut one ``(input, target)`` grayscale pair into ``window``-sized patches. + + Deliberately **not** :func:`tiling._tile_pair`. That one runs + ``qlty.cleanup.weed_sparse_classification_training_pairs_2D``, which drops + every patch containing no labelled pixel. That is right for sparse + segmentation annotations and completely wrong here: a denoiser is trained + on raw pixels, so every patch is valid training data and a weeded run would + silently discard the entire dataset (an all-zero "label" is, to the weeder, + an unlabelled patch). + + Input and target are unstitched **independently through the same + ``NCYXQuilt``**, so patch *k* of one is the exact same window of the image + as patch *k* of the other — spatial alignment comes from shared geometry + rather than from trusting a paired helper. The quilt is + :func:`tiling._quilt`, the same geometry + :func:`tiling.denoise_image_tiled` uses at inference time, so the model is + trained on and applied to identically-shaped windows. + + Note that an image smaller than one window is zero-padded up to it (see + :func:`tiling.pad_to_min`); the padded region is constant 0 in both input + and target, so it is trivially satisfiable rather than misleading — but it + is real loss mass, so tiling an undersized image is worth avoiding. + """ + import torch # noqa: PLC0415 + + import tiling # noqa: PLC0415 — optional (qlty) dependency + + if inp.shape != tgt.shape: + raise ValueError(f"Denoiser input/target shapes differ: {inp.shape} vs {tgt.shape}") + + img = tiling.pad_to_min(inp, window, window, fill=0) + tar = tiling.pad_to_min(tgt, window, window, fill=0) + quilt = tiling._quilt(img.shape[0], img.shape[1], window) + + # (1, 1, H, W) — one image, one channel. uint8 all the way to the batch's + # to_tensor_fn, same as the segmentation patch cache: 4x less memory than + # float32 for what can be tens of thousands of patches. + patches_in = quilt.unstitch(torch.from_numpy(np.ascontiguousarray(img))[None, None]) + patches_tgt = quilt.unstitch(torch.from_numpy(np.ascontiguousarray(tar))[None, None]) + + return [ + ( + np.ascontiguousarray(patches_in[k, 0].numpy()), + np.ascontiguousarray(patches_tgt[k, 0].numpy()), + ) + for k in range(len(patches_in)) + ] + + +def tile_denoise_datasets( + datasets: dict[str, list[tuple[np.ndarray, np.ndarray]]], + window: int, + progress_cb: Callable[[str], None] | None = None, + cancel_cb: Callable[[], bool] | None = None, +) -> dict[str, list[tuple[np.ndarray, np.ndarray]]] | None: + """Replace each full-resolution pair with its window-sized patches. + + Mirrors :func:`tiling.tile_datasets`' contract exactly — same in/out dict + shape, same "return ``None`` if *cancel_cb* asked to stop partway" — but + routes through :func:`tile_denoise_pair` (no weeding, single channel). + """ + out: dict[str, list[tuple[np.ndarray, np.ndarray]]] = {} + for split, pairs in datasets.items(): + patches: list[tuple[np.ndarray, np.ndarray]] = [] + for i, (inp, tgt) in enumerate(pairs): + if cancel_cb is not None and cancel_cb(): + return None + before = len(patches) + patches.extend(tile_denoise_pair(inp, tgt, window)) + if progress_cb is not None: + progress_cb(f"{split} item {i + 1}/{len(pairs)}: {len(patches) - before} patch(es)") + out[split] = patches + if progress_cb is not None: + progress_cb(f"{split}: {len(pairs)} item(s) → {len(patches)} patch(es) of {window}px") + return out + + +# --------------------------------------------------------------------------- +# Geometry for the non-tiled path +# --------------------------------------------------------------------------- + + +def letterbox_denoise_pair(inp: np.ndarray, tgt: np.ndarray, size: int) -> tuple[np.ndarray, np.ndarray]: + """Resize-keep-aspect + pad an ``(input, target)`` grayscale pair to + ``size`` x ``size``. + + :func:`train_common.letterbox` cannot be used here: it pads the *label* + with :data:`train_common.IGNORE_INDEX` (255) and nearest-resizes it as + class indices. For a regression target 255 is not "ignore", it is + *maximum brightness* — a bright frame the network would be trained to + reproduce. This variant treats both sides as images: identical bilinear + resize, identical zero padding, so they stay pixel-aligned. + + With ``tiling=True`` (the default) every patch already arrives exactly + ``size`` x ``size`` and this short-circuits, exactly as + :func:`train_common.letterbox` does. The non-tiled path is handled + explicitly rather than left to rely on that short-circuit never being + missed — but it is still the worse option for a denoiser, because + resampling a noisy image correlates neighbouring pixels and so weakens the + per-pixel noise independence both schemes are built on. The caller logs a + warning; see :func:`run_denoise_training_loop`. + """ + from PIL import Image as PILImage # noqa: PLC0415 + + if inp.shape != tgt.shape: + raise ValueError(f"Denoiser input/target shapes differ: {inp.shape} vs {tgt.shape}") + h, w = inp.shape[:2] + if (h, w) == (size, size): + return inp, tgt + + scale = min(size / h, size / w) + nh, nw = max(1, round(h * scale)), max(1, round(w * scale)) + top, left = (size - nh) // 2, (size - nw) // 2 + + out: list[np.ndarray] = [] + for arr in (inp, tgt): + resized = np.asarray(PILImage.fromarray(arr).resize((nw, nh), PILImage.Resampling.BILINEAR)) + canvas = np.zeros((size, size), dtype=np.uint8) + canvas[top : top + nh, left : left + nw] = resized # noqa: E203 + out.append(canvas) + return out[0], out[1] + + +# --------------------------------------------------------------------------- +# Noise2Void blind-spot masking +# --------------------------------------------------------------------------- + + +def _mirror_offset(centre: np.ndarray, offset: np.ndarray, n: int) -> np.ndarray: + """``centre + offset`` along one axis, folded back inside ``[0, n)`` by + flipping the offset's sign rather than by clamping the index. + + Both of the obvious alternatives silently break the blind spot at the image + border, by mapping some neighbour offset back onto the pixel itself: + + * clipping — ``centre=0, offset=-1`` clips to ``0``; + * index reflection (``-idx``) — ``centre=1, offset=-2`` reflects ``-1`` to + ``1``. + + Flipping the offset instead gives ``centre - offset``, which differs from + ``centre`` for every non-zero offset. It stays in range whenever the axis + is at least ``2 * radius + 1`` long (true of any real patch); the final + clip is a guard for degenerately small arrays, where the caller's + donor-is-not-the-centre resampling catches the leftover case. + """ + raw = centre + offset + folded = np.where((raw < 0) | (raw >= n), centre - offset, raw) + return np.clip(folded, 0, n - 1) + + +def n2v_mask_and_replace( + patch: np.ndarray, + rng: np.random.Generator, + *, + fraction: float = DEFAULT_N2V_MASK_FRACTION, + neighbourhood: int = DEFAULT_N2V_NEIGHBOURHOOD, +) -> tuple[np.ndarray, np.ndarray]: + """Apply Noise2Void masking to one ``(H, W)`` patch. + + Returns ``(masked_patch, mask)`` — a **copy** of *patch* in which each + selected pixel has been overwritten with the value of a random neighbour + from a ``neighbourhood`` x ``neighbourhood`` window, and the boolean mask + of the selected coordinates. + + The blind spot is the whole point, and it is enforced twice over: a donor + is never the masked pixel itself, and never another masked pixel. Without + the second condition a masked pixel's original value could still reach the + input by being copied into some *other* masked pixel's position — a leak + the model's receptive field is wide enough to exploit, at which point it + learns the identity function and denoises nothing. + + The count is exact (``round(H * W * fraction)``, at least 1) rather than a + per-pixel coin flip. Beyond making the fraction testable, it guarantees a + non-empty mask: ``dlsia``'s ``MSELossMasked`` divides by ``masks.sum()``, + so an all-false mask would return NaN and poison the run. + """ + if not 0.0 < fraction < 1.0: + raise ValueError(f"n2v mask fraction must be in (0, 1), got {fraction}") + if neighbourhood < 3 or neighbourhood % 2 == 0: + raise ValueError(f"n2v neighbourhood must be an odd size >= 3, got {neighbourhood}") + + h, w = patch.shape[:2] + n_pixels = h * w + n_masked = min(max(1, int(round(n_pixels * fraction))), n_pixels - 1) + + flat = rng.choice(n_pixels, size=n_masked, replace=False, shuffle=False) + ys, xs = np.divmod(flat, w) + mask = np.zeros((h, w), dtype=bool) + mask[ys, xs] = True + + radius = neighbourhood // 2 + donor_y = np.empty(n_masked, dtype=np.int64) + donor_x = np.empty(n_masked, dtype=np.int64) + # Positions in ys/xs still without a usable donor. An explicit index array + # rather than a boolean `todo[todo] = ...`, which would index an array with + # itself while writing to it. + pending = np.arange(n_masked) + + for _ in range(_N2V_MAX_RESAMPLE): + if pending.size == 0: + break + dy = rng.integers(-radius, radius + 1, size=pending.size) + dx = rng.integers(-radius, radius + 1, size=pending.size) + cand_y = _mirror_offset(ys[pending], dy, h) + cand_x = _mirror_offset(xs[pending], dx, w) + donor_y[pending] = cand_y + donor_x[pending] = cand_x + # Unusable if it landed on the pixel itself, or on another masked pixel. + unusable = ((cand_y == ys[pending]) & (cand_x == xs[pending])) | mask[cand_y, cand_x] + pending = pending[unusable] + + if pending.size: + # Degenerate only at absurd mask fractions, where a 5x5 window can be + # entirely masked. Fall back to a uniformly-drawn unmasked pixel from + # anywhere in the patch: a worse donor (no locality) but one that keeps + # the blind-spot guarantee, which correctness depends on and locality + # does not. `fraction < 1` guarantees the candidate set is non-empty. + unmasked = np.flatnonzero(~mask.ravel()) + picked = rng.choice(unmasked, size=pending.size, replace=True) + donor_y[pending], donor_x[pending] = np.divmod(picked, w) + logger.debug("n2v: %d pixel(s) fell back to a non-local donor", pending.size) + + masked = patch.copy() + masked[ys, xs] = patch[donor_y, donor_x] + return masked, mask + + +# --------------------------------------------------------------------------- +# Training loop +# --------------------------------------------------------------------------- + +# Key of the returned validation metric. Named for what it actually measures — +# Pearson correlation between the prediction and the (still noisy) target — so +# it can't be read as an image-quality score. It is NOT PSNR and must not be +# presented as one: for Noise2Noise the target is a different noisy slice, and +# for Noise2Void it is the original noisy pixel, so perfect agreement with it +# would mean the model had learned to reproduce noise. Rising early then +# plateauing is the convergence signal to read it for. +VAL_METRIC_KEY = "val_noisy_target_pearson" + + +def _pearson(pred: "Any", target: "Any") -> float: + """dlsia's regression metric, coerced to a plain finite float. + + ``torch.corrcoef`` returns NaN when either input has zero variance (a + constant patch — an all-padding window, or a model that has collapsed to a + constant early in training). Reporting NaN would propagate into + ``metrics.json`` and render as a broken value in the runs list, so it + degrades to 0.0, which is also the honest reading: no measured agreement. + """ + from dlsia.core.train_scripts import regression_metrics # noqa: PLC0415 + + if target.numel() < 2: + return 0.0 + value = float(regression_metrics(pred, target)) + return value if math.isfinite(value) else 0.0 + + +def _prep_batch( + pairs: list[tuple[np.ndarray, np.ndarray]], + indices: Sequence[int], + *, + image_size: int, + scheme: str, + to_tensor_fn: Callable[[Any], Any], + rng: np.random.Generator, + flip_augment: bool, + mask_fraction: float, + neighbourhood: int, +) -> tuple[Any, Any, Any]: + """Assemble one batch as ``(inputs, targets, mask_or_None)`` CPU tensors. + + For ``"n2v"`` the mask is drawn fresh here, per patch and per epoch, so a + patch seen ten times supervises ten different pixel subsets. ``"ae"`` needs no + per-epoch randomisation — its objective is fixed (reconstruct this patch), + and the bottleneck, not a fresh perturbation, is what stops it cheating. + """ + import torch # noqa: PLC0415 + + imgs, tgts, masks = [], [], [] + for i in indices: + inp, tgt = pairs[int(i)] + inp, tgt = letterbox_denoise_pair(inp, tgt, image_size) + flip = bool(flip_augment and rng.random() < 0.5) + if flip: + inp = np.ascontiguousarray(inp[:, ::-1]) + if scheme == "ae": + # Pure self-reconstruction: the target IS the input. Safe only + # because the 'cnn_ae' architecture has no skip connections, so the + # bottleneck cannot pass the input through unchanged (the schema + # refuses this scheme on TUNet for exactly that reason). Like n2v, + # an ae item is a slice paired with itself, so the pair's `tgt` is + # ignored — the target must be the same flip `inp` just took. + imgs.append(to_tensor_fn(inp)) + tgts.append(to_tensor_fn(inp)) + elif scheme == "n2v": + # Target is the ORIGINAL patch; the model input is the masked copy. + # `tgt` is ignored — an n2v item is a slice paired with itself, and + # the target has to be the same flip of it that `inp` just took. + masked, mask = n2v_mask_and_replace( + inp, rng, fraction=mask_fraction, neighbourhood=neighbourhood + ) + imgs.append(to_tensor_fn(masked)) + tgts.append(to_tensor_fn(inp)) + masks.append(torch.from_numpy(mask).unsqueeze(0)) + else: + if flip: + tgt = np.ascontiguousarray(tgt[:, ::-1]) + imgs.append(to_tensor_fn(inp)) + tgts.append(to_tensor_fn(tgt)) + + batch_mask = torch.stack(masks) if masks else None + return torch.stack(imgs), torch.stack(tgts), batch_mask + + +def _denoise_loss(pred: Any, target: Any, mask: Any, criterion: Any) -> Any: + """Masked MSE for Noise2Void, plain MSE for Noise2Noise.""" + return criterion(pred, target) if mask is None else criterion(pred, target, mask) + + +def evaluate_denoise( + val_pairs: list[tuple[np.ndarray, np.ndarray]], + *, + image_size: int, + scheme: str, + to_tensor_fn: Callable[[Any], Any], + forward_fn: Callable[[Any], Any], + device: str, + criterion: Any, + batch_size: int, + seed: int, + mask_fraction: float, + neighbourhood: int, +) -> tuple[float, float]: + """Average validation loss and target-correlation over *val_pairs*. + + Mirrors :func:`train_common._evaluate`'s batching and its per-sample + weighting of a batch-mean loss, with mIoU replaced by + :data:`VAL_METRIC_KEY`'s Pearson correlation — and, for Noise2Void, the + correlation restricted to the masked coordinates, which are the only + positions the model was ever asked to predict. + + The mask RNG is re-seeded from *seed* on every call, so epoch-to-epoch + changes in the metric come from the model rather than from a different + random subset of pixels being scored each time. + """ + import torch # noqa: PLC0415 + + rng = np.random.default_rng(seed) + total_loss = 0.0 + total_corr = 0.0 + n = len(val_pairs) + with torch.no_grad(): + for start in range(0, n, batch_size): + chunk_idx = range(start, min(start + batch_size, n)) + imgs, tgts, mask = _prep_batch( + val_pairs, + list(chunk_idx), + image_size=image_size, + scheme=scheme, + to_tensor_fn=to_tensor_fn, + rng=rng, + flip_augment=False, + mask_fraction=mask_fraction, + neighbourhood=neighbourhood, + ) + imgs, tgts = imgs.to(device), tgts.to(device) + mask = mask.to(device) if mask is not None else None + + pred = forward_fn(imgs) + chunk_n = len(chunk_idx) + total_loss += float(_denoise_loss(pred, tgts, mask, criterion).item()) * chunk_n + if mask is None: + total_corr += _pearson(pred, tgts) * chunk_n + else: + total_corr += _pearson(pred[mask], tgts[mask]) * chunk_n + return total_loss / n, total_corr / n + + +def run_denoise_training_loop( + *, + train_pairs: list[tuple[np.ndarray, np.ndarray]], + val_pairs: list[tuple[np.ndarray, np.ndarray]], + image_size: int, + training_scheme: str, + epochs: int, + batch_size: int, + seed: int, + flip_augment: bool, + to_tensor_fn: Callable[[Any], Any], + forward_fn: Callable[[Any], Any], + trainable_params: list[Any], + lr: float, + device: str, + n2v_mask_fraction: float = DEFAULT_N2V_MASK_FRACTION, + n2v_neighbourhood: int = DEFAULT_N2V_NEIGHBOURHOOD, + make_optimizer_fn: Callable[[list[Any], float], Any] | None = None, + on_batch: Callable[[], bool] | None = None, + on_epoch: Callable[[int, float, float | None, float | None], bool] | None = None, + set_train_mode: Callable[[bool], None] | None = None, +) -> dict[str, Any]: + """Epoch/batch loop for the self-supervised denoiser families. + + The regression counterpart of :func:`train_common.run_training_loop`, with + the same optimizer (AdamW + cosine annealing over ``epochs``), the same + cooperative-cancel protocol (``on_batch``/``on_epoch`` return ``True`` to + stop; a cancelled run still returns its partial metrics so the caller can + save a checkpoint), and the same ``set_train_mode`` handling around + validation — TUNet uses BatchNorm2d, whose running stats must not be + updated while validating. + + ``training_scheme`` selects the objective: + + * ``"n2n"`` — plain ``nn.MSELoss`` between the input slice's prediction and + the paired adjacent slice. + * ``"n2v"`` — ``dlsia.core.custom_losses.MSELossMasked`` evaluated only at + the freshly-masked blind-spot coordinates (see + :func:`n2v_mask_and_replace`). dlsia's implementation is used rather than + a hand-rolled ``(err * mask).sum() / mask.sum()`` so the normalisation + convention matches the rest of the dlsia stack. + + Returns :func:`train_common.run_training_loop`'s dict shape — + ``{"epochs_completed", "final_train_loss", "final_val_loss", "cancelled"}`` + — with mIoU replaced by :data:`VAL_METRIC_KEY`. Read that field as a + convergence signal, never as image quality: both schemes' targets are + themselves noisy. + """ + import torch # noqa: PLC0415 + import torch.nn as nn # noqa: PLC0415 + + if not train_pairs: + raise ValueError("No training data: the denoiser needs at least one slice in scope") + if training_scheme not in {"n2n", "n2v", "ae"}: + raise ValueError(f"Unknown denoiser training scheme: {training_scheme!r}") + + sample_shape = train_pairs[0][0].shape[:2] + if sample_shape != (image_size, image_size): + # Not fatal — letterbox_denoise_pair handles it, and it is the expected + # state with tiling disabled — but resampling noisy data undermines the + # per-pixel noise independence both schemes assume, so say so once. + logger.warning( + "Denoiser training on %s items that are not already %dpx: they will be letterboxed " + "(bilinear resize + zero pad), which correlates neighbouring noise. Enable tiling to " + "train on native-resolution windows instead.", + sample_shape, + image_size, + ) + + if training_scheme == "n2v": + from dlsia.core.custom_losses import MSELossMasked # noqa: PLC0415 + + criterion: Any = MSELossMasked() + else: + criterion = nn.MSELoss() + + rng = np.random.default_rng(seed) + optimizer = (make_optimizer_fn or (lambda params, lr_: torch.optim.AdamW(params, lr=lr_, weight_decay=0.01)))( + trainable_params, lr + ) + scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max(1, epochs)) + + cancelled = False + final_train_loss = 0.0 + final_val_loss: float | None = None + val_metric: float | None = None + epochs_completed = 0 + + for epoch in range(epochs): + order = rng.permutation(len(train_pairs)) + epoch_loss = 0.0 + n_batches = 0 + for start in range(0, len(order), batch_size): + batch_idx = order[start : start + batch_size] # noqa: E203 + imgs, tgts, mask = _prep_batch( + train_pairs, + batch_idx, + image_size=image_size, + scheme=training_scheme, + to_tensor_fn=to_tensor_fn, + rng=rng, + flip_augment=flip_augment, + mask_fraction=n2v_mask_fraction, + neighbourhood=n2v_neighbourhood, + ) + imgs, tgts = imgs.to(device), tgts.to(device) + mask = mask.to(device) if mask is not None else None + + pred = forward_fn(imgs) + loss = _denoise_loss(pred, tgts, mask, criterion) + optimizer.zero_grad() + loss.backward() + optimizer.step() + + epoch_loss += float(loss.detach().item()) + n_batches += 1 + if on_batch is not None and on_batch(): + cancelled = True + break + scheduler.step() + final_train_loss = epoch_loss / max(1, n_batches) + epochs_completed = epoch + 1 + + if val_pairs and not cancelled: + if set_train_mode is not None: + set_train_mode(False) + final_val_loss, val_metric = evaluate_denoise( + val_pairs, + image_size=image_size, + scheme=training_scheme, + to_tensor_fn=to_tensor_fn, + forward_fn=forward_fn, + device=device, + criterion=criterion, + batch_size=batch_size, + seed=seed, + mask_fraction=n2v_mask_fraction, + neighbourhood=n2v_neighbourhood, + ) + if set_train_mode is not None: + set_train_mode(True) + + if on_epoch is not None and on_epoch(epochs_completed, final_train_loss, final_val_loss, val_metric): + cancelled = True + if cancelled: + break + + return { + "epochs_completed": epochs_completed, + "final_train_loss": final_train_loss, + "final_val_loss": final_val_loss, + VAL_METRIC_KEY: val_metric, + "cancelled": cancelled, + } diff --git a/backend/dlsia_runtime.py b/backend/dlsia_runtime.py new file mode 100644 index 0000000..1cce2d6 --- /dev/null +++ b/backend/dlsia_runtime.py @@ -0,0 +1,99 @@ +"""dlsia TUNet model building for the Train tab. + +A TUNet trains from scratch — no pretrained checkpoint, no license +restriction beyond dlsia's own (BSD). dlsia is an optional dependency +(``backend/pyproject.toml``'s ``ml`` extra, alongside ``torch``) — this +module is import-safe without it installed; only +:func:`build_model`/:func:`load_model` actually import it, guarded by +``train_common.dlsia_available()``. + +dlsia's own ``train_scripts.train_segmentation`` runs a whole training call +synchronously with no per-batch/per-epoch hook, so it can't report progress +or honour a cooperative cancel — this app calls dlsia only for the ``TUNet`` +model class (and its own save/load helpers) and drives training through +``train_common.run_training_loop`` instead. +""" + +from __future__ import annotations + +from typing import Any + + +def build_model( + n_classes: int, image_size: int, depth: int, base_channels: int, growth_rate: float, device: str +) -> Any: + """Construct a fresh (untrained) dlsia TUNet for *n_classes* output channels. + + ``image_shape`` is fixed to ``(image_size, image_size)`` at construction — + dlsia's TUNet precomputes exact per-layer tensor sizes from it, so every + training and inference image must be letterboxed to this same size. + """ + from dlsia.core.networks.tunet import TUNet # noqa: PLC0415 — optional dependency + + model = TUNet( + image_shape=(image_size, image_size), + in_channels=3, + out_channels=n_classes, + depth=depth, + base_channels=base_channels, + growth_rate=growth_rate, + ) + return model.to(device) + + +def network_dict(model: Any) -> dict[str, Any]: + """The full ``{topo_dict, state_dict}`` dlsia uses to reconstruct a TUNet + (see ``TUNet.save_network_parameters`` / ``TUNetwork_from_file``). Stored + as-is in this run's ``adapter.pt`` — TUNet trains from scratch, so there's + no base/delta split the way LoRA has.""" + return model.save_network_parameters(name=None) + + +def load_model(state: dict[str, Any], device: str) -> Any: + """Reconstruct a TUNet from a saved :func:`network_dict`.""" + from dlsia.core.networks.tunet import TUNet # noqa: PLC0415 + + model = TUNet(**state["topo_dict"]) + model.load_state_dict(state["state_dict"]) + return model.to(device) + + +def make_forward_fn(model: Any): + """Return ``forward(batch_images) -> logits`` for :func:`train_common.run_training_loop`. + + TUNet's transposed-conv decoder is symmetric with its encoder, so output + spatial size already matches input size — no resizing needed. + """ + + def _forward(batch_images: Any) -> Any: + return model(batch_images) + + return _forward + + +def make_set_train_mode_fn(model: Any): + """Return ``set_train_mode(is_training)`` for :func:`train_common.run_training_loop`. + + TUNet defaults to ``nn.BatchNorm2d`` — its running mean/var should only + update during training, not while computing validation metrics. + """ + + def _set_train_mode(is_training: bool) -> None: + model.train(is_training) + + return _set_train_mode + + +def make_to_tensor_fn(): + """Return ``(rgb_uint8_hwc) -> float CPU tensor (3,H,W)`` scaled to [0, 1]. + + No ImageNet normalisation — the model trains from scratch on this data, + so there's no pretrained-backbone convention to match. + """ + import numpy as np + import torch + + def _to_tensor(rgb_uint8: "np.ndarray") -> Any: + return torch.from_numpy(np.ascontiguousarray(rgb_uint8)).permute(2, 0, 1).float() / 255.0 + + return _to_tensor diff --git a/backend/export_jobs.py b/backend/export_jobs.py index 315fe13..852a05e 100644 --- a/backend/export_jobs.py +++ b/backend/export_jobs.py @@ -7,7 +7,9 @@ """ from __future__ import annotations +import os import threading +import time import uuid from typing import Any @@ -15,6 +17,11 @@ _lock = threading.Lock() _MAX_LOG = 200 +# Per-stage timing is opt-in (off by default in production) — see log()'s own +# doc for why. Read once at import: this flag is meant for a deliberate +# profiling run (locally or in CI), not something toggled mid-process. +_PROFILE = os.getenv("PROFILE_JOBS") == "1" + def new_job(dataset_path: str) -> str: """Register a new export job and return its id.""" @@ -30,10 +37,40 @@ def new_job(dataset_path: str) -> str: "error": None, "dataset_path": dataset_path, "zip_path": None, # internal; surfaced as result.zip_available + "cancel_requested": False, + "_last_log_at": time.monotonic(), # internal; PROFILE_JOBS timing only } return jid +def request_cancel(jid: str) -> bool: + """Ask a running job to stop; returns False if the job is unknown. + + Cooperative: this only sets a flag. Long-running jobs poll + :func:`cancel_requested` between units of work and stop at a clean boundary, + which is what lets a half-finished dataset be discarded rather than left + looking complete. + """ + with _lock: + job = _jobs.get(jid) + if not job: + return False + job["cancel_requested"] = True + return True + + +def cancel_requested(jid: str) -> bool: + """True if cancellation was requested for *jid*. + + Reads the flag directly under the lock rather than through :func:`get_job`, + which copies the whole job dict (log lines included) just to read one bool — + and this is polled on every slice of a running job. + """ + with _lock: + job = _jobs.get(jid) + return bool(job and job.get("cancel_requested")) + + def get_job(jid: str) -> dict | None: """Return a snapshot copy of the job (without the internal zip_path).""" with _lock: @@ -42,6 +79,7 @@ def get_job(jid: str) -> dict | None: return None snap = dict(job) snap.pop("zip_path", None) + snap.pop("_last_log_at", None) return snap @@ -70,10 +108,26 @@ def bump(jid: str, n: int = 1) -> None: def log(jid: str, message: str) -> None: + """Append a log line, called at every natural stage boundary by every job + type (iPred batch jobs, dlsia train/infer). + + When ``PROFILE_JOBS=1`` is set, each line gets an elapsed-since-previous- + log suffix (e.g. ``"Wrote foo (+1.23s)"``) — a cheap per-stage timing + signal with zero new call sites, since every job already calls this at + the boundaries that matter. Off by default: real users/deployments should + see zero timing overhead, not "probably negligible" — this is meant to + run in the test suite (real job round-trips with profiling on) and local + benchmarking, not every production request. + """ with _lock: job = _jobs.get(jid) if not job: return + if _PROFILE: + now = time.monotonic() + elapsed = now - job.get("_last_log_at", now) + job["_last_log_at"] = now + message = f"{message} (+{elapsed:.2f}s)" job["log"].append(message) if len(job["log"]) > _MAX_LOG: del job["log"][: len(job["log"]) - _MAX_LOG] diff --git a/backend/images.py b/backend/images.py index 1ee0723..0f2aefb 100644 --- a/backend/images.py +++ b/backend/images.py @@ -71,31 +71,29 @@ def _minmax(i: int) -> tuple[float, float]: return result -def render_slice( +def normalize_scalar_unit( arr: np.ndarray, opts: dict[str, Any], global_range: tuple[float, float] | None = None, ) -> np.ndarray: - """Render a 2-D or H×W×C array slice to uint8 RGB. + """Map a 2-D scalar array to ``[0, 1]`` via the scale transform + vmin/vmax + normalisation ``render_slice``'s grayscale branch uses, stopping short of + the final colormap/uint8 step. + + Extracted so :mod:`denoise_train` (``_slice_to_gray_uint8``) can train a + denoiser on the exact same intensity pipeline the 2-D canvas renders, + without duplicating this logic and risking the two drifting apart. Args: - arr: 2-D grayscale or H×W×(3|4) colour array. + arr: 2-D grayscale array. opts: :class:`~schemas.RenderOpts`-compatible dict with keys - ``norm``, ``scale``, ``vmin_pct``, ``vmax_pct``, ``cmap``. + ``norm``, ``scale``, ``vmin_pct``, ``vmax_pct``. global_range: ``(vmin, vmax)`` used when ``opts["norm"] == "global"``. Ignored in slice-norm mode. Returns: - uint8 RGB array of shape ``(H, W, 3)``. + ``float64`` array the same shape as *arr*, values in ``[0, 1]``. """ - is_rgb = arr.ndim == 3 - if is_rgb: - rgb = arr[:, :, :3].astype(np.float64) - mn, mx = float(rgb.min()), float(rgb.max()) - if mx > mn: - rgb = (rgb - mn) / (mx - mn) * 255.0 - return np.clip(rgb, 0, 255).astype(np.uint8) - data = arr.astype(np.float64) data = np.nan_to_num(data, nan=0.0, posinf=0.0, neginf=0.0) @@ -122,7 +120,35 @@ def render_slice( data = (data - vmin_abs) / (vmax_abs - vmin_abs) else: data = np.zeros_like(data) - data = np.clip(data, 0.0, 1.0) + return np.clip(data, 0.0, 1.0) + + +def render_slice( + arr: np.ndarray, + opts: dict[str, Any], + global_range: tuple[float, float] | None = None, +) -> np.ndarray: + """Render a 2-D or H×W×C array slice to uint8 RGB. + + Args: + arr: 2-D grayscale or H×W×(3|4) colour array. + opts: :class:`~schemas.RenderOpts`-compatible dict with keys + ``norm``, ``scale``, ``vmin_pct``, ``vmax_pct``, ``cmap``. + global_range: ``(vmin, vmax)`` used when ``opts["norm"] == "global"``. + Ignored in slice-norm mode. + + Returns: + uint8 RGB array of shape ``(H, W, 3)``. + """ + is_rgb = arr.ndim == 3 + if is_rgb: + rgb = arr[:, :, :3].astype(np.float64) + mn, mx = float(rgb.min()), float(rgb.max()) + if mx > mn: + rgb = (rgb - mn) / (mx - mn) * 255.0 + return np.clip(rgb, 0, 255).astype(np.uint8) + + data = normalize_scalar_unit(arr, opts, global_range) cmap = opts.get("cmap", "gray") if cmap == "viridis": diff --git a/backend/infer_jobs.py b/backend/infer_jobs.py new file mode 100644 index 0000000..c9ee6bc --- /dev/null +++ b/backend/infer_jobs.py @@ -0,0 +1,528 @@ +"""Inference job orchestration for the Train tab. + +Runs a saved fine-tuned dlsia TUNet run over requested slices, producing +editable polygon annotations, a colourised overlay preview, and an optional +push of the raw label maps into Tiled (reusing the existing mask-sync +writer). Progress/cancellation share the same :mod:`export_jobs` registry as +training and export jobs. + +DINOv3 LoRA inference is out of scope here (see Phase 5.5) — only +``dlsia_tunet`` runs are supported; ``dlsia_denoiser`` runs are explicitly +refused (a denoiser is 1->1 regression, not classification). + +Prediction label-map convention matches ``tiled_mask_sync``'s semantic masks +exactly (0 = no confident prediction/background, 1..n = predicted class index ++ 1) — NOT the training-time ``IGNORE_INDEX=255`` convention, which is a +different concept (unannotated ground truth, not a model's confidence gate). +""" + +from __future__ import annotations + +import concurrent.futures +import io +import logging +import os +import threading +from typing import Any + +import numpy as np +from fastapi import HTTPException +from PIL import Image as PILImage + +import export_jobs +import train_common +from coco_export import _mask_to_polygons +from schemas import InferRequest + +logger = logging.getLogger(__name__) + +# job_id -> {"run_id", "classes", "source", "kind", "server_uri", "height", +# "width", "label_pngs": {slice_idx: png_bytes}} +# Bounded to the newest few jobs — previews are only needed for the session +# that just ran inference, not forever. +_MAX_CACHED_JOBS = 4 +_cache: dict[str, dict[str, Any]] = {} +_cache_order: list[str] = [] +_cache_lock = threading.Lock() + + +def _cache_put(job_id: str, entry: dict[str, Any]) -> None: + with _cache_lock: + _cache[job_id] = entry + _cache_order.append(job_id) + while len(_cache_order) > _MAX_CACHED_JOBS: + oldest = _cache_order.pop(0) + _cache.pop(oldest, None) + + +def _cache_get(job_id: str) -> dict[str, Any] | None: + with _cache_lock: + return _cache.get(job_id) + + +def _hex_to_rgb(color: str | None) -> tuple[int, int, int]: + if not color or not color.startswith("#"): + return (255, 0, 0) + hex_part = color[1:] + if len(hex_part) in (3, 4): + hex_part = "".join(c * 2 for c in hex_part[:3]) + try: + return tuple(int(hex_part[i : i + 2], 16) for i in (0, 2, 4)) # type: ignore[return-value] # noqa: E203 + except ValueError: + return (255, 0, 0) + + +def _mask_to_polygons_padded(component: np.ndarray) -> list[list[float]]: + """Like ``coco_export._mask_to_polygons``, but safe for a region that + touches the array border. + + ``skimage.measure.find_contours`` only traces a transition *within* the + array — a mask that's ``True`` all the way to an edge (e.g. a dominant + class spanning the whole slice) has no such transition there, so it finds + zero contours (confirmed: a full-frame True mask yields 0 contours, + verified against skimage directly). Padding with a 1px False border + guarantees a transition exists everywhere the shape meets the canvas edge; + the offset is subtracted back out of the resulting coordinates. + """ + padded = np.pad(component, pad_width=1, mode="constant", constant_values=False) + polygons = _mask_to_polygons(padded) + return [[v - 1 for v in flat] for flat in polygons] + + +def _enclosed_area(ring: np.ndarray) -> float: + """Shoelace area enclosed by a closed (n, 2) ring, ignoring winding.""" + x, y = ring[:, 0], ring[:, 1] + return float(abs(np.dot(x, np.roll(y, 1)) - np.dot(y, np.roll(x, 1))) / 2.0) + + +def _vectorize_label_map( + label_map: np.ndarray, + run_classes: list[dict[str, Any]], + min_area: int, + simplify_tol: float, + run_id: str, + slice_idx: int, +) -> list[dict[str, Any]]: + """Connected-component-vectorize a semantic label map (0=bg, 1..n=class) into + polygon shape dicts, one per component per class. + + Holes matter here, not just cosmetically. A label map assigns each pixel to + exactly one class, so a component that surrounds another class (a background + region between grains, a ring around a bore) traces an outer contour *and* + one contour per enclosed region. Emitting those inner contours as separate + solid same-class polygons — as this used to — double-claims those pixels and + makes the outer polygon paint straight over whichever class actually sits + there: shapes are drawn in list order, so a later class's filled outer + contour hides the earlier ones entirely (the label-map preview PNG never + showed this, since it colours one class per pixel by construction). + Encoding them as real ``holes`` keeps every pixel claimed exactly once, so + the imported annotations match the preview regardless of draw order. + """ + from skimage import measure + + shapes: list[dict[str, Any]] = [] + counter = 0 + for c, cls in enumerate(run_classes): + binary = label_map == (c + 1) + if not binary.any(): + continue + labeled = measure.label(binary, connectivity=2) + for region in measure.regionprops(labeled): + if region.area < min_area: + continue + component = labeled == region.label + rings: list[np.ndarray] = [] + for flat_points in _mask_to_polygons_padded(component): + coords = np.asarray(flat_points, dtype=np.float64).reshape(-1, 2) + if simplify_tol > 0: + coords = measure.approximate_polygon(coords, tolerance=simplify_tol) + if len(coords) < 3: + continue + rings.append(coords) + if not rings: + continue + # A connected region has a single outer boundary, so the widest ring + # is it and everything else it traced is enclosed by it. + rings.sort(key=_enclosed_area, reverse=True) + counter += 1 + shape: dict[str, Any] = { + "id": f"pred_{run_id[:8]}_{slice_idx}_{counter}", + "classId": cls["classId"], + "kind": "polygon", + "points": [round(float(v), 2) for v in rings[0].ravel().tolist()], + } + if len(rings) > 1: + shape["holes"] = [ + [round(float(v), 2) for v in ring.ravel().tolist()] for ring in rings[1:] + ] + shapes.append(shape) + return shapes + + +#: Slices processed concurrently in one inference job. Only the actual model +#: forward call is serialized (train_common.GPU_FORWARD_LOCK, held briefly +#: inside _predict_one_slice) — I/O, rendering, and vectorization for +#: different slices run genuinely in parallel across this pool, so CPU-bound +#: work overlaps with the GPU instead of the whole per-slice pipeline +#: serializing behind one lock (the old behavior, when ML_LOCK itself was +#: held for the entire job). Mirrors ipred_batch_jobs.py's pool shape. +_DEFAULT_INFER_CONCURRENCY = 4 + + +def _predict_one_slice( + slice_idx: int, + *, + jid: str, + node: Any, + meta: dict[str, Any], + h: int, + w: int, + render: dict[str, Any], + global_range: Any, + render_slice_fn: Any, + tiled: bool, + image_size: int, + forward_fn: Any, + to_tensor_fn: Any, + device: Any, + request: InferRequest, + run_classes: list[dict[str, Any]], + tiling_logged: threading.Event, +) -> tuple[int, bytes | None, list[dict[str, Any]] | None, str | None]: + """Read + render + (GPU-locked) predict + vectorize one slice on a worker + thread. Returns ``(slice_idx, label_png_bytes, shapes, error)`` — never + raises, so a pool of these can be driven with plain ``future.result()`` + and no per-future try/except at the call site (mirrors + ``ipred_batch_jobs._apply_one_slice``'s contract). + + A ``None`` label_png/shapes pair with no error means the tiled path was + cancelled mid-slice (``tiling.predict_label_map_tiled`` returns ``None`` + when its own ``cancel_cb`` fires) — the caller's overall cancellation + check already covers stopping the job, so this slice simply contributes + nothing rather than needing its own error to report. + """ + import torch + import torch.nn.functional as F + + import arrays as arrays_mod + + try: + arr = arrays_mod.read_slice(node, meta, slice_idx) + rgb = render_slice_fn(arr, render, global_range) + + with train_common.GPU_FORWARD_LOCK, torch.no_grad(): + if tiled: + import tiling + + # Logged once across the whole job, not once per slice — every + # slice shares the same tiling geometry. Race-free because + # this whole block already runs under GPU_FORWARD_LOCK: only + # one worker is ever inside here at a time, so check-then-set + # on tiling_logged can't interleave between two threads. + first_to_log = not tiling_logged.is_set() + if first_to_log: + tiling_logged.set() + label_map = tiling.predict_label_map_tiled( + rgb, + forward_fn=forward_fn, + to_tensor_fn=to_tensor_fn, + window=image_size, + min_confidence=request.min_confidence, + device=device, + cancel_cb=lambda: export_jobs.cancel_requested(jid), + progress_cb=(lambda msg: export_jobs.log(jid, f"tiling: {msg}")) if first_to_log else None, + ) + if label_map is None: # cancelled mid-slice + return slice_idx, None, None, None + else: + img_l, _ = train_common.letterbox(rgb, np.zeros((h, w), dtype=np.uint8), image_size) + batch = to_tensor_fn(img_l).unsqueeze(0).to(device) + + logits = forward_fn(batch)[0] + probs = F.softmax(logits, dim=0) + confidence, pred_class = probs.max(dim=0) + pred_np = pred_class.cpu().numpy() + conf_np = confidence.cpu().numpy() + label_letterboxed = np.where(conf_np >= request.min_confidence, pred_np + 1, 0).astype(np.uint8) + label_map = train_common.unletterbox(label_letterboxed, h, w, image_size) + + shapes = _vectorize_label_map( + label_map, run_classes, request.min_area, request.simplify_tol, request.run_id, slice_idx, + ) + buf = io.BytesIO() + PILImage.fromarray(label_map, mode="L").save(buf, format="PNG") + return slice_idx, buf.getvalue(), shapes, None + except Exception as exc: # noqa: BLE001 — one bad slice must not abort the job + return slice_idx, None, None, str(exc) + + +def run_infer_job(jid: str, request: InferRequest) -> None: + """Background worker: predict + vectorize each requested slice. + + ``train_common.ML_LOCK`` is still held for the whole job (unchanged — + this is what keeps a training run, a denoise bake, or another inference + job from contending for the GPU at the same time as this one). Within + the job, slices are processed by a small bounded worker pool + (``_DEFAULT_INFER_CONCURRENCY``) instead of one at a time: only the + actual model forward call is serialized (``train_common.GPU_FORWARD_LOCK``, + a separate, finer-grained lock — see its own doc), so I/O, rendering, and + vectorization for different slices overlap with the GPU instead of the + old behavior of serializing the entire per-slice pipeline behind + ``ML_LOCK``. + """ + if not train_common.ML_LOCK.acquire(blocking=False): + export_jobs.update( + jid, + state="error", + phase="error", + error="Another training or inference job is already running", + ) + return + try: + export_jobs.update(jid, state="running", phase="loading") + device = train_common.pick_device() + if device is None: + raise RuntimeError("torch is not installed on this server") + + config = train_common.load_run_config(request.run_id) + adapter_state = train_common.load_adapter_state(request.run_id) + run_classes: list[dict[str, Any]] = config["classes"] + image_size = int(config["image_size"]) + render = request.render.model_dump() if request.render is not None else config["render"] + # Predict with the geometry this run was TRAINED with — reading the flag off + # the run (not the request) keeps runs saved before tiling existed on the + # original whole-slice-rescale path, where their weights are valid. + tiled = bool((config.get("hyperparams") or {}).get("tiling", False)) + # Input denoising is read off the RUN, never off the request — same rule + # as `tiled` above. The model was trained on these exact pixels, so + # letting a caller choose differently at predict time would be a silent + # distribution shift: no error, just quietly worse predictions. Runs + # saved before this field existed have no "denoise" key and get plain + # render_slice, which is exactly what they were trained with. + render_slice_fn = train_common.denoising_render_slice_fn(config.get("denoise")) + if config.get("denoise"): + export_jobs.log( + jid, + f"Applying the run's {config['denoise'].get('method')} input denoising " + "(recorded at training time).", + ) + if tiled: + import tiling + + if not tiling.qlty_available(): + raise RuntimeError( + "This run was trained with tiling and needs the 'qlty' package to predict, " + "which is not installed on this server" + ) + + if config["model_family"] == "dlsia_tunet": + import dlsia_runtime as fam + + if not train_common.dlsia_available(): + raise RuntimeError("dlsia is not installed on this server") + model = fam.load_model(adapter_state, device) + model.eval() + forward_fn = fam.make_forward_fn(model) + to_tensor_fn = fam.make_to_tensor_fn() + elif config["model_family"] == "dlsia_denoiser": + # A denoiser is a 1->1 regression model with no class channels, so + # the label-map path below (softmax/argmax over n_classes, then + # vectorising into shapes) is meaningless for it. Refuse clearly + # instead of loading it through dlsia_runtime, which is what an + # unconditional catch-all `else` would do — that would produce a + # shape mismatch deep in the forward pass rather than an explanation. + raise RuntimeError( + "This is a denoiser run, not a segmentation model — it produces a denoised " + "image, not labelled regions. Apply it from the Annotate tab's Denoise panel." + ) + else: + raise RuntimeError(f"Unsupported model family: {config['model_family']!r}") + + import arrays as arrays_mod + import images as images_mod + + node = arrays_mod.resolve_array(request.source, request.kind, request.server_uri) + meta = arrays_mod.array_shape_meta(node) + h, w = meta["height"], meta["width"] + global_range = images_mod._sample_global_stats(node, meta) if render.get("norm") == "global" else None + + export_jobs.set_total(jid, len(request.slice_indices)) + export_jobs.update(jid, phase="predicting") + + # A single mutable dict, referenced (not copied) by the cache entry + # below — workers below only ever add their OWN slice_idx key, so + # concurrent inserts never conflict, and every `preview_png` call + # from another request thread sees results as soon as they land, no + # further _cache_put calls needed (see #15's live-preview ask: this + # is what makes `GET /api/train/infer/preview/{job_id}/{slice_index}` + # servable for early slices while the job is still `running`). + label_pngs: dict[int, bytes] = {} + slices_result: dict[str, list[dict[str, Any]]] = {} + n_shapes = 0 + cancelled = False + tiling_logged = threading.Event() + + _cache_put( + jid, + { + "run_id": request.run_id, + "classes": run_classes, + "source": request.source, + "kind": request.kind, + "server_uri": request.server_uri, + "height": h, + "width": w, + "label_pngs": label_pngs, + }, + ) + + def _publish_result(*, done: bool) -> None: + """Update the job's `result` after every completed slice (not + just once at the end) so the frontend's poll of + `GET /api/train/infer/status/{job_id}` can render a growing + preview slider while state is still `running` — see #15.""" + export_jobs.update( + jid, + state="done" if done else "running", + phase="done" if done else "predicting", + result={ + "run_id": request.run_id, + "classes": run_classes, + "slices": dict(slices_result), + "n_shapes": n_shapes, + "preview_slices": sorted(label_pngs.keys()), + "cancelled": cancelled, + }, + ) + + pool_size = max(1, int(os.getenv("DLSIA_INFER_CONCURRENCY", _DEFAULT_INFER_CONCURRENCY))) + with concurrent.futures.ThreadPoolExecutor(max_workers=pool_size) as pool: + pending = iter(request.slice_indices) + in_flight: dict[concurrent.futures.Future, int] = {} + + def submit_next() -> bool: + si = next(pending, None) + if si is None: + return False + fut = pool.submit( + _predict_one_slice, + si, + jid=jid, node=node, meta=meta, h=h, w=w, render=render, + global_range=global_range, render_slice_fn=render_slice_fn, + tiled=tiled, image_size=image_size, forward_fn=forward_fn, + to_tensor_fn=to_tensor_fn, device=device, request=request, + run_classes=run_classes, tiling_logged=tiling_logged, + ) + in_flight[fut] = si + return True + + for _ in range(pool_size): + if not submit_next(): + break + + while in_flight: + done_futs, _ = concurrent.futures.wait( + in_flight, return_when=concurrent.futures.FIRST_COMPLETED + ) + for fut in done_futs: + del in_flight[fut] + slice_idx, png_bytes, shapes, error = fut.result() + if error is not None: + logger.warning("Inference job %s: slice %d failed (%s)", jid, slice_idx, error) + export_jobs.log(jid, f"slice {slice_idx}: failed ({error})") + elif png_bytes is not None and shapes is not None: + n_shapes += len(shapes) + slices_result[str(slice_idx)] = shapes + label_pngs[slice_idx] = png_bytes + export_jobs.log(jid, f"slice {slice_idx}: {len(shapes)} region(s)") + # else: cancelled mid-slice (tiled path) — nothing to record. + export_jobs.bump(jid, 1) + _publish_result(done=False) + + if export_jobs.cancel_requested(jid): + cancelled = True + # Already-submitted work can't be un-submitted — just stop + # refilling the pool so it drains rather than growing. + continue + for _ in done_futs: + submit_next() + + _publish_result(done=True) + export_jobs.log(jid, "Inference cancelled; partial results kept." if cancelled else "Inference complete.") + except Exception as exc: # noqa: BLE001 — reported as a job error, never a crash + logger.error("Inference job %s failed: %s", jid, exc) + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + finally: + train_common.ML_LOCK.release() + + +def preview_png(job_id: str, slice_index: int) -> bytes: + """Colourise a cached predicted label map into an RGBA overlay PNG. + + Raises: + HTTPException: 404 if the job or slice isn't cached. + """ + entry = _cache_get(job_id) + if entry is None: + raise HTTPException(404, "No cached inference results for this job (they expire after a few jobs)") + png_bytes = entry["label_pngs"].get(slice_index) + if png_bytes is None: + raise HTTPException(404, f"Slice {slice_index} was not part of this inference job") + + label = np.asarray(PILImage.open(io.BytesIO(png_bytes))) + rgba = np.zeros((*label.shape, 4), dtype=np.uint8) + for c, cls in enumerate(entry["classes"]): + mask = label == (c + 1) + if not mask.any(): + continue + r, g, b = _hex_to_rgb(cls.get("color")) + rgba[mask] = (r, g, b, 180) + + buf = io.BytesIO() + PILImage.fromarray(rgba, mode="RGBA").save(buf, format="PNG") + return buf.getvalue() + + +def run_write_tiled_job(jid: str, infer_job_id: str) -> None: + """Background worker: push a completed inference job's label maps into + Tiled as a ``__masks`` sibling container (same writer the + manual "sync masks to Tiled" flow uses).""" + try: + export_jobs.update(jid, state="running", phase="writing") + entry = _cache_get(infer_job_id) + if entry is None: + raise ValueError("No cached inference results for this job (they expire after a few jobs)") + if entry["kind"] != "tiled": + raise ValueError("Inference source was not a Tiled array — nothing to write back") + + import tiled_mask_sync + + classes = entry["classes"] + slice_indices = sorted(entry["label_pngs"].keys()) + semantic = np.stack( + [np.asarray(PILImage.open(io.BytesIO(entry["label_pngs"][i]))) for i in slice_indices], + axis=0, + ) + class_vols = {cls["label"]: ((semantic == (c + 1)) * 255).astype(np.uint8) for c, cls in enumerate(classes)} + legend = [{"id": c + 1, "name": cls["label"], "color": cls.get("color")} for c, cls in enumerate(classes)] + volumes = { + "semantic": semantic, + "class_vols": class_vols, + "slice_indices": slice_indices, + "legend": legend, + } + + export_jobs.set_total(jid, 1) + # "_deep" keeps this in its own container, separate from whatever the + # manual "sync masks to Tiled" action (iPred's fast results) has + # written for the same source — letting both be loaded as independent + # mask layers in the 3-D viewer instead of merging into one. + info = tiled_mask_sync.write_masks_to_tiled( + entry["source"], entry["server_uri"], volumes, classes, container_suffix="_deep", + ) + export_jobs.bump(jid, 1) + export_jobs.update(jid, state="done", phase="done", result=info) + export_jobs.log(jid, f"Wrote predicted masks to {info['path']}.") + except Exception as exc: # noqa: BLE001 + logger.error("Write-to-Tiled job %s failed: %s", jid, exc) + export_jobs.update(jid, state="error", phase="error", error=str(exc)) diff --git a/backend/ingest.py b/backend/ingest.py index efc4777..26612c3 100644 --- a/backend/ingest.py +++ b/backend/ingest.py @@ -28,13 +28,17 @@ from __future__ import annotations import logging +import os import re +import shutil +import tempfile import threading import uuid from pathlib import Path from typing import Any import numpy as np +from fastapi import HTTPException from tiled_clients import api_key_for_uri, get_tiled_client @@ -124,6 +128,38 @@ def _read_array(path: Path) -> np.ndarray: raise ValueError(f"unsupported extension {suffix!r}") +def validate_container_path(container_path: str) -> list[str]: + """Return safe Tiled key segments beneath the configured ingest root. + + A write confinement check, not a formatting one: callers take a destination + from the client, so without this a request could name any container in the + catalog — including one holding unrelated data. Rejects traversal segments, + control characters, and anything that is not a strict descendant of + ``TILED_INGEST_ROOT`` (default ``browse``). + """ + raw = (container_path or "").strip() + if not raw or raw != raw.strip("/"): + raise ValueError("container_path must be a relative canonical Tiled path") + parts = raw.split("/") + if any( + not part + or part in {".", ".."} + or len(part) > 128 + or any(ord(char) < 32 for char in part) + for part in parts + ): + raise ValueError("container_path contains an unsafe segment") + + root_parts = [ + part for part in os.getenv("TILED_INGEST_ROOT", "browse").strip("/").split("/") if part + ] + # `len(parts) <= len(root_parts)` rejects the root itself: a dataset must be + # written *inside* it, never over it. + if not root_parts or parts[: len(root_parts)] != root_parts or len(parts) <= len(root_parts): + raise ValueError("container_path must be a child of the configured ingest root") + return parts + + def _ensure_container(client: Any, parts: list[str]) -> Any: """Navigate to ``client[parts...]``, creating containers as needed. @@ -333,6 +369,7 @@ def run_ingest_job( container_meta: dict[str, Any] = { "sample_name": parts[-1] if parts else "", "n_images": len(temp_files), + "source_format": IMAGE_STACK_SOURCE_FORMAT, } if description: container_meta["description"] = description @@ -403,3 +440,188 @@ def run_ingest_job( tmp.unlink(missing_ok=True) except Exception: # noqa: BLE001 pass + + +def node_source_kind(node: Any) -> str | None: + """The ``source_format`` tag on *node*'s metadata, or None if absent/unreadable. + + Used to tell whether an existing Tiled node at a candidate key is the SAME + kind of registration a scan is about to (re-)create, vs. an unrelated + dataset that happens to share the same stem name (e.g. a raw image folder + and its own already-registered Zarr reconstruction) — so a scan can report + that collision honestly as "shadowed" instead of a misleading "already + registered". Shared by :mod:`zarr_source`'s own scan function too. + """ + try: + return dict(getattr(node, "metadata", {}) or {}).get("source_format") + except Exception: # noqa: BLE001 — best-effort + return None + + +#: Tag written on every container this module ingests into — lets a later scan +#: (of any kind) recognize "this key already holds a real per-slice ingest" +#: and distinguish it from an external Zarr/tiff-stack-3d registration sharing +#: the same stem name. +IMAGE_STACK_SOURCE_FORMAT = "image-stack" + + +def _is_image_stack_dir(path: Path) -> bool: + """True if *path* is a directory holding 2+ directly-supported image files. + + A single image isn't a "stack" worth its own ingest — mirrors + :func:`tiff_stack_source.inspect_tiff_stack`'s own ``len(files) < 2`` rule, + generalized to every :data:`IMAGE_EXTS` type, not just TIFF. + """ + if not path.is_dir(): + return False + count = 0 + for child in path.iterdir(): + if child.is_file() and child.suffix.lower() in IMAGE_EXTS: + count += 1 + if count >= 2: + return True + return False + + +def scan_and_register_image_stacks( + server_uri: str | None, + scan_root: str, + container_path: str = "browse", + on_conflict: str = "skip", + renames: dict[str, str] | None = None, +) -> dict[str, Any]: + """Walk *scan_root* for folders of image slices and ingest each not already + present, as fast per-slice registration — no 3-D pyramid is built here. + + Deliberately the "quick" path: each slice is copied into Tiled as its own + 2-D node (exactly what dropping the folder onto the Connect page's + dropzone already does), with no multiscale volume generated. Building a + 3-D pyramid is comparatively expensive (reads every slice at least once) + and is left to the existing, on-demand "Build 3D volume" button on the 3D + page (:mod:`volume_build`) — a discovery scan should not silently take + minutes per folder. + + Only top-level subdirectories of *scan_root* are candidates — one level, + matching :func:`zarr_source.scan_and_register_zarrs`'s own non-recursive + scope. A `.zarr` directory is never a candidate here even if it happens to + also contain 2+ loose image files at its root (it shouldn't, but the check + is cheap insurance against double-registering the same data two ways). + + A raw image folder and an unrelated dataset (e.g. a Zarr reconstruction of + the same acquisition) commonly share the same stem name — in which case + they'd derive the identical Tiled key. Rather than silently treating that + as "already registered" (misleading — the raw slices were never actually + ingested) or blindly replacing someone else's data, a same-key collision + with a DIFFERENT kind of registration (per the ``source_format`` tag; see + :func:`node_source_kind`) is reported as **shadowed**, distinctly from a + same-kind ``skipped`` match from a previous run of this same scan. + + Args: + server_uri: Connected Tiled server URI. + scan_root: Absolute directory to scan. + container_path: Parent container each discovered folder registers + under, keyed by its own (sanitized) folder name — matching how a + dropped folder is named today. + on_conflict: ``"skip"`` (default, safe to re-run) or ``"replace"`` — + applies only to a same-kind (not shadowed) collision. + renames: Optional ``{folder_name: alternate_key}`` override, so a + shadowed candidate can be retried under a different key without + re-scanning everything else. + + Returns: + Dict with ``scanned``, ``registered`` (``name``/``key``/``tiled_path``), + ``skipped`` (names), ``shadowed`` (``name``/``key``/``existing_kind``/ + ``suggested_key`` — a different-kind collision, nothing registered), + and ``errors`` (``name``/``error`` pairs). + + Raises: + HTTPException: 400/404 if *scan_root* itself is unusable. + """ + from tiled.client.register import Settings + + import zarr_source + + if on_conflict not in ("skip", "replace"): + on_conflict = "skip" + renames = renames or {} + + root = Path(scan_root).expanduser() + if not root.is_absolute(): + raise HTTPException(400, f"Path must be absolute: {scan_root!r}") + if not root.is_dir(): + raise HTTPException(404, f"No such directory: {root}") + + candidates = sorted( + (p for p in root.iterdir() if not zarr_source._is_zarr_dir(p) and _is_image_stack_dir(p)), + key=lambda p: p.name, + ) + + client = get_tiled_client(server_uri, api_key_for_uri(server_uri)) + parts = [p for p in container_path.strip("/").split("/") if p] + target = _ensure_container(client, parts) + + registered: list[dict[str, Any]] = [] + skipped: list[str] = [] + shadowed: list[dict[str, str]] = [] + errors: list[dict[str, str]] = [] + for candidate in candidates: + default_key = Settings.init().key_from_filename(candidate.name) + key = renames.get(candidate.name, default_key) + existing_keys = _child_keys(target) + if key in existing_keys: + existing_kind = node_source_kind(target[key]) + if existing_kind != IMAGE_STACK_SOURCE_FORMAT: + shadowed.append( + { + "name": candidate.name, + "key": key, + "existing_kind": existing_kind or "unknown", + "suggested_key": f"{key}_images", + } + ) + continue + if on_conflict == "skip": + skipped.append(key) + continue + try: + target.delete_contents(key, recursive=True, external_only=False) + except Exception as exc: # noqa: BLE001 — one bad entry must not sink the scan + errors.append({"name": candidate.name, "error": str(exc)}) + continue + + # Copy each slice to a REAL temp file before ingesting — run_ingest_job + # unlinks its inputs when done (correct for its normal caller, the + # upload route's own tempfile.mkdtemp()); pointing it at scan_root's + # actual files would delete the user's real source data. + tmp_dir = Path(tempfile.mkdtemp(prefix="scan_ingest_")) + try: + temp_files: list[tuple[str, Path]] = [] + for index, src in enumerate( + sorted(p for p in candidate.iterdir() if p.is_file() and p.suffix.lower() in IMAGE_EXTS) + ): + dest = tmp_dir / f"{index:06d}{src.suffix.lower()}" + shutil.copy2(src, dest) + temp_files.append((src.name, dest)) + + jid = new_job(len(temp_files), server_uri, f"{container_path}/{key}".strip("/")) + run_ingest_job(jid, server_uri, f"{container_path}/{key}".strip("/"), temp_files, on_conflict=on_conflict) + job = get_job(jid) or {} + if job.get("state") == "error": + errs = job.get("errors") or [{"error": "ingest failed"}] + errors.append({"name": candidate.name, "error": errs[0].get("error", "ingest failed")}) + else: + registered.append( + {"name": candidate.name, "key": key, "tiled_path": f"{container_path}/{key}".strip("/")} + ) + except Exception as exc: # noqa: BLE001 — one bad folder must not sink the scan + errors.append({"name": candidate.name, "error": str(exc)}) + finally: + shutil.rmtree(tmp_dir, ignore_errors=True) + + return { + "scanned": len(candidates), + "registered": registered, + "skipped": skipped, + "shadowed": shadowed, + "errors": errors, + } diff --git a/backend/ipred_batch_jobs.py b/backend/ipred_batch_jobs.py new file mode 100644 index 0000000..e490cd6 --- /dev/null +++ b/backend/ipred_batch_jobs.py @@ -0,0 +1,276 @@ +"""Background jobs for sample-scale iPred operations: multi-slice train and +whole-volume apply. + +Both dispatch onto the shared ``export_jobs`` registry — mirrors +``denoise_bake.py``'s worker pattern exactly (thread target, cooperative +cancel, bump/log per unit of work). Poll ``GET /api/export/status/{job_id}``; +the routes that spawn these live in ``ipred_routes.py``. + +The apply job deliberately stops at "run inference per slice" and does NOT +vectorize the conformal commit/status maps into polygon shapes server-side — +that stays client-side (``AnnotatePage.tsx`` + ``pixelClf.ts``), reusing the +already-proven, already-tested polygon tracer instead of porting it to Python. +""" + +from __future__ import annotations + +import concurrent.futures +import logging +import os +from typing import Any + +import httpx + +import export_jobs +import ipred_client as ipred_client_mod + +logger = logging.getLogger(__name__) + + +def run_ipred_multi_train_job( + jid: str, + *, + session_id: str, + per_slice_shapes: dict[int, list[dict[str, Any]]], + composition_id: str | None, + feature_setup_id: str | None, + trainer_id: str, + config: dict[str, Any] | None, +) -> None: + """Preprocess every given slice (feature banks are cache-aware, so an + already-computed slice is cheap), then train ONE model pooling all their + labeled pixels via ``POST /train/multi``. + + A slice's preprocess failing aborts the whole job rather than being + skipped — unlike the apply job below, a pooled model that silently + dropped one of the slices the user asked for isn't the model they asked + to train, so this fails loudly instead. + """ + try: + export_jobs.update(jid, state="running", phase="preprocessing") + slice_indices = sorted(per_slice_shapes) + # +1 unit for the final pooled-train phase, so the bar doesn't read + # 100% while the (potentially slow) pooled fit is still running. + export_jobs.set_total(jid, len(slice_indices) + 1) + + feature_ids: dict[str, str] = {} + for slice_index in slice_indices: + if export_jobs.cancel_requested(jid): + export_jobs.update(jid, state="error", phase="cancelled", error="Cancelled.") + return + bank = ipred_client_mod.preprocess( + session_id=session_id, + feature_setup_id=feature_setup_id, + composition_id=composition_id, + slice_index=slice_index, + ) + feature_ids[str(slice_index)] = bank["feature_id"] + export_jobs.bump(jid, 1) + cache_note = "cached" if bank.get("cache_hit") else "computed" + export_jobs.log(jid, f"slice {slice_index}: features ready ({cache_note})") + + if export_jobs.cancel_requested(jid): + export_jobs.update(jid, state="error", phase="cancelled", error="Cancelled.") + return + + export_jobs.update(jid, phase="training") + result = ipred_client_mod.train_multi( + session_id=session_id, + slices={str(k): v for k, v in per_slice_shapes.items()}, + feature_ids=feature_ids, + trainer_id=trainer_id, + config=config, + ) + export_jobs.bump(jid, 1) + export_jobs.update(jid, state="done", phase="done", result=result) + export_jobs.log( + jid, + f"Trained on {len(slice_indices)} slice(s), " + f"{result.get('n_samples', 0):,} pooled samples.", + ) + except Exception as exc: # noqa: BLE001 — reported as a job error, never a crash + logger.exception("ipred multi-slice train job %s failed", jid) + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + + +#: Slices processed concurrently in a volume-wide apply job. Each in-flight +#: slice holds ~1 GB (its feature bank, released immediately after use — see +#: `Catalog.delete_feature_bank`) resident on the ipred service process, so +#: this directly multiplies peak memory there. Overridable per machine +#: without a code change via IPRED_APPLY_CONCURRENCY; kept conservative by +#: default. `preprocess`/`infer` are sync FastAPI routes, which Starlette +#: already runs in its own thread pool even under a single uvicorn process — +#: nothing on the ipred service side needs to change to make this safe (see +#: the encoder's own `_session_lock`/`_model_lock`, which already serialize +#: concurrent ONNX access correctly). +_DEFAULT_APPLY_CONCURRENCY = 4 + + +def _apply_one_slice( + slice_index: int, + *, + client: httpx.Client, + session_id: str, + model_id: str, + composition_id: str | None, + feature_setup_id: str | None, + alpha: float, +) -> tuple[int, str | None, str | None]: + """Preprocess + infer one slice on a worker thread. Returns + ``(slice_index, run_id, error)`` — exactly one of the last two is set — + never raises, so a pool of these can be driven with plain + ``future.result()`` and no per-future try/except at the call site.""" + bank: dict[str, Any] | None = None + try: + bank = ipred_client_mod.preprocess( + session_id=session_id, + feature_setup_id=feature_setup_id, + composition_id=composition_id, + slice_index=slice_index, + client=client, + ) + run = ipred_client_mod.infer( + session_id=session_id, + model_id=model_id, + feature_id=bank["feature_id"], + alpha=alpha, + # Volume-apply never reads a run's proba.npy back (only + # commit.png, client-side) — see run_infer's own doc for why + # this is safe to skip here but not for interactive infer. + store_probabilities=False, + client=client, + ) + return slice_index, run["run_id"], None + except Exception as exc: # noqa: BLE001 — one bad slice must not abort the volume + return slice_index, None, str(exc) + finally: + # Release the feature bank whether infer succeeded or failed — + # a failed infer still leaves an orphaned ~1 GB bank behind + # otherwise. Best-effort: a cleanup failure must not fail the + # slice that already predicted successfully. + if bank is not None: + try: + ipred_client_mod.delete_feature_bank(bank["feature_id"], client=client) + except Exception as exc: # noqa: BLE001 + logger.warning( + "ipred volume apply: could not release feature bank for slice %d (%s)", + slice_index, exc, + ) + + +def run_ipred_volume_apply_job( + jid: str, + *, + session_id: str, + model_id: str, + slice_indices: list[int], + composition_id: str | None, + feature_setup_id: str | None, + alpha: float, +) -> None: + """Run inference on every given slice, ensuring each has a feature bank + first. Stores ``result.runs = {slice_index: run_id}`` for the frontend to + fetch commit/status PNGs from and vectorize per slice. + + Unlike the train job, one slice failing does not abort the rest — a + volume-wide apply that stops at the first unreadable slice would be far + more disruptive than a job that finishes with a handful of gaps reported + in ``result.errors``. + + Slices are processed by a small bounded thread pool (see + ``_DEFAULT_APPLY_CONCURRENCY``) sharing one persistent HTTP connection + (``ipred_client_mod.new_shared_client``) instead of one-at-a-time with a + fresh connection per call — see the plan doc for why each piece of + shared state here (export_jobs' lock, Catalog's per-call connections, + CatBoost's read-only predict, the encoder's existing session locks) is + already safe under this concurrency. + """ + try: + export_jobs.update(jid, state="running", phase="predicting") + export_jobs.set_total(jid, len(slice_indices)) + + runs: dict[str, str] = {} + errors: list[dict[str, Any]] = [] + cancelled = False + pool_size = max(1, int(os.getenv("IPRED_APPLY_CONCURRENCY", _DEFAULT_APPLY_CONCURRENCY))) + + def _publish_result() -> None: + """Update the job's `result` after every completed slice, not + just once at the end — otherwise a UI polling this job while it + runs sees `result: null` the entire time, no matter how many + slices have actually already predicted (same class of bug fixed + in infer_jobs.py's dlsia inference job for the same reason).""" + export_jobs.update(jid, result={"runs": dict(runs), "errors": list(errors), "cancelled": cancelled}) + + with ipred_client_mod.new_shared_client() as client, \ + concurrent.futures.ThreadPoolExecutor(max_workers=pool_size) as pool: + pending = iter(slice_indices) + in_flight: dict[concurrent.futures.Future, int] = {} + + def submit_next() -> bool: + si = next(pending, None) + if si is None: + return False + fut = pool.submit( + _apply_one_slice, + si, + client=client, + session_id=session_id, + model_id=model_id, + composition_id=composition_id, + feature_setup_id=feature_setup_id, + alpha=alpha, + ) + in_flight[fut] = si + return True + + for _ in range(pool_size): + if not submit_next(): + break + + while in_flight: + done, _ = concurrent.futures.wait( + in_flight, return_when=concurrent.futures.FIRST_COMPLETED + ) + for fut in done: + del in_flight[fut] + slice_index, run_id, error = fut.result() + if error is not None: + logger.warning("ipred volume apply: slice %d failed (%s)", slice_index, error) + errors.append({"slice": slice_index, "error": error}) + else: + runs[str(slice_index)] = run_id # type: ignore[assignment] + export_jobs.log(jid, f"slice {slice_index}: predicted (run {run_id[:8]})") + export_jobs.bump(jid, 1) + _publish_result() + + if export_jobs.cancel_requested(jid): + cancelled = True + # Already-submitted work can't be un-submitted (Python + # threads have no preemptive cancel) — just stop + # refilling the pool so it drains rather than growing. + continue + # Refill one slot per completed future, keeping ~pool_size + # in flight rather than waiting for a full batch to finish. + for _ in done: + submit_next() + + result = {"runs": runs, "errors": errors, "cancelled": cancelled} + if not runs: + export_jobs.update( + jid, + state="error", + phase="error", + error="No slices could be predicted.", + result=result, + ) + return + export_jobs.update(jid, state="done", phase="done", result=result) + export_jobs.log( + jid, + f"{'Cancelled after' if cancelled else 'Applied to'} " + f"{len(runs)}/{len(slice_indices)} slice(s).", + ) + except Exception as exc: # noqa: BLE001 — reported as a job error, never a crash + logger.exception("ipred volume apply job %s failed", jid) + export_jobs.update(jid, state="error", phase="error", error=str(exc)) diff --git a/backend/ipred_client.py b/backend/ipred_client.py new file mode 100644 index 0000000..24c7799 --- /dev/null +++ b/backend/ipred_client.py @@ -0,0 +1,359 @@ +"""HTTP client for the disjoint ipred backend (no Python imports of ``ipred``).""" + +from __future__ import annotations + +import logging +import os +from contextlib import contextmanager +from typing import Any, Iterator + +import httpx + +logger = logging.getLogger(__name__) + +DEFAULT_IPRED_URL = "http://127.0.0.1:8003" + + +def ipred_url() -> str: + """Base URL for the iterative prediction backend.""" + # Prefer IPRED_URL; accept legacy CLF_ENGINE_URL during transition. + return os.getenv( + "IPRED_URL", + os.getenv("CLF_ENGINE_URL", DEFAULT_IPRED_URL), + ).rstrip("/") + + +def _client(timeout: float = 300.0) -> httpx.Client: + return httpx.Client(base_url=ipred_url(), timeout=timeout) + + +def new_shared_client(timeout: float = 300.0) -> httpx.Client: + """A `httpx.Client` a caller owns and reuses across many calls (e.g. one + per slice in a volume-wide batch-apply job), instead of the fresh + open/close-per-call `_client()` every other function here defaults to. + `httpx.Client` is documented safe for concurrent use from multiple + threads, so this is also what a thread-pooled batch job should share. + Caller is responsible for closing it (a `with` block or `.close()`). + """ + return httpx.Client(base_url=ipred_url(), timeout=timeout) + + +@contextmanager +def _use_client(client: httpx.Client | None, timeout: float = 300.0) -> Iterator[httpx.Client]: + """Yield `client` if given (never closing it — the caller owns its + lifecycle), else open-and-close a fresh one exactly like every call site + here did before `client=` params existed. Keeps every function's + single-call default behavior unchanged for callers that don't pass one.""" + if client is not None: + yield client + return + with _client(timeout=timeout) as owned: + yield owned + + +def health() -> dict[str, Any]: + """GET /health.""" + with _client(timeout=5.0) as client: + r = client.get("/health") + r.raise_for_status() + return r.json() + + +def open_session( + *, + kind: str, + source: str, + server_uri: str | None = None, + root: str | None = None, +) -> dict[str, Any]: + """POST /sessions.""" + with _client() as client: + r = client.post( + "/sessions", + json={ + "kind": kind, + "source": source, + "server_uri": server_uri, + "root": root, + }, + ) + r.raise_for_status() + return r.json() + + +def list_setups() -> list[dict[str, Any]]: + """GET /setups.""" + with _client() as client: + r = client.get("/setups") + r.raise_for_status() + return list(r.json().get("setups") or []) + + +def get_setup(setup_id: str) -> dict[str, Any]: + """GET /setups/{id}.""" + with _client() as client: + r = client.get(f"/setups/{setup_id}") + r.raise_for_status() + return r.json() + + +def upsert_setup(payload: dict[str, Any]) -> dict[str, Any]: + """POST /setups.""" + with _client() as client: + r = client.post("/setups", json=payload) + r.raise_for_status() + return r.json() + + +def list_trainers() -> list[str]: + """GET /trainers.""" + with _client() as client: + r = client.get("/trainers") + r.raise_for_status() + return list(r.json().get("trainers") or []) + + +def list_modules() -> list[dict[str, Any]]: + """GET /modules.""" + with _client() as client: + r = client.get("/modules") + r.raise_for_status() + return list(r.json().get("modules") or []) + + +def list_compositions() -> list[dict[str, Any]]: + """GET /compositions.""" + with _client() as client: + r = client.get("/compositions") + r.raise_for_status() + return list(r.json().get("compositions") or []) + + +def get_composition(composition_id: str) -> dict[str, Any]: + """GET /compositions/{id}.""" + with _client() as client: + r = client.get(f"/compositions/{composition_id}") + r.raise_for_status() + return r.json() + + +def upsert_composition(payload: dict[str, Any]) -> dict[str, Any]: + """POST /compositions.""" + with _client() as client: + r = client.post("/compositions", json=payload) + r.raise_for_status() + return r.json() + + +def preview_composition(payload: dict[str, Any]) -> dict[str, Any]: + """POST /compositions/preview.""" + with _client() as client: + r = client.post("/compositions/preview", json=payload) + r.raise_for_status() + return r.json() + + +def upload_session_array(session_id: str, payload: dict[str, Any]) -> dict[str, Any]: + """POST /sessions/{id}/arrays.""" + with _client() as client: + r = client.post(f"/sessions/{session_id}/arrays", json=payload) + r.raise_for_status() + return r.json() + + +def preprocess( + *, + session_id: str, + feature_setup_id: str | None = None, + composition_id: str | None = None, + slice_index: int = 0, + array_ref: str | None = None, + client: httpx.Client | None = None, +) -> dict[str, Any]: + """POST /preprocess. Pass `client` (see `new_shared_client`) to reuse a + connection across many calls instead of opening a fresh one each time.""" + body: dict[str, Any] = { + "session_id": session_id, + "slice_index": slice_index, + } + if composition_id: + body["composition_id"] = composition_id + if feature_setup_id: + body["feature_setup_id"] = feature_setup_id + if array_ref: + body["array_ref"] = array_ref + with _use_client(client) as c: + r = c.post("/preprocess", json=body) + r.raise_for_status() + return r.json() + + +def delete_feature_bank(feature_id: str, *, client: httpx.Client | None = None) -> dict[str, Any]: + """DELETE /features/{id} — release a feature bank's disk blob once nothing + later in the current job needs it (see ipred_batch_jobs.py's volume-apply + job, the only caller).""" + with _use_client(client) as c: + r = c.delete(f"/features/{feature_id}") + r.raise_for_status() + return r.json() + + +def feature_channel_bytes(feature_id: str, index: int) -> bytes: + """GET /features/{id}/channels/{index}.""" + with _client() as client: + r = client.get(f"/features/{feature_id}/channels/{index}") + r.raise_for_status() + return r.content + + +def train( + *, + session_id: str, + shapes: list[dict[str, Any]], + feature_id: str | None = None, + trainer_id: str = "catboost", + config: dict[str, Any] | None = None, +) -> dict[str, Any]: + """POST /train.""" + with _client() as client: + r = client.post( + "/train", + json={ + "session_id": session_id, + "shapes": shapes, + "feature_id": feature_id, + "trainer_id": trainer_id, + "config": config or {}, + }, + ) + r.raise_for_status() + return r.json() + + +def train_multi( + *, + session_id: str, + slices: dict[str, list[dict[str, Any]]], + feature_ids: dict[str, str], + trainer_id: str = "catboost", + config: dict[str, Any] | None = None, +) -> dict[str, Any]: + """POST /train/multi — pool labeled pixels across multiple slices into one model.""" + with _client() as client: + r = client.post( + "/train/multi", + json={ + "session_id": session_id, + "slices": slices, + "feature_ids": feature_ids, + "trainer_id": trainer_id, + "config": config or {}, + }, + ) + r.raise_for_status() + return r.json() + + +def infer( + *, + session_id: str, + model_id: str | None = None, + feature_id: str | None = None, + alpha: float = 0.05, + store_probabilities: bool = True, + client: httpx.Client | None = None, +) -> dict[str, Any]: + """POST /infer. `store_probabilities=False` skips persisting the run's + full probability array — see `train_infer.run_infer`'s doc for why a + volume-wide batch-apply job (the only caller that passes `False`) never + needs it. Pass `client` (see `new_shared_client`) to reuse a connection + across many calls instead of opening a fresh one each time.""" + with _use_client(client) as c: + r = c.post( + "/infer", + json={ + "session_id": session_id, + "model_id": model_id, + "feature_id": feature_id, + "alpha": alpha, + "store_probabilities": store_probabilities, + }, + ) + r.raise_for_status() + return r.json() + + +def rethreshold( + *, + session_id: str, + alpha: float, + run_id: str | None = None, +) -> dict[str, Any]: + """POST /rethreshold.""" + with _client() as client: + r = client.post( + "/rethreshold", + json={ + "session_id": session_id, + "alpha": alpha, + "run_id": run_id, + }, + ) + r.raise_for_status() + return r.json() + + +def run_commit_png(run_id: str) -> bytes: + """GET /runs/{id}/commit.png.""" + with _client() as client: + r = client.get(f"/runs/{run_id}/commit.png") + r.raise_for_status() + return r.content + + +def run_status_png(run_id: str) -> bytes: + """GET /runs/{id}/status.png.""" + with _client() as client: + r = client.get(f"/runs/{run_id}/status.png") + r.raise_for_status() + return r.content + + +def run_proba_png(run_id: str, class_index: int) -> bytes: + """GET /runs/{id}/proba/{class_index}.png.""" + with _client() as client: + r = client.get(f"/runs/{run_id}/proba/{int(class_index)}.png") + r.raise_for_status() + return r.content + + +def threshold_class_map( + run_id: str, + *, + class_id: int, + threshold: float, +) -> dict[str, Any]: + """POST /runs/{id}/threshold-class.""" + with _client() as client: + r = client.post( + f"/runs/{run_id}/threshold-class", + json={"class_id": int(class_id), "threshold": float(threshold)}, + ) + r.raise_for_status() + return r.json() + + +def manifold_sample(payload: dict[str, Any]) -> dict[str, Any]: + """POST /manifold/sample.""" + with _client() as client: + r = client.post("/manifold/sample", json=payload) + r.raise_for_status() + return r.json() + + +def manifold_heatmap_png(sample_id: str) -> bytes: + """GET /manifold/{id}/heatmap.png.""" + with _client() as client: + r = client.get(f"/manifold/{sample_id}/heatmap.png") + r.raise_for_status() + return r.content diff --git a/backend/ipred_routes.py b/backend/ipred_routes.py new file mode 100644 index 0000000..08b8697 --- /dev/null +++ b/backend/ipred_routes.py @@ -0,0 +1,456 @@ +"""Proxy routes for the disjoint ipred service (GUI never calls port 8003 directly). + +All routes live under ``/api/ipred/*`` and are thin pass-throughs into +``ipred_client``, which talks to the standalone iPred FastAPI service over +HTTP. A connection failure surfaces as 503 rather than a 500/toast, so the +frontend can show a clear "not running" state. +""" + +from __future__ import annotations + +import threading +from typing import Any, Optional + +import httpx +from fastapi import APIRouter, HTTPException +from fastapi.responses import Response +from pydantic import BaseModel + +import export_jobs +import ipred_batch_jobs +import ipred_client as ipred_client_mod + +router = APIRouter(prefix="/api/ipred") + + +class IpredSessionRequest(BaseModel): + """Open an ipred session for one image/stack project.""" + + kind: str + source: str + server_uri: Optional[str] = None + root: Optional[str] = None + + +class IpredPreprocessRequest(BaseModel): + """Proxy featurize request.""" + + session_id: str + feature_setup_id: Optional[str] = None + composition_id: Optional[str] = None + slice_index: int = 0 + array_ref: Optional[str] = None + + +class IpredCompositionUpsertRequest(BaseModel): + """Create/update composition on ipred.""" + + name: str + nodes: list[dict[str, Any]] + outputs: list[str] + composition_id: Optional[str] = None + builtin: bool = False + + +class IpredArrayUploadRequest(BaseModel): + """Upload float array for remote ipred preprocess.""" + + session_id: str + shape: list[int] + dtype: str = "float32" + data_b64: str + array_ref: Optional[str] = None + + +class IpredTrainRequest(BaseModel): + """Proxy train request.""" + + session_id: str + shapes: list[dict[str, Any]] + feature_id: Optional[str] = None + trainer_id: str = "catboost" + config: Optional[dict[str, Any]] = None + + +class IpredInferRequest(BaseModel): + """Proxy infer request.""" + + session_id: str + model_id: Optional[str] = None + feature_id: Optional[str] = None + alpha: float = 0.05 + + +class IpredRethresholdRequest(BaseModel): + """Proxy rethreshold request.""" + + session_id: str + alpha: float + run_id: Optional[str] = None + + +class IpredSetupUpsertRequest(BaseModel): + """Create or update a Feature Setup on ipred.""" + + name: str + kind: str + procedure_id: Optional[str] = None + params: Optional[dict[str, Any]] = None + encoder_setup_id: Optional[str] = None + weights_path: Optional[str] = None + weights_format: Optional[str] = None + inference: Optional[dict[str, Any]] = None + setup_id: Optional[str] = None + + +class IpredManifoldSampleRequest(BaseModel): + """Suggest Labels via ipred feature bank.""" + + feature_id: str + k: int = 24 + box_size: Optional[int] = None + stride: Optional[int] = None + pca_dims: int = 16 + shapes: Optional[list[dict[str, Any]]] = None + + +class IpredThresholdClassRequest(BaseModel): + class_id: int + threshold: float = 0.5 + + +class IpredBatchTrainRequest(BaseModel): + """Train one model pooling labeled pixels across multiple slices.""" + + session_id: str + slices: dict[str, list[dict[str, Any]]] + composition_id: Optional[str] = None + feature_setup_id: Optional[str] = None + trainer_id: str = "catboost" + config: Optional[dict[str, Any]] = None + + +class IpredBatchApplyRequest(BaseModel): + """Run inference across a set of slices (e.g. the whole volume).""" + + session_id: str + model_id: str + slice_indices: list[int] + composition_id: Optional[str] = None + feature_setup_id: Optional[str] = None + alpha: float = 0.05 + + +def _ipred_http_error(exc: Exception) -> HTTPException: + """Translate an ipred_client exception into the right HTTP response.""" + if isinstance(exc, httpx.HTTPStatusError): + detail: Any = exc.response.text + try: + detail = exc.response.json() + except Exception: + pass + return HTTPException(exc.response.status_code, detail) + if isinstance(exc, httpx.ConnectError): + return HTTPException(503, f"ipred unreachable at {ipred_client_mod.ipred_url()}") + return HTTPException(500, str(exc)) + + +@router.get("/health") +async def ipred_health() -> dict: + """Liveness of the disjoint ipred.""" + try: + return ipred_client_mod.health() + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/sessions") +async def ipred_open_session(body: IpredSessionRequest) -> dict: + """Open/create an engine session for a project.""" + try: + return ipred_client_mod.open_session( + kind=body.kind, + source=body.source, + server_uri=body.server_uri, + root=body.root, + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/setups") +async def ipred_list_setups() -> dict: + """List Feature Setups from the ipred.""" + try: + return {"setups": ipred_client_mod.list_setups()} + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/setups/{setup_id}") +async def ipred_get_setup(setup_id: str) -> dict: + """Get one Feature Setup from ipred.""" + try: + return ipred_client_mod.get_setup(setup_id) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/setups") +async def ipred_upsert_setup(body: IpredSetupUpsertRequest) -> dict: + """Create or update a Feature Setup on ipred.""" + try: + return ipred_client_mod.upsert_setup(body.model_dump(exclude_none=True)) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/trainers") +async def ipred_list_trainers() -> dict: + """List trainer plugin ids from ipred.""" + try: + return {"trainers": ipred_client_mod.list_trainers()} + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/modules") +async def ipred_list_modules() -> dict: + """List composable feature modules (each reports `runtime` + `ready`).""" + try: + return {"modules": ipred_client_mod.list_modules()} + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/compositions") +async def ipred_list_compositions() -> dict: + """List feature compositions.""" + try: + return {"compositions": ipred_client_mod.list_compositions()} + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/compositions/{composition_id}") +async def ipred_get_composition(composition_id: str) -> dict: + """Get one composition.""" + try: + return ipred_client_mod.get_composition(composition_id) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/compositions") +async def ipred_upsert_composition(body: IpredCompositionUpsertRequest) -> dict: + """Create or update a composition.""" + try: + return ipred_client_mod.upsert_composition(body.model_dump(exclude_none=True)) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/compositions/preview") +async def ipred_preview_composition(body: IpredCompositionUpsertRequest) -> dict: + """Preview concat labels for a composition draft.""" + try: + return ipred_client_mod.preview_composition(body.model_dump(exclude_none=True)) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/sessions/{session_id}/arrays") +async def ipred_upload_array(session_id: str, body: IpredArrayUploadRequest) -> dict: + """Upload an array blob to ipred for remote preprocess.""" + try: + payload = body.model_dump(exclude_none=True) + payload["session_id"] = session_id + return ipred_client_mod.upload_session_array(session_id, payload) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/manifold/sample") +async def ipred_manifold_sample(body: IpredManifoldSampleRequest) -> dict: + """Suggest Labels on an ipred feature bank.""" + try: + return ipred_client_mod.manifold_sample(body.model_dump(exclude_none=True)) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/manifold/{sample_id}/heatmap.png") +async def ipred_manifold_heatmap(sample_id: str) -> Response: + """Proxy manifold residual heatmap PNG.""" + try: + return Response( + content=ipred_client_mod.manifold_heatmap_png(sample_id), + media_type="image/png", + headers={"Cache-Control": "private, max-age=300"}, + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/preprocess") +async def ipred_preprocess(body: IpredPreprocessRequest) -> dict: + """Cache-aware featurize via ipred.""" + try: + return ipred_client_mod.preprocess( + session_id=body.session_id, + feature_setup_id=body.feature_setup_id, + composition_id=body.composition_id, + slice_index=body.slice_index, + array_ref=body.array_ref, + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/features/{feature_id}/channels/{index}") +async def ipred_feature_channel(feature_id: str, index: int) -> Response: + """Proxy a feature-channel PNG from the ipred.""" + try: + data = ipred_client_mod.feature_channel_bytes(feature_id, index) + return Response(content=data, media_type="image/png") + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/train") +async def ipred_train(body: IpredTrainRequest) -> dict: + """Train via ipred trainer plugin.""" + try: + return ipred_client_mod.train( + session_id=body.session_id, + shapes=body.shapes, + feature_id=body.feature_id, + trainer_id=body.trainer_id, + config=body.config, + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/infer") +async def ipred_infer(body: IpredInferRequest) -> dict: + """Infer + conformal products via ipred.""" + try: + return ipred_client_mod.infer( + session_id=body.session_id, + model_id=body.model_id, + feature_id=body.feature_id, + alpha=body.alpha, + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/rethreshold") +async def ipred_rethreshold(body: IpredRethresholdRequest) -> dict: + """Rethreshold from cached proba via ipred.""" + try: + return ipred_client_mod.rethreshold( + session_id=body.session_id, + alpha=body.alpha, + run_id=body.run_id, + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/runs/{run_id}/commit.png") +async def ipred_run_commit(run_id: str) -> Response: + """Proxy commit map PNG.""" + try: + return Response( + content=ipred_client_mod.run_commit_png(run_id), + media_type="image/png", + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/runs/{run_id}/status.png") +async def ipred_run_status(run_id: str) -> Response: + """Proxy status map PNG.""" + try: + return Response( + content=ipred_client_mod.run_status_png(run_id), + media_type="image/png", + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.get("/runs/{run_id}/proba/{class_index}.png") +async def ipred_run_proba(run_id: str, class_index: int) -> Response: + """Proxy softmax class heatmap PNG from ipred.""" + try: + return Response( + content=ipred_client_mod.run_proba_png(run_id, class_index), + media_type="image/png", + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc + + +@router.post("/batch/train") +async def ipred_batch_train(body: IpredBatchTrainRequest) -> dict: + """Train across multiple slices in the background; poll via export_jobs. + + Mirrors the `/api/denoise/bake` job-spawning pattern: create the job, + hand it to a daemon thread, return its id immediately. + """ + if not body.slices: + raise HTTPException(422, "at least one slice with shapes is required") + jid = export_jobs.new_job(body.session_id) + threading.Thread( + target=ipred_batch_jobs.run_ipred_multi_train_job, + kwargs=dict( + jid=jid, + session_id=body.session_id, + per_slice_shapes={int(k): v for k, v in body.slices.items()}, + composition_id=body.composition_id, + feature_setup_id=body.feature_setup_id, + trainer_id=body.trainer_id, + config=body.config, + ), + daemon=True, + ).start() + return {"job_id": jid} + + +@router.post("/batch/apply") +async def ipred_batch_apply(body: IpredBatchApplyRequest) -> dict: + """Run inference across many slices in the background; poll via export_jobs.""" + if not body.slice_indices: + raise HTTPException(422, "at least one slice_index is required") + jid = export_jobs.new_job(body.session_id) + threading.Thread( + target=ipred_batch_jobs.run_ipred_volume_apply_job, + kwargs=dict( + jid=jid, + session_id=body.session_id, + model_id=body.model_id, + slice_indices=list(body.slice_indices), + composition_id=body.composition_id, + feature_setup_id=body.feature_setup_id, + alpha=body.alpha, + ), + daemon=True, + ).start() + return {"job_id": jid} + + +@router.post("/runs/{run_id}/threshold-class") +async def ipred_threshold_class(run_id: str, body: IpredThresholdClassRequest) -> dict: + """Threshold one softmax class into a dense label map via ipred.""" + try: + return ipred_client_mod.threshold_class_map( + run_id, + class_id=body.class_id, + threshold=body.threshold, + ) + except Exception as exc: + raise _ipred_http_error(exc) from exc diff --git a/backend/local_fs.py b/backend/local_fs.py index 4845287..1990833 100644 --- a/backend/local_fs.py +++ b/backend/local_fs.py @@ -28,6 +28,16 @@ _DEFAULT_ROOT: Path = Path(os.getenv("LOCAL_DATA_ROOT", "~/data")).expanduser().resolve() +def default_root() -> str: + """Return the default browse root (``LOCAL_DATA_ROOT``) as an absolute path. + + Used by callers that need to construct an absolute path from a directory + listing's root-relative entries (``list_dir`` returns paths relative to + the root, not absolute ones) when no explicit *root* was granted. + """ + return str(_DEFAULT_ROOT) + + def _resolve_root(root: str | None) -> Path: """Return the granted browse root as an absolute, resolved Path. diff --git a/backend/mask_pyramid.py b/backend/mask_pyramid.py new file mode 100644 index 0000000..6b932ff --- /dev/null +++ b/backend/mask_pyramid.py @@ -0,0 +1,162 @@ +"""Register a class-id mask/annotation volume as a real OME-NGFF multiscale +Tiled node, so the volume viewer's ``loadMask()``/``openOmeZarr()`` can +actually open it. + +``tiled_mask_sync.write_masks_to_tiled`` used to write its ``semantic`` array +with a plain ``container.write_array()`` call, which sets no ``multiscales`` +metadata — Tiled then serves it as an ordinary array node, and the viewer's +``openOmeZarr()`` rejects it with "missing multiscales" (see +``volume_nodes.has_multiscales``, which exists specifically to distinguish the +two). This module gives masks the same real multiscale registration the +primary volume already gets from ``volume_build.build_volume``. + +Differs from ``volume_build.build_volume`` in two ways: + - Downsamples by majority vote (mode), never by mean — an averaged class id + is meaningless, the same reasoning the viewer's own ``mask-texture.ts`` + gives for using ``r8uint``/``textureLoad`` instead of ``r8unorm``/sampled. + - Every level is computed directly from the native-resolution array rather + than cascaded level-to-level: the whole mask is already fully in memory + (unlike the primary volume, which is built by streaming a stack that may + be tens of GB), so there is no incremental-IO reason to cascade. +""" +from __future__ import annotations + +import asyncio +import logging +from typing import Any + +import numpy as np + +import tiff_stack_source as tss + +logger = logging.getLogger("mask_pyramid") + + +def majority_downsample(volume: np.ndarray, factor: list[int]) -> np.ndarray: + """Mode-downsample a uint8 class-id volume by the per-axis *factor*. + + Vectorized over the small set of class ids actually present in *volume* + (real annotation taxonomies are a handful of classes even though uint8 + allows 256), rather than per-voxel. Ties favor the lower class id: + ``np.unique`` is sorted ascending and a later id only overwrites the + current winner on a strictly higher count. + """ + fz, fy, fx = (int(f) for f in factor) + if fz < 1 or fy < 1 or fx < 1: + raise ValueError("downsample factors must be >= 1") + z = (volume.shape[0] // fz) * fz + y = (volume.shape[1] // fy) * fy + x = (volume.shape[2] // fx) * fx + trimmed = volume[:z, :y, :x] + out_shape = (z // fz, y // fy, x // fx) + + class_ids = np.unique(trimmed) + if class_ids.size <= 1: + fill = int(class_ids[0]) if class_ids.size else 0 + return np.full(out_shape, fill, dtype=np.uint8) + + reshaped = trimmed.reshape(z // fz, fz, y // fy, fy, x // fx, fx) + best_count = np.zeros(out_shape, dtype=np.int64) + best_id = np.zeros(out_shape, dtype=np.uint8) + for class_id in class_ids: + count = (reshaped == class_id).sum(axis=(1, 3, 5)) + better = count > best_count + best_count = np.where(better, count, best_count) + best_id = np.where(better, class_id, best_id) + return best_id + + +def build_mask_pyramid(semantic: np.ndarray) -> tuple[dict[str, np.ndarray], list[dict[str, Any]]]: + """Compute ``{"scale0": semantic, "scale1": ..., ...}`` for a class-id volume. + + Returns the level dict plus the *generated* (non-scale0) plan entries from + ``tiff_stack_source.pyramid_plan`` — the same shape ``volume_build.py`` + passes to ``multiscales_metadata``. + """ + shape = tuple(int(s) for s in semantic.shape) + generated = tss.pyramid_plan(shape) + levels: dict[str, np.ndarray] = {"scale0": semantic.astype(np.uint8, copy=False)} + for level in generated: + levels[level["path"]] = majority_downsample(semantic, level["factor"]) + return levels, generated + + +def register_mask_pyramid( + semantic: np.ndarray, key: str, container: Any, *, cache_key: str, +) -> dict[str, Any]: + """Build and register *semantic* (a ``(z, y, x)`` uint8 class-id volume) as + a real OME-NGFF multiscale node named *key* inside *container* (the + ``__masks`` container ``tiled_mask_sync`` already manages). + + Reuses ``volume_build.build_volume``'s own write-pyramid-sidecar + + ``register_single_item`` + ``multiscales_metadata`` machinery + (``tiff_stack_source.write_pyramid_store``/``pyramid_plan``/ + ``multiscales_metadata``/``PYRAMID_KEY``) so Tiled serves this exactly the + way it already serves the primary volume's pyramid — just downsampled by + majority vote instead of mean. + + *key* only names the Tiled sub-node under *container* (always + ``"semantic"`` today) — it is NOT unique across datasets or producers. + *cache_key* is what actually scopes the on-disk sidecar + ``write_pyramid_store`` writes to (``pyramid_cache_root() / cache_key``), + and MUST be globally unique per (source, mask producer) — e.g. the + already-unique ``__masks`` container name. Passing a + constant here (this function's own bug, previously always ``"semantic"``) + means every mask sync for every dataset and every producer (iPred vs + dlsia) silently overwrites the same on-disk pyramid file: Tiled's + per-container metadata stays correct (it lives in Tiled's own DB), but + the actual pixel data every such container's registration points at is + the same shared file, so whichever sync ran last "wins" for everyone — + exactly like ``build_volume`` already keys its own sidecar by the + source's stem, which this must match to avoid the same class of bug. + + Replaces any previous registration for *key* — a re-sync must refresh the + pyramid, not fail or accumulate duplicates, matching ``build_volume``'s + same rule for the primary volume. + """ + from tiled.client.register import Settings, register_single_item + + import ingest as ingest_mod + + levels, generated = build_mask_pyramid(semantic) + sidecar = tss.write_pyramid_store(cache_key, levels) + + if key in ingest_mod._child_keys(container): + container.delete_contents(key, recursive=True, external_only=False) + node = container.create_container(key=key, metadata={}) + + try: + asyncio.run( + register_single_item(node, sidecar, is_directory=True, settings=Settings.init()) + ) + except Exception as exc: # noqa: BLE001 — surfaced to the caller, same as build_volume + logger.warning("mask pyramid registration failed for %s: %s", sidecar, exc) + raise + + pyramid_node = ingest_mod._walk(node, [tss.PYRAMID_KEY]) + if pyramid_node is None or not list(pyramid_node): + raise RuntimeError( + f"Tiled registered no levels for mask pyramid {key!r} from {sidecar}. " + "The usual cause is that this path is not in the Tiled server's " + "`readable_storage` (see tiled/config.yml)." + ) + + # scale0 is generated by THIS module rather than living in a pre-existing + # per-slice sibling (the primary volume's case) — every level, scale0 + # included, is written into the one self-contained pyramid store above, so + # it takes the SAME `f"{PYRAMID_KEY}/{level['path']}"` nesting the + # generated levels get. Passing include_scale0=False and folding scale0 + # into `plan` (rather than True, which would emit it as a bare top-level + # sibling path) is what keeps the metadata pointing at where the data + # actually landed. + plan = [{"path": "scale0", "factor": [1, 1, 1], "shape": list(semantic.shape)}, *generated] + node.update_metadata(metadata=tss.multiscales_metadata(key, plan, include_scale0=False)) + return {"key": key, "full_shape": list(semantic.shape), "pyramid_plan": generated} + + +def read_mask_scale0(container: Any, key: str) -> np.ndarray: + """Read back the native-resolution level of a mask pyramid registered by + :func:`register_mask_pyramid` — the merge path's source of truth (the + downsampled levels are viewer-only conveniences, never re-read).""" + scale0 = container[key][tss.PYRAMID_KEY]["scale0"] + return np.asarray(scale0[...]) diff --git a/backend/pyproject.toml b/backend/pyproject.toml index b3ffeed..a1ebda4 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -23,13 +23,27 @@ dependencies = [ "matplotlib>=3.8", "pycocotools>=2.0.7", "scikit-image>=0.22", + # denoise.py imports scipy.ndimage directly (gaussian/median filters, and the + # convolution behind its noise estimator). scikit-image pulls scipy in + # transitively, but a direct import earns a direct pin. + "scipy>=1.11", "tifffile>=2024.0", "imagecodecs", + # ipred_client.py talks to the standalone iPred service over plain HTTP — + # pulled in transitively today via tiled[client], but the direct import + # (ipred_client.py, ipred_routes.py) earns a direct pin. + "httpx>=0.27", ] [project.optional-dependencies] dev = ["flake8", "flake8-isort", "black"] test = ["pytest>=8", "pytest-asyncio>=0.23", "httpx>=0.27"] +# Multi-GB (torch) — not installed in the Docker image; start_all.sh's +# ensure_ml_env() installs it locally. Every /api/train/* route (and +# denoise_bake.py's "model" denoise method) degrades to a clear error via +# train_common.torch_available()/dlsia_available() rather than crashing on +# import when this extra isn't installed. +ml = ["torch>=2.4", "dlsia>=0.3", "qlty>=1.5"] [tool.setuptools] py-modules = [] diff --git a/backend/schemas.py b/backend/schemas.py index eaf50ea..93c3470 100644 --- a/backend/schemas.py +++ b/backend/schemas.py @@ -18,9 +18,9 @@ from __future__ import annotations -from typing import Any, Literal +from typing import Annotated, Any, Literal, Self, Union -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, field_validator, model_validator class BrushStroke(BaseModel): @@ -151,6 +151,24 @@ class RenderOpts(BaseModel): cmap: Literal["gray", "viridis"] = "gray" +class PredictedSlicePointer(BaseModel): + """An un-vectorized iPred volume-apply result for one slice — the frontend's + ``predictedRasterStore`` pointer, sent straight through instead of first + fetching+tracing the commit PNG into polygon shapes client-side (see the + lazy-vectorization plan item this supports). + + Attributes: + run_id: ipred run id — ``tiled_mask_sync`` fetches its commit.png + directly from the ipred service via ``ipred_client.run_commit_png``. + class_ids: Frontend class ids the commit PNG's pixel values are drawn + from directly (NOT a 1-based sequential index — matches the + frontend's own ``labelMapToPolygonShapes`` convention). + """ + + run_id: str + class_ids: list[int] + + class ExportSourceItem(BaseModel): """A single annotated sample within a multi-source export. @@ -161,6 +179,11 @@ class ExportSourceItem(BaseModel): slices: Mapping of slice key → list of serialised shape dicts. split_by_slice: Mapping of slice key → dataset split name. negative_slices: Slice keys included as negative (unannotated) examples. + predicted_slices: Mapping of slice key → an un-vectorized iPred result + pointer (see :class:`PredictedSlicePointer`). Only consulted by + ``tiled_mask_sync.build_mask_volumes`` (the "Push to Tiled" path) + for slice keys absent from ``slices`` — a slice with real shapes + always wins, matching the frontend's own precedence. """ kind: Literal["tiled", "local"] @@ -169,6 +192,7 @@ class ExportSourceItem(BaseModel): slices: dict[str, list[dict[str, Any]]] = Field(default_factory=dict) split_by_slice: dict[str, str] = Field(default_factory=dict) negative_slices: list[str] = Field(default_factory=list) + predicted_slices: dict[str, PredictedSlicePointer] = Field(default_factory=dict) class ExportRequest(BaseModel): @@ -281,6 +305,120 @@ class IngestPreflightRequest(BaseModel): server_uri: str | None = None +class ZarrRegisterRequest(BaseModel): + """Request body for registering an on-disk Zarr volume with Tiled. + + No bytes are uploaded: the volume stays where it is and Tiled reads it in + place, so this carries a server-side path rather than file content. + + Attributes: + path: Absolute path to the ``.zarr`` directory on the server. + container_path: Target container, e.g. ``browse``. + description: Optional keyword(s) stored on the node (same treatment as + the drag-and-drop ingest, so the volume is filterable in Browse). + on_conflict: ``"fail"``, ``"replace"`` or ``"skip"`` when the key exists. + server_uri: Target Tiled server URI; ``None`` uses the default server. + """ + + path: str + container_path: str = "browse" + description: str = "" + on_conflict: str = "fail" + server_uri: str | None = None + + +class ZarrScanRequest(BaseModel): + """Request body for scanning a directory and registering every Zarr store found. + + For pointing a mounted directory of already-reconstructed volumes at Tiled + in bulk, rather than registering each one individually via + :class:`ZarrRegisterRequest`. Non-recursive: only immediate subdirectories + of ``scan_root`` that look like a Zarr store are considered. + + Attributes: + scan_root: Absolute directory to scan (e.g. a bind-mounted host folder). + container_path: Target container every discovered store registers into. + on_conflict: ``"skip"`` (default, safe to re-run) or ``"replace"`` for + a same-kind entry that already exists — ``"fail"`` makes no sense + here since one conflicting store shouldn't abort the whole scan. + server_uri: Target Tiled server URI; ``None`` uses the default server. + renames: Optional ``{folder_name: alternate_key}`` override, so a + candidate reported as "shadowed" (its natural key collides with + an unrelated, different-kind registration) can be retried under a + different key without re-scanning everything else. + """ + + scan_root: str + container_path: str = "browse" + on_conflict: str = "skip" + renames: dict[str, str] = Field(default_factory=dict) + server_uri: str | None = None + + +class TiffStackRegisterRequest(ZarrRegisterRequest): + """Request body for registering a TIFF directory as a 3-D volume. + + Same fields as :class:`ZarrRegisterRequest` — ``path`` is a directory of 2-D + TIFF slices rather than a ``.zarr`` store. Kept as its own type so the two + endpoints document themselves and can diverge without a breaking change. + + Unlike the Zarr path this does real work: the full-resolution slices are + registered in place, but the downsampled pyramid levels the 3-D viewer + actually renders have to be computed, so registration runs as a job. + """ + + +class DenoiseBakeRequest(BaseModel): + """Request body to denoise a whole volume into a new Tiled dataset. + + The Annotate preview is non-destructive — it changes what you see, not the + data, and exports still use the original pixels. This is the other half: it + writes a denoised copy as a first-class dataset you can open, annotate and + export. + + Attributes: + source: Tiled path of the volume to denoise. + server_uri: Tiled server URI; ``None`` uses the default server. + method: A ``denoise.ALL_METHODS`` entry other than ``"none"``. + strength: 0..1, mapped onto the method's native parameter. + target_path: Destination Tiled path; defaults to ``_denoised``, a + sibling so it lands next to its source in Browse. Must sit beneath + the configured ingest root. + description: Optional comma-separated tags, treated exactly as ingest + treats them — searchable in Browse. + """ + + source: str + server_uri: str | None = None + method: str + strength: float = Field(default=0.5, ge=0.0, le=1.0) + target_path: str | None = None + description: str = "" + run_id: str | None = None + """Saved denoiser run, for the trained-model path. Not yet wired up here.""" + + +class VolumeBuildRequest(BaseModel): + """Request body for building a 3-D volume from a stack already in Tiled. + + Deliberately minimal: the slices are already in the catalog, so the only + thing needed is which dataset. No source path, because requiring one would + mean asking the user to re-supply data the app already holds. + + Attributes: + source: Tiled path of the per-slice dataset. + kind: Source kind; only ``"tiled"`` has a catalog to register into. + container_path: Where to place the volume sidecar; defaults to the + dataset's own parent container. + server_uri: Target Tiled server URI; ``None`` uses the default server. + """ + + source: str + kind: str = "tiled" + container_path: str | None = None + server_uri: str | None = None + + class GuideClass(BaseModel): """One class entry in an annotation guide. @@ -319,8 +457,27 @@ class ImageMeta(BaseModel): dtype: NumPy dtype string (e.g. ``"float32"``). is_rgb: ``True`` if the array has a colour channel dimension. value_range: ``[min, max]`` of the first slice. + global_value_range: ``[vmin, vmax]`` actually used by + ``GET /api/image/slice``'s default ``norm="global"`` rendering — + the 1st/99th percentile (by default) sampled across up to 64 + slices spanning the whole volume (see ``images._sample_global_stats``), + cached 5 minutes. Distinct from ``value_range`` above (which is + just slice 0's raw min/max): this is the real contrast window the + 2D canvas's pixel bytes are normalized against, needed by anything + that must convert a displayed 0-255 value back to a physical + intensity (e.g. the Sampler-fitted band sent to the 3D viewer). keywords: Dataset tags stored at ingest; each is pre-created as an annotation class in the Annotate tab. + level_key: For a multiscale Zarr volume, the pyramid level being read + (e.g. ``"scale2"``). ``height``/``width``/``n_slices`` always + describe the FINEST level, since annotations are stored in + full-resolution coordinates; these fields describe what is actually + being displayed underneath them. + level_index: Index of that level, finest = 0. + level_count: Number of levels in the pyramid. + level_height / level_width / level_n_slices: The open level's own shape. + z_downsample: Finest-z / level-z. When > 1 the level can only address + every f-th full-resolution slice. """ n_slices: int @@ -329,4 +486,283 @@ class ImageMeta(BaseModel): dtype: str is_rgb: bool value_range: list[float] + global_value_range: list[float] | None = None keywords: list[str] = Field(default_factory=list) + level_key: str | None = None + level_index: int | None = None + level_count: int | None = None + level_height: int | None = None + level_width: int | None = None + level_n_slices: int | None = None + + +# --------------------------------------------------------------------------- +# Train tab (Phase 5) — dlsia TUNet segmentation + the dlsia/autoencoder +# denoiser family only. DINOv3 LoRA is deferred to Phase 5.5 and has no +# schema here at all, not even a stubbed-out variant. +# --------------------------------------------------------------------------- + + +class TunetHyperParams(BaseModel): + """Bounded training hyperparameters for a dlsia TUNet trained from scratch. + + Attributes: + epochs: Number of training epochs (from-scratch training typically + needs more epochs than fine-tuning a pretrained backbone). + lr: Learning rate for the whole network. + depth: U-Net depth (encoder/decoder stages). + base_channels: Channel count of the first conv stage. + growth_rate: Channel growth factor per depth level. + batch_size: Training batch size. + image_size: Square side length images/labels are letterboxed to. + dlsia's TUNet fixes its layer sizes to ``image_shape`` at + construction, so train and inference images must match exactly. + With ``tiling`` on this is the tile window — every window is exactly + this size, so the constraint still holds. + seed: Seed for shuffling and augmentation. + flip_augment: Whether to randomly flip image+label together. + tiling: Train on native-resolution ``image_size`` windows cut from each + slice, instead of rescaling the whole slice down to ``image_size`` + (see ``tiling.py``). Recorded on the run, so inference reproduces + whichever geometry the run was trained with. + """ + + epochs: int = Field(default=60, ge=1, le=1000) + lr: float = Field(default=1e-3, gt=0, le=1) + depth: int = Field(default=4, ge=2, le=6) + base_channels: int = Field(default=8, ge=1, le=128) + growth_rate: float = Field(default=1.5, gt=0, le=4) + batch_size: int = Field(default=4, ge=1, le=32) + image_size: int = Field(default=512, ge=64, le=2048) + seed: int = 1234 + flip_augment: bool = True + tiling: bool = True + + @field_validator("image_size") + @classmethod + def validate_image_size(cls, value: int) -> int: + """A TUNet at ``depth`` stages needs the side divisible by 2**depth + so every downsample/upsample step lands on a whole-pixel size.""" + if value % 64 != 0: + raise ValueError("image_size must be a multiple of 64 (covers depth up to 6)") + return value + + +class DlsiaTunetConfig(BaseModel): + """Model-family config: dlsia tunable U-Net trained from scratch. + + Attributes: + model_family: Discriminator literal. + hyperparams: Training hyperparameters. + """ + + model_family: Literal["dlsia_tunet"] = "dlsia_tunet" + hyperparams: TunetHyperParams = Field(default_factory=TunetHyperParams) + + +class DlsiaDenoiserConfig(BaseModel): + """Model-family config: a single-channel image denoiser trained + self-supervised, instead of a multi-class segmenter. + + Two architectures share this family. Keeping them under one + ``model_family`` is deliberate: the frontend partitions runs with + ``isSegmentationRun == !isDenoiserRun`` (``lib/runCompatibility.ts``), so a + *new* family value would be silently treated as segmentation and offered in + the fine-tune / apply / inference pickers, which key off a class list a + denoiser does not have. An ``architecture`` discriminator inside the family + avoids that entirely. + + ``TunetHyperParams`` is reused verbatim rather than defining a parallel + hyperparameter class — ``depth``/``base_channels`` mean downsampling levels + and first-conv width for both architectures. Architecture-specific knobs + (``ae_compression``) live here on the config, not there. + + Attributes: + model_family: Discriminator literal. + architecture: Which network to build — ``"tunet"`` (dlsia TUNet with + ``in_channels=1, out_channels=1``; see ``denoise_runtime.py``) or + ``"cnn_ae"`` (a plain convolutional autoencoder with an explicit + latent bottleneck and NO skip connections; see + ``autoencoder_runtime.py``). Defaults to ``"tunet"`` so runs saved + before this field existed keep their meaning. + hyperparams: Training hyperparameters (reuses the segmentation + family's TUNet knobs — depth/base_channels/growth_rate/etc.). + ae_compression: How much the ``"cnn_ae"`` bottleneck compresses, as a + ratio of input values to latent values. Higher removes more noise + but also discards more real detail. Only meaningful for + ``"cnn_ae"``. + training_scheme: Self-supervised training objective — ``"n2n"`` + (Noise2Noise: paired noisy/noisy training), ``"n2v"`` + (Noise2Void: blind-spot training from single noisy images), or + ``"ae"`` (pure self-reconstruction: target IS the input, and the + bottleneck is what forces noise out). + + ``"ae"`` is valid ONLY with ``architecture="cnn_ae"``, enforced + below. On a skip-connected network like TUNet, training on + ``target == input`` makes ``f(x) = x`` trivially learnable: it + converges to copying the input, removes no noise whatsoever, and + still reports a falling loss. Without skips, the bottleneck cannot + pass the input through unchanged, so reconstruction becomes a real + denoising objective. That pairing is a correctness constraint, not + a convenience. + """ + + model_family: Literal["dlsia_denoiser"] = "dlsia_denoiser" + architecture: Literal["tunet", "cnn_ae"] = "tunet" + hyperparams: TunetHyperParams = Field(default_factory=TunetHyperParams) + training_scheme: Literal["n2n", "n2v", "ae"] + ae_compression: int = Field(default=16, ge=4, le=64) + + @model_validator(mode="after") + def _check_scheme_matches_architecture(self) -> "DlsiaDenoiserConfig": + """Keep scheme and architecture to the combinations that make sense. + + ``ae`` + ``tunet`` is the identity-collapse footgun described above and + must be impossible to request. The reverse (``cnn_ae`` with a masking or + paired scheme) is not unsound in principle, just untested here, so it is + refused rather than silently shipped. + """ + if self.training_scheme == "ae" and self.architecture != "cnn_ae": + raise ValueError( + "training_scheme='ae' (pure self-reconstruction) requires " + "architecture='cnn_ae'. On a skip-connected network it would just learn to " + "copy its input and remove no noise." + ) + if self.architecture == "cnn_ae" and self.training_scheme != "ae": + raise ValueError( + "architecture='cnn_ae' is only supported with training_scheme='ae'; " + f"got {self.training_scheme!r}." + ) + return self + + +ModelConfig = Annotated[ + Union[DlsiaTunetConfig, DlsiaDenoiserConfig], + Field(discriminator="model_family"), +] + + +class BatchProbeRequest(BaseModel): + """Request body to measure the largest batch size a model config can fit. + + Deliberately not a :class:`TrainRequest`: the probe feeds synthetic tensors, so + it needs no sources, and requiring them would force the caller to invent data + just to ask a question about memory. + + Attributes: + model: Model-family configuration to size (patch size and the per-family + hyperparameters come from its ``hyperparams``). + n_classes: Segmentation head output channels. Affects memory only + marginally, so it defaults to a typical value — meaning a batch size + can be estimated before any classes have been defined. + """ + + model: ModelConfig + n_classes: int = Field(default=2, ge=1) + + +class DenoiseTrainOpts(BaseModel): + """Denoising applied to a model's INPUT pixels, at training and inference alike. + + Distinct from the Annotate tab's denoise preview, which is display-only and + never reaches a model. When this is set on a training request it is recorded + on the saved run, and inference reads it back off the run rather than off the + request — the same rule ``tiling`` already follows, and for the same reason: + a model must see the same pixel distribution it was trained on. Letting the + two be chosen independently would produce a silent distribution shift with + no error, just quietly worse predictions. + + Attributes: + method: A ``denoise.ALL_METHODS`` entry other than ``"none"``/``"model"`` + (a learned denoiser as a preprocessor for another model is not + supported — it would need its own run and GPU pass per slice). + strength: 0..1, mapped onto the method's native parameter. + """ + + method: str + strength: float = Field(default=0.5, ge=0, le=1) + + +class TrainRequest(BaseModel): + """Request body to start a fine-tuning job. + + Attributes: + task: What the trained model is for — ``"segmentation"`` (requires at + least one class) or ``"denoising"`` (a self-supervised + Noise2Noise/Noise2Void/autoencoder denoiser, which has no class + taxonomy at all). + sources: Annotated samples to train on (each with its own slices/splits). + classes: Annotation classes/taxonomy shared across every source. + Required (at least one) when ``task == "segmentation"``; may be + empty when ``task == "denoising"``. + render: Render options used to rasterise training images. + denoise: Optional denoising applied to the model's INPUT pixels. Recorded + on the run and reapplied automatically at inference — see + :class:`DenoiseTrainOpts`. Ignored when ``task == "denoising"``: a + denoiser learns to remove noise, so pre-cleaning its input would + defeat the point. + auto_split: Auto-split configuration for slices without an explicit split. + model: Model-family configuration (discriminated on ``model_family``). + run_name: Optional human-readable label for the resulting run. + resume_from_run_id: Continue fine-tuning from this saved run's weights + instead of starting from scratch. The saved run's architecture-defining + settings win over anything sent in ``model`` — a resume that changed + them could not load the saved weights at all. Result is always a NEW + run; the parent is never modified. + """ + + task: Literal["segmentation", "denoising"] = "segmentation" + sources: list[ExportSourceItem] = Field(min_length=1) + classes: list[AnnotationClass] = Field(default_factory=list) + render: RenderOpts = Field(default_factory=RenderOpts) + denoise: DenoiseTrainOpts | None = None + auto_split: dict[str, Any] = Field( + default_factory=lambda: {"ratios": [0.8, 0.1, 0.1], "seed": 1234} + ) + model: ModelConfig + run_name: str | None = None + resume_from_run_id: str | None = None + + @model_validator(mode="after") + def validate_train_taxonomy(self) -> Self: + """A self-supervised denoiser has no classes at all — still required + (and still validated) for segmentation, allowed empty for denoising.""" + if self.task == "segmentation" and len(self.classes) < 1: + raise ValueError("Segmentation training requires at least one class") + return self + + +class InferRequest(BaseModel): + """Request body to run inference with a saved fine-tuned run. + + Attributes: + run_id: Identifier of a previously trained run. + kind: Source kind — ``"tiled"`` or ``"local"``. + source: Tiled path or local relative path to run inference on. + server_uri: Tiled server URI (kind == "tiled" only). + slice_indices: Zero-based slice indices to run inference on. + render: Render options; ``None`` reuses the run's stored render options. + min_area: Minimum connected-component pixel area kept per predicted region. + simplify_tol: Polygon simplification tolerance (pixels). + min_confidence: Softmax confidence below which a pixel is treated as + background (no class) — training never sees an explicit background + class, since unannotated pixels are the ignore index, not a label. + """ + + run_id: str + kind: Literal["tiled", "local"] + source: str + server_uri: str | None = None + slice_indices: list[int] = Field(min_length=1) + render: RenderOpts | None = None + min_area: int = Field(default=64, ge=0) + simplify_tol: float = Field(default=1.5, ge=0, le=50) + min_confidence: float = Field(default=0.5, ge=0, le=1) + + @field_validator("slice_indices") + @classmethod + def validate_slice_indices(cls, value: list[int]) -> list[int]: + if len(value) != len(set(value)): + raise ValueError("slice_indices must be unique") + return value + z_downsample: float | None = None diff --git a/backend/tests/test_annotation_server_browse_routes.py b/backend/tests/test_annotation_server_browse_routes.py new file mode 100644 index 0000000..f531d9b --- /dev/null +++ b/backend/tests/test_annotation_server_browse_routes.py @@ -0,0 +1,177 @@ +"""HTTP-level tests for annotation_server.py's /api/browse/* routes. + +Only `get_tiled_client` is faked (returns a lightweight fake Tiled root node, +same duck-typed FakeNode as test_browse_helpers.py) — the real +`get_browse_container_for` runs against it unmocked, since it's pure +navigation logic (`node[k]` per path segment). Every test supplies an +explicit container_path so get_browse_container_for never falls through to +the heuristic root-discovery path (get_browse_container), which is out of +scope here. The underlying field-mapping/distinct/search logic itself is +already thoroughly covered directly in test_browse_helpers.py — these tests +are about routing, query-param parsing, and error handling, not re-proving +that logic. +""" +from __future__ import annotations + +from collections import Counter + +import pytest +from httpx import ASGITransport, AsyncClient + +import annotation_server + +tiled_queries = pytest.importorskip("tiled.queries") + + +class FakeNode: + def __init__(self, children=None, metadata=None, is_container=True, search_raises=False): + self._children = children or {} + self.metadata = metadata or {} + self.structure_family = "container" if is_container else "array" + self._search_raises = search_raises + + def __iter__(self): + return iter(self._children) + + def __getitem__(self, key): + return self._children[key] + + def __len__(self): + return len(self._children) + + def search(self, query): + if self._search_raises: + raise RuntimeError("search exploded") + matched = {} + for k, child in self._children.items(): + meta = child.metadata or {} + if isinstance(query, tiled_queries.Contains): + val = meta.get(query.key) + if isinstance(val, (list, tuple)) and query.value in val: + matched[k] = child + else: # Eq + if meta.get(query.key) == query.value: + matched[k] = child + return FakeNode(children=matched, metadata=self.metadata) + + def distinct(self, key, counts=True): + counter: Counter = Counter() + for child in self._children.values(): + val = (child.metadata or {}).get(key) + if val is None: + continue + counter[val] += 1 + return {"metadata": {key: [{"value": v, "count": n} for v, n in counter.items()]}} + + +def _sample(metadata): + return FakeNode(metadata=metadata, is_container=False) + + +@pytest.fixture() +def browse_root(): + """browse/sample1 (PI=Smith), browse/sample2 (PI=Jones) — enough to + exercise facet discovery (>=2 distinct PI values), column distinct, and + item search/filtering.""" + samples = { + "sample1": _sample({"PI": "Smith", "sample_name": "s1", "technique": "GIWAXS"}), + "sample2": _sample({"PI": "Jones", "sample_name": "s2", "technique": "GIWAXS"}), + } + browse = FakeNode(children=samples) + return FakeNode(children={"browse": browse}) + + +@pytest.fixture(autouse=True) +def fake_tiled_client(monkeypatch: pytest.MonkeyPatch, browse_root): + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda *a, **k: browse_root) + # Field-mapping / column / item caches are module-level and would leak + # a fake result from one test's fake node into the next test's assertions + # (same server_uri/technique/container_path -> same cache key). + annotation_server._field_mapping_cache._store.clear() + annotation_server._column_cache._store.clear() + annotation_server._items_cache._store.clear() + + +@pytest.fixture() +async def client(): + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as c: + yield c + + +app = annotation_server.app + + +@pytest.mark.asyncio +async def test_browse_facets_finds_multi_valued_field(client): + response = await client.get("/api/browse/facets", params={"container_path": "browse"}) + assert response.status_code == 200 + assert "PI" in response.json()["facets"] + + +@pytest.mark.asyncio +async def test_browse_facets_returns_empty_lists_on_internal_error(client, monkeypatch): + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda *a, **k: (_ for _ in ()).throw(RuntimeError("down"))) + response = await client.get("/api/browse/facets", params={"container_path": "browse"}) + assert response.status_code == 200 + assert response.json() == {"facets": [], "techniques": []} + + +@pytest.mark.asyncio +async def test_browse_column_returns_distinct_values(client): + response = await client.get( + "/api/browse/column", params={"field": "PI", "container_path": "browse"}, + ) + assert response.status_code == 200 + values = {v["value"] for v in response.json()["values"]} + assert values == {"Smith", "Jones"} + + +@pytest.mark.asyncio +async def test_browse_column_502s_on_internal_error(client, monkeypatch): + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda *a, **k: (_ for _ in ()).throw(RuntimeError("down"))) + response = await client.get( + "/api/browse/column", params={"field": "PI", "container_path": "browse"}, + ) + assert response.status_code == 502 + + +@pytest.mark.asyncio +async def test_browse_items_returns_matching_samples(client): + response = await client.get( + "/api/browse/items", + params={"container_path": "browse", "filters": '{"PI": "Smith"}'}, + ) + assert response.status_code == 200 + body = response.json() + assert body["total"] == 1 + assert body["items"][0]["sample"] == "sample1" + + +@pytest.mark.asyncio +async def test_browse_items_no_filters_returns_all(client): + response = await client.get("/api/browse/items", params={"container_path": "browse"}) + assert response.status_code == 200 + assert response.json()["total"] == 2 + + +@pytest.mark.asyncio +async def test_browse_items_malformed_filter_json_treated_as_no_filter(client): + response = await client.get( + "/api/browse/items", params={"container_path": "browse", "filters": "not json"}, + ) + assert response.status_code == 200 + assert response.json()["total"] == 2 + + +@pytest.mark.asyncio +async def test_browse_items_502s_on_internal_error(client, monkeypatch): + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda *a, **k: (_ for _ in ()).throw(RuntimeError("down"))) + response = await client.get("/api/browse/items", params={"container_path": "browse"}) + assert response.status_code == 502 + + +@pytest.mark.asyncio +async def test_browse_slices_lists_container_children(client): + response = await client.get("/api/browse/slices", params={"path": "browse"}) + assert response.status_code == 200 + assert response.json()["total"] == 2 diff --git a/backend/tests/test_annotation_server_data_routes.py b/backend/tests/test_annotation_server_data_routes.py new file mode 100644 index 0000000..e7b0e20 --- /dev/null +++ b/backend/tests/test_annotation_server_data_routes.py @@ -0,0 +1,477 @@ +"""HTTP-level tests for annotation_server.py's data-source registration +(zarr/tiff-stack), denoise, volume, and remaining train-route wiring. + +The underlying modules (zarr_source, tiff_stack_source, volume_build, +volume_nodes, denoise_bake, infer_jobs, train_jobs) already have their own +direct unit tests — these tests are about routing, request validation, and +background-job wiring, so the delegate functions are monkeypatched rather +than re-proven here. Background threads are made synchronous via a fake +`threading.Thread` so job results are observable without polling/sleeping. +""" +from __future__ import annotations + +import types + +import pytest +from httpx import ASGITransport, AsyncClient + +import annotation_server +import denoise as denoise_mod +import export_jobs +import infer_jobs +import ingest as ingest_mod +import tiff_stack_source +import train_common +import volume_build +import volume_nodes +import zarr_source + +app = annotation_server.app + + +@pytest.fixture() +async def client(): + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as c: + yield c + + +class ImmediateThread: + """Runs its target synchronously instead of spawning a real thread, so a + background-job route's result is observable right after the request + returns, with no polling or sleeping needed.""" + + def __init__(self, target=None, args=(), kwargs=None, daemon=None): + self._target = target + self._args = args + self._kwargs = kwargs or {} + + def start(self): + self._target(*self._args, **self._kwargs) + + +@pytest.fixture() +def sync_threads(monkeypatch: pytest.MonkeyPatch): + # Replace the NAME `threading` inside annotation_server's own module + # namespace (not the shared stdlib module object) — otherwise this would + # also intercept unrelated internal thread creation, e.g. inside + # asyncio.to_thread's own ThreadPoolExecutor, which every route here uses. + monkeypatch.setattr(annotation_server, "threading", types.SimpleNamespace(Thread=ImmediateThread)) + + +# --------------------------------------------------------------------------- +# /api/zarr/* +# --------------------------------------------------------------------------- + +class TestZarrRoutes: + @pytest.mark.asyncio + async def test_inspect_delegates_to_zarr_source(self, client, monkeypatch): + monkeypatch.setattr(zarr_source, "inspect_zarr", lambda path: {"path": path, "levels": 3}) + response = await client.get("/api/zarr/inspect", params={"path": "/data/x.zarr"}) + assert response.status_code == 200 + assert response.json() == {"path": "/data/x.zarr", "levels": 3} + + @pytest.mark.asyncio + async def test_preflight_delegates_with_request_fields(self, client, monkeypatch): + calls = [] + monkeypatch.setattr( + zarr_source, "preflight_zarr", + lambda server_uri, path, container_path: calls.append((server_uri, path, container_path)) or {"ok": True}, + ) + response = await client.post( + "/api/zarr/preflight", + json={"path": "/data/x.zarr", "container_path": "browse/x"}, + ) + assert response.status_code == 200 + assert calls == [(None, "/data/x.zarr", "browse/x")] + + @pytest.mark.asyncio + async def test_register_rejects_bad_on_conflict(self, client): + response = await client.post( + "/api/zarr/register", + json={"path": "/data/x.zarr", "on_conflict": "explode"}, + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_register_delegates_on_valid_request(self, client, monkeypatch): + monkeypatch.setattr(zarr_source, "register_zarr", lambda *a: {"key": "x"}) + response = await client.post( + "/api/zarr/register", + json={"path": "/data/x.zarr", "on_conflict": "replace"}, + ) + assert response.status_code == 200 + assert response.json() == {"key": "x"} + + @pytest.mark.asyncio + async def test_scan_delegates_with_request_fields(self, client, monkeypatch): + calls = [] + monkeypatch.setattr( + zarr_source, "scan_and_register_zarrs", + lambda server_uri, scan_root, container_path, on_conflict, renames: ( + calls.append((server_uri, scan_root, container_path, on_conflict, renames)) + or {"scanned": 2, "registered": [], "skipped": [], "shadowed": [], "errors": []} + ), + ) + response = await client.post( + "/api/zarr/scan", + json={"scan_root": "/data/processed", "container_path": "browse"}, + ) + assert response.status_code == 200 + assert response.json() == {"scanned": 2, "registered": [], "skipped": [], "shadowed": [], "errors": []} + assert calls == [(None, "/data/processed", "browse", "skip", {})] + + @pytest.mark.asyncio + async def test_ingest_scan_delegates_with_request_fields(self, client, monkeypatch): + calls = [] + monkeypatch.setattr( + ingest_mod, "scan_and_register_image_stacks", + lambda server_uri, scan_root, container_path, on_conflict, renames: ( + calls.append((server_uri, scan_root, container_path, on_conflict, renames)) + or {"scanned": 1, "registered": [], "skipped": [], "shadowed": [], "errors": []} + ), + ) + response = await client.post( + "/api/ingest/scan", + json={"scan_root": "/data/processed", "container_path": "browse"}, + ) + assert response.status_code == 200 + assert response.json() == {"scanned": 1, "registered": [], "skipped": [], "shadowed": [], "errors": []} + assert calls == [(None, "/data/processed", "browse", "skip", {})] + + @pytest.mark.asyncio + async def test_scan_datasets_merges_zarr_and_image_scan_results(self, client, monkeypatch): + monkeypatch.setattr( + zarr_source, "scan_and_register_zarrs", + lambda *a: { + "scanned": 2, + "registered": [{"name": "a.zarr", "key": "a", "tiled_path": "browse/a"}], + "skipped": ["b"], + "shadowed": [{"name": "e", "key": "e", "existing_kind": "image-stack", "suggested_key": "e_zarr"}], + "errors": [], + }, + ) + monkeypatch.setattr( + ingest_mod, "scan_and_register_image_stacks", + lambda *a: { + "scanned": 1, + "registered": [{"name": "c", "key": "c", "tiled_path": "browse/c"}], + "skipped": [], + "shadowed": [], + "errors": [{"name": "d", "error": "boom"}], + }, + ) + response = await client.post( + "/api/scan-datasets", + json={"scan_root": "/data/processed", "container_path": "browse"}, + ) + assert response.status_code == 200 + body = response.json() + assert body["scanned"] == 3 + assert sorted(r["name"] for r in body["registered"]) == ["a.zarr", "c"] + assert body["skipped"] == ["b"] + assert body["shadowed"] == [ + {"name": "e", "key": "e", "existing_kind": "image-stack", "suggested_key": "e_zarr"} + ] + assert body["errors"] == [{"name": "d", "error": "boom"}] + + +# --------------------------------------------------------------------------- +# /api/tiff-stack/* +# --------------------------------------------------------------------------- + +class TestTiffStackRoutes: + @pytest.mark.asyncio + async def test_inspect_delegates(self, client, monkeypatch): + monkeypatch.setattr(tiff_stack_source, "inspect_tiff_stack", lambda path: {"slices_to_read": 5}) + response = await client.get("/api/tiff-stack/inspect", params={"path": "/data/stack"}) + assert response.status_code == 200 + assert response.json() == {"slices_to_read": 5} + + @pytest.mark.asyncio + async def test_preflight_delegates(self, client, monkeypatch): + monkeypatch.setattr( + tiff_stack_source, "preflight_tiff_stack", + lambda server_uri, path, container_path: {"collision": False}, + ) + response = await client.post( + "/api/tiff-stack/preflight", + json={"path": "/data/stack", "container_path": "browse/s"}, + ) + assert response.status_code == 200 + assert response.json() == {"collision": False} + + @pytest.mark.asyncio + async def test_register_rejects_bad_on_conflict(self, client): + response = await client.post( + "/api/tiff-stack/register", + json={"path": "/data/stack", "on_conflict": "nope"}, + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_register_runs_job_and_reports_done(self, client, monkeypatch, sync_threads): + monkeypatch.setattr(tiff_stack_source, "inspect_tiff_stack", lambda path: {"slices_to_read": 3}) + monkeypatch.setattr( + tiff_stack_source, "register_tiff_stack", + lambda server_uri, path, container_path, description, on_conflict, progress: {"key": "stack1"}, + ) + response = await client.post( + "/api/tiff-stack/register", + json={"path": "/data/stack", "on_conflict": "fail"}, + ) + assert response.status_code == 200 + jid = response.json()["job_id"] + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"] == {"key": "stack1"} + + @pytest.mark.asyncio + async def test_register_job_reports_error_on_exception(self, client, monkeypatch, sync_threads): + monkeypatch.setattr(tiff_stack_source, "inspect_tiff_stack", lambda path: {"slices_to_read": 3}) + + def boom(*a, **k): + raise RuntimeError("registration blew up") + + monkeypatch.setattr(tiff_stack_source, "register_tiff_stack", boom) + response = await client.post( + "/api/tiff-stack/register", + json={"path": "/data/stack", "on_conflict": "fail"}, + ) + jid = response.json()["job_id"] + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "registration blew up" in job["error"] + + +# --------------------------------------------------------------------------- +# /api/denoise/* +# --------------------------------------------------------------------------- + +class TestDenoiseMethods: + @pytest.mark.asyncio + async def test_lists_described_methods(self, client, monkeypatch): + monkeypatch.setattr(denoise_mod, "describe_methods", lambda: [{"id": "median"}]) + response = await client.get("/api/denoise/methods") + assert response.json() == {"methods": [{"id": "median"}]} + + +class TestDenoiseAuto: + @pytest.mark.asyncio + async def test_unknown_method_is_400(self, client): + response = await client.get( + "/api/denoise/auto", params={"source": "x", "kind": "local", "method": "not_real"}, + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_suggests_strength_for_a_real_slice(self, client, monkeypatch): + import numpy as np + + import arrays as arrays_mod + + arr = np.random.default_rng(0).random((8, 8)).astype(np.float32) + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri, root: "node") + monkeypatch.setattr(arrays_mod, "pyramid_info", lambda source, kind, server_uri, root: None) + monkeypatch.setattr(arrays_mod, "array_shape_meta", lambda node, pyramid: {"shape_kind": "HW"}) + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: arr) + response = await client.get( + "/api/denoise/auto", params={"source": "x", "kind": "local", "method": "median"}, + ) + assert response.status_code == 200 + body = response.json() + assert "strength" in body + assert "noise_sigma" in body + + +class TestDenoiseBakeRoute: + @pytest.mark.asyncio + async def test_method_none_is_422(self, client): + response = await client.post( + "/api/denoise/bake", json={"source": "browse/s", "method": "none"}, + ) + assert response.status_code == 422 + + @pytest.mark.asyncio + async def test_unknown_method_is_422(self, client): + response = await client.post( + "/api/denoise/bake", json={"source": "browse/s", "method": "not_real"}, + ) + assert response.status_code == 422 + + @pytest.mark.asyncio + async def test_unavailable_method_is_422(self, client, monkeypatch): + monkeypatch.setattr(denoise_mod, "available_methods", lambda: ("median",)) + response = await client.post( + "/api/denoise/bake", json={"source": "browse/s", "method": "wavelet"}, + ) + assert response.status_code == 422 + + @pytest.mark.asyncio + async def test_invalid_target_path_is_422(self, client): + response = await client.post( + "/api/denoise/bake", + json={"source": "browse/s", "method": "median", "target_path": "../escape"}, + ) + assert response.status_code == 422 + + @pytest.mark.asyncio + async def test_valid_request_starts_a_job(self, client, monkeypatch, sync_threads): + import denoise_bake as denoise_bake_mod + + monkeypatch.setattr(denoise_bake_mod, "run_denoise_bake_job", lambda jid, payload: export_jobs.update(jid, state="done")) + response = await client.post( + "/api/denoise/bake", json={"source": "browse/s", "method": "median"}, + ) + assert response.status_code == 200 + body = response.json() + assert "job_id" in body + assert body["target_path"] == "browse/s_denoised" + assert export_jobs.get_job(body["job_id"])["state"] == "done" + + +# --------------------------------------------------------------------------- +# /api/train/start — remaining guard paths not covered by test_train_routes.py +# --------------------------------------------------------------------------- + +class TestTrainStartGuards: + def _payload(self, **overrides): + base = { + "task": "segmentation", + "sources": [{"kind": "local", "source": "fake.tif", "slices": {"0": []}}], + "classes": [{"classId": 1, "label": "a", "color": "#ff0000"}], + "model": {"model_family": "dlsia_tunet"}, + } + base.update(overrides) + return base + + @pytest.mark.asyncio + async def test_torch_unavailable_is_503(self, client, monkeypatch): + monkeypatch.setattr(train_common, "torch_available", lambda: False) + response = await client.post("/api/train/start", json=self._payload()) + assert response.status_code == 503 + + @pytest.mark.asyncio + async def test_dlsia_unavailable_for_tunet_is_503(self, client, monkeypatch): + monkeypatch.setattr(train_common, "torch_available", lambda: True) + monkeypatch.setattr(train_common, "dlsia_available", lambda: False) + response = await client.post("/api/train/start", json=self._payload()) + assert response.status_code == 503 + + @pytest.mark.asyncio + async def test_busy_ml_lock_is_409(self, client, monkeypatch): + monkeypatch.setattr(train_common, "torch_available", lambda: True) + monkeypatch.setattr(train_common, "dlsia_available", lambda: True) + train_common.ML_LOCK.acquire() + try: + response = await client.post("/api/train/start", json=self._payload()) + assert response.status_code == 409 + finally: + train_common.ML_LOCK.release() + + @pytest.mark.asyncio + async def test_resume_incompatible_is_400(self, client, monkeypatch): + import train_jobs + + monkeypatch.setattr(train_common, "torch_available", lambda: True) + monkeypatch.setattr(train_common, "dlsia_available", lambda: True) + monkeypatch.setattr(train_common, "load_run_config", lambda run_id: {"some": "config"}) + + def raise_incompatible(parent_config, payload): + raise ValueError("architectures differ") + + monkeypatch.setattr(train_jobs, "check_resume_compatible", raise_incompatible) + response = await client.post( + "/api/train/start", json=self._payload(resume_from_run_id="parent1"), + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_success_creates_a_job(self, client, monkeypatch, sync_threads): + import train_jobs + + monkeypatch.setattr(train_common, "torch_available", lambda: True) + monkeypatch.setattr(train_common, "dlsia_available", lambda: True) + monkeypatch.setattr(train_jobs, "run_train_job", lambda jid, payload, run_id: export_jobs.update(jid, state="done")) + response = await client.post("/api/train/start", json=self._payload()) + assert response.status_code == 200 + body = response.json() + assert "job_id" in body and "run_id" in body + assert export_jobs.get_job(body["job_id"])["state"] == "done" + + +class TestTrainInferPreviewAndWriteTiled: + @pytest.mark.asyncio + async def test_preview_delegates_to_infer_jobs(self, client, monkeypatch): + monkeypatch.setattr(infer_jobs, "preview_png", lambda job_id, slice_index: b"pngbytes") + response = await client.get("/api/train/infer/preview/job1/3") + assert response.status_code == 200 + assert response.content == b"pngbytes" + assert response.headers["content-type"] == "image/png" + + @pytest.mark.asyncio + async def test_write_tiled_starts_a_new_job(self, client, monkeypatch, sync_threads): + monkeypatch.setattr( + infer_jobs, "run_write_tiled_job", + lambda write_jid, infer_job_id: export_jobs.update(write_jid, state="done", result={"source": infer_job_id}), + ) + response = await client.post("/api/train/infer/write-tiled/infer-job-1") + assert response.status_code == 200 + jid = response.json()["job_id"] + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"] == {"source": "infer-job-1"} + + +# --------------------------------------------------------------------------- +# /api/volume/* +# --------------------------------------------------------------------------- + +class TestVolumeRoutes: + @pytest.mark.asyncio + async def test_resolve_delegates(self, client, monkeypatch): + monkeypatch.setattr(volume_nodes, "resolve_volume", lambda server_uri, source: {"key": "vol1"}) + response = await client.get("/api/volume/resolve", params={"source": "browse/s"}) + assert response.json() == {"key": "vol1"} + + @pytest.mark.asyncio + async def test_build_inspect_delegates(self, client, monkeypatch): + monkeypatch.setattr( + volume_build, "inspect_volume_build", + lambda source, kind, server_uri: {"slices_to_read": 10}, + ) + response = await client.get("/api/volume/build/inspect", params={"source": "browse/s"}) + assert response.json() == {"slices_to_read": 10} + + @pytest.mark.asyncio + async def test_build_start_runs_job_and_reports_done(self, client, monkeypatch, sync_threads): + monkeypatch.setattr( + volume_build, "inspect_volume_build", + lambda source, kind, server_uri: {"slices_to_read": 4}, + ) + monkeypatch.setattr( + volume_build, "build_volume", + lambda source, kind, server_uri, container_path, progress: {"key": "vol1"}, + ) + response = await client.post("/api/volume/build", json={"source": "browse/s"}) + assert response.status_code == 200 + body = response.json() + assert body["slices_to_read"] == 4 + job = export_jobs.get_job(body["job_id"]) + assert job["state"] == "done" + assert job["result"] == {"key": "vol1"} + + @pytest.mark.asyncio + async def test_build_start_job_reports_error_on_exception(self, client, monkeypatch, sync_threads): + monkeypatch.setattr( + volume_build, "inspect_volume_build", + lambda source, kind, server_uri: {"slices_to_read": 4}, + ) + + def boom(*a, **k): + raise RuntimeError("build blew up") + + monkeypatch.setattr(volume_build, "build_volume", boom) + response = await client.post("/api/volume/build", json={"source": "browse/s"}) + job = export_jobs.get_job(response.json()["job_id"]) + assert job["state"] == "error" + assert "build blew up" in job["error"] diff --git a/backend/tests/test_annotation_server_local_image_routes.py b/backend/tests/test_annotation_server_local_image_routes.py new file mode 100644 index 0000000..e43a014 --- /dev/null +++ b/backend/tests/test_annotation_server_local_image_routes.py @@ -0,0 +1,371 @@ +"""HTTP-level tests for annotation_server.py's local-filesystem, connection +summary, Tiled listing, and image meta/slice routes. Continues the pattern +established in test_annotation_server_routes.py/test_annotation_server_browse_routes.py: +httpx.AsyncClient + ASGITransport, monkeypatched I/O boundaries +(local_fs/get_tiled_client/arrays_mod/images_mod), no real Tiled server or +filesystem beyond what a test itself creates.""" +from __future__ import annotations + +import numpy as np +import pytest +from httpx import ASGITransport, AsyncClient + +import annotation_server +import arrays as arrays_mod +import images as images_mod +import local_fs + +app = annotation_server.app + + +@pytest.fixture() +async def client(): + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as c: + yield c + + +class FakeNode: + def __init__(self, children=None, is_container=True): + self._children = children or {} + self.structure_family = "container" if is_container else "array" + + def __iter__(self): + return iter(self._children) + + def __getitem__(self, key): + return self._children[key] + + def keys(self): + return list(self._children.keys()) + + def __len__(self): + return len(self._children) + + +# --------------------------------------------------------------------------- +# /api/local/root +# --------------------------------------------------------------------------- + +class TestLocalRoot: + @pytest.mark.asyncio + async def test_returns_the_default_root(self, client, monkeypatch): + monkeypatch.setattr(local_fs, "default_root", lambda: "/data/raw") + response = await client.get("/api/local/root") + assert response.status_code == 200 + assert response.json() == {"root": "/data/raw"} + + +# --------------------------------------------------------------------------- +# /api/local/list, /api/local/samples +# --------------------------------------------------------------------------- + +class TestLocalList: + @pytest.mark.asyncio + async def test_lists_directory_entries(self, client, monkeypatch): + monkeypatch.setattr(local_fs, "list_dir", lambda rel, root: [{"name": "a", "is_dir": True}]) + response = await client.get("/api/local/list", params={"rel": ""}) + assert response.status_code == 200 + assert response.json() == [{"name": "a", "is_dir": True}] + + @pytest.mark.asyncio + async def test_passes_rel_and_root_through(self, client, monkeypatch): + calls = [] + monkeypatch.setattr(local_fs, "list_dir", lambda rel, root: calls.append((rel, root)) or []) + await client.get("/api/local/list", params={"rel": "sub", "root": "/data"}) + assert calls == [("sub", "/data")] + + +class TestLocalSamples: + @pytest.mark.asyncio + async def test_returns_items_and_total(self, client, monkeypatch): + monkeypatch.setattr( + local_fs, "list_image_files", + lambda rel, root: [{"name": "a.tif", "path": "a.tif"}, {"name": "b.tif", "path": "b.tif"}], + ) + response = await client.get("/api/local/samples", params={"rel": "folder"}) + assert response.status_code == 200 + body = response.json() + assert body["total"] == 2 + assert len(body["items"]) == 2 + + @pytest.mark.asyncio + async def test_requires_rel_query_param(self, client): + response = await client.get("/api/local/samples") + assert response.status_code == 422 + + +# --------------------------------------------------------------------------- +# /api/connect/summary +# --------------------------------------------------------------------------- + +class TestConnectSummaryLocal: + @pytest.mark.asyncio + async def test_counts_local_files(self, client, monkeypatch): + monkeypatch.setattr(local_fs, "count_image_files", lambda rel, root: 42) + response = await client.get( + "/api/connect/summary", params={"kind": "local", "rel": "sub", "root": "/data"}, + ) + assert response.status_code == 200 + body = response.json() + assert body["sample_count"] == 42 + assert body["label"] == "/data/sub" + assert body["kind"] == "local" + + @pytest.mark.asyncio + async def test_label_falls_back_to_default_when_no_root(self, client, monkeypatch): + monkeypatch.setattr(local_fs, "count_image_files", lambda rel, root: 0) + response = await client.get("/api/connect/summary", params={"kind": "local"}) + assert response.json()["label"] == "Local Data Root" + + +class TestConnectSummaryTiled: + @pytest.mark.asyncio + async def test_counts_container_via_len(self, client, monkeypatch): + root = FakeNode({"a": 1, "b": 2, "c": 3}) + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda uri: root) + monkeypatch.setattr(annotation_server, "get_browse_container_for", lambda client, path: (root, "browse")) + monkeypatch.setattr(annotation_server, "get_tiled_servers", lambda: {}) + response = await client.get("/api/connect/summary", params={"kind": "tiled", "server_uri": "http://x"}) + assert response.status_code == 200 + body = response.json() + assert body["sample_count"] == 3 + assert body["kind"] == "tiled" + + @pytest.mark.asyncio + async def test_label_uses_configured_server_name(self, client, monkeypatch): + root = FakeNode({}) + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda uri: root) + monkeypatch.setattr(annotation_server, "get_browse_container_for", lambda client, path: (root, "browse")) + monkeypatch.setattr( + annotation_server, "get_tiled_servers", + lambda: {"local": {"uri": "http://x:1", "name": "My Server"}}, + ) + response = await client.get("/api/connect/summary", params={"kind": "tiled", "server_uri": "http://x:1"}) + assert response.json()["label"] == "My Server" + + @pytest.mark.asyncio + async def test_container_path_appended_to_label(self, client, monkeypatch): + root = FakeNode({}) + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda uri: root) + monkeypatch.setattr(annotation_server, "get_browse_container_for", lambda client, path: (root, path or "")) + monkeypatch.setattr(annotation_server, "get_tiled_servers", lambda: {}) + response = await client.get( + "/api/connect/summary", + params={"kind": "tiled", "server_uri": "http://x", "container_path": "browse/sample"}, + ) + assert "browse/sample" in response.json()["label"] + + @pytest.mark.asyncio + async def test_count_failure_falls_back_to_zero(self, client, monkeypatch): + def boom(uri): + raise RuntimeError("down") + + monkeypatch.setattr(annotation_server, "get_tiled_client", boom) + monkeypatch.setattr(annotation_server, "get_tiled_servers", lambda: {}) + response = await client.get("/api/connect/summary", params={"kind": "tiled", "server_uri": "http://x"}) + assert response.status_code == 200 + assert response.json()["sample_count"] == 0 + + @pytest.mark.asyncio + async def test_falls_back_to_search_when_len_unsupported(self, client, monkeypatch): + class NoLenNode(FakeNode): + def __len__(self): + raise TypeError("no len") + + root = NoLenNode({"a": 1}) + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda uri: root) + monkeypatch.setattr(annotation_server, "get_browse_container_for", lambda client, path: (root, "browse")) + monkeypatch.setattr(annotation_server, "get_tiled_servers", lambda: {}) + monkeypatch.setattr(annotation_server, "tiled_search_items", lambda container, filters, limit: {"total": 7}) + response = await client.get("/api/connect/summary", params={"kind": "tiled", "server_uri": "http://x"}) + assert response.json()["sample_count"] == 7 + + +class TestConnectSummaryUnknownKind: + @pytest.mark.asyncio + async def test_unknown_kind_is_400(self, client): + response = await client.get("/api/connect/summary", params={"kind": "weird"}) + assert response.status_code == 400 + + +# --------------------------------------------------------------------------- +# /api/tiled/list +# --------------------------------------------------------------------------- + +class TestTiledList: + @pytest.mark.asyncio + async def test_lists_children_sorted_containers_first(self, client, monkeypatch): + root = FakeNode({ + "z_array": FakeNode(is_container=False), + "a_container": FakeNode({}), + }) + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda uri: root) + response = await client.get("/api/tiled/list", params={"path": ""}) + assert response.status_code == 200 + body = response.json() + assert body[0]["name"] == "a_container" + assert body[0]["is_dir"] is True + assert body[1]["name"] == "z_array" + assert body[1]["is_array"] is True + + @pytest.mark.asyncio + async def test_missing_path_segment_is_404(self, client, monkeypatch): + root = FakeNode({}) + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda uri: root) + response = await client.get("/api/tiled/list", params={"path": "nope"}) + assert response.status_code == 404 + + @pytest.mark.asyncio + async def test_leaf_array_node_is_400(self, client, monkeypatch): + class LeafNoKeys(FakeNode): + def keys(self): + raise RuntimeError("no children") + + root = FakeNode({"leaf": LeafNoKeys(is_container=False)}) + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda uri: root) + response = await client.get("/api/tiled/list", params={"path": "leaf"}) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_nested_path_joins_correctly(self, client, monkeypatch): + root = FakeNode({"a": FakeNode({"b": FakeNode(is_container=False)})}) + monkeypatch.setattr(annotation_server, "get_tiled_client", lambda uri: root) + response = await client.get("/api/tiled/list", params={"path": "a"}) + assert response.json()[0]["path"] == "a/b" + + @pytest.mark.asyncio + async def test_unexpected_error_is_500(self, client, monkeypatch): + def boom(uri): + raise RuntimeError("connection refused") + + monkeypatch.setattr(annotation_server, "get_tiled_client", boom) + response = await client.get("/api/tiled/list", params={"path": ""}) + assert response.status_code == 500 + + +# --------------------------------------------------------------------------- +# /api/image/meta, /api/image/slice +# --------------------------------------------------------------------------- + +@pytest.fixture() +def fake_image_source(monkeypatch: pytest.MonkeyPatch): + arr = np.arange(100, dtype=np.float32).reshape(10, 10) + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri, root: "node") + monkeypatch.setattr(arrays_mod, "pyramid_info", lambda source, kind, server_uri, root: None) + monkeypatch.setattr(arrays_mod, "array_shape_meta", lambda node, pyramid: { + "n_slices": 1, "height": 10, "width": 10, "dtype": "float32", "is_rgb": False, "shape_kind": "HW", + }) + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: arr) + monkeypatch.setattr(arrays_mod, "node_keywords", lambda node: ["tag1"]) + return arr + + +class TestImageMeta: + @pytest.mark.asyncio + async def test_returns_shape_and_value_range(self, client, fake_image_source): + response = await client.get( + "/api/image/meta", params={"source": "local:foo.tif", "kind": "local"}, + ) + assert response.status_code == 200 + body = response.json() + assert body["n_slices"] == 1 + assert body["height"] == 10 + assert body["width"] == 10 + assert body["value_range"] == [0.0, 99.0] + assert body["keywords"] == ["tag1"] + + @pytest.mark.asyncio + async def test_includes_global_value_range_used_by_slice_rendering(self, client, fake_image_source, monkeypatch): + monkeypatch.setattr(images_mod, "_sample_global_stats", lambda node, meta: (-73.0, 71.3)) + response = await client.get( + "/api/image/meta", params={"source": "local:foo.tif", "kind": "local"}, + ) + assert response.status_code == 200 + assert response.json()["global_value_range"] == [-73.0, 71.3] + + @pytest.mark.asyncio + async def test_resolve_failure_is_500(self, client, monkeypatch): + monkeypatch.setattr( + arrays_mod, "resolve_array", + lambda source, kind, server_uri, root: (_ for _ in ()).throw(RuntimeError("boom")), + ) + response = await client.get( + "/api/image/meta", params={"source": "local:foo.tif", "kind": "local"}, + ) + assert response.status_code == 500 + + @pytest.mark.asyncio + async def test_http_exception_is_passed_through(self, monkeypatch, client): + from fastapi import HTTPException + + def raise_404(source, kind, server_uri, root): + raise HTTPException(404, "not found") + + monkeypatch.setattr(arrays_mod, "resolve_array", raise_404) + response = await client.get( + "/api/image/meta", params={"source": "local:foo.tif", "kind": "local"}, + ) + assert response.status_code == 404 + + +class TestImageSlice: + @pytest.mark.asyncio + async def test_renders_a_png(self, client, fake_image_source, monkeypatch): + monkeypatch.setattr(images_mod, "_sample_global_stats", lambda node, meta: (0.0, 99.0)) + response = await client.get( + "/api/image/slice", params={"source": "local:foo.tif", "kind": "local"}, + ) + assert response.status_code == 200 + assert response.headers["content-type"] == "image/png" + + @pytest.mark.asyncio + async def test_unknown_denoise_method_is_400(self, client, fake_image_source): + response = await client.get( + "/api/image/slice", + params={"source": "local:foo.tif", "kind": "local", "denoise_method": "not_a_method"}, + ) + assert response.status_code == 400 + + @pytest.mark.asyncio + async def test_denoised_result_is_cached_across_requests(self, client, fake_image_source, monkeypatch): + calls = [] + real_denoised = annotation_server._denoised_slice + + def spy_denoised(node, meta, idx, method, strength, crop): + calls.append(1) + return real_denoised(node, meta, idx, method, strength, crop) + + monkeypatch.setattr(annotation_server, "_denoised_slice", spy_denoised) + monkeypatch.setattr(images_mod, "_sample_global_stats", lambda node, meta: (0.0, 99.0)) + params = {"source": "local:foo.tif", "kind": "local", "denoise_method": "median"} + r1 = await client.get("/api/image/slice", params=params) + r2 = await client.get("/api/image/slice", params=params) + assert r1.status_code == 200 and r2.status_code == 200 + assert r1.content == r2.content + assert len(calls) == 1 # second request served from cache + + @pytest.mark.asyncio + async def test_render_failure_is_500(self, client, fake_image_source, monkeypatch): + monkeypatch.setattr( + images_mod, "render_slice", + lambda sl, opts, gr: (_ for _ in ()).throw(RuntimeError("render broke")), + ) + monkeypatch.setattr(images_mod, "_sample_global_stats", lambda node, meta: (0.0, 99.0)) + response = await client.get( + "/api/image/slice", params={"source": "local:foo.tif", "kind": "local"}, + ) + assert response.status_code == 500 + + @pytest.mark.asyncio + async def test_slice_norm_skips_global_stats(self, client, fake_image_source, monkeypatch): + calls = [] + monkeypatch.setattr( + images_mod, "_sample_global_stats", + lambda node, meta: calls.append(1) or (0.0, 99.0), + ) + response = await client.get( + "/api/image/slice", + params={"source": "local:foo.tif", "kind": "local", "norm": "slice"}, + ) + assert response.status_code == 200 + assert calls == [] diff --git a/backend/tests/test_annotation_server_routes.py b/backend/tests/test_annotation_server_routes.py new file mode 100644 index 0000000..6fccab4 --- /dev/null +++ b/backend/tests/test_annotation_server_routes.py @@ -0,0 +1,217 @@ +"""HTTP-level tests for annotation_server.py's session-persistence and +export-job-status routes — mirrors the httpx.AsyncClient + ASGITransport +pattern already established in test_ipred_routes.py/test_train_routes.py. + +Scoped to the routes testable with local filesystem + in-memory job registry +state (drafts, guide, measure, export status/cancel/download) — the +biggest, most tractable slice of this file's route coverage. Browse/ingest/ +zarr/volume/tiff-stack/denoise routes need a real or heavily-mocked Tiled +client and are left for a further pass. +""" +from __future__ import annotations + +import numpy as np +import pytest +from httpx import ASGITransport, AsyncClient + +import arrays as arrays_mod +import drafts +import export_jobs +import guides +from annotation_server import app + + +@pytest.fixture() +def local_data_root(tmp_path, monkeypatch: pytest.MonkeyPatch): + """drafts.py/guides.py compute their storage dir at IMPORT time + (module-level `_DRAFT_DIR`), so setting LOCAL_DATA_ROOT alone has no + effect on an already-imported module — patch the attribute directly.""" + draft_dir = tmp_path / ".drafts" + monkeypatch.setattr(drafts, "_DRAFT_DIR", draft_dir) + monkeypatch.setattr(guides, "_DRAFT_DIR", draft_dir) + return draft_dir + + +@pytest.fixture() +async def client(): + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as c: + yield c + + +# --------------------------------------------------------------------------- +# Drafts +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_get_draft_404_when_none_saved(client, local_data_root): + response = await client.get("/api/annotations/draft", params={"source_key": "local:x.tif"}) + assert response.status_code == 404 + + +@pytest.mark.asyncio +async def test_put_then_get_draft_round_trips(client, local_data_root): + payload = {"classes": [{"classId": 1, "label": "A", "color": "#f00"}], "slices": {"0": []}} + put_res = await client.put( + "/api/annotations/draft", params={"source_key": "local:x.tif"}, json=payload, + ) + assert put_res.status_code == 200 + + get_res = await client.get("/api/annotations/draft", params={"source_key": "local:x.tif"}) + assert get_res.status_code == 200 + body = get_res.json() + assert body["source_key"] == "local:x.tif" + assert body["payload"]["classes"][0]["label"] == "A" + + +@pytest.mark.asyncio +async def test_list_drafts_includes_saved_ones(client, local_data_root): + await client.put( + "/api/annotations/draft", params={"source_key": "local:x.tif"}, + json={"classes": [], "slices": {}}, + ) + response = await client.get("/api/annotations/drafts") + assert response.status_code == 200 + keys = [d["source_key"] for d in response.json()] + assert "local:x.tif" in keys + + +# --------------------------------------------------------------------------- +# Guide +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_get_guide_404_when_none_saved(client, local_data_root): + response = await client.get("/api/guide", params={"source_key": "local:x.tif"}) + assert response.status_code == 404 + + +@pytest.mark.asyncio +async def test_put_then_get_guide_round_trips(client, local_data_root): + payload = {"classes": [{"label": "A", "color": "#f00", "description": "an example"}], "notes": "n"} + put_res = await client.put("/api/guide", params={"source_key": "local:x.tif"}, json=payload) + assert put_res.status_code == 200 + + get_res = await client.get("/api/guide", params={"source_key": "local:x.tif"}) + assert get_res.status_code == 200 + assert get_res.json()["guide"]["notes"] == "n" + + +# --------------------------------------------------------------------------- +# Measure +# --------------------------------------------------------------------------- + +@pytest.fixture() +def fake_measure_source(monkeypatch: pytest.MonkeyPatch): + arr = np.zeros((10, 10), dtype=np.float32) + # shape_to_mask's rectangle rasterization is inclusive of both endpoints + # (skimage.draw.rectangle "like mlex") — x=2,y=2,w=3,h=3 covers columns/ + # rows 2..5 inclusive, a 4x4=16 region, not the 3x3=9 a half-open range + # would give. Filled to match exactly so every measured pixel is 10.0. + arr[2:6, 2:6] = 10.0 + + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri: arr) + monkeypatch.setattr( + arrays_mod, "array_shape_meta", + lambda node, pyramid=None: {"height": 10, "width": 10, "n_slices": 1}, + ) + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: node) + return arr + + +@pytest.mark.asyncio +async def test_measure_returns_stats_for_shape_region(client, fake_measure_source): + shape = {"id": "a", "kind": "rectangle", "classId": 1, "x": 2, "y": 2, "w": 3, "h": 3} + response = await client.post( + "/api/measure", params={"source_key": "local:x.tif"}, + json={"slice_index": 0, "shapes": [shape]}, + ) + assert response.status_code == 200 + body = response.json() + assert body["pixel_count"] == 16 + assert body["min"] == 10.0 + assert body["max"] == 10.0 + assert body["mean"] == 10.0 + + +@pytest.mark.asyncio +async def test_measure_with_no_shapes_returns_empty_stats(client, fake_measure_source): + response = await client.post( + "/api/measure", params={"source_key": "local:x.tif"}, + json={"slice_index": 0, "shapes": []}, + ) + assert response.status_code == 200 + body = response.json() + assert body["pixel_count"] == 0 + assert body["min"] is None + + +@pytest.mark.asyncio +async def test_measure_ignores_a_shape_that_fails_to_rasterize(client, fake_measure_source): + bad_shape = {"id": "bad", "kind": "polygon", "classId": 1, "points": []} + good_shape = {"id": "good", "kind": "rectangle", "classId": 1, "x": 2, "y": 2, "w": 3, "h": 3} + response = await client.post( + "/api/measure", params={"source_key": "local:x.tif"}, + json={"slice_index": 0, "shapes": [bad_shape, good_shape]}, + ) + assert response.status_code == 200 + assert response.json()["pixel_count"] == 16 + + +# --------------------------------------------------------------------------- +# Export job status/cancel/download +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_export_status_unknown_job_404s(client): + response = await client.get("/api/export/status/does-not-exist") + assert response.status_code == 404 + + +@pytest.mark.asyncio +async def test_export_status_returns_real_job_state(client): + jid = export_jobs.new_job("some/path") + export_jobs.update(jid, state="running", phase="working") + export_jobs.set_total(jid, 10) + export_jobs.bump(jid, 3) + + response = await client.get(f"/api/export/status/{jid}") + assert response.status_code == 200 + body = response.json() + assert body["state"] == "running" + assert body["done"] == 3 + assert body["total"] == 10 + + +@pytest.mark.asyncio +async def test_export_cancel_unknown_job_404s(client): + response = await client.post("/api/export/cancel/does-not-exist") + assert response.status_code == 404 + + +@pytest.mark.asyncio +async def test_export_cancel_sets_flag_on_real_job(client): + jid = export_jobs.new_job("x") + response = await client.post(f"/api/export/cancel/{jid}") + assert response.status_code == 200 + assert response.json() == {"cancelled": True} + assert export_jobs.cancel_requested(jid) is True + + +@pytest.mark.asyncio +async def test_export_download_404s_when_zip_not_ready(client): + jid = export_jobs.new_job("x") + response = await client.get(f"/api/export/download/{jid}") + assert response.status_code == 404 + + +@pytest.mark.asyncio +async def test_export_download_streams_the_real_zip_file(client, tmp_path): + jid = export_jobs.new_job("x") + zip_path = tmp_path / "out.zip" + zip_path.write_bytes(b"PK\x03\x04fakezip") + export_jobs.update(jid, zip_path=str(zip_path)) + + response = await client.get(f"/api/export/download/{jid}") + assert response.status_code == 200 + assert response.content == b"PK\x03\x04fakezip" + assert response.headers["content-type"] == "application/zip" diff --git a/backend/tests/test_annotation_thumbnails.py b/backend/tests/test_annotation_thumbnails.py new file mode 100644 index 0000000..ddd6eb9 --- /dev/null +++ b/backend/tests/test_annotation_thumbnails.py @@ -0,0 +1,334 @@ +"""Tests for annotation_thumbnails.py — real PIL/numpy rendering, arrays.py +resolve/meta/read functions monkeypatched (same fake-array-source pattern as +test_annotation_server_routes.py's measure fixture); Tiled upload path uses a +lightweight duck-typed fake container.""" +from __future__ import annotations + +import base64 +from io import BytesIO + +import numpy as np +import pytest +from PIL import Image as PILImage + +import annotation_thumbnails +import arrays as arrays_mod + +# --------------------------------------------------------------------------- +# decode_thumbnail_base64 +# --------------------------------------------------------------------------- + + +class TestDecodeThumbnailBase64: + def test_empty_string_returns_none(self): + assert annotation_thumbnails.decode_thumbnail_base64("") is None + + def test_plain_base64_decodes(self): + raw = base64.b64encode(b"hello").decode() + assert annotation_thumbnails.decode_thumbnail_base64(raw) == b"hello" + + def test_data_url_prefix_is_stripped(self): + raw = base64.b64encode(b"pngdata").decode() + data_url = f"data:image/png;base64,{raw}" + assert annotation_thumbnails.decode_thumbnail_base64(data_url) == b"pngdata" + + def test_invalid_base64_returns_none(self): + assert annotation_thumbnails.decode_thumbnail_base64("not-@@base64!!") is None + + +# --------------------------------------------------------------------------- +# _stride_downsample +# --------------------------------------------------------------------------- + +class TestStrideDownsample: + def test_small_array_returned_unchanged(self): + arr = np.zeros((10, 10)) + out = annotation_thumbnails._stride_downsample(arr, 512) + assert out is arr + + def test_large_2d_array_is_strided(self): + arr = np.zeros((1000, 1000)) + out = annotation_thumbnails._stride_downsample(arr, 500) + assert max(out.shape) <= 500 + + def test_large_3d_array_preserves_channel_dim(self): + arr = np.zeros((1000, 1000, 3)) + out = annotation_thumbnails._stride_downsample(arr, 500) + assert out.shape[2] == 3 + assert max(out.shape[:2]) <= 500 + + +# --------------------------------------------------------------------------- +# _hex_to_rgba +# --------------------------------------------------------------------------- + +class TestHexToRgba: + def test_full_hex_with_hash(self): + assert annotation_thumbnails._hex_to_rgba("#ff0080", 100) == (255, 0, 128, 100) + + def test_hex_without_hash(self): + assert annotation_thumbnails._hex_to_rgba("00ff00", 200) == (0, 255, 0, 200) + + def test_short_hex_is_padded(self): + # "abc" -> ljust(6, "0") -> "abc000" + assert annotation_thumbnails._hex_to_rgba("#abc", 50) == (0xAB, 0xC0, 0x00, 50) + + +# --------------------------------------------------------------------------- +# _draw_shapes (via a real PIL ImageDraw) +# --------------------------------------------------------------------------- + +class TestDrawShapes: + def _draw(self, shapes, class_colors=None): + img = PILImage.new("RGBA", (50, 50), (0, 0, 0, 0)) + draw = annotation_thumbnails.ImageDraw.Draw(img) + annotation_thumbnails._draw_shapes(draw, shapes, class_colors or {}, scale=1.0) + return np.array(img) + + def test_polygon_is_drawn(self): + shape = {"kind": "polygon", "classId": 1, "points": [5, 5, 20, 5, 20, 20, 5, 20]} + pixels = self._draw([shape], {1: "#ff0000"}) + assert pixels[10, 10, 3] > 0 # alpha channel non-zero inside the fill + + def test_polygon_with_too_few_points_is_skipped(self): + shape = {"kind": "polygon", "classId": 1, "points": [5, 5]} + pixels = self._draw([shape]) + assert np.all(pixels[:, :, 3] == 0) + + def test_rectangle_is_drawn(self): + shape = {"kind": "rectangle", "classId": 1, "x": 5, "y": 5, "w": 10, "h": 10} + pixels = self._draw([shape], {1: "#00ff00"}) + assert pixels[10, 10, 3] > 0 + + def test_ellipse_is_drawn(self): + shape = {"kind": "ellipse", "classId": 1, "cx": 25, "cy": 25, "rx": 10, "ry": 10} + pixels = self._draw([shape], {1: "#0000ff"}) + assert pixels[25, 25, 3] > 0 + + def test_brush_stroke_is_drawn(self): + shape = { + "kind": "brush", + "classId": 1, + "strokes": [{"mode": "paint", "points": [5, 5, 30, 30], "radius": 3}], + } + pixels = self._draw([shape], {1: "#ffffff"}) + assert pixels.sum() > 0 + + def test_brush_erase_stroke_is_skipped(self): + shape = { + "kind": "brush", + "classId": 1, + "strokes": [{"mode": "erase", "points": [5, 5, 30, 30], "radius": 3}], + } + pixels = self._draw([shape]) + assert np.all(pixels[:, :, 3] == 0) + + def test_brush_stroke_with_one_point_is_skipped(self): + shape = { + "kind": "brush", + "classId": 1, + "strokes": [{"mode": "paint", "points": [5, 5], "radius": 3}], + } + pixels = self._draw([shape]) + assert np.all(pixels[:, :, 3] == 0) + + def test_unknown_class_id_falls_back_to_default_color(self): + shape = {"kind": "rectangle", "classId": 999, "x": 0, "y": 0, "w": 5, "h": 5} + pixels = self._draw([shape], {1: "#00ff00"}) + assert pixels[2, 2, 3] > 0 + + +# --------------------------------------------------------------------------- +# render_annotated_thumbnail +# --------------------------------------------------------------------------- + +@pytest.fixture() +def fake_array_source(monkeypatch: pytest.MonkeyPatch): + arr = np.random.default_rng(0).integers(0, 255, size=(20, 20), dtype=np.uint8).astype(np.float32) + + def fake_read_slice(node, meta, idx): + return arr + idx # differ per slice so "best slice" pick is checkable + + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri: "node") + monkeypatch.setattr(arrays_mod, "array_shape_meta", lambda node: {"n_slices": 5}) + monkeypatch.setattr(arrays_mod, "read_slice", fake_read_slice) + return arr + + +class TestRenderAnnotatedThumbnail: + def test_returns_png_bytes_for_grayscale_source(self, fake_array_source): + payload = {"classes": [], "slices": {}} + png = annotation_thumbnails.render_annotated_thumbnail("local:foo.tif", payload) + assert png is not None + decoded = np.array(PILImage.open(BytesIO(png))) + assert decoded.shape[:2] == (20, 20) + + def test_picks_slice_with_most_shapes(self, fake_array_source, monkeypatch): + calls = [] + real_read_slice = arrays_mod.read_slice + + def spy_read_slice(node, meta, idx): + calls.append(idx) + return real_read_slice(node, meta, idx) + + monkeypatch.setattr(arrays_mod, "read_slice", spy_read_slice) + payload = { + "classes": [{"classId": 1, "color": "#ff0000"}], + "slices": { + "0": [{"kind": "rectangle", "classId": 1, "x": 0, "y": 0, "w": 2, "h": 2}], + "3": [ + {"kind": "rectangle", "classId": 1, "x": 0, "y": 0, "w": 2, "h": 2}, + {"kind": "rectangle", "classId": 1, "x": 5, "y": 5, "w": 2, "h": 2}, + ], + }, + } + png = annotation_thumbnails.render_annotated_thumbnail("local:foo.tif", payload) + assert png is not None + assert calls == [3] + + def test_best_slice_clamped_to_valid_range(self, fake_array_source): + payload = {"classes": [], "slices": {"999": [{"kind": "rectangle", "x": 0, "y": 0, "w": 1, "h": 1}]}} + png = annotation_thumbnails.render_annotated_thumbnail("local:foo.tif", payload) + assert png is not None + + def test_non_integer_slice_key_is_ignored(self, fake_array_source): + payload = {"classes": [], "slices": {"not-a-number": [{"kind": "rectangle"}]}} + png = annotation_thumbnails.render_annotated_thumbnail("local:foo.tif", payload) + assert png is not None + + def test_missing_arrays_module_returns_none(self, monkeypatch): + import builtins + real_import = builtins.__import__ + + def fake_import(name, *args, **kwargs): + if name == "arrays": + raise ImportError("no arrays") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", fake_import) + assert annotation_thumbnails.render_annotated_thumbnail("local:x", {}) is None + + def test_resolve_array_failure_returns_none(self, monkeypatch): + monkeypatch.setattr( + arrays_mod, "resolve_array", + lambda source, kind, server_uri: (_ for _ in ()).throw(RuntimeError("boom")), + ) + assert annotation_thumbnails.render_annotated_thumbnail("local:x", {}) is None + + def test_unsupported_array_shape_returns_none(self, monkeypatch): + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri: "node") + monkeypatch.setattr(arrays_mod, "array_shape_meta", lambda node: {"n_slices": 1}) + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: np.zeros(10)) + assert annotation_thumbnails.render_annotated_thumbnail("local:x", {}) is None + + def test_read_slice_failure_returns_none(self, monkeypatch): + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri: "node") + monkeypatch.setattr(arrays_mod, "array_shape_meta", lambda node: {"n_slices": 1}) + + def boom(node, meta, idx): + raise RuntimeError("read failed") + + monkeypatch.setattr(arrays_mod, "read_slice", boom) + assert annotation_thumbnails.render_annotated_thumbnail("local:x", {}) is None + + def test_rgb_source_uses_prepare_rgb_path(self, monkeypatch): + rgb_arr = np.zeros((15, 15, 3), dtype=np.uint8) + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri: "node") + monkeypatch.setattr(arrays_mod, "array_shape_meta", lambda node: {"n_slices": 1}) + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: rgb_arr) + png = annotation_thumbnails.render_annotated_thumbnail("local:x", {"classes": [], "slices": {}}) + assert png is not None + decoded = np.array(PILImage.open(BytesIO(png))) + assert decoded.shape[:2] == (15, 15) + + +# --------------------------------------------------------------------------- +# upload_thumbnail_to_tiled +# --------------------------------------------------------------------------- + +class FakeArrayContainer: + def __init__(self): + self.written = [] + + def write_array(self, arr, key, metadata): + self.written.append((key, arr, metadata)) + + +class FakeContainerNode(dict): + def __init__(self, *a, **kw): + super().__init__(*a, **kw) + self.created = [] + + def create_container(self, key, metadata): + c = FakeArrayContainer() + self[key] = c + self.created.append((key, metadata)) + return c + + +@pytest.fixture() +def fake_tiled(monkeypatch: pytest.MonkeyPatch): + import tiled_clients + + browse = FakeContainerNode() + root = FakeContainerNode({"browse": browse}) + monkeypatch.setattr(tiled_clients, "api_key_for_uri", lambda uri: None) + monkeypatch.setattr(tiled_clients, "get_tiled_client", lambda uri, key: root) + return browse + + +class TestUploadThumbnailToTiled: + def _png_bytes(self): + img = PILImage.new("RGB", (4, 4), (10, 20, 30)) + buf = BytesIO() + img.save(buf, format="PNG") + return buf.getvalue() + + def test_local_source_is_a_no_op(self, fake_tiled): + annotation_thumbnails.upload_thumbnail_to_tiled("local:foo.tif", 1, "2024-01-01", self._png_bytes()) + assert fake_tiled.created == [] + + def test_creates_container_and_writes_array_on_first_version(self, fake_tiled): + annotation_thumbnails.upload_thumbnail_to_tiled( + "tiled::browse/sample1", 3, "2024-01-01T00:00:00", self._png_bytes(), + ) + assert len(fake_tiled.created) == 1 + key, metadata = fake_tiled.created[0] + assert key == "sample1__v_thumbs" + container = fake_tiled[key] + assert container.written[0][0] == "v0003" + assert container.written[0][2]["version"] == 3 + + def test_reuses_existing_container_on_subsequent_versions(self, fake_tiled): + annotation_thumbnails.upload_thumbnail_to_tiled( + "tiled::browse/sample1", 1, "t1", self._png_bytes(), + ) + annotation_thumbnails.upload_thumbnail_to_tiled( + "tiled::browse/sample1", 2, "t2", self._png_bytes(), + ) + assert len(fake_tiled.created) == 1 + container = fake_tiled["sample1__v_thumbs"] + assert [w[0] for w in container.written] == ["v0001", "v0002"] + + def test_missing_tiled_clients_module_is_a_no_op(self, monkeypatch): + import builtins + real_import = builtins.__import__ + + def fake_import(name, *args, **kwargs): + if name == "tiled_clients": + raise ImportError("no tiled_clients") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", fake_import) + # Should not raise. + annotation_thumbnails.upload_thumbnail_to_tiled("tiled::browse/s", 1, "t", self._png_bytes()) + + def test_tiled_error_is_swallowed(self, monkeypatch): + import tiled_clients + monkeypatch.setattr(tiled_clients, "api_key_for_uri", lambda uri: None) + monkeypatch.setattr( + tiled_clients, "get_tiled_client", + lambda uri, key: (_ for _ in ()).throw(RuntimeError("down")), + ) + # Should not raise even though the Tiled call blows up. + annotation_thumbnails.upload_thumbnail_to_tiled("tiled::browse/s", 1, "t", self._png_bytes()) diff --git a/backend/tests/test_arrays.py b/backend/tests/test_arrays.py new file mode 100644 index 0000000..59f9fb8 --- /dev/null +++ b/backend/tests/test_arrays.py @@ -0,0 +1,516 @@ +"""Tests for arrays.py — the shape-dispatch/slice-reading core shared by +almost every route. Fake Tiled containers are plain duck-typed classes (same +minimal-surface pattern as test_browse_helpers.py's FakeNode); leaf arrays +are either real numpy arrays or a thin FakeArrayNode wrapper exposing +.shape/.dtype/.metadata/__array__/__getitem__, matching what a real Tiled +array client node looks like.""" +from __future__ import annotations + +import numpy as np +import pytest +from fastapi import HTTPException + +import arrays +import local_fs + + +class FakeArrayNode: + def __init__(self, arr, metadata=None): + self._arr = np.asarray(arr) + self.shape = self._arr.shape + self.dtype = self._arr.dtype + self.metadata = metadata or {} + + def __array__(self, dtype=None): + return self._arr if dtype is None else self._arr.astype(dtype) + + def __getitem__(self, idx): + return self._arr[idx] + + +class FakeContainer: + def __init__(self, children=None, metadata=None): + self.structure_family = "container" + self._children = dict(children or {}) + self.metadata = metadata or {} + + def __iter__(self): + return iter(self._children) + + def __getitem__(self, key): + return self._children[key] + + def __len__(self): + return len(self._children) + + +@pytest.fixture(autouse=True) +def clear_node_cache(): + arrays._node_cache.clear() + yield + arrays._node_cache.clear() + + +# --------------------------------------------------------------------------- +# array_shape_meta — direct array shape dispatch +# --------------------------------------------------------------------------- + +class TestArrayShapeMetaDirect: + def test_2d_array_is_hw(self): + meta = arrays.array_shape_meta(np.zeros((10, 20), dtype=np.uint8)) + assert meta == { + "n_slices": 1, "height": 10, "width": 20, + "dtype": "uint8", "is_rgb": False, "shape_kind": "HW", + } + + def test_3d_rgb_array_is_hwc(self): + meta = arrays.array_shape_meta(np.zeros((10, 20, 3), dtype=np.uint8)) + assert meta["shape_kind"] == "HWC" + assert meta["is_rgb"] is True + assert meta["n_slices"] == 1 + + def test_3d_rgba_array_is_hwc(self): + meta = arrays.array_shape_meta(np.zeros((10, 20, 4), dtype=np.uint8)) + assert meta["shape_kind"] == "HWC" + + def test_3d_stack_is_nhw(self): + meta = arrays.array_shape_meta(np.zeros((5, 10, 20), dtype=np.float32)) + assert meta["shape_kind"] == "NHW" + assert meta["n_slices"] == 5 + assert meta["height"] == 10 + assert meta["width"] == 20 + + def test_4d_rgb_stack_is_nhwc(self): + meta = arrays.array_shape_meta(np.zeros((5, 10, 20, 3), dtype=np.uint8)) + assert meta["shape_kind"] == "NHWC" + assert meta["n_slices"] == 5 + assert meta["is_rgb"] is True + + def test_1d_array_raises(self): + with pytest.raises(HTTPException) as exc: + arrays.array_shape_meta(np.zeros(10)) + assert exc.value.status_code == 422 + + def test_5d_array_raises(self): + with pytest.raises(HTTPException) as exc: + arrays.array_shape_meta(np.zeros((1, 2, 3, 4, 5))) + assert exc.value.status_code == 422 + + def test_array_like_without_shape_attr_goes_through_asarray(self): + class ArrayLikeNoShape: + def __array__(self, dtype=None): + return np.zeros((4, 4), dtype=np.uint8) + + meta = arrays.array_shape_meta(ArrayLikeNoShape()) + assert meta["shape_kind"] == "HW" + + def test_pyramid_kwarg_reports_finest_geometry(self): + pyramid = { + "full_shape": [100, 512, 512], "z_downsample": 4.0, + "level_key": "scale2", "level_index": 2, "level_count": 3, + } + meta = arrays.array_shape_meta(np.zeros((25, 128, 128)), pyramid=pyramid) + assert meta["n_slices"] == 100 + assert meta["height"] == 512 + assert meta["width"] == 512 + assert meta["z_downsample"] == 4.0 + assert meta["level_height"] == 128 + assert meta["level_n_slices"] == 25 + + +# --------------------------------------------------------------------------- +# array_shape_meta — container/stack dispatch +# --------------------------------------------------------------------------- + +class TestArrayShapeMetaStack: + def test_container_of_2d_arrays_is_stack(self): + node = FakeContainer({"b": np.zeros((4, 4)), "a": np.ones((4, 4))}) + meta = arrays.array_shape_meta(node) + assert meta["shape_kind"] == "STACK" + assert meta["n_slices"] == 2 + assert meta["keys"] == ["a", "b"] # sorted lexically + assert meta["is_rgb"] is False + + def test_container_of_rgb_arrays_is_rgb_stack(self): + node = FakeContainer({"a": np.zeros((4, 4, 3))}) + meta = arrays.array_shape_meta(node) + assert meta["is_rgb"] is True + + def test_empty_container_raises(self): + with pytest.raises(HTTPException) as exc: + arrays.array_shape_meta(FakeContainer({})) + assert exc.value.status_code == 422 + + def test_container_with_unsupported_child_shape_raises(self): + node = FakeContainer({"a": np.zeros((4, 4, 5))}) + with pytest.raises(HTTPException) as exc: + arrays.array_shape_meta(node) + assert exc.value.status_code == 422 + + +# --------------------------------------------------------------------------- +# read_slice +# --------------------------------------------------------------------------- + +class TestReadSliceDirect: + def test_hw_ignores_index(self): + arr = np.arange(9).reshape(3, 3) + meta = {"shape_kind": "HW"} + out = arrays.read_slice(arr, meta, 5) + assert np.array_equal(out, arr) + + def test_hwc_ignores_index(self): + arr = np.zeros((3, 3, 3)) + out = arrays.read_slice(arr, {"shape_kind": "HWC"}, 2) + assert out.shape == (3, 3, 3) + + def test_nhw_indexes_by_slice(self): + node = FakeArrayNode(np.arange(2 * 4 * 4).reshape(2, 4, 4)) + out = arrays.read_slice(node, {"shape_kind": "NHW"}, 1) + assert np.array_equal(out, node._arr[1]) + + def test_nhw_maps_full_res_index_through_z_downsample(self): + node = FakeArrayNode(np.arange(3 * 2 * 2).reshape(3, 2, 2)) + meta = {"shape_kind": "NHW", "z_downsample": 3.0, "level_n_slices": 3} + # full-res idx 6 -> round(6/3)=2, clamped to level_n_slices-1=2 + out = arrays.read_slice(node, meta, 6) + assert np.array_equal(out, node._arr[2]) + + def test_nhw_z_downsample_of_one_is_a_no_op(self): + node = FakeArrayNode(np.arange(3 * 2 * 2).reshape(3, 2, 2)) + meta = {"shape_kind": "NHW", "z_downsample": 1.0} + out = arrays.read_slice(node, meta, 1) + assert np.array_equal(out, node._arr[1]) + + def test_unrecognized_shape_kind_raises(self): + with pytest.raises(HTTPException) as exc: + arrays.read_slice(np.zeros((2, 2)), {"shape_kind": "WEIRD"}, 0) + assert exc.value.status_code == 422 + + +class TestReadSliceStack: + def test_reads_by_sorted_key_order(self): + node = FakeContainer({"b": np.full((2, 2), 2.0), "a": np.full((2, 2), 1.0)}) + meta = {"shape_kind": "STACK", "keys": ["a", "b"]} + out = arrays.read_slice(node, meta, 1) + assert np.all(out == 2.0) + + def test_out_of_range_index_falls_back_to_first_key(self): + node = FakeContainer({"a": np.full((2, 2), 1.0), "b": np.full((2, 2), 2.0)}) + meta = {"shape_kind": "STACK", "keys": ["a", "b"]} + out = arrays.read_slice(node, meta, 99) + assert np.all(out == 1.0) + + def test_missing_keys_in_meta_recomputes_from_node(self): + node = FakeContainer({"a": np.full((2, 2), 1.0)}) + meta = {"shape_kind": "STACK"} + out = arrays.read_slice(node, meta, 0) + assert np.all(out == 1.0) + + def test_empty_stack_raises(self): + node = FakeContainer({}) + with pytest.raises(HTTPException) as exc: + arrays.read_slice(node, {"shape_kind": "STACK", "keys": []}, 0) + assert exc.value.status_code == 422 + + +# --------------------------------------------------------------------------- +# _is_container_node / multiscale_levels / _level_array / _descend_to_stack +# --------------------------------------------------------------------------- + +class TestIsContainerNode: + def test_plain_array_is_not_a_container(self): + assert arrays._is_container_node(np.zeros((2, 2))) is False + + def test_fake_container_is_a_container(self): + assert arrays._is_container_node(FakeContainer({})) is True + + def test_enum_like_structure_family_value_is_read(self): + class FakeEnum: + value = "container" + + class Node: + structure_family = FakeEnum() + + assert arrays._is_container_node(Node()) is True + + +class TestMultiscaleLevels: + def test_non_container_returns_none(self): + assert arrays.multiscale_levels(np.zeros((2, 2))) is None + + def test_container_without_scale_children_returns_none(self): + assert arrays.multiscale_levels(FakeContainer({"foo": 1, "bar": 2})) is None + + def test_single_scale_child_returns_none(self): + assert arrays.multiscale_levels(FakeContainer({"scale0": 1})) is None + + def test_multiple_scale_children_sorted_numerically(self): + node = FakeContainer({"scale10": 1, "scale2": 2, "scale0": 3}) + assert arrays.multiscale_levels(node) == ["scale0", "scale2", "scale10"] + + def test_non_enumerable_container_returns_none(self): + class BrokenContainer: + structure_family = "container" + + def __iter__(self): + raise RuntimeError("boom") + + assert arrays.multiscale_levels(BrokenContainer()) is None + + +class TestLevelArray: + def test_level_wrapping_a_single_array_child(self): + arr = np.zeros((2, 2)) + level = FakeContainer({"image": arr}) + node = FakeContainer({"scale0": level}) + assert arrays._level_array(node, "scale0") is arr + + def test_level_that_is_already_an_array(self): + arr = np.zeros((2, 2)) + node = FakeContainer({"scale0": arr}) + assert arrays._level_array(node, "scale0") is arr + + def test_missing_level_key_returns_none(self): + node = FakeContainer({}) + assert arrays._level_array(node, "scale0") is None + + def test_empty_level_container_returns_none(self): + node = FakeContainer({"scale0": FakeContainer({})}) + assert arrays._level_array(node, "scale0") is None + + +class TestDescendToStack: + def test_non_container_returned_unchanged(self): + arr = np.zeros((2, 2)) + assert arrays._descend_to_stack(arr) is arr + + def test_descends_through_wrapper_container_to_array_stack(self): + stack = FakeContainer({"slice_0": np.zeros((2, 2))}) + wrapper = FakeContainer({"dataset": stack}) + assert arrays._descend_to_stack(wrapper) is stack + + def test_stops_at_container_of_arrays(self): + stack = FakeContainer({"slice_0": np.zeros((2, 2)), "slice_1": np.zeros((2, 2))}) + assert arrays._descend_to_stack(stack) is stack + + def test_empty_container_returned_as_is(self): + empty = FakeContainer({}) + assert arrays._descend_to_stack(empty) is empty + + def test_max_depth_stops_infinite_wrapper_chain(self): + node = FakeContainer({}) + for _ in range(20): + node = FakeContainer({"only": node}) + result = arrays._descend_to_stack(node, max_depth=3) + # Should not infinite-loop or raise; just stop after max_depth hops. + assert arrays._is_container_node(result) + + def test_multiscale_volume_resolves_to_finest_level_array(self): + finest = np.zeros((3, 3)) + level0 = FakeContainer({"image": finest}) + level1 = FakeContainer({"image": np.zeros((2, 2))}) + node = FakeContainer({"scale0": level0, "scale1": level1}) + assert arrays._descend_to_stack(node) is finest + + +# --------------------------------------------------------------------------- +# node_keywords +# --------------------------------------------------------------------------- + +class TestNodeKeywords: + def test_no_metadata_returns_empty(self): + assert arrays.node_keywords(np.zeros((2, 2))) == [] + + def test_list_keywords_on_node_itself(self): + node = FakeArrayNode(np.zeros((2, 2)), metadata={"keywords": ["a", "b"]}) + assert arrays.node_keywords(node) == ["a", "b"] + + def test_string_keyword_is_wrapped_in_a_list(self): + node = FakeArrayNode(np.zeros((2, 2)), metadata={"keywords": "solo"}) + assert arrays.node_keywords(node) == ["solo"] + + def test_blank_string_values_filtered_out(self): + node = FakeArrayNode(np.zeros((2, 2)), metadata={"keywords": ["a", " ", ""]}) + assert arrays.node_keywords(node) == ["a"] + + def test_falls_back_to_first_child_of_container(self): + child = FakeArrayNode(np.zeros((2, 2)), metadata={"keywords": ["child-tag"]}) + container = FakeContainer({"a": child}) + assert arrays.node_keywords(container) == ["child-tag"] + + def test_container_metadata_wins_over_child(self): + child = FakeArrayNode(np.zeros((2, 2)), metadata={"keywords": ["child-tag"]}) + container = FakeContainer({"a": child}, metadata={"keywords": ["container-tag"]}) + assert arrays.node_keywords(container) == ["container-tag"] + + def test_container_with_no_tags_anywhere_returns_empty(self): + container = FakeContainer({"a": FakeArrayNode(np.zeros((2, 2)))}) + assert arrays.node_keywords(container) == [] + + def test_error_reading_child_metadata_is_swallowed(self): + class BrokenContainer(FakeContainer): + def __getitem__(self, key): + raise RuntimeError("boom") + + container = BrokenContainer({"a": 1}) + assert arrays.node_keywords(container) == [] + + +# --------------------------------------------------------------------------- +# resolve_array / resolve_container — kind dispatch + caching +# --------------------------------------------------------------------------- + +class TestResolveArrayLocal: + def test_local_kind_delegates_to_local_fs(self, monkeypatch: pytest.MonkeyPatch, tmp_path): + arr = np.zeros((3, 3)) + monkeypatch.setattr(local_fs, "open_array", lambda source, root: arr) + result = arrays.resolve_array("foo.npy", "local", root=str(tmp_path)) + assert result is arr + + def test_unknown_kind_raises_422(self): + with pytest.raises(HTTPException) as exc: + arrays.resolve_array("x", "weird") + assert exc.value.status_code == 422 + + +class TestResolveArrayTiled: + def test_resolves_nested_path_and_descends(self, monkeypatch: pytest.MonkeyPatch): + leaf = np.zeros((2, 2)) + stack = FakeContainer({"slice_0": leaf}) + root = FakeContainer({"browse": FakeContainer({"sample1": stack})}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + result = arrays.resolve_array("browse/sample1", "tiled") + assert result is stack + + def test_missing_path_raises_404(self, monkeypatch: pytest.MonkeyPatch): + root = FakeContainer({}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + with pytest.raises(HTTPException) as exc: + arrays.resolve_array("does/not/exist", "tiled") + assert exc.value.status_code == 404 + + def test_result_is_cached_by_key(self, monkeypatch: pytest.MonkeyPatch): + calls = [] + leaf = np.zeros((2, 2)) + root = FakeContainer({"a": leaf}) + + def fake_get_client(uri, key): + calls.append(uri) + return root + + monkeypatch.setattr("arrays.get_tiled_client", fake_get_client) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + arrays.resolve_array("a", "tiled", server_uri="http://x") + arrays.resolve_array("a", "tiled", server_uri="http://x") + assert len(calls) == 1 + + def test_different_server_uri_is_a_different_cache_key(self, monkeypatch: pytest.MonkeyPatch): + calls = [] + leaf = np.zeros((2, 2)) + root = FakeContainer({"a": leaf}) + + def fake_get_client(uri, key): + calls.append(uri) + return root + + monkeypatch.setattr("arrays.get_tiled_client", fake_get_client) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + arrays.resolve_array("a", "tiled", server_uri="http://x") + arrays.resolve_array("a", "tiled", server_uri="http://y") + assert len(calls) == 2 + + +class TestResolveContainer: + def test_non_tiled_kind_raises(self): + with pytest.raises(HTTPException) as exc: + arrays.resolve_container("x", "local") + assert exc.value.status_code == 422 + + def test_resolves_without_descending(self, monkeypatch: pytest.MonkeyPatch): + stack = FakeContainer({"slice_0": np.zeros((2, 2))}) + root = FakeContainer({"browse": FakeContainer({"sample1": stack})}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + result = arrays.resolve_container("browse/sample1", "tiled") + assert result is stack # NOT descended into slice_0 + + def test_missing_path_raises_404(self, monkeypatch: pytest.MonkeyPatch): + root = FakeContainer({}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + with pytest.raises(HTTPException) as exc: + arrays.resolve_container("nope", "tiled") + assert exc.value.status_code == 404 + + def test_blank_path_segments_are_skipped(self, monkeypatch: pytest.MonkeyPatch): + leaf = FakeContainer({}) + root = FakeContainer({"a": leaf}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + result = arrays.resolve_container("/a//", "tiled") + assert result is leaf + + +# --------------------------------------------------------------------------- +# pyramid_info +# --------------------------------------------------------------------------- + +class TestPyramidInfo: + def test_non_tiled_kind_returns_none(self): + assert arrays.pyramid_info("x", "local") is None + + def test_path_not_addressing_a_scale_level_returns_none(self, monkeypatch: pytest.MonkeyPatch): + assert arrays.pyramid_info("browse/sample1", "tiled") is None + + def test_path_addressing_a_scale_level_directly(self, monkeypatch: pytest.MonkeyPatch): + finest = np.zeros((100, 512, 512)) + coarse = np.zeros((25, 128, 128)) + volume = FakeContainer({"scale0": finest, "scale1": coarse}) + root = FakeContainer({"browse": FakeContainer({"sample1": volume})}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + + info = arrays.pyramid_info("browse/sample1/scale1", "tiled") + assert info is not None + assert info["level_key"] == "scale1" + assert info["level_index"] == 1 + assert info["level_count"] == 2 + assert info["full_shape"] == [100, 512, 512] + assert info["z_downsample"] == 4.0 + + def test_path_addressing_a_scale_levels_image_child(self, monkeypatch: pytest.MonkeyPatch): + finest_arr = np.zeros((10, 20, 20)) + finest_level = FakeContainer({"image": finest_arr}) + coarse_arr = np.zeros((5, 10, 10)) + coarse_level = FakeContainer({"image": coarse_arr}) + volume = FakeContainer({"scale0": finest_level, "scale1": coarse_level}) + root = FakeContainer({"browse": FakeContainer({"sample1": volume})}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + + info = arrays.pyramid_info("browse/sample1/scale1/image", "tiled") + assert info is not None + assert info["level_key"] == "scale1" + + def test_nonexistent_parent_path_returns_none(self, monkeypatch: pytest.MonkeyPatch): + root = FakeContainer({}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + assert arrays.pyramid_info("nope/scale0", "tiled") is None + + def test_parent_without_multiscale_levels_returns_none(self, monkeypatch: pytest.MonkeyPatch): + volume = FakeContainer({"scale0": np.zeros((2, 2, 2))}) # only one level + root = FakeContainer({"sample1": volume}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + assert arrays.pyramid_info("sample1/scale0", "tiled") is None + + def test_non_3d_level_shape_returns_none(self, monkeypatch: pytest.MonkeyPatch): + volume = FakeContainer({"scale0": np.zeros((10, 10)), "scale1": np.zeros((5, 5))}) + root = FakeContainer({"sample1": volume}) + monkeypatch.setattr("arrays.get_tiled_client", lambda uri, key: root) + monkeypatch.setattr("arrays.api_key_for_uri", lambda uri: None) + assert arrays.pyramid_info("sample1/scale1", "tiled") is None diff --git a/backend/tests/test_batch_probe.py b/backend/tests/test_batch_probe.py new file mode 100644 index 0000000..8011c82 --- /dev/null +++ b/backend/tests/test_batch_probe.py @@ -0,0 +1,271 @@ +"""Tests for batch_probe.py — pure helpers directly, plus real (no mocks) +CPU forward+backward+optimizer-step round trips via train_common.build_family +(same tiny-model pattern as test_train_e2e_real_ml.py), and the job-level +guard/OOM/cancellation paths via a monkeypatched _try_batch.""" +from __future__ import annotations + +import concurrent.futures + +import pytest + +torch = pytest.importorskip("torch") + +import batch_probe # noqa: E402 +import export_jobs # noqa: E402 +import train_common # noqa: E402 +from schemas import BatchProbeRequest, DlsiaTunetConfig # noqa: E402 + + +def _request(image_size=64, batch_size=2, depth=2, base_channels=4): + return BatchProbeRequest( + model=DlsiaTunetConfig( + hyperparams={ + "depth": depth, "base_channels": base_channels, + "image_size": image_size, "batch_size": batch_size, + }, + ), + n_classes=2, + ) + + +# --------------------------------------------------------------------------- +# Pure helpers +# --------------------------------------------------------------------------- + +class TestSchemaBatchCap: + def test_returns_the_schema_upper_bound_when_lower_than_default(self): + hp = _request().model.hyperparams + cap = batch_probe.schema_batch_cap(hp, default=1000) + assert cap <= 1000 + + def test_default_used_when_no_le_constraint_found(self): + class NoConstraints: + model_fields = {} + + assert batch_probe.schema_batch_cap(NoConstraints(), default=42) == 42 + + def test_broken_metadata_falls_back_to_default(self): + class Broken: + pass + + assert batch_probe.schema_batch_cap(Broken(), default=7) == 7 + + +class TestAttemptCount: + def test_cap_one_is_one_attempt(self): + assert batch_probe._attempt_count(1) == 1 + + def test_higher_caps_scale_with_bit_length(self): + assert batch_probe._attempt_count(64) == 64 .bit_length() + + def test_never_below_one(self): + assert batch_probe._attempt_count(0) >= 1 + + +class TestSuggest: + def test_applies_safety_factor(self): + assert batch_probe._suggest(10, cap=64) == int(10 * batch_probe._SAFETY_FACTOR) + + def test_never_below_one(self): + assert batch_probe._suggest(1, cap=64) >= 1 + + def test_never_exceeds_cap(self): + assert batch_probe._suggest(1000, cap=8) <= 8 + + +class TestIsOom: + @pytest.mark.parametrize("message", [ + "CUDA out of memory. Tried to allocate...", + "MPS backend out of memory", + "Insufficient Memory!", + "can't allocate memory", + "Cannot allocate 4GB", + ]) + def test_recognizes_oom_messages_case_insensitively(self, message): + assert batch_probe._is_oom(RuntimeError(message)) is True + + def test_unrelated_error_is_not_oom(self): + assert batch_probe._is_oom(RuntimeError("shape mismatch")) is False + + +class TestDeviceMemoryGib: + def test_cpu_reports_nothing(self): + assert batch_probe._device_memory_gib("cpu") is None + + def test_unknown_device_reports_nothing(self): + assert batch_probe._device_memory_gib("weird") is None + + +class TestRelease: + def test_cpu_does_not_raise(self): + batch_probe._release("cpu") # gc.collect() only; no torch cache to clear + + +# --------------------------------------------------------------------------- +# _try_batch — real forward + backward + optimizer step on CPU +# --------------------------------------------------------------------------- + +class TestTryBatchReal: + def test_completes_a_real_step_without_raising(self): + req = _request() + built = train_common.build_family(req.model, req.n_classes, "cpu", lambda msg: None) + batch_probe._try_batch( + 2, image_size=16, forward_fn=built.forward_fn, trainable=built.trainable_params, device="cpu", + ) + + def test_updates_model_parameters(self): + req = _request() + built = train_common.build_family(req.model, req.n_classes, "cpu", lambda msg: None) + before = [p.clone() for p in built.trainable_params] + batch_probe._try_batch( + 2, image_size=16, forward_fn=built.forward_fn, trainable=built.trainable_params, device="cpu", + ) + after = built.trainable_params + assert any(not torch.equal(b, a) for b, a in zip(before, after)) + + +# --------------------------------------------------------------------------- +# _run_attempt_cancellable +# --------------------------------------------------------------------------- + +class TestRunAttemptCancellable: + def test_returns_ok_on_success(self): + # A real trainable module, not an identity lambda: `_try_batch` calls + # `loss.backward()`, which needs a real grad_fn in the graph. + conv = torch.nn.Conv2d(3, 3, kernel_size=1) + jid = export_jobs.new_job("x") + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + outcome = batch_probe._run_attempt_cancellable( + executor, jid, 1, image_size=8, forward_fn=conv, trainable=list(conv.parameters()), device="cpu", + ) + assert outcome == "ok" + + def test_returns_oom_when_try_batch_raises_oom_message(self, monkeypatch): + def boom(*a, **k): + raise RuntimeError("CUDA out of memory") + + monkeypatch.setattr(batch_probe, "_try_batch", boom) + jid = export_jobs.new_job("x") + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + outcome = batch_probe._run_attempt_cancellable( + executor, jid, 1, image_size=8, forward_fn=lambda x: x, trainable=[], device="cpu", + ) + assert outcome == "oom" + + def test_reraises_non_oom_exceptions(self, monkeypatch): + def boom(*a, **k): + raise RuntimeError("something else broke") + + monkeypatch.setattr(batch_probe, "_try_batch", boom) + jid = export_jobs.new_job("x") + with pytest.raises(RuntimeError, match="something else broke"): + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + batch_probe._run_attempt_cancellable( + executor, jid, 1, image_size=8, forward_fn=lambda x: x, trainable=[], device="cpu", + ) + + def test_polls_for_cancellation_while_attempt_runs(self, monkeypatch): + import time + + def slow(*a, **k): + time.sleep(1.5) + + monkeypatch.setattr(batch_probe, "_try_batch", slow) + monkeypatch.setattr(batch_probe, "_ATTEMPT_POLL_SECONDS", 0.1) + jid = export_jobs.new_job("x") + export_jobs.request_cancel(jid) + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + outcome = batch_probe._run_attempt_cancellable( + executor, jid, 1, image_size=8, forward_fn=lambda x: x, trainable=[], device="cpu", + ) + assert outcome == "cancelled" + + +# --------------------------------------------------------------------------- +# run_probe_job — guard paths and outcome-shaping logic +# --------------------------------------------------------------------------- + +class TestRunProbeJobGuards: + def test_busy_ml_lock_reports_error_not_a_crash(self): + train_common.ML_LOCK.acquire() + try: + jid = export_jobs.new_job("x") + batch_probe.run_probe_job(jid, _request()) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "already running" in job["error"] + finally: + train_common.ML_LOCK.release() + + def test_no_device_reports_error(self, monkeypatch): + monkeypatch.setattr(train_common, "pick_device", lambda: None) + jid = export_jobs.new_job("x") + batch_probe.run_probe_job(jid, _request()) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "torch is not installed" in job["error"] + assert train_common.ML_LOCK.locked() is False # released even on early failure + + +class TestRunProbeJobOutcomes: + def test_real_probe_succeeds_and_releases_lock(self): + jid = export_jobs.new_job("x") + batch_probe.run_probe_job(jid, _request()) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["suggested_batch_size"] >= 1 + assert job["result"]["largest_ok"] >= 1 + assert job["result"]["cancelled"] is False + assert train_common.ML_LOCK.locked() is False + + def test_oom_on_first_attempt_reports_no_suggestion(self, monkeypatch): + def always_oom(*a, **k): + raise RuntimeError("out of memory") + + monkeypatch.setattr(batch_probe, "_try_batch", always_oom) + jid = export_jobs.new_job("x") + batch_probe.run_probe_job(jid, _request()) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "Even batch size 1" in job["error"] + + def test_oom_after_some_success_suggests_a_smaller_batch(self, monkeypatch): + calls = [] + + def sometimes_oom(batch_size, **k): + calls.append(batch_size) + if batch_size > 2: + raise RuntimeError("out of memory") + + monkeypatch.setattr(batch_probe, "_try_batch", sometimes_oom) + jid = export_jobs.new_job("x") + batch_probe.run_probe_job(jid, _request(batch_size=32)) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["largest_ok"] == 2 + assert job["result"]["first_failure"] == 4 + assert job["result"]["suggested_batch_size"] <= 2 + + def test_cancelled_before_any_success_reports_done_not_error(self, monkeypatch): + monkeypatch.setattr(export_jobs, "cancel_requested", lambda jid: True) + jid = export_jobs.new_job("x") + batch_probe.run_probe_job(jid, _request()) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["cancelled"] is True + assert job["result"]["suggested_batch_size"] is None + + def test_cancelled_after_a_success_suggests_from_partial_progress(self, monkeypatch): + calls = {"n": 0} + + def cancel_after_first(jid): + calls["n"] += 1 + return calls["n"] > 1 + + monkeypatch.setattr(export_jobs, "cancel_requested", cancel_after_first) + jid = export_jobs.new_job("x") + batch_probe.run_probe_job(jid, _request(batch_size=32)) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["cancelled"] is True + assert job["result"]["largest_ok"] >= 1 diff --git a/backend/tests/test_browse_helpers.py b/backend/tests/test_browse_helpers.py new file mode 100644 index 0000000..1fa40c1 --- /dev/null +++ b/backend/tests/test_browse_helpers.py @@ -0,0 +1,370 @@ +"""Tests for browse_helpers.py using lightweight fake Tiled nodes — duck-typed +to the subset of the Tiled client API this module actually uses +(__iter__/__getitem__/.metadata/.structure_family/.search()/.distinct()), so +no real Tiled server is needed. search() is driven by REAL tiled.queries.Eq/ +Contains objects (what Key(...) == val / Contains(...) actually produce), +not a further mock of the query layer itself. +""" +from __future__ import annotations + +from collections import Counter + +import pytest + +import browse_helpers as bh + +tiled_queries = pytest.importorskip("tiled.queries") + + +class FakeNode: + def __init__(self, children=None, metadata=None, is_container=True, search_raises=False): + self._children = children or {} + self.metadata = metadata or {} + self.structure_family = "container" if is_container else "array" + self._search_raises = search_raises + + def __iter__(self): + return iter(self._children) + + def __getitem__(self, key): + return self._children[key] + + def __len__(self): + return len(self._children) + + def search(self, query): + if self._search_raises: + raise RuntimeError("search exploded") + matched = {} + for k, child in self._children.items(): + meta = child.metadata or {} + if isinstance(query, tiled_queries.Contains): + val = meta.get(query.key) + if isinstance(val, (list, tuple)) and query.value in val: + matched[k] = child + else: # Eq + if meta.get(query.key) == query.value: + matched[k] = child + return FakeNode(children=matched, metadata=self.metadata) + + def distinct(self, key, counts=True): + counter: Counter = Counter() + for child in self._children.values(): + val = (child.metadata or {}).get(key) + if val is None: + continue + counter[val] += 1 + return {"metadata": {key: [{"value": v, "count": n} for v, n in counter.items()]}} + + +# --------------------------------------------------------------------------- +# Field discovery +# --------------------------------------------------------------------------- + +class TestDisplayName: + def test_strips_thinfilm_prefix(self): + assert bh._display_name("thinfilm_AnnealingTemp") == "Temp" + + def test_thinfilm_prefix_without_alias_returns_stripped(self): + assert bh._display_name("thinfilm_Custom") == "Custom" + + def test_studio_key_alias(self): + assert bh._display_name("studio_annotated") == "Annotated" + + def test_unknown_key_passthrough(self): + assert bh._display_name("foo") == "foo" + + +class TestIsRichEnough: + def test_thinfilm_key_alone_is_enough(self): + assert bh._is_rich_enough(["thinfilm_x"]) is True + + def test_important_keys_with_min_count(self): + keys = ["PI", "sample_name", "a", "b"] + assert bh._is_rich_enough(keys) is True + + def test_important_key_without_min_count_not_enough(self): + assert bh._is_rich_enough(["PI"]) is False + + def test_plain_key_richness_threshold(self): + assert bh._is_rich_enough([f"k{i}" for i in range(10)]) is True + assert bh._is_rich_enough([f"k{i}" for i in range(9)]) is False + + +class TestBuildFieldMapping: + def test_finds_first_rich_enough_sample(self): + sparse = FakeNode(metadata={"a": 1}) + rich = FakeNode(metadata={f"k{i}": i for i in range(10)}) + node = FakeNode(children={"s0": sparse, "s1": rich}) + mapping = bh.build_field_mapping(node) + assert "k0" in mapping.raw_to_display + + def test_falls_back_to_last_scanned_if_none_rich_enough(self): + node = FakeNode(children={"s0": FakeNode(metadata={"a": 1})}) + mapping = bh.build_field_mapping(node) + assert "a" in mapping.raw_to_display + + def test_skips_samples_that_error_on_open(self): + class ExplodingNode: + @property + def metadata(self): + raise RuntimeError("boom") + + node = FakeNode(children={"bad": ExplodingNode(), "good": FakeNode(metadata={f"k{i}": i for i in range(10)})}) + mapping = bh.build_field_mapping(node) + assert "k0" in mapping.raw_to_display + + def test_always_injects_studio_and_ingest_keys(self): + node = FakeNode(children={"s0": FakeNode(metadata={})}) + mapping = bh.build_field_mapping(node) + assert "Annotated" in mapping.all_display_keys + assert "studio_annotated" in mapping.raw_to_display + + def test_empty_container_does_not_raise(self): + mapping = bh.build_field_mapping(FakeNode(children={})) + assert "studio_annotated" in mapping.raw_to_display + + +# --------------------------------------------------------------------------- +# Distinct values +# --------------------------------------------------------------------------- + +class TestScopedMetadataRows: + def test_reads_each_childs_metadata(self): + node = FakeNode(children={"a": FakeNode(metadata={"x": 1}), "b": FakeNode(metadata={"x": 2})}) + rows = bh.scoped_metadata_rows(node) + assert rows == [{"x": 1}, {"x": 2}] + + def test_respects_limit(self): + node = FakeNode(children={str(i): FakeNode(metadata={"x": i}) for i in range(10)}) + rows = bh.scoped_metadata_rows(node, limit=3) + assert len(rows) == 3 + + def test_skips_children_that_fail_to_open(self): + class ExplodingNode: + @property + def metadata(self): + raise RuntimeError("boom") + + node = FakeNode(children={"bad": ExplodingNode(), "good": FakeNode(metadata={"x": 1})}) + assert bh.scoped_metadata_rows(node) == [{"x": 1}] + + +class TestDistinctFromRows: + def test_tallies_scalar_values(self): + rows = [{"k": "a"}, {"k": "a"}, {"k": "b"}] + result = bh.distinct_from_rows(rows, "k") + by_value = {e["value"]: e["count"] for e in result} + assert by_value == {"a": 2, "b": 1} + + def test_explodes_list_valued_metadata(self): + rows = [{"k": ["x", "y"]}, {"k": ["x"]}] + result = bh.distinct_from_rows(rows, "k") + by_value = {e["value"]: e["count"] for e in result} + assert by_value == {"x": 2, "y": 1} + + def test_missing_key_ignored(self): + assert bh.distinct_from_rows([{"other": 1}], "k") == [] + + +class TestTiledDistinctValues: + def _node(self): + return FakeNode(children={ + "a": FakeNode(metadata={"studio_annotated": "yes"}), + "b": FakeNode(metadata={"studio_annotated": "yes"}), + "c": FakeNode(metadata={"studio_annotated": "no"}), + }) + + def test_global_distinct_uses_node_distinct(self): + result = bh.tiled_distinct_values(self._node(), "studio_annotated") + by_value = {e["value"]: e["count"] for e in result["values"]} + assert by_value == {"yes": 2, "no": 1} + assert result["total"] == 2 + + def test_scoped_distinct_uses_iteration(self): + result = bh.tiled_distinct_values(self._node(), "studio_annotated", scoped=True) + by_value = {e["value"]: e["count"] for e in result["values"]} + assert by_value == {"yes": 2, "no": 1} + + def test_invalid_values_filtered_out(self): + node = FakeNode(children={ + "a": FakeNode(metadata={"k": "real"}), + "b": FakeNode(metadata={"k": "NaN"}), + }) + result = bh.tiled_distinct_values(node, "k") + assert result["values"] == [{"value": "real", "count": 1, "sample_paths": []}] + + def test_field_mapping_translates_display_key_to_raw(self): + mapping = bh.FieldMapping( + display_to_raw={}, raw_to_display={"studio_annotated": "Annotated"}, all_display_keys=[], + ) + result = bh.tiled_distinct_values(self._node(), "studio_annotated", field_mapping=mapping) + assert result["field"] == "Annotated" + + def test_filter_narrows_node_before_distinct(self): + mapping = bh.FieldMapping( + display_to_raw={"Annotated": "studio_annotated"}, raw_to_display={}, all_display_keys=[], + ) + result = bh.tiled_distinct_values( + self._node(), "studio_annotated", filters={"Annotated": "yes"}, field_mapping=mapping, + ) + assert result["total"] == 1 # only "yes" remains after filtering to studio_annotated == "yes" + + def test_distinct_call_failure_returns_empty(self): + node = FakeNode(children={}) + node.distinct = lambda *a, **k: (_ for _ in ()).throw(RuntimeError("boom")) + result = bh.tiled_distinct_values(node, "k") + assert result == {"values": [], "total": 0, "field": "k"} + + +# --------------------------------------------------------------------------- +# Item search +# --------------------------------------------------------------------------- + +class TestTiledSearchItemsDispatch: + def test_no_array_only_filters_uses_container_only_path(self): + node = FakeNode(children={"a": FakeNode(metadata={"PI": "smith"}, is_container=False)}) + result = bh.tiled_search_items(node, filters={"PI": "smith"}) + assert result["total"] == 1 + assert result["items"][0]["sample"] == "a" + + def test_only_array_only_filters_uses_array_only_path(self): + # bar's metadata is a real int here, matching how Tiled actually stores + # this field — _apply_filters coerces the filter string "1" to int 1 + # via _typed_query_value, so a string-valued fixture would (correctly) + # never match and silently fail the test for the wrong reason. + node = FakeNode(children={"a": FakeNode(metadata={"bar": 1}, is_container=False)}) + result = bh.tiled_search_items(node, filters={"bar": "1"}) + assert result["total"] == 1 + + def test_mixed_filters_uses_mixed_path(self): + array_child = FakeNode(metadata={"bar": 1}, is_container=False) + sample = FakeNode(children={"arr": array_child}, metadata={"PI": "smith"}) + node = FakeNode(children={"s0": sample}) + result = bh.tiled_search_items(node, filters={"PI": "smith", "bar": "1"}) + assert result["total"] == 1 + assert result["items"][0]["sample"] == "s0" + + def test_no_filters_returns_everything(self): + node = FakeNode(children={"a": FakeNode(metadata={}), "b": FakeNode(metadata={})}) + result = bh.tiled_search_items(node) + assert result["total"] == 2 + + +class TestSearchContainerOnly: + def test_reports_n_slices_for_flat_array_stack(self): + stack = FakeNode(children={ + str(i): FakeNode(metadata={}, is_container=False) for i in range(5) + }, metadata={}) + node = FakeNode(children={"sample1": stack}) + result = bh._search_container_only(node, [], "", 500) + assert result["items"][0]["n_slices"] == 5 + + def test_borrows_first_childs_metadata_when_container_has_none(self): + stack = FakeNode(children={"0": FakeNode(metadata={"x": 1}, is_container=False)}, metadata={}) + node = FakeNode(children={"sample1": stack}) + result = bh._search_container_only(node, [], "", 500) + assert result["items"][0]["metadata"] == {"x": 1} + + def test_leaf_array_reports_n_slices_one(self): + node = FakeNode(children={"a": FakeNode(metadata={}, is_container=False)}) + result = bh._search_container_only(node, [], "", 500) + assert result["items"][0]["n_slices"] == 1 + + def test_path_prefix_applied(self): + node = FakeNode(children={"a": FakeNode(metadata={}, is_container=False)}) + result = bh._search_container_only(node, [], "parent", 500) + assert result["items"][0]["path"] == "parent/a" + + def test_children_that_fail_to_open_are_skipped(self): + class Exploding: + metadata = property(lambda self: (_ for _ in ()).throw(RuntimeError())) + + node = FakeNode(children={"bad": Exploding(), "good": FakeNode(metadata={}, is_container=False)}) + result = bh._search_container_only(node, [], "", 500) + assert result["total"] == 1 + + +# --------------------------------------------------------------------------- +# Small helpers +# --------------------------------------------------------------------------- + +class TestRawFilters: + def test_translates_display_to_raw_and_stringifies(self): + out = bh._raw_filters({"Temp": 5}, {"Temp": "thinfilm_AnnealingTemp"}) + assert out == [("thinfilm_AnnealingTemp", "5")] + + def test_none_values_dropped(self): + assert bh._raw_filters({"Temp": None}, {}) == [] + + def test_unmapped_display_key_passthrough(self): + assert bh._raw_filters({"raw_key": "v"}, {}) == [("raw_key", "v")] + + +class TestMatchesContainerFilters: + def test_scalar_case_insensitive_match(self): + assert bh._matches_container_filters({"PI": "Smith"}, [("PI", "smith")]) is True + + def test_scalar_mismatch(self): + assert bh._matches_container_filters({"PI": "Jones"}, [("PI", "smith")]) is False + + def test_missing_key_fails(self): + assert bh._matches_container_filters({}, [("PI", "smith")]) is False + + def test_list_valued_membership_match(self): + assert bh._matches_container_filters({"keywords": ["Alpha", "Beta"]}, [("keywords", "alpha")]) is True + + def test_list_valued_membership_mismatch(self): + assert bh._matches_container_filters({"keywords": ["Beta"]}, [("keywords", "alpha")]) is False + + +class TestTypedQueryValue: + def test_int_roundtrip(self): + assert bh._typed_query_value("5") == 5 + + def test_float_roundtrip(self): + assert bh._typed_query_value("5.5") == 5.5 + + def test_non_numeric_passthrough(self): + assert bh._typed_query_value("abc") == "abc" + + def test_leading_zero_not_coerced_to_int(self): + # "007" != str(int("007")) == "7", so it must stay a string. + assert bh._typed_query_value("007") == "007" + + +class TestIsValidValue: + @pytest.mark.parametrize("value", [None, "", "None", "NaN", "nan", " "]) + def test_invalid_values(self, value): + assert bh._is_valid_value(value) is False + + def test_valid_value(self): + assert bh._is_valid_value("real") is True + + +class TestJoin: + def test_with_prefix(self): + assert bh._join("a", "b") == "a/b" + + def test_without_prefix(self): + assert bh._join("", "b") == "b" + + +class TestApplyFilters: + def test_eq_filter_narrows_node(self): + node = FakeNode(children={"a": FakeNode(metadata={"k": "v"}), "b": FakeNode(metadata={"k": "other"})}) + result = bh._apply_filters(node, [("k", "v")]) + assert list(result) == ["a"] + + def test_contains_filter_for_list_valued_key(self): + node = FakeNode(children={ + "a": FakeNode(metadata={"keywords": ["x", "y"]}), + "b": FakeNode(metadata={"keywords": ["z"]}), + }) + result = bh._apply_filters(node, [("keywords", "x")]) + assert list(result) == ["a"] + + def test_search_failure_is_swallowed_and_node_unchanged(self): + node = FakeNode(children={"a": FakeNode(metadata={})}, search_raises=True) + result = bh._apply_filters(node, [("k", "v")]) + assert result is node diff --git a/backend/tests/test_coco_export.py b/backend/tests/test_coco_export.py index 72c4c05..0d6b225 100644 --- a/backend/tests/test_coco_export.py +++ b/backend/tests/test_coco_export.py @@ -235,3 +235,45 @@ def test_polygon_hole_carved_out() -> None: assert mask[2, 2], "frame corner must be filled" # Area ~= 40*40 - 20*20 = 1200, allow boundary slack. assert 1050 < int(mask.sum()) < 1350 + + +def test_resolve_split_single_slice_goes_to_train() -> None: + """A lone annotated slice at the default 80/10/10 ratio must not be + floored out of every split — that leaves nothing to train on.""" + from coco_export import _resolve_split + + result = _resolve_split(["0"], {}, {"ratios": [0.8, 0.1, 0.1], "seed": 1234}) + assert result == {"0": "train"} + + +def test_resolve_split_two_slices_still_gets_a_train_slice() -> None: + from coco_export import _resolve_split + + result = _resolve_split(["0", "1"], {}, {"ratios": [0.8, 0.1, 0.1], "seed": 1234}) + assert "train" in result.values() + + +def test_resolve_split_borrows_from_valid_before_test() -> None: + """When train floors to zero but valid has slices to spare, borrow from + valid rather than test, so the split still roughly tracks the ratios.""" + from coco_export import _resolve_split + + result = _resolve_split( + [str(i) for i in range(3)], {}, {"ratios": [0.1, 0.8, 0.1], "seed": 1234} + ) + counts = {"train": 0, "valid": 0, "test": 0} + for v in result.values(): + counts[v] += 1 + assert counts["train"] == 1 + assert counts["valid"] == 1 + assert counts["test"] == 1 + + +def test_resolve_split_respects_explicit_assignments() -> None: + """Slices already given an explicit (non-'auto') split are left alone.""" + from coco_export import _resolve_split + + result = _resolve_split( + ["0", "1"], {"0": "valid"}, {"ratios": [0.8, 0.1, 0.1], "seed": 1234} + ) + assert result["0"] == "valid" diff --git a/backend/tests/test_denoise.py b/backend/tests/test_denoise.py new file mode 100644 index 0000000..45eb058 --- /dev/null +++ b/backend/tests/test_denoise.py @@ -0,0 +1,189 @@ +"""Tests for the classical denoise filters. + +Ported alongside :mod:`denoise` itself. The properties pinned here are the ones +a user would notice if they broke: that a filter actually reduces noise without +destroying the edges this app exists to annotate, that shape and dtype survive +the round trip through normalized units, and that "Auto" never hands back a +setting that makes the image worse. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import numpy as np +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import denoise # noqa: E402 + + +def step_phantom(h=192, w=192, noise=12.0, seed=3): + """A hard vertical edge plus additive Gaussian noise. + + The standard shape for asking "did it denoise, or did it just blur?" — the + edge is what a bad filter destroys. + """ + rng = np.random.default_rng(seed) + clean = np.zeros((h, w), np.float32) + clean[:, w // 2:] = 100.0 + return clean, clean + rng.normal(0, noise, clean.shape).astype(np.float32) + + +def snr_db(clean: np.ndarray, test: np.ndarray) -> float: + return float(10 * np.log10(clean.var() / ((test - clean) ** 2).mean())) + + +class TestDenoiseSlice: + @pytest.mark.parametrize("method", ["gaussian", "median", "bilateral", "tv", "nlm"]) + def test_improves_snr_on_a_noisy_edge(self, method): + clean, noisy = step_phantom() + out = denoise.denoise_slice(noisy, method, 0.5).astype(np.float32) + assert snr_db(clean, out) > snr_db(clean, noisy) + + @pytest.mark.parametrize("method", ["gaussian", "median", "tv"]) + @pytest.mark.parametrize("dtype", [np.uint8, np.uint16, np.float32]) + def test_preserves_shape_and_dtype(self, method, dtype): + # Filtering happens in normalized [0,1] floats; the round trip back to + # the source dtype is where clipping and rounding bugs would live. + rng = np.random.default_rng(1) + arr = (rng.random((48, 64)) * 200).astype(dtype) + out = denoise.denoise_slice(arr, method, 0.5) + assert out.shape == arr.shape + assert out.dtype == arr.dtype + + def test_none_is_an_identity(self): + _, noisy = step_phantom(h=32, w=32) + assert denoise.denoise_slice(noisy, "none", 1.0) is noisy + + def test_flat_slice_is_returned_untouched(self): + # Zero dynamic range: normalizing would divide by zero. + flat = np.full((32, 32), 7, np.uint16) + assert np.array_equal(denoise.denoise_slice(flat, "tv", 0.5), flat) + + def test_stronger_settings_smooth_more(self): + _, noisy = step_phantom() + weak = denoise.denoise_slice(noisy, "gaussian", 0.1).astype(np.float32) + strong = denoise.denoise_slice(noisy, "gaussian", 0.9).astype(np.float32) + assert strong.std() < weak.std() + + def test_rejects_a_3d_method(self): + _, noisy = step_phantom(h=16, w=16) + with pytest.raises(ValueError, match="denoise_stack"): + denoise.denoise_slice(noisy, "gaussian3d", 0.5) + + def test_rejects_an_unknown_method(self): + _, noisy = step_phantom(h=16, w=16) + with pytest.raises(ValueError, match="unknown"): + denoise.denoise_slice(noisy, "bogus", 0.5) + + def test_rejects_a_3d_array(self): + with pytest.raises(ValueError, match="2-D"): + denoise.denoise_slice(np.zeros((3, 8, 8), np.float32), "tv", 0.5) + + +class TestDenoiseStack: + def test_uses_z_neighbours_to_beat_the_2d_equivalent(self): + # The whole argument for the 3-D filters: adjacent slices share + # structure while their noise is independent, so averaging along z buys + # noise reduction at a far lower cost in real detail than in-plane blur. + rng = np.random.default_rng(5) + clean2d = np.zeros((96, 96), np.float32) + clean2d[:, 48:] = 100.0 + clean = np.repeat(clean2d[None], 5, axis=0) + noisy = clean + rng.normal(0, 15, clean.shape).astype(np.float32) + + out3d = denoise.denoise_stack(noisy, "gaussian3d", 0.5).astype(np.float32)[2] + out2d = denoise.denoise_slice(noisy[2], "gaussian", 0.5).astype(np.float32) + assert snr_db(clean2d, out3d) > snr_db(clean2d, out2d) + + @pytest.mark.parametrize("method", ["gaussian3d", "median3d"]) + def test_preserves_shape_and_dtype(self, method): + rng = np.random.default_rng(2) + stack = (rng.random((5, 32, 32)) * 500).astype(np.uint16) + out = denoise.denoise_stack(stack, method, 0.5) + assert out.shape == stack.shape + assert out.dtype == stack.dtype + + def test_rejects_a_2d_method(self): + with pytest.raises(ValueError, match="not a 3-D"): + denoise.denoise_stack(np.zeros((3, 8, 8), np.float32), "tv", 0.5) + + def test_rejects_a_2d_array(self): + with pytest.raises(ValueError, match=r"\(z, y, x\)"): + denoise.denoise_stack(np.zeros((8, 8), np.float32), "gaussian3d", 0.5) + + +class TestNoiseEstimate: + def test_tracks_the_true_sigma(self): + # Immerkaer's estimator, on a flat field where the answer is known. + rng = np.random.default_rng(9) + for sigma in (0.01, 0.05, 0.1): + field = rng.normal(0.5, sigma, (256, 256)) + estimate = denoise.estimate_noise_sigma(field) + assert abs(estimate - sigma) / sigma < 0.1 + + def test_is_not_fooled_by_structure(self): + # A clean step edge is structure, not noise; the estimate must stay low + # or "Auto" would recommend smoothing a noiseless image. + clean, _ = step_phantom(noise=0.0) + unit, _, span = denoise._to_unit(clean) + assert denoise.estimate_noise_sigma(unit) < 0.02 + + def test_degenerate_input_returns_zero(self): + assert denoise.estimate_noise_sigma(np.zeros((2, 2))) == 0.0 + + +class TestAutoStrength: + def test_suggests_more_for_a_noisier_image(self): + _, quiet = step_phantom(noise=2.0) + _, loud = step_phantom(noise=30.0) + assert denoise.auto_strength(loud, "tv") > denoise.auto_strength(quiet, "tv") + + def test_never_makes_a_clean_image_worse(self): + # The failure this guards: "Auto" on a clean slice returning a strength + # that visibly softens it. + clean, _ = step_phantom(noise=0.0) + assert denoise.auto_strength(clean, "tv") == pytest.approx(0.0, abs=0.05) + + def test_damps_gaussian_below_the_edge_preserving_methods(self): + # Measured: at equal nominal strength a plain Gaussian can LOWER SNR + # where TV raises it, so Auto must not push it as hard. + _, noisy = step_phantom(noise=25.0) + assert denoise.auto_strength(noisy, "gaussian") < denoise.auto_strength(noisy, "tv") + + def test_none_is_zero(self): + _, noisy = step_phantom() + assert denoise.auto_strength(noisy, "none") == 0.0 + + def test_stays_in_range(self): + for noise in (0.0, 1.0, 50.0, 500.0): + _, noisy = step_phantom(noise=noise) + for method in denoise.available_methods(): + assert 0.0 <= denoise.auto_strength(noisy, method) <= 1.0 + + +class TestCapabilities: + def test_describes_every_method(self): + described = {m["method"] for m in denoise.describe_methods()} + assert described == set(denoise.ALL_METHODS) + + def test_availability_is_probed_not_assumed(self): + # denoise_wavelet imports fine without PyWavelets and only fails when + # called, so the menu has to probe rather than trust the import. + for entry in denoise.describe_methods(): + if entry["method"] == "wavelet": + assert entry["available"] == denoise.wavelet_available() + + def test_only_3d_methods_need_z_neighbours(self): + for entry in denoise.describe_methods(): + expected = entry["method"] in denoise.METHODS_3D + assert (entry["z_radius"] > 0) is expected + + def test_unavailable_methods_are_excluded_from_available_methods(self): + usable = set(denoise.available_methods()) + assert usable <= set(denoise.ALL_METHODS) + if not denoise.wavelet_available(): + assert "wavelet" not in usable diff --git a/backend/tests/test_denoise_bake.py b/backend/tests/test_denoise_bake.py new file mode 100644 index 0000000..fbdbeff --- /dev/null +++ b/backend/tests/test_denoise_bake.py @@ -0,0 +1,337 @@ +"""Tests for denoise_bake.py — real job-registry (export_jobs), fake array +source (same monkeypatch pattern as test_infer_jobs.py/test_denoise_train.py) +and a lightweight duck-typed fake Tiled container supporting compound +slash-separated __getitem__ (matching how client[target_path] is used +directly in run_denoise_bake_job).""" +from __future__ import annotations + +import numpy as np +import pytest + +torch = pytest.importorskip("torch") + +import arrays as arrays_mod # noqa: E402 +import denoise as denoise_mod # noqa: E402 +import denoise_bake # noqa: E402 +import export_jobs # noqa: E402 +import schemas # noqa: E402 +import train_common # noqa: E402 + + +class FakeTiledContainer: + def __init__(self): + self._children: dict[str, "FakeTiledContainer"] = {} + self.metadata: dict = {} + self.written: list[dict] = [] + self.updated_metadata: dict | None = None + + def __getitem__(self, key): + node = self + for part in str(key).split("/"): + node = node._children[part] + return node + + def create_container(self, key, metadata): + child = FakeTiledContainer() + child.metadata = metadata + self._children[key] = child + return child + + def write_array(self, arr, key, dims=None, metadata=None): + self.written.append({"key": key, "arr": arr, "dims": dims, "metadata": metadata}) + + def update_metadata(self, metadata): + self.updated_metadata = metadata + + +@pytest.fixture() +def fake_client(monkeypatch: pytest.MonkeyPatch): + client = FakeTiledContainer() + client.create_container("browse", {}) + monkeypatch.setattr(denoise_bake, "get_tiled_client", lambda uri: client) + return client + + +@pytest.fixture() +def fake_array_source(monkeypatch: pytest.MonkeyPatch): + """A 4-slice volume of small 2D frames, distinguishable per slice.""" + frames = [np.full((6, 6), fill_value=float(i * 10), dtype=np.float32) for i in range(4)] + + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri: "node") + monkeypatch.setattr(arrays_mod, "array_shape_meta", lambda node: {"n_slices": len(frames)}) + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: frames[idx]) + return frames + + +def _request(**overrides): + defaults = dict( + source="browse/sample", server_uri=None, method="median", + strength=0.5, target_path="browse/sample_denoised", description="", + ) + defaults.update(overrides) + return schemas.DenoiseBakeRequest(**defaults) + + +# --------------------------------------------------------------------------- +# Pure helpers +# --------------------------------------------------------------------------- + +class TestDefaultTargetPath: + def test_appends_suffix_to_last_segment(self): + assert denoise_bake.default_target_path("browse/sample") == "browse/sample_denoised" + + def test_custom_suffix(self): + assert denoise_bake.default_target_path("browse/sample", suffix="clean") == "browse/sample_clean" + + def test_strips_whitespace_and_slashes(self): + assert denoise_bake.default_target_path(" /browse/sample/ ") == "browse/sample_denoised" + + def test_empty_path_raises(self): + with pytest.raises(ValueError): + denoise_bake.default_target_path(" /// ") + + +class TestSliceKey: + def test_zero_padded(self): + assert denoise_bake._slice_key(7) == "slice_0007" + + def test_large_index(self): + assert denoise_bake._slice_key(12345) == "slice_12345" + + +# --------------------------------------------------------------------------- +# run_denoise_bake_job — guard/validation paths +# --------------------------------------------------------------------------- + +class TestGuardPaths: + def test_method_none_reports_error(self): + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="none")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "method" in job["error"].lower() + + def test_model_method_without_run_id_reports_error(self): + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="model", run_id=None)) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "run_id" in job["error"] + + def test_unavailable_method_reports_error(self): + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="not_a_real_method")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "unavailable" in job["error"] + + def test_invalid_target_path_reports_error(self, fake_client): + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(target_path="../escape")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + + def test_existing_target_path_reports_error(self, fake_client): + fake_client["browse"].create_container("sample_denoised", {}) + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request()) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "already exists" in job["error"] + + def test_unexpected_exception_is_reported_not_raised(self, fake_client, monkeypatch): + monkeypatch.setattr( + arrays_mod, "resolve_array", + lambda source, kind, server_uri: (_ for _ in ()).throw(RuntimeError("boom")), + ) + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request()) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "boom" in job["error"] + + +# --------------------------------------------------------------------------- +# run_denoise_bake_job — real classical-filter round trips +# --------------------------------------------------------------------------- + +class TestClassicalBakeRoundTrip: + def test_2d_method_writes_every_slice(self, fake_client, fake_array_source): + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="median")) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["n_slices"] == 4 + target = fake_client["browse"]["sample_denoised"] + assert len(target.written) == 4 + assert target.updated_metadata["n_images"] == 4 + assert target.updated_metadata["denoise_method"] == "median" + + def test_3d_method_uses_z_window_and_stacks(self, fake_client, fake_array_source, monkeypatch): + calls = [] + real_stack = denoise_mod.denoise_stack + + def spy_stack(stack, method, strength): + calls.append(stack.shape[0]) + return real_stack(stack, method, strength) + + monkeypatch.setattr(denoise_mod, "denoise_stack", spy_stack) + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="median3d")) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + # radius=1: middle slices see a 3-frame window, edges see 2. + assert calls == [2, 3, 3, 2] + + def test_per_slice_failure_falls_back_to_unfiltered_copy(self, fake_client, fake_array_source, monkeypatch): + real_denoise_slice = denoise_mod.denoise_slice + + def flaky_denoise_slice(arr, method, strength): + if arr[0, 0] == 10.0: # slice index 1 + raise RuntimeError("filter blew up") + return real_denoise_slice(arr, method, strength) + + monkeypatch.setattr(denoise_mod, "denoise_slice", flaky_denoise_slice) + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="median")) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["n_slices"] == 4 + assert len(job["result"]["errors"]) == 1 + assert job["result"]["errors"][0]["slice"] == 1 + target = fake_client["browse"]["sample_denoised"] + # The failed slice was still written (as an unfiltered copy of source). + written_slice1 = next(w for w in target.written if w["key"] == "slice_0001") + assert np.array_equal(written_slice1["arr"], fake_array_source[1]) + + def test_totally_unreadable_slice_is_skipped_not_written(self, fake_client, fake_array_source, monkeypatch): + monkeypatch.setattr(denoise_mod, "denoise_slice", lambda arr, m, s: (_ for _ in ()).throw(RuntimeError("f"))) + + def flaky_read_slice(node, meta, idx): + if idx == 2: + raise RuntimeError("source also unreadable") + return fake_array_source[idx] + + monkeypatch.setattr(arrays_mod, "read_slice", flaky_read_slice) + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="median")) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["n_slices"] == 3 # slice 2 was skipped entirely + slice2_error = next(e for e in job["result"]["errors"] if e["slice"] == 2) + assert "also unreadable" in slice2_error["error"] + + def test_cancellation_mid_job_stops_early(self, fake_client, fake_array_source, monkeypatch): + monkeypatch.setattr(export_jobs, "cancel_requested", lambda jid: True) + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="median")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "nothing was written" in job["error"] + + def test_all_slices_failing_reports_error(self, fake_client, fake_array_source, monkeypatch): + monkeypatch.setattr(denoise_mod, "denoise_slice", lambda arr, m, s: (_ for _ in ()).throw(RuntimeError("f"))) + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: (_ for _ in ()).throw(RuntimeError("g"))) + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="median")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "nothing was written" in job["error"] + + def test_description_and_keywords_propagate_to_metadata(self, fake_client, fake_array_source): + jid = export_jobs.new_job("x") + denoise_bake.run_denoise_bake_job(jid, _request(method="median", description="alpha, beta")) + target = fake_client["browse"]["sample_denoised"] + assert target.updated_metadata["description"] == "alpha, beta" + assert set(target.updated_metadata["keywords"]) == {"alpha", "beta"} + + +# --------------------------------------------------------------------------- +# _ModelDenoiser — real forward pass via a minimal nn.Conv2d stand-in +# --------------------------------------------------------------------------- + +class _FakeDenoiserRuntime: + """Stand-in for denoise_runtime/autoencoder_runtime — generic over + forward_fn/to_tensor_fn like the real modules, using a real nn.Conv2d.""" + + @staticmethod + def load_model(state, device): + return torch.nn.Conv2d(1, 1, kernel_size=3, padding=1) + + @staticmethod + def make_forward_fn(model): + def forward_fn(batch): + return model(batch) + return forward_fn + + @staticmethod + def make_to_tensor_fn(): + def to_tensor_fn(arr): + return torch.from_numpy(np.ascontiguousarray(arr)).float().unsqueeze(0) / 255.0 + return to_tensor_fn + + +@pytest.fixture() +def fake_model_run(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr( + train_common, "load_run_config", + lambda run_id: {"model_family": "dlsia_denoiser", "task": "denoising", "image_size": 16, "render": {}}, + ) + monkeypatch.setattr(train_common, "denoiser_runtime_for", lambda config: _FakeDenoiserRuntime) + monkeypatch.setattr(train_common, "denoiser_needs_dlsia", lambda config: False) + monkeypatch.setattr(train_common, "dlsia_available", lambda: True) + monkeypatch.setattr(train_common, "pick_device", lambda: "cpu") + monkeypatch.setattr(train_common, "load_adapter_state", lambda run_id: {}) + # _ModelDenoiser.__init__ always computes global stats via arrays.read_slice + # regardless of which guard path a test is exercising. + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: np.zeros((4, 4))) + + +class TestModelDenoiser: + def test_rejects_run_that_is_not_a_denoiser(self, monkeypatch): + monkeypatch.setattr(train_common, "load_run_config", lambda run_id: {"model_family": "dlsia", "task": "segmentation"}) + with pytest.raises(ValueError, match="not a denoiser"): + denoise_bake._ModelDenoiser("r1", node=object(), meta={"n_slices": 1}) + + def test_missing_dlsia_when_needed_raises(self, fake_model_run, monkeypatch): + monkeypatch.setattr(train_common, "denoiser_needs_dlsia", lambda config: True) + monkeypatch.setattr(train_common, "dlsia_available", lambda: False) + with pytest.raises(ValueError, match="dlsia is not installed"): + denoise_bake._ModelDenoiser("r1", node=object(), meta={"n_slices": 1}) + + def test_missing_device_raises(self, fake_model_run, monkeypatch): + monkeypatch.setattr(train_common, "pick_device", lambda: None) + with pytest.raises(ValueError, match="torch is not installed"): + denoise_bake._ModelDenoiser("r1", node=object(), meta={"n_slices": 1}) + + def test_busy_lock_raises_and_does_not_deadlock(self, fake_model_run): + train_common.ML_LOCK.acquire() + try: + with pytest.raises(ValueError, match="busy"): + denoise_bake._ModelDenoiser("r1", node=object(), meta={"n_slices": 1}) + finally: + train_common.ML_LOCK.release() + + def test_real_forward_pass_and_lock_lifecycle(self, fake_model_run, monkeypatch): + gray = np.full((16, 16), 128, dtype=np.uint8) + monkeypatch.setattr("denoise_train._slice_to_gray_uint8", lambda node, meta, idx, opts, gr: gray) + monkeypatch.setattr("images._sample_global_stats", lambda node, meta: (0.0, 255.0)) + + assert train_common.ML_LOCK.locked() is False + md = denoise_bake._ModelDenoiser("r1", node=object(), meta={"n_slices": 1}) + assert train_common.ML_LOCK.locked() is True + try: + out = md.denoise(0) + assert out.shape == (16, 16) + assert out.dtype == np.uint8 + finally: + md.close() + assert train_common.ML_LOCK.locked() is False + # Idempotent: closing again must not raise or double-release. + md.close() + + def test_load_failure_releases_lock(self, fake_model_run, monkeypatch): + monkeypatch.setattr(train_common, "load_adapter_state", lambda run_id: (_ for _ in ()).throw(RuntimeError("nope"))) + with pytest.raises(RuntimeError, match="nope"): + denoise_bake._ModelDenoiser("r1", node=object(), meta={"n_slices": 1}) + assert train_common.ML_LOCK.locked() is False diff --git a/backend/tests/test_denoise_runtime.py b/backend/tests/test_denoise_runtime.py new file mode 100644 index 0000000..fa1c8be --- /dev/null +++ b/backend/tests/test_denoise_runtime.py @@ -0,0 +1,93 @@ +"""Real (no mocks) build/save/load/forward round trip for denoise_runtime.py +— mirrors test_train_e2e_real_ml.py's pattern for the sibling dlsia_runtime.py +segmentation family. +""" +from __future__ import annotations + +import numpy as np +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("dlsia") + +import denoise_runtime # noqa: E402 + +IMAGE_SIZE = 64 + + +def _model(): + return denoise_runtime.build_model( + image_size=IMAGE_SIZE, depth=2, base_channels=4, growth_rate=1.2, device="cpu", + ) + + +class TestBuildModel: + def test_forward_pass_produces_single_channel_output(self): + model = _model() + model.eval() + forward_fn = denoise_runtime.make_forward_fn(model) + to_tensor_fn = denoise_runtime.make_to_tensor_fn() + gray = np.zeros((IMAGE_SIZE, IMAGE_SIZE), dtype=np.uint8) + batch = to_tensor_fn(gray).unsqueeze(0) + with torch.no_grad(): + out = forward_fn(batch) + assert out.shape == (1, 1, IMAGE_SIZE, IMAGE_SIZE) + + def test_final_activation_is_pinned_to_none(self): + # Checked directly on the model rather than by hoping an untrained + # (effectively random) network happens to produce an out-of-[0,1] + # value — that was flaky (not guaranteed within a handful of seeds). + # final_activation=None must hold: a squashed output would mean a + # future dlsia default silently added an activation the denoiser + # regression head must never have. + model = _model() + assert model.final_activation is None + + +class TestSaveLoadRoundTrip: + def test_reloaded_model_reproduces_output(self): + model = _model() + model.eval() + forward_fn = denoise_runtime.make_forward_fn(model) + to_tensor_fn = denoise_runtime.make_to_tensor_fn() + gray = np.random.default_rng(0).integers(0, 256, size=(IMAGE_SIZE, IMAGE_SIZE), dtype=np.uint8) + batch = to_tensor_fn(gray).unsqueeze(0) + + with torch.no_grad(): + original_out = forward_fn(batch) + + state = denoise_runtime.network_dict(model) + assert "topo_dict" in state and "state_dict" in state + + reloaded = denoise_runtime.load_model(state, "cpu") + reloaded.eval() + reloaded_forward = denoise_runtime.make_forward_fn(reloaded) + with torch.no_grad(): + reloaded_out = reloaded_forward(batch) + + torch.testing.assert_close(original_out, reloaded_out) + + +class TestSetTrainMode: + def test_toggles_module_training_flag(self): + model = _model() + set_train_mode = denoise_runtime.make_set_train_mode_fn(model) + set_train_mode(False) + assert model.training is False + set_train_mode(True) + assert model.training is True + + +class TestToTensorFn: + def test_scales_uint8_to_unit_range(self): + to_tensor_fn = denoise_runtime.make_to_tensor_fn() + gray = np.full((8, 8), 255, dtype=np.uint8) + t = to_tensor_fn(gray) + assert t.shape == (1, 8, 8) + assert torch.allclose(t, torch.ones_like(t)) + + def test_accepts_single_channel_hwc(self): + to_tensor_fn = denoise_runtime.make_to_tensor_fn() + gray = np.full((8, 8, 1), 255, dtype=np.uint8) + t = to_tensor_fn(gray) + assert t.shape == (1, 8, 8) diff --git a/backend/tests/test_denoise_train.py b/backend/tests/test_denoise_train.py new file mode 100644 index 0000000..da027a6 --- /dev/null +++ b/backend/tests/test_denoise_train.py @@ -0,0 +1,292 @@ +"""Tests for denoise_train.py: pure sampling/geometry/masking logic directly, +plus a real (no mocks) run_denoise_training_loop round trip using a minimal +trainable model — the loop itself is generic over forward_fn/to_tensor_fn/ +trainable_params, so it doesn't need a real dlsia model to exercise for real. +""" +from __future__ import annotations + +import numpy as np +import pytest + +torch = pytest.importorskip("torch") + +import arrays as arrays_mod # noqa: E402 +import denoise_train # noqa: E402 + + +class FakeItem: + def __init__(self, source="fake.tif", kind="local", server_uri=None, slices=None): + self.source = source + self.kind = kind + self.server_uri = server_uri + self.slices = slices or {} + + +# --------------------------------------------------------------------------- +# selected_slice_indices +# --------------------------------------------------------------------------- + +class TestSelectedSliceIndices: + def test_empty_slices_returns_whole_range(self): + assert denoise_train.selected_slice_indices(FakeItem(slices={}), 5) == [0, 1, 2, 3, 4] + + def test_explicit_keys_sorted_and_deduped(self): + item = FakeItem(slices={"3": [], "1": [], "3": []}) # noqa: F601 + assert denoise_train.selected_slice_indices(item, 10) == [1, 3] + + def test_out_of_range_keys_dropped(self): + item = FakeItem(slices={"1": [], "99": [], "-1": []}) + assert denoise_train.selected_slice_indices(item, 5) == [1] + + def test_non_integer_keys_ignored(self): + item = FakeItem(slices={"abc": [], "2": []}) + assert denoise_train.selected_slice_indices(item, 5) == [2] + + def test_all_keys_out_of_range_falls_back_to_whole_range(self): + item = FakeItem(slices={"99": []}) + assert denoise_train.selected_slice_indices(item, 3) == [0, 1, 2] + + +# --------------------------------------------------------------------------- +# noise2noise_pairs +# --------------------------------------------------------------------------- + +class TestNoise2NoisePairs: + def test_adjacent_pairs_both_directions(self): + pairs = denoise_train.noise2noise_pairs([0, 1, 2], stride=1, both_directions=True) + assert set(pairs) == {(0, 1), (1, 0), (1, 2), (2, 1)} + + def test_single_direction(self): + pairs = denoise_train.noise2noise_pairs([0, 1, 2], stride=1, both_directions=False) + assert set(pairs) == {(0, 1), (1, 2)} + + def test_last_slice_has_no_partner_not_clamped(self): + pairs = denoise_train.noise2noise_pairs([0, 1], stride=1, both_directions=False) + assert (1, 1) not in pairs + + def test_stride_skips_non_adjacent(self): + pairs = denoise_train.noise2noise_pairs([0, 1, 2], stride=2, both_directions=False) + assert pairs == [(0, 2)] + + def test_no_qualifying_pairs_returns_empty(self): + assert denoise_train.noise2noise_pairs([0], stride=1) == [] + + def test_stride_below_one_raises(self): + with pytest.raises(ValueError, match="stride must be >= 1"): + denoise_train.noise2noise_pairs([0, 1], stride=0) + + +# --------------------------------------------------------------------------- +# _mirror_offset +# --------------------------------------------------------------------------- + +class TestMirrorOffset: + def test_in_range_offset_unchanged(self): + out = denoise_train._mirror_offset(np.array([5]), np.array([2]), 10) + assert out[0] == 7 + + def test_out_of_range_flips_sign_instead_of_clamping(self): + # centre=0, offset=-1 would clip to 0 (== centre) under naive clamping; + # flipping gives centre - offset = 1, which differs from centre. + out = denoise_train._mirror_offset(np.array([0]), np.array([-1]), 10) + assert out[0] == 1 + + def test_result_always_in_bounds(self): + centre = np.array([0, 9]) + offset = np.array([-5, 5]) + out = denoise_train._mirror_offset(centre, offset, 10) + assert np.all((out >= 0) & (out < 10)) + + +# --------------------------------------------------------------------------- +# n2v_mask_and_replace +# --------------------------------------------------------------------------- + +class TestN2vMaskAndReplace: + def test_mask_count_matches_fraction(self): + patch = (np.arange(400) % 256).astype(np.uint8).reshape(20, 20) + rng = np.random.default_rng(0) + _, mask = denoise_train.n2v_mask_and_replace(patch, rng, fraction=0.05, neighbourhood=5) + assert mask.sum() == round(400 * 0.05) + + def test_donor_never_the_masked_pixel_itself(self): + patch = (np.arange(400) % 256).astype(np.uint8).reshape(20, 20) + rng = np.random.default_rng(1) + masked, mask = denoise_train.n2v_mask_and_replace(patch, rng, fraction=0.1, neighbourhood=5) + # A masked pixel that kept its original value would mean its donor was itself. + assert not np.array_equal(masked[mask], patch[mask]) + + def test_at_least_one_pixel_masked_for_tiny_fraction(self): + patch = np.zeros((10, 10), dtype=np.uint8) + rng = np.random.default_rng(0) + _, mask = denoise_train.n2v_mask_and_replace(patch, rng, fraction=0.0001, neighbourhood=3) + assert mask.sum() >= 1 + + def test_unmasked_pixels_untouched(self): + patch = (np.arange(400) % 256).astype(np.uint8).reshape(20, 20) + rng = np.random.default_rng(2) + masked, mask = denoise_train.n2v_mask_and_replace(patch, rng, fraction=0.05, neighbourhood=5) + assert np.array_equal(masked[~mask], patch[~mask]) + + @pytest.mark.parametrize("fraction", [0.0, 1.0, -0.1, 1.5]) + def test_fraction_out_of_range_raises(self, fraction): + with pytest.raises(ValueError, match="fraction must be in"): + denoise_train.n2v_mask_and_replace(np.zeros((10, 10), dtype=np.uint8), np.random.default_rng(0), fraction=fraction) + + @pytest.mark.parametrize("neighbourhood", [2, 4, 1]) + def test_neighbourhood_must_be_odd_and_at_least_three(self, neighbourhood): + if neighbourhood >= 3 and neighbourhood % 2 == 1: + return + with pytest.raises(ValueError, match="neighbourhood must be an odd size"): + denoise_train.n2v_mask_and_replace( + np.zeros((10, 10), dtype=np.uint8), np.random.default_rng(0), neighbourhood=neighbourhood, + ) + + +# --------------------------------------------------------------------------- +# letterbox_denoise_pair +# --------------------------------------------------------------------------- + +class TestLetterboxDenoisePair: + def test_already_correct_size_short_circuits(self): + inp = np.zeros((32, 32), dtype=np.uint8) + tgt = np.ones((32, 32), dtype=np.uint8) + out_inp, out_tgt = denoise_train.letterbox_denoise_pair(inp, tgt, 32) + assert out_inp is inp and out_tgt is tgt + + def test_resizes_and_pads_to_target_size(self): + inp = np.full((16, 32), 100, dtype=np.uint8) + tgt = np.full((16, 32), 200, dtype=np.uint8) + out_inp, out_tgt = denoise_train.letterbox_denoise_pair(inp, tgt, 64) + assert out_inp.shape == (64, 64) + assert out_tgt.shape == (64, 64) + # Corners should be zero-padded (aspect-ratio letterboxing). + assert out_inp[0, 0] == 0 + + def test_shape_mismatch_raises(self): + with pytest.raises(ValueError, match="shapes differ"): + denoise_train.letterbox_denoise_pair( + np.zeros((10, 10), dtype=np.uint8), np.zeros((10, 20), dtype=np.uint8), 32, + ) + + +# --------------------------------------------------------------------------- +# prepare_noise2noise_datasets / prepare_noise2void_datasets (real images_mod, +# fake arrays_mod — same pattern as test_infer_jobs.py's fake array source) +# --------------------------------------------------------------------------- + +@pytest.fixture() +def fake_array_source(monkeypatch: pytest.MonkeyPatch): + rng = np.random.default_rng(1) + volume = rng.integers(0, 256, size=(4, 16, 16), dtype=np.uint8) + + monkeypatch.setattr(arrays_mod, "resolve_array", lambda source, kind, server_uri: volume) + monkeypatch.setattr( + arrays_mod, "array_shape_meta", + lambda node, pyramid=None: {"height": 16, "width": 16, "n_slices": 4}, + ) + monkeypatch.setattr(arrays_mod, "read_slice", lambda node, meta, idx: node[idx]) + return volume + + +class TestPrepareDatasets: + def test_noise2noise_produces_pairs_from_adjacent_slices(self, fake_array_source): + result = denoise_train.prepare_noise2noise_datasets([FakeItem()], render={}) + assert result["val"] == [] + assert len(result["train"]) > 0 + for inp, tgt in result["train"]: + assert inp.dtype == np.uint8 and inp.shape == (16, 16) + + def test_noise2noise_raises_when_only_one_slice_in_scope(self, fake_array_source): + item = FakeItem(slices={"0": []}) + with pytest.raises(ValueError, match="Noise2Noise needs at least two slices"): + denoise_train.prepare_noise2noise_datasets([item], render={}) + + def test_noise2void_pairs_each_slice_with_itself(self, fake_array_source): + result = denoise_train.prepare_noise2void_datasets([FakeItem()], render={}) + assert len(result["train"]) == 4 + for inp, tgt in result["train"]: + assert np.array_equal(inp, tgt) + + def test_noise2void_raises_with_no_slices_in_scope(self, monkeypatch, fake_array_source): + monkeypatch.setattr( + arrays_mod, "array_shape_meta", + lambda node, pyramid=None: {"height": 16, "width": 16, "n_slices": 0}, + ) + with pytest.raises(ValueError, match="Noise2Void needs at least one slice"): + denoise_train.prepare_noise2void_datasets([FakeItem()], render={}) + + +# --------------------------------------------------------------------------- +# run_denoise_training_loop — real round trip, minimal trainable model +# --------------------------------------------------------------------------- + +def _tiny_model(): + """A minimal real nn.Module — the loop is generic over forward_fn/ + to_tensor_fn/trainable_params, so it doesn't need dlsia's actual TUNet.""" + model = torch.nn.Conv2d(1, 1, kernel_size=3, padding=1) + + def forward_fn(batch): + return model(batch) + + def to_tensor_fn(arr): + return torch.from_numpy(np.ascontiguousarray(arr)).float().unsqueeze(0) / 255.0 + + return forward_fn, to_tensor_fn, list(model.parameters()) + + +def _pairs(n, size=8, seed=0): + rng = np.random.default_rng(seed) + return [ + (rng.integers(0, 256, size=(size, size), dtype=np.uint8), rng.integers(0, 256, size=(size, size), dtype=np.uint8)) + for _ in range(n) + ] + + +class TestRunDenoiseTrainingLoop: + @pytest.mark.parametrize("scheme", ["n2n", "n2v", "ae"]) + def test_real_round_trip_all_schemes(self, scheme): + forward_fn, to_tensor_fn, params = _tiny_model() + metrics = denoise_train.run_denoise_training_loop( + train_pairs=_pairs(6), val_pairs=_pairs(2, seed=1), + image_size=8, training_scheme=scheme, epochs=2, batch_size=2, seed=0, + flip_augment=True, to_tensor_fn=to_tensor_fn, forward_fn=forward_fn, + trainable_params=params, lr=1e-3, device="cpu", + ) + assert metrics["epochs_completed"] == 2 + assert metrics["cancelled"] is False + assert np.isfinite(metrics["final_train_loss"]) + assert np.isfinite(metrics["final_val_loss"]) + assert metrics[denoise_train.VAL_METRIC_KEY] is not None + + def test_cancellation_via_on_epoch_stops_early_with_partial_metrics(self): + forward_fn, to_tensor_fn, params = _tiny_model() + metrics = denoise_train.run_denoise_training_loop( + train_pairs=_pairs(4), val_pairs=[], + image_size=8, training_scheme="n2n", epochs=5, batch_size=2, seed=0, + flip_augment=False, to_tensor_fn=to_tensor_fn, forward_fn=forward_fn, + trainable_params=params, lr=1e-3, device="cpu", + on_epoch=lambda epoch, *_: epoch == 1, + ) + assert metrics["cancelled"] is True + assert metrics["epochs_completed"] == 1 + + def test_no_training_data_raises(self): + forward_fn, to_tensor_fn, params = _tiny_model() + with pytest.raises(ValueError, match="No training data"): + denoise_train.run_denoise_training_loop( + train_pairs=[], val_pairs=[], image_size=8, training_scheme="n2n", + epochs=1, batch_size=2, seed=0, flip_augment=False, + to_tensor_fn=to_tensor_fn, forward_fn=forward_fn, trainable_params=params, + lr=1e-3, device="cpu", + ) + + def test_unknown_scheme_raises(self): + forward_fn, to_tensor_fn, params = _tiny_model() + with pytest.raises(ValueError, match="Unknown denoiser training scheme"): + denoise_train.run_denoise_training_loop( + train_pairs=_pairs(2), val_pairs=[], image_size=8, training_scheme="bogus", + epochs=1, batch_size=2, seed=0, flip_augment=False, + to_tensor_fn=to_tensor_fn, forward_fn=forward_fn, trainable_params=params, + lr=1e-3, device="cpu", + ) diff --git a/backend/tests/test_drafts.py b/backend/tests/test_drafts.py index 8b841e2..1e8f123 100644 --- a/backend/tests/test_drafts.py +++ b/backend/tests/test_drafts.py @@ -1,58 +1,197 @@ -"""Tests for session-draft persistence.""" - +"""Tests for drafts.py — real filesystem (tmp_path), no mocks needed. Sets +`drafts._DRAFT_DIR` directly rather than the LOCAL_DATA_ROOT env var, since +it's computed once at import time (see MEMORY.md's import-time-constant +gotcha) — the pattern already used in test_annotation_server_routes.py.""" from __future__ import annotations -import importlib -import os +import json + +import pytest + +import drafts + + +@pytest.fixture(autouse=True) +def draft_dir(tmp_path, monkeypatch: pytest.MonkeyPatch): + d = tmp_path / ".drafts" + monkeypatch.setattr(drafts, "_DRAFT_DIR", d) + return d + + +# --------------------------------------------------------------------------- +# save_draft / load_draft +# --------------------------------------------------------------------------- + +class TestSaveAndLoadDraft: + def test_round_trips_payload(self): + result = drafts.save_draft("local:foo.tif", {"classes": [], "slices": {}}) + assert "saved_at" in result + assert "path" in result + + loaded = drafts.load_draft("local:foo.tif") + assert loaded["source_key"] == "local:foo.tif" + assert loaded["payload"] == {"classes": [], "slices": {}} + def test_missing_draft_returns_none(self): + assert drafts.load_draft("nope") is None -def test_save_and_load_draft(tmp_path) -> None: - """A saved draft should be loadable with the correct source_key.""" - with __import__("unittest.mock", fromlist=["patch"]).patch.dict( - os.environ, {"LOCAL_DATA_ROOT": str(tmp_path)} - ): - import drafts + def test_different_source_keys_do_not_collide(self): + drafts.save_draft("a", {"v": 1}) + drafts.save_draft("b", {"v": 2}) + assert drafts.load_draft("a")["payload"] == {"v": 1} + assert drafts.load_draft("b")["payload"] == {"v": 2} - importlib.reload(drafts) - meta = drafts.save_draft("test-source", {"classes": [], "slices": {}}) - assert "saved_at" in meta - loaded = drafts.load_draft("test-source") - assert loaded is not None - assert loaded["source_key"] == "test-source" + def test_overwrites_previous_draft_for_same_key(self): + drafts.save_draft("a", {"v": 1}) + drafts.save_draft("a", {"v": 2}) + assert drafts.load_draft("a")["payload"] == {"v": 2} + def test_corrupt_draft_file_returns_none(self, draft_dir): + draft_dir.mkdir(parents=True, exist_ok=True) + path = drafts._draft_path("bad") + path.write_text("{not valid json") + assert drafts.load_draft("bad") is None -def test_load_missing_draft(tmp_path) -> None: - """Loading a draft that was never saved should return None.""" - with __import__("unittest.mock", fromlist=["patch"]).patch.dict( - os.environ, {"LOCAL_DATA_ROOT": str(tmp_path)} - ): - import drafts + def test_write_is_atomic_no_leftover_tmp_file(self, draft_dir): + drafts.save_draft("a", {"v": 1}) + tmp_files = list(draft_dir.glob("*.tmp")) + assert tmp_files == [] - importlib.reload(drafts) - result = drafts.load_draft("nonexistent-key") - assert result is None +class TestDraftPath: + def test_same_key_produces_same_path(self): + assert drafts._draft_path("x") == drafts._draft_path("x") -def test_list_drafts_returns_saved(tmp_path) -> None: - """list_drafts should include a previously saved draft.""" - with __import__("unittest.mock", fromlist=["patch"]).patch.dict( - os.environ, {"LOCAL_DATA_ROOT": str(tmp_path)} - ): - import drafts + def test_different_keys_produce_different_paths(self): + assert drafts._draft_path("x") != drafts._draft_path("y") - importlib.reload(drafts) - drafts.save_draft("my-source", {"classes": [{"classId": 1, "label": "cell", "color": "#f00", "isVisible": True}]}) - all_drafts = drafts.list_drafts() - assert any(d["source_key"] == "my-source" for d in all_drafts) +# --------------------------------------------------------------------------- +# list_drafts +# --------------------------------------------------------------------------- -def test_list_drafts_empty_when_none(tmp_path) -> None: - """list_drafts should return an empty list when no drafts exist.""" - with __import__("unittest.mock", fromlist=["patch"]).patch.dict( - os.environ, {"LOCAL_DATA_ROOT": str(tmp_path)} - ): - import drafts +class TestListDrafts: + def test_empty_when_dir_does_not_exist(self): + assert drafts.list_drafts() == [] - importlib.reload(drafts) + def test_lists_saved_drafts(self): + drafts.save_draft("a", {"slices": {}}) + drafts.save_draft("b", {"slices": {}}) result = drafts.list_drafts() - assert result == [] + assert len(result) == 2 + keys = {d["source_key"] for d in result} + assert keys == {"a", "b"} + + def test_has_annotations_true_when_a_slice_has_shapes(self): + drafts.save_draft("a", {"slices": {"0": [{"id": "s1"}]}}) + result = drafts.list_drafts() + assert result[0]["has_annotations"] is True + + def test_has_annotations_false_when_all_slices_empty(self): + drafts.save_draft("a", {"slices": {"0": []}}) + result = drafts.list_drafts() + assert result[0]["has_annotations"] is False + + def test_corrupt_draft_file_is_skipped_not_raised(self, draft_dir): + drafts.save_draft("a", {"slices": {}}) + draft_dir.mkdir(parents=True, exist_ok=True) + (draft_dir / "corrupt.json").write_text("{not json") + result = drafts.list_drafts() + assert len(result) == 1 + + +# --------------------------------------------------------------------------- +# save_version / list_versions / get_version +# --------------------------------------------------------------------------- + +class TestSaveVersion: + def test_first_version_is_1(self): + result = drafts.save_version("a", {"classes": [], "slices": {}}) + assert result["version"] == 1 + + def test_versions_increment(self): + drafts.save_version("a", {"classes": [], "slices": {}}) + result2 = drafts.save_version("a", {"classes": [], "slices": {}}) + assert result2["version"] == 2 + + def test_shape_count_and_class_count_computed(self): + payload = { + "classes": [{"classId": 1}, {"classId": 2}], + "slices": {"0": [{"id": "s1"}, {"id": "s2"}], "1": [{"id": "s3"}]}, + } + result = drafts.save_version("a", payload) + assert result["shape_count"] == 3 + + def test_also_updates_the_crash_recovery_draft(self): + drafts.save_version("a", {"classes": [], "slices": {"0": [{"id": "s1"}]}}) + loaded = drafts.load_draft("a") + assert loaded["payload"]["slices"] == {"0": [{"id": "s1"}]} + + def test_annotated_by_and_notes_are_trimmed_and_none_when_blank(self): + result_blank = drafts.save_version("a", {"classes": [], "slices": {}}, annotated_by=" ", notes="") + vdoc = json.loads((drafts._versions_dir("a") / "v0001.json").read_text()) + assert vdoc["annotated_by"] is None + assert vdoc["notes"] is None + + drafts.save_version("a", {"classes": [], "slices": {}}, annotated_by=" Alice ", notes=" hi ") + vdoc2 = json.loads((drafts._versions_dir("a") / "v0002.json").read_text()) + assert vdoc2["annotated_by"] == "Alice" + assert vdoc2["notes"] == "hi" + assert result_blank["version"] == 1 + + +class TestListVersions: + def test_empty_when_no_versions_dir(self): + assert drafts.list_versions("nope") == [] + + def test_sorted_oldest_first(self): + drafts.save_version("a", {"classes": [], "slices": {}}) + drafts.save_version("a", {"classes": [], "slices": {}}) + result = drafts.list_versions("a") + assert [v["version"] for v in result] == [1, 2] + + def test_has_thumbnail_reflects_saved_thumbnail(self): + drafts.save_version("a", {"classes": [], "slices": {}}) + result = drafts.list_versions("a") + assert result[0]["has_thumbnail"] is False + drafts.save_version_thumbnail("a", 1, b"pngbytes") + result2 = drafts.list_versions("a") + assert result2[0]["has_thumbnail"] is True + + def test_corrupt_version_file_is_skipped(self): + drafts.save_version("a", {"classes": [], "slices": {}}) + (drafts._versions_dir("a") / "vbad.json").write_text("not json at all {") + result = drafts.list_versions("a") + assert len(result) == 1 + + +class TestGetVersion: + def test_returns_full_document(self): + drafts.save_version("a", {"classes": [{"classId": 1}], "slices": {}}, notes="test") + doc = drafts.get_version("a", 1) + assert doc["payload"]["classes"] == [{"classId": 1}] + assert doc["notes"] == "test" + + def test_missing_version_returns_none(self): + drafts.save_version("a", {"classes": [], "slices": {}}) + assert drafts.get_version("a", 99) is None + + def test_corrupt_version_returns_none(self): + vdir = drafts._versions_dir("a") + vdir.mkdir(parents=True, exist_ok=True) + (vdir / "v0001.json").write_text("{ broken") + assert drafts.get_version("a", 1) is None + + +class TestVersionThumbnails: + def test_missing_thumbnail_returns_none(self): + assert drafts.get_version_thumbnail("a", 1) is None + + def test_round_trips_thumbnail_bytes(self): + drafts.save_version_thumbnail("a", 3, b"\x89PNGdata") + assert drafts.get_version_thumbnail("a", 3) == b"\x89PNGdata" + + def test_thumbnail_write_is_atomic_no_leftover_tmp(self): + drafts.save_version_thumbnail("a", 1, b"data") + vdir = drafts._versions_dir("a") + assert list(vdir.glob("*.tmp")) == [] diff --git a/backend/tests/test_export_jobs.py b/backend/tests/test_export_jobs.py new file mode 100644 index 0000000..7f9bc38 --- /dev/null +++ b/backend/tests/test_export_jobs.py @@ -0,0 +1,56 @@ +"""Tests for export_jobs' shared job registry, focused on PROFILE_JOBS timing.""" +import importlib +import os + +import pytest + +import export_jobs + + +@pytest.fixture(autouse=True) +def _restore_default_module_state(): + """PROFILE_JOBS is read once at import — reloading the module to flip it + (see _reload_with_profile) mutates process-global state that outlives + monkeypatch's own env var cleanup, regardless of fixture teardown order. + Explicitly clear the env var and reload back to the un-profiled default + after every test so other test files never see a leftover reload here.""" + yield + os.environ.pop("PROFILE_JOBS", None) + importlib.reload(export_jobs) + + +def _reload_with_profile(monkeypatch, enabled: bool): + """PROFILE_JOBS is read once at import time — reload the module under a + patched env var rather than mutating a private flag directly, so the test + exercises the real startup path.""" + monkeypatch.setenv("PROFILE_JOBS", "1" if enabled else "0") + module = importlib.reload(export_jobs) + return module + + +def test_log_has_no_timing_suffix_by_default(monkeypatch): + module = _reload_with_profile(monkeypatch, enabled=False) + jid = module.new_job("some/path") + module.log(jid, "starting") + job = module.get_job(jid) + assert job["log"] == ["starting"] + + +def test_log_adds_elapsed_suffix_when_profiling_enabled(monkeypatch): + module = _reload_with_profile(monkeypatch, enabled=True) + jid = module.new_job("some/path") + module.log(jid, "starting") + module.log(jid, "done") + job = module.get_job(jid) + assert len(job["log"]) == 2 + for line in job["log"]: + assert "(+" in line and line.endswith("s)") + + +def test_internal_last_log_at_field_never_leaks_into_snapshot(monkeypatch): + module = _reload_with_profile(monkeypatch, enabled=True) + jid = module.new_job("some/path") + module.log(jid, "starting") + job = module.get_job(jid) + assert "_last_log_at" not in job + assert "zip_path" not in job diff --git a/backend/tests/test_infer_jobs.py b/backend/tests/test_infer_jobs.py new file mode 100644 index 0000000..b0cbff2 --- /dev/null +++ b/backend/tests/test_infer_jobs.py @@ -0,0 +1,513 @@ +"""Real (not mocked) round-trip tests for infer_jobs.run_infer_job, focused on +this session's rework: the per-slice worker pool, the GPU_FORWARD_LOCK/ML_LOCK +split, and #15's live-preview (incremental result/cache updates). + +Skips automatically if the `ml` extra (torch/dlsia) isn't installed — see +test_train_e2e_real_ml.py for the sibling test proving the underlying +train/save/load contract. This file builds on the same real-model pattern but +drives it through infer_jobs.run_infer_job itself, with a fake array source +(no real Tiled/local dataset needed) so it stays fast and self-contained. +""" + +from __future__ import annotations + +import io +import sys +import threading +import time +from pathlib import Path + +import numpy as np +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +pytest.importorskip("torch") +pytest.importorskip("dlsia") + +import arrays as arrays_mod # noqa: E402 +import export_jobs # noqa: E402 +import infer_jobs # noqa: E402 +import train_common # noqa: E402 +from schemas import DlsiaTunetConfig, InferRequest # noqa: E402 + +IMAGE_SIZE = 64 +N_SLICES = 6 +N_CLASSES = 2 + + +@pytest.fixture() +def runs_dir(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Path: + monkeypatch.setenv("DINO_RUNS_DIR", str(tmp_path / "runs")) + return tmp_path / "runs" + + +@pytest.fixture() +def trained_run_id(runs_dir) -> str: + """A real, tiny, trained dlsia_tunet run — same pattern as + test_train_e2e_real_ml.py, sized down for test speed.""" + model_cfg = DlsiaTunetConfig( + hyperparams={ + "epochs": 1, "depth": 2, "base_channels": 4, "growth_rate": 1.2, + "batch_size": 2, "image_size": IMAGE_SIZE, "tiling": False, + } + ) + built = train_common.build_family(model_cfg, N_CLASSES, "cpu", lambda _msg: None) + + rng = np.random.default_rng(0) + train_pairs = [ + ( + rng.integers(0, 256, size=(IMAGE_SIZE, IMAGE_SIZE, 3), dtype=np.uint8), + rng.integers(0, N_CLASSES, size=(IMAGE_SIZE, IMAGE_SIZE), dtype=np.uint8), + ) + for _ in range(4) + ] + train_common.run_training_loop( + train_pairs=train_pairs, val_pairs=[], image_size=IMAGE_SIZE, n_classes=N_CLASSES, + epochs=1, batch_size=2, seed=0, flip_augment=False, + to_tensor_fn=built.to_tensor_fn, forward_fn=built.forward_fn, + trainable_params=built.trainable_params, lr=1e-3, device="cpu", + set_train_mode=built.set_train_mode, + ) + + run_id = "test-infer-job-run" + train_common.save_run( + run_id, + model_family=model_cfg.model_family, + model_config=built.model_config_snapshot, + classes=[{"classId": 1, "label": "a", "color": "#f00"}, {"classId": 2, "label": "b", "color": "#0f0"}], + render={}, + image_size=IMAGE_SIZE, + hyperparams=model_cfg.hyperparams.model_dump(), + source_keys=["local:fake.tif"], + adapter_state=built.adapter_state_fn(), + metrics={"epochs_completed": 1, "cancelled": False}, + ) + return run_id + + +@pytest.fixture() +def fake_array_source(monkeypatch: pytest.MonkeyPatch): + """Serve a small synthetic in-memory volume in place of a real Tiled/ + local dataset — infer_jobs only ever calls these three arrays.py + functions, so patching them is enough to drive run_infer_job for real + without any actual dataset on disk. + + A short sleep per read_slice call gives concurrent slice-workers a real + window to overlap in (proving the pool actually parallelizes I/O across + slices, not just that it doesn't crash) and gives a background-thread + poller time to observe partial (live-preview) results before the job + finishes. + """ + rng = np.random.default_rng(1) + volume = rng.integers(0, 256, size=(N_SLICES, IMAGE_SIZE, IMAGE_SIZE), dtype=np.uint8) + + def fake_resolve_array(source, kind, server_uri): + return volume + + def fake_array_shape_meta(node, pyramid=None): + return {"height": IMAGE_SIZE, "width": IMAGE_SIZE, "n_slices": N_SLICES} + + def fake_read_slice(node, meta, idx): + time.sleep(0.05) + return node[idx] + + monkeypatch.setattr(arrays_mod, "resolve_array", fake_resolve_array) + monkeypatch.setattr(arrays_mod, "array_shape_meta", fake_array_shape_meta) + monkeypatch.setattr(arrays_mod, "read_slice", fake_read_slice) + return volume + + +def _infer_request(run_id: str) -> InferRequest: + return InferRequest( + run_id=run_id, kind="local", source="fake.tif", + slice_indices=list(range(N_SLICES)), min_confidence=0.0, + ) + + +def test_run_infer_job_completes_with_correct_results_and_cache(trained_run_id, fake_array_source): + jid = export_jobs.new_job(trained_run_id) + infer_jobs.run_infer_job(jid, _infer_request(trained_run_id)) + + job = export_jobs.get_job(jid) + assert job["state"] == "done" + result = job["result"] + assert result["cancelled"] is False + assert sorted(int(k) for k in result["slices"]) == list(range(N_SLICES)) + assert result["preview_slices"] == list(range(N_SLICES)) + + cached = infer_jobs._cache_get(jid) + assert cached is not None + assert sorted(cached["label_pngs"].keys()) == list(range(N_SLICES)) + + # ML_LOCK must always be released, whether the job succeeds or not — + # otherwise every subsequent train/infer/bake/probe job would wrongly + # report "another job is already running" forever. + assert not train_common.ML_LOCK.locked() + + +def test_run_infer_job_publishes_partial_results_before_completion(trained_run_id, fake_array_source): + """Regression test for #15: the job's `result` must be readable and + growing WHILE state is still "running", not only once "done" — this is + what lets the frontend show a usable preview slider mid-job instead of + the old "nothing until every slice finishes" behavior.""" + jid = export_jobs.new_job(trained_run_id) + thread = threading.Thread( + target=infer_jobs.run_infer_job, args=(jid, _infer_request(trained_run_id)), daemon=True, + ) + thread.start() + + saw_partial_progress = False + deadline = time.monotonic() + 30.0 + while time.monotonic() < deadline: + job = export_jobs.get_job(jid) + if job["state"] == "running" and job.get("result"): + n = len(job["result"].get("preview_slices", [])) + if 0 < n < N_SLICES: + saw_partial_progress = True + break + if job["state"] in ("done", "error"): + break + time.sleep(0.02) + thread.join(timeout=30.0) + + assert saw_partial_progress, "expected to observe a partial (running) result with some but not all slices done" + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert len(job["result"]["preview_slices"]) == N_SLICES + + +def test_run_infer_job_cancellation_stops_early_with_partial_results(trained_run_id, fake_array_source): + jid = export_jobs.new_job(trained_run_id) + thread = threading.Thread( + target=infer_jobs.run_infer_job, args=(jid, _infer_request(trained_run_id)), daemon=True, + ) + thread.start() + time.sleep(0.06) # let at least one slice start + export_jobs.request_cancel(jid) + thread.join(timeout=30.0) + + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["cancelled"] is True + assert len(job["result"]["preview_slices"]) <= N_SLICES + assert not train_common.ML_LOCK.locked() + + +def test_predict_one_slice_serializes_gpu_forward_calls(trained_run_id, fake_array_source, monkeypatch): + """Direct test of the concurrency guarantee GPU_FORWARD_LOCK exists for: + even with several slice-workers running at once, at most one is ever + inside the model-forward critical section at a time.""" + import dlsia_runtime as fam + + config = train_common.load_run_config(trained_run_id) + adapter_state = train_common.load_adapter_state(trained_run_id) + model = fam.load_model(adapter_state, "cpu") + model.eval() + real_forward_fn = fam.make_forward_fn(model) + to_tensor_fn = fam.make_to_tensor_fn() + + concurrent_count = 0 + max_concurrent = 0 + count_lock = threading.Lock() + + def spying_forward_fn(batch): + nonlocal concurrent_count, max_concurrent + with count_lock: + concurrent_count += 1 + max_concurrent = max(max_concurrent, concurrent_count) + try: + time.sleep(0.05) # widen the window a real race would need + return real_forward_fn(batch) + finally: + with count_lock: + concurrent_count -= 1 + + node = fake_array_source + meta = {"height": IMAGE_SIZE, "width": IMAGE_SIZE} + request = _infer_request(trained_run_id) + run_classes = config["classes"] + tiling_logged = threading.Event() + + import concurrent.futures + + with concurrent.futures.ThreadPoolExecutor(max_workers=4) as pool: + futures = [ + pool.submit( + infer_jobs._predict_one_slice, + i, + jid="test-jid", node=node, meta=meta, h=IMAGE_SIZE, w=IMAGE_SIZE, + render={}, global_range=None, + render_slice_fn=train_common.denoising_render_slice_fn(None), + tiled=False, image_size=IMAGE_SIZE, forward_fn=spying_forward_fn, + to_tensor_fn=to_tensor_fn, device="cpu", request=request, + run_classes=run_classes, tiling_logged=tiling_logged, + ) + for i in range(N_SLICES) + ] + results = [f.result() for f in futures] + + for slice_idx, png_bytes, shapes, error in results: + assert error is None, f"slice {slice_idx} failed: {error}" + assert png_bytes is not None + assert shapes is not None + + assert max_concurrent == 1, "GPU_FORWARD_LOCK failed to serialize concurrent forward calls" + + +# --------------------------------------------------------------------------- +# Pure helpers +# --------------------------------------------------------------------------- + +class TestHexToRgb: + def test_full_hex(self): + assert infer_jobs._hex_to_rgb("#00ff80") == (0, 255, 128) + + def test_short_hex_is_expanded(self): + assert infer_jobs._hex_to_rgb("#0f8") == (0, 255, 136) + + def test_none_falls_back_to_red(self): + assert infer_jobs._hex_to_rgb(None) == (255, 0, 0) + + def test_missing_hash_falls_back_to_red(self): + assert infer_jobs._hex_to_rgb("00ff80") == (255, 0, 0) + + def test_invalid_hex_digits_fall_back_to_red(self): + assert infer_jobs._hex_to_rgb("#zzzzzz") == (255, 0, 0) + + +class TestVectorizeLabelMap: + def _run_classes(self): + return [{"classId": 1, "label": "a"}, {"classId": 2, "label": "b"}] + + def test_component_below_min_area_is_dropped(self): + label_map = np.zeros((20, 20), dtype=np.uint8) + label_map[0:2, 0:2] = 1 # 4px, tiny + shapes = infer_jobs._vectorize_label_map(label_map, self._run_classes(), min_area=50, simplify_tol=0.0, run_id="run12345", slice_idx=0) + assert shapes == [] + + def test_ring_shaped_component_gets_a_hole(self): + label_map = np.ones((30, 30), dtype=np.uint8) # all class 1 + label_map[10:20, 10:20] = 2 # a class-2 hole inside class 1 + shapes = infer_jobs._vectorize_label_map(label_map, self._run_classes(), min_area=1, simplify_tol=0.0, run_id="run12345", slice_idx=0) + class1_shape = next(s for s in shapes if s["classId"] == 1) + assert "holes" in class1_shape + assert len(class1_shape["holes"]) == 1 + + def test_simplify_tol_reduces_point_count(self): + label_map = np.zeros((40, 40), dtype=np.uint8) + label_map[5:35, 5:35] = 1 # a large, simple square + unsimplified = infer_jobs._vectorize_label_map(label_map, self._run_classes(), min_area=1, simplify_tol=0.0, run_id="run12345", slice_idx=0) + simplified = infer_jobs._vectorize_label_map(label_map, self._run_classes(), min_area=1, simplify_tol=5.0, run_id="run12345", slice_idx=0) + assert len(simplified[0]["points"]) <= len(unsimplified[0]["points"]) + + def test_shape_ids_are_unique_per_component(self): + label_map = np.zeros((30, 30), dtype=np.uint8) + label_map[2:6, 2:6] = 1 + label_map[20:26, 20:26] = 1 + shapes = infer_jobs._vectorize_label_map(label_map, self._run_classes(), min_area=1, simplify_tol=0.0, run_id="run12345", slice_idx=3) + ids = [s["id"] for s in shapes] + assert len(ids) == len(set(ids)) + + def test_class_with_no_pixels_is_skipped(self): + label_map = np.zeros((10, 10), dtype=np.uint8) + shapes = infer_jobs._vectorize_label_map(label_map, self._run_classes(), min_area=1, simplify_tol=0.0, run_id="run12345", slice_idx=0) + assert shapes == [] + + +# --------------------------------------------------------------------------- +# run_infer_job — guard paths that don't need a real trained run +# --------------------------------------------------------------------------- + +class TestRunInferJobGuards: + def test_busy_ml_lock_reports_error_not_a_crash(self): + train_common.ML_LOCK.acquire() + try: + jid = export_jobs.new_job("x") + infer_jobs.run_infer_job(jid, _infer_request("whatever")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "already running" in job["error"] + finally: + train_common.ML_LOCK.release() + + def test_no_device_reports_error_and_releases_lock(self, monkeypatch): + monkeypatch.setattr(train_common, "pick_device", lambda: None) + jid = export_jobs.new_job("x") + infer_jobs.run_infer_job(jid, _infer_request("whatever")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "torch is not installed" in job["error"] + assert train_common.ML_LOCK.locked() is False + + def test_denoiser_run_is_refused(self, monkeypatch): + monkeypatch.setattr( + train_common, "load_run_config", + lambda run_id: {"model_family": "dlsia_denoiser", "classes": [], "image_size": 64, "render": {}}, + ) + monkeypatch.setattr(train_common, "load_adapter_state", lambda run_id: {}) + jid = export_jobs.new_job("x") + infer_jobs.run_infer_job(jid, _infer_request("whatever")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "denoiser run" in job["error"] + + def test_unsupported_model_family_is_refused(self, monkeypatch): + monkeypatch.setattr( + train_common, "load_run_config", + lambda run_id: {"model_family": "something_weird", "classes": [], "image_size": 64, "render": {}}, + ) + monkeypatch.setattr(train_common, "load_adapter_state", lambda run_id: {}) + jid = export_jobs.new_job("x") + infer_jobs.run_infer_job(jid, _infer_request("whatever")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "Unsupported model family" in job["error"] + + def test_tiled_run_without_qlty_is_refused(self, monkeypatch): + import tiling + + monkeypatch.setattr( + train_common, "load_run_config", + lambda run_id: { + "model_family": "dlsia_tunet", "classes": [], "image_size": 64, "render": {}, + "hyperparams": {"tiling": True}, + }, + ) + monkeypatch.setattr(train_common, "load_adapter_state", lambda run_id: {}) + monkeypatch.setattr(tiling, "qlty_available", lambda: False) + jid = export_jobs.new_job("x") + infer_jobs.run_infer_job(jid, _infer_request("whatever")) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "qlty" in job["error"] + + +# --------------------------------------------------------------------------- +# preview_png +# --------------------------------------------------------------------------- + +class TestPreviewPng: + def test_missing_job_is_404(self): + with pytest.raises(Exception) as exc: + infer_jobs.preview_png("no-such-job", 0) + assert exc.value.status_code == 404 + + def test_missing_slice_is_404(self): + infer_jobs._cache_put("job-with-no-slices", {"classes": [], "label_pngs": {}}) + with pytest.raises(Exception) as exc: + infer_jobs.preview_png("job-with-no-slices", 0) + assert exc.value.status_code == 404 + + def test_colorizes_predicted_classes(self): + from PIL import Image as PILImage + + label = np.zeros((8, 8), dtype=np.uint8) + label[0:4, 0:4] = 1 + label[4:8, 4:8] = 2 + buf = io.BytesIO() + PILImage.fromarray(label, mode="L").save(buf, format="PNG") + + infer_jobs._cache_put( + "job-colorize", + { + "classes": [{"classId": 1, "color": "#ff0000"}, {"classId": 2, "color": "#00ff00"}], + "label_pngs": {0: buf.getvalue()}, + }, + ) + png = infer_jobs.preview_png("job-colorize", 0) + rgba = np.asarray(PILImage.open(io.BytesIO(png))) + assert tuple(rgba[0, 0]) == (255, 0, 0, 180) + assert tuple(rgba[4, 4]) == (0, 255, 0, 180) + assert rgba[7, 0][3] == 0 # background stays transparent + + +# --------------------------------------------------------------------------- +# run_write_tiled_job +# --------------------------------------------------------------------------- + +class TestRunWriteTiledJob: + def test_missing_cache_entry_reports_error(self): + jid = export_jobs.new_job("x") + infer_jobs.run_write_tiled_job(jid, "no-such-infer-job") + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "No cached inference results" in job["error"] + + def test_non_tiled_source_is_refused(self): + infer_jobs._cache_put("infer-local", {"kind": "local", "classes": [], "label_pngs": {}}) + jid = export_jobs.new_job("x") + infer_jobs.run_write_tiled_job(jid, "infer-local") + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "not a Tiled array" in job["error"] + + def test_writes_semantic_and_class_volumes(self, monkeypatch): + from PIL import Image as PILImage + + label0 = np.zeros((4, 4), dtype=np.uint8) + label0[0:2, 0:2] = 1 + label1 = np.zeros((4, 4), dtype=np.uint8) + label1[2:4, 2:4] = 2 + + def _png(arr): + buf = io.BytesIO() + PILImage.fromarray(arr, mode="L").save(buf, format="PNG") + return buf.getvalue() + + infer_jobs._cache_put( + "infer-tiled", + { + "kind": "tiled", "source": "browse/sample", "server_uri": "http://x", + "classes": [{"classId": 1, "label": "a", "color": "#f00"}, {"classId": 2, "label": "b", "color": "#0f0"}], + "label_pngs": {0: _png(label0), 1: _png(label1)}, + }, + ) + + captured = {} + + def fake_write_masks_to_tiled(source, server_uri, volumes, classes, container_suffix=""): + captured.update(source=source, server_uri=server_uri, volumes=volumes, container_suffix=container_suffix) + return {"path": "browse/sample__masks_deep"} + + import tiled_mask_sync + + monkeypatch.setattr(tiled_mask_sync, "write_masks_to_tiled", fake_write_masks_to_tiled) + + jid = export_jobs.new_job("x") + infer_jobs.run_write_tiled_job(jid, "infer-tiled") + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"] == {"path": "browse/sample__masks_deep"} + assert captured["container_suffix"] == "_deep" + assert captured["volumes"]["semantic"].shape == (2, 4, 4) + assert np.array_equal(captured["volumes"]["class_vols"]["a"][0], (label0 == 1) * 255) + + def test_write_failure_is_reported_as_a_job_error(self, monkeypatch): + infer_jobs._cache_put( + "infer-tiled-fail", + { + "kind": "tiled", "source": "browse/sample", "server_uri": None, + "classes": [{"classId": 1, "label": "a", "color": "#f00"}], + "label_pngs": {0: _blank_label_png()}, + }, + ) + import tiled_mask_sync + + def boom(*a, **k): + raise RuntimeError("tiled write blew up") + + monkeypatch.setattr(tiled_mask_sync, "write_masks_to_tiled", boom) + jid = export_jobs.new_job("x") + infer_jobs.run_write_tiled_job(jid, "infer-tiled-fail") + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "tiled write blew up" in job["error"] + + +def _blank_label_png() -> bytes: + from PIL import Image as PILImage + + buf = io.BytesIO() + PILImage.fromarray(np.zeros((4, 4), dtype=np.uint8), mode="L").save(buf, format="PNG") + return buf.getvalue() diff --git a/backend/tests/test_ingest_scan.py b/backend/tests/test_ingest_scan.py new file mode 100644 index 0000000..c5f5dee --- /dev/null +++ b/backend/tests/test_ingest_scan.py @@ -0,0 +1,197 @@ +"""Tests for scan_and_register_image_stacks — bulk auto-discovery of image +folders under a mounted directory, mirroring zarr_source's +scan_and_register_zarrs for the non-Zarr case. register_zarr/run_ingest_job +themselves are already covered elsewhere (test_zarr_source.py, +test_ingest_conflict.py); these tests are about candidate discovery, +temp-file safety, and result aggregation. +""" + +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import pytest +from fastapi import HTTPException +from PIL import Image + +import ingest + + +class FakeNode: + """Minimal stand-in for a Tiled container client (see test_ingest_conflict.py).""" + + def __init__(self, children: dict | None = None, metadata: dict | None = None) -> None: + self._children: dict = dict(children or {}) + self.metadata: dict = dict(metadata or {}) + self.written: dict = {} + + def __getitem__(self, key: str): + if key not in self._children: + raise KeyError(key) + return self._children[key] + + def keys(self) -> list[str]: + return list(self._children) + + def create_container(self, key: str, metadata: dict | None = None) -> "FakeNode": + node = FakeNode(metadata=metadata) + self._children[key] = node + return node + + def update_metadata(self, metadata: dict | None = None) -> None: + self.metadata = dict(metadata or {}) + + def write_array(self, arr, key=None, metadata=None, dims=None) -> None: + self._children[key] = "array" + self.written[key] = (arr, metadata) + + def delete_contents(self, keys=None, recursive: bool = False, external_only: bool = True) -> None: + for k in [keys] if isinstance(keys, str) else keys: + self._children.pop(k, None) + + +@pytest.fixture +def fake_client(monkeypatch): + root = FakeNode({"browse": FakeNode()}) + monkeypatch.setattr(ingest, "api_key_for_uri", lambda uri: None) + monkeypatch.setattr(ingest, "get_tiled_client", lambda uri, key=None: root) + return root + + +def _make_stack(root: Path, name: str, n: int = 3, ext: str = ".tif") -> Path: + """A folder of *n* tiny real image files (real bytes, no mocks).""" + stack = root / name + stack.mkdir() + for i in range(n): + Image.fromarray(np.zeros((4, 4), dtype=np.uint8)).save(stack / f"{i:03d}{ext}") + return stack + + +class TestIsImageStackDir: + def test_two_or_more_images_qualifies(self, tmp_path: Path) -> None: + stack = _make_stack(tmp_path, "s1", n=2) + assert ingest._is_image_stack_dir(stack) is True + + def test_a_single_image_does_not_qualify(self, tmp_path: Path) -> None: + stack = _make_stack(tmp_path, "s1", n=1) + assert ingest._is_image_stack_dir(stack) is False + + def test_a_plain_file_does_not_qualify(self, tmp_path: Path) -> None: + f = tmp_path / "readme.txt" + f.write_text("hi") + assert ingest._is_image_stack_dir(f) is False + + def test_unsupported_extensions_do_not_count(self, tmp_path: Path) -> None: + d = tmp_path / "junk" + d.mkdir() + (d / "a.txt").write_text("x") + (d / "b.txt").write_text("y") + assert ingest._is_image_stack_dir(d) is False + + +class TestScanAndRegisterImageStacks: + def test_finds_and_ingests_top_level_stacks_only(self, tmp_path: Path, fake_client) -> None: + _make_stack(tmp_path, "stack_a", n=3) + _make_stack(tmp_path, "stack_b", n=2) + (tmp_path / "not_a_stack").mkdir() + # A stack's OWN internals must never be treated as a second candidate. + (tmp_path / "stack_a" / "nested_stack").mkdir() + + result = ingest.scan_and_register_image_stacks(None, str(tmp_path), "browse") + assert result["scanned"] == 2 + assert sorted(r["name"] for r in result["registered"]) == ["stack_a", "stack_b"] + assert result["skipped"] == [] + assert result["errors"] == [] + + # Real per-slice ingest actually happened (not a no-op / dry run). + browse = fake_client["browse"] + assert "stack_a" in browse.keys() + assert len(browse["stack_a"].written) == 3 + + def test_a_zarr_directory_is_never_treated_as_an_image_stack(self, tmp_path: Path, fake_client) -> None: + zarr_dir = tmp_path / "vol.zarr" + zarr_dir.mkdir() + (zarr_dir / ".zarray").write_text("{}") + # Even if it happens to also hold loose image files at its root. + Image.fromarray(np.zeros((4, 4), dtype=np.uint8)).save(zarr_dir / "a.tif") + Image.fromarray(np.zeros((4, 4), dtype=np.uint8)).save(zarr_dir / "b.tif") + + result = ingest.scan_and_register_image_stacks(None, str(tmp_path), "browse") + assert result["scanned"] == 0 + assert result["registered"] == [] + + def test_does_not_delete_the_original_source_files(self, tmp_path: Path, fake_client) -> None: + stack = _make_stack(tmp_path, "stack_a", n=3) + originals = sorted(stack.iterdir()) + assert len(originals) == 3 + + ingest.scan_and_register_image_stacks(None, str(tmp_path), "browse") + + # The whole point of copying to a real temp dir first: run_ingest_job + # unlinks its inputs when done, and must never touch the originals. + assert sorted(stack.iterdir()) == originals + for f in originals: + assert f.exists() + + def test_skips_an_already_registered_stack_by_default(self, tmp_path: Path, fake_client) -> None: + _make_stack(tmp_path, "stack_a", n=2) + fake_client["browse"].create_container("stack_a", metadata={"source_format": "image-stack"}) + + result = ingest.scan_and_register_image_stacks(None, str(tmp_path), "browse") + assert result["registered"] == [] + assert result["skipped"] == ["stack_a"] + + def test_replace_on_conflict_re_ingests(self, tmp_path: Path, fake_client) -> None: + _make_stack(tmp_path, "stack_a", n=2) + fake_client["browse"].create_container("stack_a", metadata={"source_format": "image-stack"}) + + result = ingest.scan_and_register_image_stacks(None, str(tmp_path), "browse", on_conflict="replace") + assert [r["name"] for r in result["registered"]] == ["stack_a"] + assert result["skipped"] == [] + + def test_a_different_kind_collision_is_shadowed_not_skipped(self, tmp_path: Path, fake_client) -> None: + _make_stack(tmp_path, "stack_a", n=2) + # Same stem, but an unrelated Zarr registration got there first. + fake_client["browse"].create_container("stack_a", metadata={"source_format": "zarr"}) + + result = ingest.scan_and_register_image_stacks(None, str(tmp_path), "browse") + assert result["registered"] == [] + assert result["skipped"] == [] + assert result["shadowed"] == [ + {"name": "stack_a", "key": "stack_a", "existing_kind": "zarr", "suggested_key": "stack_a_images"} + ] + + def test_renames_lets_a_shadowed_candidate_ingest_under_an_alternate_key( + self, tmp_path: Path, fake_client + ) -> None: + _make_stack(tmp_path, "stack_a", n=2) + fake_client["browse"].create_container("stack_a", metadata={"source_format": "zarr"}) + + result = ingest.scan_and_register_image_stacks( + None, str(tmp_path), "browse", renames={"stack_a": "stack_a_images"} + ) + assert result["shadowed"] == [] + assert [r["key"] for r in result["registered"]] == ["stack_a_images"] + assert fake_client["browse"]["stack_a_images"].written + + def test_rejects_relative_scan_root(self) -> None: + with pytest.raises(HTTPException) as exc: + ingest.scan_and_register_image_stacks(None, "relative/dir", "browse") + assert exc.value.status_code == 400 + + def test_rejects_missing_scan_root(self, tmp_path: Path) -> None: + with pytest.raises(HTTPException) as exc: + ingest.scan_and_register_image_stacks(None, str(tmp_path / "nope"), "browse") + assert exc.value.status_code == 404 + + def test_empty_directory_scans_cleanly_with_nothing_found(self, tmp_path: Path, fake_client) -> None: + result = ingest.scan_and_register_image_stacks(None, str(tmp_path), "browse") + assert result == {"scanned": 0, "registered": [], "skipped": [], "shadowed": [], "errors": []} + + def test_invalid_on_conflict_falls_back_to_skip(self, tmp_path: Path, fake_client) -> None: + _make_stack(tmp_path, "stack_a", n=2) + fake_client["browse"].create_container("stack_a", metadata={"source_format": "image-stack"}) + + result = ingest.scan_and_register_image_stacks(None, str(tmp_path), "browse", on_conflict="bogus") + assert result["skipped"] == ["stack_a"] diff --git a/backend/tests/test_ipred_batch_jobs.py b/backend/tests/test_ipred_batch_jobs.py new file mode 100644 index 0000000..641c46e --- /dev/null +++ b/backend/tests/test_ipred_batch_jobs.py @@ -0,0 +1,185 @@ +"""Tests for the parallelized iPred volume-apply job (ipred_batch_jobs.py). + +Uses real export_jobs (in-memory registry, no I/O) and monkeypatches +ipred_client_mod so no real HTTP/ipred service is needed. +""" +from __future__ import annotations + +import threading +import time +from contextlib import contextmanager + +import pytest + +import export_jobs +import ipred_batch_jobs +import ipred_client as ipred_client_mod + + +@contextmanager +def _fake_shared_client(): + yield object() # never actually used for HTTP in these tests + + +@pytest.fixture(autouse=True) +def fake_shared_client(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "new_shared_client", _fake_shared_client) + + +def test_one_failing_slice_does_not_abort_the_rest(monkeypatch: pytest.MonkeyPatch) -> None: + def fake_preprocess(*, session_id, feature_setup_id, composition_id, slice_index, client): + if slice_index == 2: + raise RuntimeError("boom") + return {"feature_id": f"fid-{slice_index}"} + + def fake_infer(*, session_id, model_id, feature_id, alpha, store_probabilities, client): + assert store_probabilities is False # volume-apply must never persist proba.npy + return {"run_id": f"run-{feature_id}"} + + deleted: list[str] = [] + + def fake_delete(feature_id, *, client): + deleted.append(feature_id) + return {"deleted": True} + + monkeypatch.setattr(ipred_client_mod, "preprocess", fake_preprocess) + monkeypatch.setattr(ipred_client_mod, "infer", fake_infer) + monkeypatch.setattr(ipred_client_mod, "delete_feature_bank", fake_delete) + + jid = export_jobs.new_job("") + ipred_batch_jobs.run_ipred_volume_apply_job( + jid, + session_id="s1", + model_id="m1", + slice_indices=[0, 1, 2, 3, 4], + composition_id=None, + feature_setup_id=None, + alpha=0.05, + ) + + job = export_jobs.get_job(jid) + assert job["state"] == "done" + result = job["result"] + assert sorted(int(k) for k in result["runs"]) == [0, 1, 3, 4] + assert result["errors"] == [{"slice": 2, "error": "boom"}] + assert result["cancelled"] is False + # No feature bank leaked, including for the failed slice (preprocess for + # slice 2 never succeeded, so there's nothing of its to delete). + assert sorted(deleted) == ["fid-0", "fid-1", "fid-3", "fid-4"] + + +def test_feature_bank_is_released_even_when_infer_fails(monkeypatch: pytest.MonkeyPatch) -> None: + def fake_preprocess(*, session_id, feature_setup_id, composition_id, slice_index, client): + return {"feature_id": f"fid-{slice_index}"} + + def fake_infer(*, session_id, model_id, feature_id, alpha, store_probabilities, client): + raise RuntimeError("infer exploded") + + deleted: list[str] = [] + monkeypatch.setattr(ipred_client_mod, "preprocess", fake_preprocess) + monkeypatch.setattr(ipred_client_mod, "infer", fake_infer) + monkeypatch.setattr( + ipred_client_mod, "delete_feature_bank", + lambda feature_id, *, client: deleted.append(feature_id), + ) + + jid = export_jobs.new_job("") + ipred_batch_jobs.run_ipred_volume_apply_job( + jid, session_id="s1", model_id="m1", slice_indices=[0, 1], + composition_id=None, feature_setup_id=None, alpha=0.05, + ) + + job = export_jobs.get_job(jid) + # Both slices failed (infer always raises) -> job reports an error state, + # but the orphaned-bank cleanup must still have run for each. + assert job["state"] == "error" + assert sorted(deleted) == ["fid-0", "fid-1"] + + +def test_cancellation_stops_submitting_new_slices(monkeypatch: pytest.MonkeyPatch) -> None: + started: list[int] = [] + lock = threading.Lock() + + def fake_preprocess(*, session_id, feature_setup_id, composition_id, slice_index, client): + with lock: + started.append(slice_index) + # First slice requests cancellation while still "in flight", so the + # pool must not keep refilling after it's observed. + if slice_index == 0: + export_jobs.request_cancel(jid_holder["jid"]) + time.sleep(0.05) + return {"feature_id": f"fid-{slice_index}"} + + def fake_infer(*, session_id, model_id, feature_id, alpha, store_probabilities, client): + return {"run_id": f"run-{feature_id}"} + + monkeypatch.setattr(ipred_client_mod, "preprocess", fake_preprocess) + monkeypatch.setattr(ipred_client_mod, "infer", fake_infer) + monkeypatch.setattr(ipred_client_mod, "delete_feature_bank", lambda feature_id, *, client: None) + monkeypatch.setenv("IPRED_APPLY_CONCURRENCY", "1") # deterministic: one slice at a time + + jid_holder: dict[str, str] = {} + jid = export_jobs.new_job("") + jid_holder["jid"] = jid + + ipred_batch_jobs.run_ipred_volume_apply_job( + jid, session_id="s1", model_id="m1", slice_indices=[0, 1, 2, 3, 4], + composition_id=None, feature_setup_id=None, alpha=0.05, + ) + + job = export_jobs.get_job(jid) + assert job["result"]["cancelled"] is True + # With concurrency=1, cancellation is observed right after slice 0 + # completes — no later slice should ever have started. + assert started == [0] + + +def test_result_is_published_incrementally_while_running(monkeypatch: pytest.MonkeyPatch) -> None: + """Regression test: result.runs must be visible WHILE the job is still + `running`, not only once the whole volume finishes — otherwise a live + per-slice preview in the UI (Annotate's PixelClassifierPanel) has nothing + to show for slices that already predicted, no matter how far along the + job is (this is exactly what a user reported live: slice 2 of a 690-slice + apply job stayed blank even though 11 slices had already predicted).""" + def fake_preprocess(*, session_id, feature_setup_id, composition_id, slice_index, client): + time.sleep(0.03) + return {"feature_id": f"fid-{slice_index}"} + + def fake_infer(*, session_id, model_id, feature_id, alpha, store_probabilities, client): + return {"run_id": f"run-{feature_id}"} + + monkeypatch.setattr(ipred_client_mod, "preprocess", fake_preprocess) + monkeypatch.setattr(ipred_client_mod, "infer", fake_infer) + monkeypatch.setattr(ipred_client_mod, "delete_feature_bank", lambda feature_id, *, client: None) + monkeypatch.setenv("IPRED_APPLY_CONCURRENCY", "2") + + n_slices = 8 + jid = export_jobs.new_job("") + thread = threading.Thread( + target=ipred_batch_jobs.run_ipred_volume_apply_job, + kwargs=dict( + jid=jid, session_id="s1", model_id="m1", slice_indices=list(range(n_slices)), + composition_id=None, feature_setup_id=None, alpha=0.05, + ), + daemon=True, + ) + thread.start() + + saw_partial_progress = False + deadline = time.monotonic() + 10.0 + while time.monotonic() < deadline: + job = export_jobs.get_job(jid) + if job["state"] == "running" and job.get("result"): + n = len(job["result"].get("runs", {})) + if 0 < n < n_slices: + saw_partial_progress = True + break + if job["state"] in ("done", "error"): + break + time.sleep(0.005) + thread.join(timeout=10.0) + + assert saw_partial_progress, "expected to observe a partial (running) result with some but not all slices done" + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert len(job["result"]["runs"]) == n_slices diff --git a/backend/tests/test_ipred_client.py b/backend/tests/test_ipred_client.py new file mode 100644 index 0000000..a9a26c6 --- /dev/null +++ b/backend/tests/test_ipred_client.py @@ -0,0 +1,325 @@ +"""Tests for ipred_client.py — a real httpx.Client is exercised end to end +via httpx.MockTransport (no real network), so ipred_url()/_client()/_use_client() +and every wrapper function's request construction + response parsing are all +genuinely tested, not mocked away.""" +from __future__ import annotations + +import types + +import httpx +import pytest + +import ipred_client + + +class RequestLog: + def __init__(self): + self.requests: list[httpx.Request] = [] + + +def _install_handler(monkeypatch: pytest.MonkeyPatch, handler): + """Patch the `httpx` name inside ipred_client's module namespace so every + httpx.Client(...) constructed by the module routes through a MockTransport + calling `handler`, while everything else about httpx.Client stays real.""" + real_client_cls = httpx.Client + + def fake_client(**kwargs): + kwargs["transport"] = httpx.MockTransport(handler) + return real_client_cls(**kwargs) + + fake_httpx = types.SimpleNamespace(Client=fake_client) + monkeypatch.setattr(ipred_client, "httpx", fake_httpx) + + +def _json_handler(log: RequestLog, routes: dict): + def handler(request: httpx.Request) -> httpx.Response: + log.requests.append(request) + key = (request.method, request.url.path) + if key not in routes: + return httpx.Response(404, json={"detail": "no route"}) + body, status = routes[key] + if isinstance(body, (bytes, bytearray)): + return httpx.Response(status, content=body) + return httpx.Response(status, json=body) + return handler + + +# --------------------------------------------------------------------------- +# ipred_url +# --------------------------------------------------------------------------- + +class TestIpredUrl: + def test_defaults_when_no_env_vars_set(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("IPRED_URL", raising=False) + monkeypatch.delenv("CLF_ENGINE_URL", raising=False) + assert ipred_client.ipred_url() == ipred_client.DEFAULT_IPRED_URL + + def test_ipred_url_env_var_wins(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("IPRED_URL", "http://a:1/") + monkeypatch.setenv("CLF_ENGINE_URL", "http://b:2/") + assert ipred_client.ipred_url() == "http://a:1" + + def test_legacy_clf_engine_url_used_when_ipred_url_unset(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("IPRED_URL", raising=False) + monkeypatch.setenv("CLF_ENGINE_URL", "http://legacy:9/") + assert ipred_client.ipred_url() == "http://legacy:9" + + def test_trailing_slash_stripped(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("IPRED_URL", "http://x:1///") + assert ipred_client.ipred_url() == "http://x:1" + + +# --------------------------------------------------------------------------- +# health / open_session / setups / trainers / modules / compositions +# --------------------------------------------------------------------------- + +class TestSimpleGetPostWrappers: + def test_health(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/health"): ({"status": "ok"}, 200)})) + assert ipred_client.health() == {"status": "ok"} + + def test_open_session_posts_expected_body(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/sessions"): ({"session_id": "s1"}, 200)})) + result = ipred_client.open_session(kind="tiled", source="foo", server_uri="http://t", root=None) + assert result == {"session_id": "s1"} + sent = log.requests[0] + assert sent.method == "POST" + import json + body = json.loads(sent.content) + assert body == {"kind": "tiled", "source": "foo", "server_uri": "http://t", "root": None} + + def test_list_setups_defaults_to_empty_list(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/setups"): ({}, 200)})) + assert ipred_client.list_setups() == [] + + def test_list_setups_returns_setups_key(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/setups"): ({"setups": [{"id": 1}]}, 200)})) + assert ipred_client.list_setups() == [{"id": 1}] + + def test_get_setup(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/setups/abc"): ({"id": "abc"}, 200)})) + assert ipred_client.get_setup("abc") == {"id": "abc"} + + def test_upsert_setup(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/setups"): ({"id": "new"}, 200)})) + assert ipred_client.upsert_setup({"name": "x"}) == {"id": "new"} + + def test_list_trainers(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/trainers"): ({"trainers": ["catboost"]}, 200)})) + assert ipred_client.list_trainers() == ["catboost"] + + def test_list_modules(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/modules"): ({"modules": [{"id": "m1"}]}, 200)})) + assert ipred_client.list_modules() == [{"id": "m1"}] + + def test_list_compositions(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/compositions"): ({"compositions": []}, 200)})) + assert ipred_client.list_compositions() == [] + + def test_get_composition(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/compositions/c1"): ({"id": "c1"}, 200)})) + assert ipred_client.get_composition("c1") == {"id": "c1"} + + def test_upsert_composition(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/compositions"): ({"id": "c2"}, 200)})) + assert ipred_client.upsert_composition({"name": "y"}) == {"id": "c2"} + + def test_preview_composition(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/compositions/preview"): ({"preview": True}, 200)})) + assert ipred_client.preview_composition({"x": 1}) == {"preview": True} + + def test_upload_session_array(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/sessions/s1/arrays"): ({"ok": True}, 200)})) + assert ipred_client.upload_session_array("s1", {"a": 1}) == {"ok": True} + + +# --------------------------------------------------------------------------- +# preprocess / delete_feature_bank (accept an optional shared client) +# --------------------------------------------------------------------------- + +class TestSharedClientParam: + def test_preprocess_builds_minimal_body(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/preprocess"): ({"ok": 1}, 200)})) + ipred_client.preprocess(session_id="s1") + import json + body = json.loads(log.requests[0].content) + assert body == {"session_id": "s1", "slice_index": 0} + + def test_preprocess_includes_optional_fields_when_given(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/preprocess"): ({"ok": 1}, 200)})) + ipred_client.preprocess( + session_id="s1", feature_setup_id="f1", composition_id="c1", + slice_index=5, array_ref="ref1", + ) + import json + body = json.loads(log.requests[0].content) + assert body == { + "session_id": "s1", "slice_index": 5, "composition_id": "c1", + "feature_setup_id": "f1", "array_ref": "ref1", + } + + def test_preprocess_reuses_a_passed_in_client_without_closing_it(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/preprocess"): ({"ok": 1}, 200)})) + shared = ipred_client.new_shared_client() + try: + ipred_client.preprocess(session_id="s1", client=shared) + assert shared.is_closed is False + # A second call on the same shared client still works — proves it + # wasn't closed by _use_client after the first call. + ipred_client.preprocess(session_id="s2", client=shared) + assert len(log.requests) == 2 + finally: + shared.close() + + def test_delete_feature_bank(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("DELETE", "/features/f1"): ({"deleted": True}, 200)})) + assert ipred_client.delete_feature_bank("f1") == {"deleted": True} + + +# --------------------------------------------------------------------------- +# byte-returning endpoints +# --------------------------------------------------------------------------- + +class TestByteEndpoints: + def test_feature_channel_bytes(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/features/f1/channels/2"): (b"\x89PNG", 200)})) + assert ipred_client.feature_channel_bytes("f1", 2) == b"\x89PNG" + + def test_run_commit_png(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/runs/r1/commit.png"): (b"commitbytes", 200)})) + assert ipred_client.run_commit_png("r1") == b"commitbytes" + + def test_run_status_png(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/runs/r1/status.png"): (b"statusbytes", 200)})) + assert ipred_client.run_status_png("r1") == b"statusbytes" + + def test_run_proba_png(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/runs/r1/proba/3.png"): (b"probabytes", 200)})) + assert ipred_client.run_proba_png("r1", 3) == b"probabytes" + + def test_manifold_heatmap_png(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/manifold/m1/heatmap.png"): (b"heatbytes", 200)})) + assert ipred_client.manifold_heatmap_png("m1") == b"heatbytes" + + +# --------------------------------------------------------------------------- +# train / train_multi / infer / rethreshold / threshold_class_map / manifold_sample +# --------------------------------------------------------------------------- + +class TestMlWrappers: + def test_train_fills_default_trainer_and_empty_config(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/train"): ({"model_id": "m1"}, 200)})) + result = ipred_client.train(session_id="s1", shapes=[{"a": 1}]) + assert result == {"model_id": "m1"} + import json + body = json.loads(log.requests[0].content) + assert body["trainer_id"] == "catboost" + assert body["config"] == {} + + def test_train_multi(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/train/multi"): ({"model_id": "m2"}, 200)})) + result = ipred_client.train_multi( + session_id="s1", slices={"0": []}, feature_ids={"0": "f1"}, + ) + assert result == {"model_id": "m2"} + + def test_infer_defaults(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/infer"): ({"run_id": "r1"}, 200)})) + result = ipred_client.infer(session_id="s1") + assert result == {"run_id": "r1"} + import json + body = json.loads(log.requests[0].content) + assert body["store_probabilities"] is True + assert body["alpha"] == 0.05 + + def test_infer_store_probabilities_false_for_batch_apply(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/infer"): ({"run_id": "r2"}, 200)})) + ipred_client.infer(session_id="s1", store_probabilities=False) + import json + body = json.loads(log.requests[0].content) + assert body["store_probabilities"] is False + + def test_rethreshold(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/rethreshold"): ({"ok": 1}, 200)})) + assert ipred_client.rethreshold(session_id="s1", alpha=0.1, run_id="r1") == {"ok": 1} + + def test_threshold_class_map(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/runs/r1/threshold-class"): ({"ok": 1}, 200)})) + result = ipred_client.threshold_class_map("r1", class_id=2, threshold=0.5) + assert result == {"ok": 1} + import json + body = json.loads(log.requests[0].content) + assert body == {"class_id": 2, "threshold": 0.5} + + def test_manifold_sample(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("POST", "/manifold/sample"): ({"sample_id": "sm1"}, 200)})) + assert ipred_client.manifold_sample({"k": "v"}) == {"sample_id": "sm1"} + + +# --------------------------------------------------------------------------- +# error propagation +# --------------------------------------------------------------------------- + +class TestErrorPropagation: + def test_http_error_status_raises(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/health"): ({"error": "down"}, 500)})) + with pytest.raises(httpx.HTTPStatusError): + ipred_client.health() + + def test_missing_route_also_raises(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {})) + with pytest.raises(httpx.HTTPStatusError): + ipred_client.list_setups() + + +# --------------------------------------------------------------------------- +# new_shared_client / _use_client +# --------------------------------------------------------------------------- + +class TestSharedClientLifecycle: + def test_new_shared_client_is_open_and_usable(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/health"): ({"status": "ok"}, 200)})) + client = ipred_client.new_shared_client() + try: + assert client.is_closed is False + finally: + client.close() + + def test_use_client_without_explicit_client_opens_and_closes_its_own(self, monkeypatch: pytest.MonkeyPatch): + log = RequestLog() + _install_handler(monkeypatch, _json_handler(log, {("GET", "/health"): ({"status": "ok"}, 200)})) + with ipred_client._use_client(None, timeout=1.0) as c: + assert c.is_closed is False + assert c.is_closed is True diff --git a/backend/tests/test_ipred_routes.py b/backend/tests/test_ipred_routes.py new file mode 100644 index 0000000..b31a49e --- /dev/null +++ b/backend/tests/test_ipred_routes.py @@ -0,0 +1,423 @@ +"""Tests for the /api/ipred/* proxy — mocks ipred_client, never hits :8003.""" +from __future__ import annotations + +import asyncio +import time + +import httpx +import pytest +from httpx import ASGITransport, AsyncClient + +import export_jobs +import ipred_client as ipred_client_mod +from annotation_server import app + + +@pytest.mark.asyncio +async def test_health_proxies_ok(monkeypatch: pytest.MonkeyPatch) -> None: + """A healthy ipred returns its body verbatim.""" + monkeypatch.setattr(ipred_client_mod, "health", lambda: {"status": "ok"}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/health") + assert response.status_code == 200 + assert response.json() == {"status": "ok"} + + +@pytest.mark.asyncio +async def test_health_reports_503_when_ipred_unreachable(monkeypatch: pytest.MonkeyPatch) -> None: + """A down/missing ipred surfaces as 503, not a 500 or crash.""" + + def _raise() -> dict: + raise httpx.ConnectError("boom") + + monkeypatch.setattr(ipred_client_mod, "health", _raise) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/health") + assert response.status_code == 503 + assert "unreachable" in response.json()["detail"] + + +@pytest.mark.asyncio +async def test_upstream_http_error_status_is_preserved(monkeypatch: pytest.MonkeyPatch) -> None: + """A 4xx from ipred itself is forwarded with the same status + detail.""" + + def _raise() -> dict: + req = httpx.Request("GET", "http://127.0.0.1:8003/modules") + resp = httpx.Response(422, json={"detail": "bad params"}, request=req) + raise httpx.HTTPStatusError("bad", request=req, response=resp) + + monkeypatch.setattr(ipred_client_mod, "list_modules", _raise) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/modules") + assert response.status_code == 422 + assert response.json()["detail"] == {"detail": "bad params"} + + +@pytest.mark.asyncio +async def test_open_session_forwards_body(monkeypatch: pytest.MonkeyPatch) -> None: + """POST /sessions passes kind/source/server_uri/root through to the client.""" + captured: dict = {} + + def _open_session(**kwargs: object) -> dict: + captured.update(kwargs) + return {"session_id": "s1", "project_id": "p1"} + + monkeypatch.setattr(ipred_client_mod, "open_session", _open_session) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/sessions", + json={"kind": "local", "source": "foo.tiff"}, + ) + assert response.status_code == 200 + assert response.json() == {"session_id": "s1", "project_id": "p1"} + assert captured == { + "kind": "local", + "source": "foo.tiff", + "server_uri": None, + "root": None, + } + + +@pytest.mark.asyncio +async def test_list_modules_reports_ready_flags(monkeypatch: pytest.MonkeyPatch) -> None: + """Module readiness (e.g. tomojepa without weights) passes through unchanged.""" + modules = [ + {"id": "slimsam", "ready": True, "runtime": "onnx"}, + {"id": "tomojepa", "ready": False, "runtime": "torch"}, + ] + monkeypatch.setattr(ipred_client_mod, "list_modules", lambda: modules) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/modules") + assert response.status_code == 200 + assert response.json() == {"modules": modules} + + +@pytest.mark.asyncio +async def test_run_proba_png_proxies_bytes(monkeypatch: pytest.MonkeyPatch) -> None: + """Binary PNG proxies return the raw bytes with an image content-type.""" + monkeypatch.setattr(ipred_client_mod, "run_proba_png", lambda run_id, idx: b"\x89PNG\r\n") + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/runs/run1/proba/0.png") + assert response.status_code == 200 + assert response.content == b"\x89PNG\r\n" + assert response.headers["content-type"] == "image/png" + + +async def _await_job(jid: str, timeout: float = 5.0) -> dict: + """Poll export_jobs until the (thread-backed) job finishes.""" + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + job = export_jobs.get_job(jid) + assert job is not None + if job["state"] in ("done", "error"): + return job + await asyncio.sleep(0.02) + raise AssertionError(f"job {jid} did not finish within {timeout}s") + + +@pytest.mark.asyncio +async def test_batch_train_pools_slices_and_reports_done(monkeypatch: pytest.MonkeyPatch) -> None: + """POST /batch/train preprocesses every slice, then trains once, pooled.""" + preprocessed: list[int] = [] + + def _preprocess(*, session_id, feature_setup_id, composition_id, slice_index, array_ref=None): + preprocessed.append(slice_index) + return {"feature_id": f"feat-{slice_index}", "cache_hit": False} + + captured_train: dict = {} + + def _train_multi(*, session_id, slices, feature_ids, trainer_id, config): + captured_train.update( + session_id=session_id, slices=slices, feature_ids=feature_ids, trainer_id=trainer_id + ) + return {"model_id": "m1", "class_ids": [1, 2], "n_samples": 4000} + + monkeypatch.setattr(ipred_client_mod, "preprocess", _preprocess) + monkeypatch.setattr(ipred_client_mod, "train_multi", _train_multi) + + shapes = [{"kind": "rectangle", "classId": 1, "x": 0, "y": 0, "w": 5, "h": 5}] + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/batch/train", + json={ + "session_id": "s1", + "slices": {"0": shapes, "2": shapes}, + "composition_id": "comp-skimage", + "trainer_id": "catboost", + }, + ) + assert response.status_code == 200 + jid = response.json()["job_id"] + job = await _await_job(jid) + assert job["state"] == "done" + assert job["result"] == {"model_id": "m1", "class_ids": [1, 2], "n_samples": 4000} + assert sorted(preprocessed) == [0, 2] + assert set(captured_train["feature_ids"]) == {"0", "2"} + + +@pytest.mark.asyncio +async def test_batch_apply_tolerates_one_bad_slice(monkeypatch: pytest.MonkeyPatch) -> None: + """POST /batch/apply keeps going after one slice fails and reports it in result.errors.""" + + def _preprocess(*, session_id, feature_setup_id, composition_id, slice_index, array_ref=None, client=None): + if slice_index == 1: + raise RuntimeError("boom") + return {"feature_id": f"feat-{slice_index}"} + + def _infer(*, session_id, model_id, feature_id, alpha, store_probabilities=True, client=None): + return {"run_id": f"run-{feature_id}"} + + monkeypatch.setattr(ipred_client_mod, "preprocess", _preprocess) + monkeypatch.setattr(ipred_client_mod, "infer", _infer) + monkeypatch.setattr(ipred_client_mod, "delete_feature_bank", lambda feature_id, *, client=None: None) + + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/batch/apply", + json={ + "session_id": "s1", + "model_id": "m1", + "slice_indices": [0, 1, 2], + }, + ) + assert response.status_code == 200 + jid = response.json()["job_id"] + job = await _await_job(jid) + assert job["state"] == "done" + assert job["result"]["runs"] == {"0": "run-feat-0", "2": "run-feat-2"} + assert job["result"]["errors"] == [{"slice": 1, "error": "boom"}] + + +@pytest.mark.asyncio +async def test_batch_train_rejects_empty_slices() -> None: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/batch/train", json={"session_id": "s1", "slices": {}}, + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +async def test_batch_apply_rejects_empty_slice_indices() -> None: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/batch/apply", + json={"session_id": "s1", "model_id": "m1", "slice_indices": []}, + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +async def test_generic_exception_reports_500(monkeypatch: pytest.MonkeyPatch) -> None: + """An exception that is neither an HTTPStatusError nor a ConnectError + (e.g. ipred is unreachable via a different transport failure) reports 500 + with the exception's own message rather than crashing the request.""" + def _raise() -> dict: + raise ValueError("something unexpected") + + monkeypatch.setattr(ipred_client_mod, "list_trainers", _raise) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/trainers") + assert response.status_code == 500 + assert "something unexpected" in response.json()["detail"] + + +@pytest.mark.asyncio +async def test_http_status_error_falls_back_to_text_when_not_json(monkeypatch: pytest.MonkeyPatch) -> None: + def _raise() -> dict: + req = httpx.Request("GET", "http://127.0.0.1:8003/setups") + resp = httpx.Response(500, content=b"plain text error", request=req) + raise httpx.HTTPStatusError("bad", request=req, response=resp) + + monkeypatch.setattr(ipred_client_mod, "list_setups", _raise) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/setups") + assert response.status_code == 500 + assert response.json()["detail"] == "plain text error" + + +class TestSimpleProxyRoutes: + """Each of these is a thin call-through + exception translation; one happy + path per route is enough since _ipred_http_error's branches are already + covered above and shared by all of them.""" + + @pytest.mark.asyncio + async def test_get_setup(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "get_setup", lambda setup_id: {"id": setup_id}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/setups/s1") + assert response.json() == {"id": "s1"} + + @pytest.mark.asyncio + async def test_upsert_setup_excludes_none_fields(self, monkeypatch: pytest.MonkeyPatch) -> None: + captured = {} + monkeypatch.setattr(ipred_client_mod, "upsert_setup", lambda payload: captured.update(payload) or {"id": "s1"}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/setups", json={"name": "s", "kind": "onnx"}, + ) + assert response.status_code == 200 + assert "procedure_id" not in captured + assert captured["name"] == "s" + + @pytest.mark.asyncio + async def test_list_trainers(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "list_trainers", lambda: ["catboost"]) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/trainers") + assert response.json() == {"trainers": ["catboost"]} + + @pytest.mark.asyncio + async def test_list_compositions(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "list_compositions", lambda: [{"id": "c1"}]) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/compositions") + assert response.json() == {"compositions": [{"id": "c1"}]} + + @pytest.mark.asyncio + async def test_get_composition(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "get_composition", lambda cid: {"id": cid}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/compositions/c1") + assert response.json() == {"id": "c1"} + + @pytest.mark.asyncio + async def test_upsert_composition(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "upsert_composition", lambda payload: {"id": "c2"}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/compositions", + json={"name": "comp", "nodes": [], "outputs": []}, + ) + assert response.json() == {"id": "c2"} + + @pytest.mark.asyncio + async def test_preview_composition(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "preview_composition", lambda payload: {"preview": True}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/compositions/preview", + json={"name": "comp", "nodes": [], "outputs": []}, + ) + assert response.json() == {"preview": True} + + @pytest.mark.asyncio + async def test_upload_array_injects_session_id_into_payload(self, monkeypatch: pytest.MonkeyPatch) -> None: + captured = {} + + def _upload(session_id, payload): + captured.update(payload) + return {"ok": True} + + monkeypatch.setattr(ipred_client_mod, "upload_session_array", _upload) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/sessions/s1/arrays", + json={"session_id": "ignored", "shape": [2, 2], "data_b64": "AAA="}, + ) + assert response.status_code == 200 + assert captured["session_id"] == "s1" + + @pytest.mark.asyncio + async def test_manifold_sample(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "manifold_sample", lambda payload: {"sample_id": "sm1"}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/manifold/sample", json={"feature_id": "f1"}, + ) + assert response.json() == {"sample_id": "sm1"} + + @pytest.mark.asyncio + async def test_manifold_heatmap_png(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "manifold_heatmap_png", lambda sample_id: b"heatbytes") + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/manifold/sm1/heatmap.png") + assert response.content == b"heatbytes" + assert response.headers["content-type"] == "image/png" + + @pytest.mark.asyncio + async def test_preprocess(self, monkeypatch: pytest.MonkeyPatch) -> None: + captured = {} + + def _preprocess(**kwargs): + captured.update(kwargs) + return {"feature_id": "f1"} + + monkeypatch.setattr(ipred_client_mod, "preprocess", _preprocess) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/preprocess", json={"session_id": "s1", "slice_index": 3}, + ) + assert response.json() == {"feature_id": "f1"} + assert captured["slice_index"] == 3 + + @pytest.mark.asyncio + async def test_feature_channel_png(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "feature_channel_bytes", lambda fid, idx: b"chanbytes") + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/features/f1/channels/2") + assert response.content == b"chanbytes" + + @pytest.mark.asyncio + async def test_train(self, monkeypatch: pytest.MonkeyPatch) -> None: + captured = {} + + def _train(**kwargs): + captured.update(kwargs) + return {"model_id": "m1"} + + monkeypatch.setattr(ipred_client_mod, "train", _train) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/train", json={"session_id": "s1", "shapes": []}, + ) + assert response.json() == {"model_id": "m1"} + assert captured["trainer_id"] == "catboost" + + @pytest.mark.asyncio + async def test_infer(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "infer", lambda **kwargs: {"run_id": "r1"}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/infer", json={"session_id": "s1"}, + ) + assert response.json() == {"run_id": "r1"} + + @pytest.mark.asyncio + async def test_rethreshold(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "rethreshold", lambda **kwargs: {"ok": True}) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/rethreshold", json={"session_id": "s1", "alpha": 0.1}, + ) + assert response.json() == {"ok": True} + + @pytest.mark.asyncio + async def test_run_commit_png(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "run_commit_png", lambda run_id: b"commitbytes") + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/runs/r1/commit.png") + assert response.content == b"commitbytes" + + @pytest.mark.asyncio + async def test_run_status_png(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(ipred_client_mod, "run_status_png", lambda run_id: b"statusbytes") + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/ipred/runs/r1/status.png") + assert response.content == b"statusbytes" + + @pytest.mark.asyncio + async def test_threshold_class(self, monkeypatch: pytest.MonkeyPatch) -> None: + captured = {} + + def _threshold(run_id, *, class_id, threshold): + captured.update(run_id=run_id, class_id=class_id, threshold=threshold) + return {"ok": True} + + monkeypatch.setattr(ipred_client_mod, "threshold_class_map", _threshold) + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/ipred/runs/r1/threshold-class", json={"class_id": 2, "threshold": 0.7}, + ) + assert response.json() == {"ok": True} + assert captured == {"run_id": "r1", "class_id": 2, "threshold": 0.7} diff --git a/backend/tests/test_local_fs.py b/backend/tests/test_local_fs.py index dcc1444..e6bcb93 100644 --- a/backend/tests/test_local_fs.py +++ b/backend/tests/test_local_fs.py @@ -1,64 +1,208 @@ -"""Tests for local filesystem sandboxing.""" - +"""Tests for local_fs.py — real filesystem (tmp_path), no mocks. Sets +`local_fs._DEFAULT_ROOT` directly rather than the LOCAL_DATA_ROOT env var, +since it's computed once at import time (see MEMORY.md's import-time-constant +gotcha).""" from __future__ import annotations -import importlib -import os - +import numpy as np import pytest from fastapi import HTTPException +from PIL import Image as PILImage + +import local_fs + + +@pytest.fixture(autouse=True) +def default_root(tmp_path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(local_fs, "_DEFAULT_ROOT", tmp_path) + return tmp_path + + +# --------------------------------------------------------------------------- +# default_root +# --------------------------------------------------------------------------- + +class TestDefaultRoot: + def test_returns_the_default_root_as_a_string(self, default_root): + assert local_fs.default_root() == str(default_root) + + +# --------------------------------------------------------------------------- +# _resolve_root / _within / _safe +# --------------------------------------------------------------------------- + +class TestResolveRoot: + def test_none_falls_back_to_default_root(self, default_root): + assert local_fs._resolve_root(None) == default_root + def test_explicit_root_is_expanded_and_resolved(self, tmp_path): + granted = tmp_path / "granted" + granted.mkdir() + assert local_fs._resolve_root(str(granted)) == granted.resolve() -def test_path_traversal_rejected(tmp_path) -> None: - """Paths escaping LOCAL_DATA_ROOT must raise HTTP 403.""" - with __import__("unittest.mock", fromlist=["patch"]).patch.dict( - os.environ, {"LOCAL_DATA_ROOT": str(tmp_path)} - ): - import local_fs - importlib.reload(local_fs) - with pytest.raises(HTTPException) as exc_info: +class TestWithin: + def test_base_itself_is_within(self, tmp_path): + assert local_fs._within(tmp_path, tmp_path) is True + + def test_descendant_is_within(self, tmp_path): + assert local_fs._within(tmp_path, tmp_path / "a" / "b") is True + + def test_sibling_with_shared_prefix_is_not_within(self, tmp_path): + sibling = tmp_path.parent / (tmp_path.name + "2") + assert local_fs._within(tmp_path, sibling) is False + + +class TestSafe: + def test_traversal_with_explicit_root_is_403(self, tmp_path): + granted = tmp_path / "granted" + granted.mkdir() + with pytest.raises(HTTPException) as exc: + local_fs._safe("../../etc/passwd", root=str(granted)) + assert exc.value.status_code == 403 + + def test_traversal_without_root_under_default_is_403(self): + with pytest.raises(HTTPException) as exc: local_fs._safe("../../etc/passwd") - assert exc_info.value.status_code == 403 + assert exc.value.status_code == 403 + def test_relative_path_resolved_under_granted_root(self, tmp_path): + granted = tmp_path / "granted" + granted.mkdir() + result = local_fs._safe("sub/file.tif", root=str(granted)) + assert result == (granted / "sub" / "file.tif").resolve() -def test_list_dir_returns_entries(tmp_path) -> None: - """list_dir should return file entries for an existing directory.""" - (tmp_path / "sample.npy").write_bytes(b"\x93NUMPY") - with __import__("unittest.mock", fromlist=["patch"]).patch.dict( - os.environ, {"LOCAL_DATA_ROOT": str(tmp_path)} - ): - import local_fs + def test_absolute_path_without_root_returned_as_is_no_sandbox(self, tmp_path): + absolute = tmp_path / "elsewhere" / "file.tif" + result = local_fs._safe(str(absolute)) + assert result == absolute.resolve() - importlib.reload(local_fs) + def test_relative_path_without_root_resolved_under_default(self, default_root): + result = local_fs._safe("sub/file.tif") + assert result == (default_root / "sub" / "file.tif").resolve() + + +# --------------------------------------------------------------------------- +# list_dir +# --------------------------------------------------------------------------- + +class TestListDir: + def test_lists_files_and_dirs_with_sizes(self, default_root): + (default_root / "a.npy").write_bytes(b"12345") + (default_root / "sub").mkdir() entries = local_fs.list_dir("") - names = [e["name"] for e in entries] - assert "sample.npy" in names - - -def test_list_dir_missing_raises_404(tmp_path) -> None: - """list_dir on a non-existent path should raise HTTP 404.""" - with __import__("unittest.mock", fromlist=["patch"]).patch.dict( - os.environ, {"LOCAL_DATA_ROOT": str(tmp_path)} - ): - import local_fs - - importlib.reload(local_fs) - with pytest.raises(HTTPException) as exc_info: - local_fs.list_dir("does_not_exist") - assert exc_info.value.status_code == 404 - - -def test_open_array_unsupported_type_raises_422(tmp_path) -> None: - """open_array on an unsupported extension should raise HTTP 422.""" - bad_file = tmp_path / "data.csv" - bad_file.write_text("a,b\n1,2\n") - with __import__("unittest.mock", fromlist=["patch"]).patch.dict( - os.environ, {"LOCAL_DATA_ROOT": str(tmp_path)} - ): - import local_fs - - importlib.reload(local_fs) - with pytest.raises(HTTPException) as exc_info: + by_name = {e["name"]: e for e in entries} + assert by_name["a.npy"]["is_dir"] is False + assert by_name["a.npy"]["size"] == 5 + assert by_name["sub"]["is_dir"] is True + assert by_name["sub"]["size"] is None + + def test_missing_root_itself_returns_empty_not_404(self, tmp_path, monkeypatch): + monkeypatch.setattr(local_fs, "_DEFAULT_ROOT", tmp_path / "does_not_exist") + assert local_fs.list_dir("") == [] + assert local_fs.list_dir(".") == [] + + def test_missing_subpath_raises_404(self): + with pytest.raises(HTTPException) as exc: + local_fs.list_dir("no/such/dir") + assert exc.value.status_code == 404 + + def test_file_instead_of_dir_raises_400(self, default_root): + (default_root / "afile.npy").write_bytes(b"x") + with pytest.raises(HTTPException) as exc: + local_fs.list_dir("afile.npy") + assert exc.value.status_code == 400 + + def test_paths_are_relative_to_root(self, default_root): + (default_root / "sub").mkdir() + (default_root / "sub" / "f.npy").write_bytes(b"x") + entries = local_fs.list_dir("sub") + assert entries[0]["path"] == "sub/f.npy" + + +# --------------------------------------------------------------------------- +# count_image_files / list_image_files +# --------------------------------------------------------------------------- + +class TestCountImageFiles: + def test_counts_only_supported_extensions_recursively(self, default_root): + (default_root / "a.tif").write_bytes(b"x") + (default_root / "b.npy").write_bytes(b"x") + (default_root / "c.txt").write_bytes(b"x") + (default_root / "sub").mkdir() + (default_root / "sub" / "d.png").write_bytes(b"x") + assert local_fs.count_image_files("") == 3 + + def test_nonexistent_path_returns_zero(self): + assert local_fs.count_image_files("nope") == 0 + + def test_file_instead_of_dir_returns_zero(self, default_root): + (default_root / "a.tif").write_bytes(b"x") + assert local_fs.count_image_files("a.tif") == 0 + + def test_extension_matching_is_case_insensitive(self, default_root): + (default_root / "A.TIF").write_bytes(b"x") + assert local_fs.count_image_files("") == 1 + + +class TestListImageFiles: + def test_returns_sorted_flat_list(self, default_root): + (default_root / "sub").mkdir() + (default_root / "sub" / "b.tif").write_bytes(b"x") + (default_root / "a.tif").write_bytes(b"x") + (default_root / "readme.md").write_bytes(b"x") + entries = local_fs.list_image_files("") + assert [e["path"] for e in entries] == ["a.tif", "sub/b.tif"] + + def test_missing_path_raises_404(self): + with pytest.raises(HTTPException) as exc: + local_fs.list_image_files("nope") + assert exc.value.status_code == 404 + + def test_file_instead_of_dir_raises_400(self, default_root): + (default_root / "a.tif").write_bytes(b"x") + with pytest.raises(HTTPException) as exc: + local_fs.list_image_files("a.tif") + assert exc.value.status_code == 400 + + +# --------------------------------------------------------------------------- +# open_array +# --------------------------------------------------------------------------- + +class TestOpenArray: + def test_missing_file_raises_404(self): + with pytest.raises(HTTPException) as exc: + local_fs.open_array("nope.npy") + assert exc.value.status_code == 404 + + def test_unsupported_extension_raises_422(self, default_root): + (default_root / "data.csv").write_text("a,b\n1,2\n") + with pytest.raises(HTTPException) as exc: local_fs.open_array("data.csv") - assert exc_info.value.status_code == 422 + assert exc.value.status_code == 422 + + def test_opens_npy_file(self, default_root): + arr = np.arange(12).reshape(3, 4).astype(np.float32) + np.save(default_root / "a.npy", arr) + result = local_fs.open_array("a.npy") + assert np.array_equal(np.asarray(result), arr) + + def test_opens_png_file(self, default_root): + img = PILImage.new("L", (5, 5), color=128) + img.save(default_root / "a.png") + result = local_fs.open_array("a.png") + assert result.shape == (5, 5) + + def test_opens_tif_file(self, default_root): + tifffile = pytest.importorskip("tifffile") + arr = np.arange(16).reshape(4, 4).astype(np.uint16) + tifffile.imwrite(default_root / "a.tif", arr) + result = local_fs.open_array("a.tif") + assert np.array_equal(np.asarray(result), arr) + + def test_corrupt_npy_raises_500(self, default_root): + (default_root / "bad.npy").write_bytes(b"not a real npy file") + with pytest.raises(HTTPException) as exc: + local_fs.open_array("bad.npy") + assert exc.value.status_code == 500 diff --git a/backend/tests/test_mask_pyramid.py b/backend/tests/test_mask_pyramid.py new file mode 100644 index 0000000..182eef8 --- /dev/null +++ b/backend/tests/test_mask_pyramid.py @@ -0,0 +1,98 @@ +"""Unit tests for mask_pyramid's pure downsampling (no Tiled I/O).""" +import numpy as np + +import tiff_stack_source as tss +from mask_pyramid import build_mask_pyramid, majority_downsample + + +def test_majority_downsample_picks_the_more_frequent_class(): + # A 2x4x4 volume, factor (1,2,2): each output voxel covers a 1x2x2 block. + # Top-left block is all class 1 except one class-2 voxel — 1 wins. + block = np.zeros((1, 4, 4), dtype=np.uint8) + block[0, 0:2, 0:2] = 1 + block[0, 0, 0] = 2 # one dissenting voxel inside that block + out = majority_downsample(block, [1, 2, 2]) + assert out.shape == (1, 2, 2) + assert out[0, 0, 0] == 1 + + +def test_majority_downsample_uniform_block(): + volume = np.full((2, 4, 4), 3, dtype=np.uint8) + out = majority_downsample(volume, [2, 2, 2]) + assert out.shape == (1, 2, 2) + assert np.all(out == 3) + + +def test_majority_downsample_never_averages(): + # Two classes, 50/50 split within a block: averaging would invent a value + # (e.g. 1.5 -> rounds to 2), which must never happen — the winner must be + # one of the actual present class ids. + block = np.zeros((1, 2, 2), dtype=np.uint8) + block[0, 0, 0] = 5 + block[0, 0, 1] = 9 + out = majority_downsample(block, [1, 1, 2]) + assert out[0, 0, 0] in (5, 9) + + +def test_majority_downsample_tie_favors_lower_class_id(): + block = np.zeros((1, 1, 2), dtype=np.uint8) + block[0, 0, 0] = 7 + block[0, 0, 1] = 3 + out = majority_downsample(block, [1, 1, 2]) + assert out[0, 0, 0] == 3 + + +def test_majority_downsample_rejects_factor_below_one(): + import pytest + + with pytest.raises(ValueError): + majority_downsample(np.zeros((2, 2, 2), dtype=np.uint8), [0, 1, 1]) + + +def test_build_mask_pyramid_scale0_is_the_original_array(): + semantic = np.zeros((4, 8, 8), dtype=np.uint8) + semantic[:, 2:6, 2:6] = 1 + levels, generated = build_mask_pyramid(semantic) + assert "scale0" in levels + assert np.array_equal(levels["scale0"], semantic) + # 8x8 is already tiny — pyramid_plan should generate nothing further. + assert generated == [] + + +def test_build_mask_pyramid_generates_coarser_levels_for_a_large_volume(): + semantic = np.zeros((4, 4096, 4096), dtype=np.uint8) + semantic[:, :2048, :2048] = 1 + levels, generated = build_mask_pyramid(semantic) + assert len(generated) > 0 + for level in generated: + assert level["path"] in levels + assert levels[level["path"]].dtype == np.uint8 + # Downsampled level must only contain class ids actually present. + assert set(np.unique(levels[level["path"]])).issubset({0, 1}) + + +def test_write_pyramid_store_scopes_by_key_not_shared(monkeypatch, tmp_path): + """Regression test for the bug where every mask sync (any dataset, iPred + or dlsia) wrote to the same on-disk path because register_mask_pyramid + always passed the literal "semantic" as write_pyramid_store's key — + Tiled's per-container metadata differed but the actual pixel data on + disk was shared and silently clobbered by whichever sync ran last. + """ + monkeypatch.setenv("VOLUME_CACHE_DIR", str(tmp_path)) + + fast = np.zeros((1, 2, 2), dtype=np.uint8) + fast[:] = 1 + deep = np.zeros((1, 2, 2), dtype=np.uint8) + deep[:] = 2 + + fast_path = tss.write_pyramid_store("sample__masks", {"scale0": fast}) + deep_path = tss.write_pyramid_store("sample__masks_deep", {"scale0": deep}) + + assert fast_path != deep_path + + import zarr + + fast_readback = np.asarray(zarr.open_group(str(fast_path), mode="r")["scale0"][:]) + deep_readback = np.asarray(zarr.open_group(str(deep_path), mode="r")["scale0"][:]) + assert np.all(fast_readback == 1) + assert np.all(deep_readback == 2) diff --git a/backend/tests/test_source_keys.py b/backend/tests/test_source_keys.py new file mode 100644 index 0000000..2bab808 --- /dev/null +++ b/backend/tests/test_source_keys.py @@ -0,0 +1,91 @@ +"""Tests for source_keys.parse_source_key.""" +from __future__ import annotations + +import pytest + +import source_keys +import tiled_config + + +@pytest.fixture() +def two_servers(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr( + tiled_config, "get_tiled_servers", + lambda: { + "local": {"uri": "http://127.0.0.1:8010"}, + "remote": {"uri": "http://example.com:9000/"}, + }, + ) + monkeypatch.setattr(source_keys, "get_tiled_servers", tiled_config.get_tiled_servers) + + +class TestLocalKeys: + def test_parses_local_prefix(self): + assert source_keys.parse_source_key("local:foo/bar.tif") == { + "kind": "local", "server_uri": None, "path": "foo/bar.tif", + } + + def test_empty_path_after_prefix(self): + assert source_keys.parse_source_key("local:") == { + "kind": "local", "server_uri": None, "path": "", + } + + +class TestTiledKeys: + def test_matches_known_server_uri(self, two_servers): + result = source_keys.parse_source_key("tiled:http://127.0.0.1:8010:browse/sample") + assert result == {"kind": "tiled", "server_uri": "http://127.0.0.1:8010", "path": "browse/sample"} + + def test_trailing_slash_on_configured_uri_is_normalized(self, two_servers): + result = source_keys.parse_source_key("tiled:http://example.com:9000:browse/x") + assert result["server_uri"] == "http://example.com:9000" + assert result["path"] == "browse/x" + + def test_prefers_longest_matching_uri(self, monkeypatch: pytest.MonkeyPatch): + # "http://x:1" is a prefix of "http://x:1/extra" — the longer match + # must win so the path split lands in the right place. + monkeypatch.setattr( + tiled_config, "get_tiled_servers", + lambda: { + "short": {"uri": "http://x:1"}, + "long": {"uri": "http://x:1/extra"}, + }, + ) + monkeypatch.setattr(source_keys, "get_tiled_servers", tiled_config.get_tiled_servers) + result = source_keys.parse_source_key("tiled:http://x:1/extra:browse/s") + assert result == {"kind": "tiled", "server_uri": "http://x:1/extra", "path": "browse/s"} + + def test_unmatched_uri_falls_back_to_default_server_with_colon(self, two_servers): + result = source_keys.parse_source_key("tiled::browse/sample") + assert result == {"kind": "tiled", "server_uri": None, "path": "browse/sample"} + + def test_unmatched_uri_falls_back_to_default_server_without_colon(self, two_servers): + result = source_keys.parse_source_key("tiled:browse/sample") + assert result == {"kind": "tiled", "server_uri": None, "path": "browse/sample"} + + def test_no_configured_servers_falls_back_to_default(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(tiled_config, "get_tiled_servers", lambda: {}) + monkeypatch.setattr(source_keys, "get_tiled_servers", tiled_config.get_tiled_servers) + result = source_keys.parse_source_key("tiled:browse/sample") + assert result == {"kind": "tiled", "server_uri": None, "path": "browse/sample"} + + def test_server_with_no_uri_configured_is_skipped(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr( + tiled_config, "get_tiled_servers", + lambda: {"broken": {}, "ok": {"uri": "http://ok:1"}}, + ) + monkeypatch.setattr(source_keys, "get_tiled_servers", tiled_config.get_tiled_servers) + result = source_keys.parse_source_key("tiled:http://ok:1:browse/s") + assert result["server_uri"] == "http://ok:1" + + +class TestUnknownKeys: + def test_unrecognized_prefix_returns_unknown_kind(self): + assert source_keys.parse_source_key("weird:thing") == { + "kind": "unknown", "server_uri": None, "path": "weird:thing", + } + + def test_empty_string(self): + assert source_keys.parse_source_key("") == { + "kind": "unknown", "server_uri": None, "path": "", + } diff --git a/backend/tests/test_thumbnails.py b/backend/tests/test_thumbnails.py new file mode 100644 index 0000000..3f1e295 --- /dev/null +++ b/backend/tests/test_thumbnails.py @@ -0,0 +1,144 @@ +"""Tests for thumbnails.py — real numpy/PIL, no mocks needed since the whole +module is pure array->PNG logic.""" +from __future__ import annotations + +from io import BytesIO + +import numpy as np +import pytest +from PIL import Image as PILImage + +import thumbnails + + +class FakeArrayNode: + def __init__(self, arr): + self._arr = arr + + def read(self): + return self._arr + + +class FakeContainer: + def __init__(self, children): + self._children = children + + def __iter__(self): + return iter(self._children) + + def __getitem__(self, key): + return self._children[key] + + +def _decode(png_bytes): + return np.array(PILImage.open(BytesIO(png_bytes))) + + +class TestRenderThumbnailRgb: + def test_uint8_rgb_array_round_trips(self): + arr = np.zeros((10, 10, 3), dtype=np.uint8) + arr[:, :, 0] = 200 + png = thumbnails.render_thumbnail(FakeArrayNode(arr), size=32) + assert png is not None + decoded = _decode(png) + assert decoded.shape[:2] == (10, 10) + + def test_rgba_array_drops_alpha(self): + arr = np.zeros((8, 8, 4), dtype=np.uint8) + png = thumbnails.render_thumbnail(FakeArrayNode(arr), size=32) + assert png is not None + decoded = _decode(png) + assert decoded.shape[:2] == (8, 8) + + def test_non_uint8_rgb_is_normalized(self): + arr = np.zeros((4, 4, 3), dtype=np.float32) + arr[0, 0, :] = 100.0 + arr[1, 1, :] = 50.0 + png = thumbnails.render_thumbnail(FakeArrayNode(arr), size=32) + assert png is not None + + def test_constant_rgb_array_yields_all_zero(self): + rgb = thumbnails._prepare_rgb(np.full((4, 4, 3), 7.0, dtype=np.float32)) + assert np.all(rgb == 0) + + +class TestRenderThumbnailIntensity: + def test_2d_array_gets_colormapped(self): + arr = np.random.default_rng(0).random((16, 16)) * 1000 + png = thumbnails.render_thumbnail(FakeArrayNode(arr), size=32) + assert png is not None + decoded = _decode(png) + assert decoded.shape[:2] == (16, 16) + + def test_constant_intensity_array_does_not_crash(self): + arr = np.full((4, 4), 5.0) + png = thumbnails.render_thumbnail(FakeArrayNode(arr), size=16) + assert png is not None + + def test_nan_and_inf_are_sanitized(self): + arr = np.array([[np.nan, np.inf], [-np.inf, 1.0]]) + png = thumbnails.render_thumbnail(FakeArrayNode(arr), size=16) + assert png is not None + + def test_negative_values_are_clamped_before_log(self): + arr = np.array([[-5.0, -1.0], [0.0, 10.0]]) + png = thumbnails.render_thumbnail(FakeArrayNode(arr), size=16) + assert png is not None + + def test_grayscale_fallback_when_no_viridis(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(thumbnails, "_VIRIDIS", None) + rgb = thumbnails._prepare_intensity(np.array([[0.0, 1.0], [2.0, 3.0]])) + assert rgb.shape == (2, 2, 3) + assert np.all(rgb[:, :, 0] == rgb[:, :, 1]) + assert np.all(rgb[:, :, 1] == rgb[:, :, 2]) + + +class TestRenderThumbnailUnsupportedShapes: + def test_1d_array_returns_none(self): + assert thumbnails.render_thumbnail(FakeArrayNode(np.zeros(10)), size=32) is None + + def test_4d_array_returns_none_after_squeeze(self): + arr = np.zeros((1, 5, 5, 2)) + assert thumbnails.render_thumbnail(FakeArrayNode(arr), size=32) is None + + def test_no_array_child_returns_none(self): + assert thumbnails.render_thumbnail(FakeContainer({}), size=32) is None + + +class TestResolveArrayNode: + def test_direct_array_node_returned_as_is(self): + node = FakeArrayNode(np.zeros((2, 2))) + assert thumbnails._resolve_array_node(node) is node + + def test_finds_first_array_child(self): + child = FakeArrayNode(np.zeros((2, 2))) + container = FakeContainer({"data": child}) + assert thumbnails._resolve_array_node(container) is child + + def test_skips_qmap_suffixed_children(self): + qmap_child = FakeArrayNode(np.zeros((2, 2))) + real_child = FakeArrayNode(np.ones((3, 3))) + container = FakeContainer({"foo_qmap": qmap_child, "bar": real_child}) + assert thumbnails._resolve_array_node(container) is real_child + + def test_non_iterable_node_returns_none(self): + assert thumbnails._resolve_array_node(object()) is None + + def test_container_with_no_array_children_returns_none(self): + container = FakeContainer({"nested": FakeContainer({})}) + assert thumbnails._resolve_array_node(container) is None + + +class TestEncodePng: + def test_downscales_to_fit_size(self): + rgb = np.zeros((100, 200, 3), dtype=np.uint8) + png = thumbnails._encode_png(rgb, size=50) + decoded = _decode(png) + assert decoded.shape[0] <= 50 + assert decoded.shape[1] <= 50 + + def test_never_upscales(self): + rgb = np.zeros((10, 10, 3), dtype=np.uint8) + png = thumbnails._encode_png(rgb, size=256) + decoded = _decode(png) + assert decoded.shape[:2] == (10, 10) diff --git a/backend/tests/test_tiff_stack_source.py b/backend/tests/test_tiff_stack_source.py new file mode 100644 index 0000000..9d405ee --- /dev/null +++ b/backend/tests/test_tiff_stack_source.py @@ -0,0 +1,429 @@ +"""Tests for exposing a TIFF directory as a streamable 3-D Zarr volume. + +The registration path needs a live Tiled server, so what is pinned here is +everything that decides *correctness* and can be checked without one: the +pyramid plan, the downsampling arithmetic, the OME-NGFF metadata shape, and the +inspection guards that stop a shuffled or unusable stack from being registered. +Fixtures write tiny TIFFs on the fly, so these run anywhere. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import numpy as np +import pytest +from fastapi import HTTPException + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import tiff_stack_source as tss # noqa: E402 + +tifffile = pytest.importorskip("tifffile") + + +def write_stack( + root: Path, n: int = 8, h: int = 16, w: int = 16, pad: int = 4, dtype=np.uint16 +) -> Path: + """A directory of zero-padded 2-D TIFFs whose value encodes the slice index.""" + root.mkdir(parents=True, exist_ok=True) + for i in range(n): + frame = np.full((h, w), i, dtype=dtype) + tifffile.imwrite(str(root / f"img_{i:0{pad}d}.tif"), frame) + return root + + +class TestPyramidPlan: + def test_no_levels_when_already_small(self): + # Nothing to generate: scale0 alone already fits a GPU texture. + assert tss.pyramid_plan((10, 64, 64), target_dim=384) == [] + + def test_first_level_brings_every_axis_under_target(self): + plan = tss.pyramid_plan((2000, 3232, 3232), target_dim=384) + assert plan, "a 3232-wide stack must generate levels" + first = plan[0] + assert max(first["shape"]) <= 384 + # Power-of-two factors keep the block mean exact. + for f in first["factor"]: + assert f & (f - 1) == 0 + + def test_factor_is_the_smallest_that_fits(self): + # 3232/8 = 404 > 384, so 8 is not enough and 16 is the answer. + assert tss.pyramid_plan((2000, 3232, 3232), target_dim=384)[0]["factor"] == [8, 16, 16] + + def test_levels_double_and_are_named_in_order(self): + plan = tss.pyramid_plan((2000, 3232, 3232), target_dim=384, levels=3) + assert [p["path"] for p in plan] == ["scale1", "scale2", "scale3"] + assert [p["factor"] for p in plan] == [[8, 16, 16], [16, 32, 32], [32, 64, 64]] + + def test_anisotropic_stack_still_gets_levels(self): + # 4 slices of 4096x4096 is an ordinary tomography shape. A single shared + # factor picked for the wide axes would ask for 4//16 = 0 slices and drop + # every level, leaving nothing renderable — which looks exactly like a + # broken viewer. Per-axis factors keep z alive. + plan = tss.pyramid_plan((4, 4096, 4096), target_dim=384, levels=3) + assert plan, "the xy axes still need downsampling" + for level in plan: + assert min(level["shape"]) >= 1 + # z has nowhere to go, so it is left alone rather than collapsed. + assert all(level["shape"][0] >= 1 for level in plan) + assert plan[0]["factor"][0] == 1 + + def test_does_not_emit_a_level_identical_to_the_one_above(self): + # Once every axis has bottomed out, another level would cost bytes and + # add no detail. + plan = tss.pyramid_plan((2, 1024, 1024), target_dim=384, levels=6) + shapes = [tuple(level["shape"]) for level in plan] + assert len(shapes) == len(set(shapes)) + + def test_degenerate_shape_yields_nothing(self): + assert tss.pyramid_plan((0, 512, 512)) == [] + + +class TestBlockMean: + def test_averages_over_the_whole_block(self): + block = np.stack([np.full((4, 4), 10.0), np.full((4, 4), 20.0)]) + out = tss.block_mean(block, 2, 2) + assert out.shape == (2, 2) + # Mean over z as well as over each 2x2 tile. + assert np.allclose(out, 15.0) + + def test_drops_partial_tiles_rather_than_averaging_them(self): + # 5 columns with fx=2: the last column cannot fill a tile. Including it + # would make that output column brighter purely because of where the + # edge fell. + block = np.ones((1, 4, 5), dtype=np.float32) + assert tss.block_mean(block, 2, 2).shape == (2, 2) + + def test_does_not_overflow_integer_inputs(self): + # uint16 sums overflow almost immediately; the mean must be exact. + block = np.full((4, 4, 4), 60000, dtype=np.uint16) + assert np.allclose(tss.block_mean(block, 2, 2), 60000.0) + + def test_rejects_zero_factor(self): + with pytest.raises(ValueError): + tss.block_mean(np.ones((1, 2, 2), dtype=np.float32), 0, 1) + + +class TestMultiscalesMetadata: + def test_nests_under_attributes(self): + # Tiled's .zattrs route returns metadata["attributes"] verbatim; anywhere + # else and the viewer reports "missing multiscales". + meta = tss.multiscales_metadata("sample", []) + assert "multiscales" in meta["attributes"] + + def test_always_includes_scale0_first(self): + datasets = tss.multiscales_metadata("s", [])["attributes"]["multiscales"][0]["datasets"] + assert [d["path"] for d in datasets] == ["scale0"] + + def test_lists_every_generated_level_with_a_scaled_transform(self): + plan = tss.pyramid_plan((2000, 3232, 3232), target_dim=384) + ms = tss.multiscales_metadata("s", plan)["attributes"]["multiscales"][0] + # scale0 is the in-place TIFF sequence; the generated levels live in the + # registered sidecar sub-group, so their paths are nested. + assert [d["path"] for d in ms["datasets"]] == ["scale0"] + [ + f"{tss.PYRAMID_KEY}/{p['path']}" for p in plan + ] + # A level downsampled by f on an axis covers f times as much space per + # voxel on that axis — anisotropic factors must survive into the transform. + for dataset, level in zip(ms["datasets"][1:], plan): + scale = dataset["coordinateTransformations"][0]["scale"] + assert scale == [float(f) for f in level["factor"]] + + def test_declares_zyx_axes(self): + ms = tss.multiscales_metadata("s", [])["attributes"]["multiscales"][0] + assert [a["name"] for a in ms["axes"]] == ["z", "y", "x"] + + +class TestInspect: + def test_describes_a_well_formed_stack(self, tmp_path): + info = tss.inspect_tiff_stack(str(write_stack(tmp_path / "scan", n=8, h=16, w=16))) + assert info["full_shape"] == [8, 16, 16] + assert info["dtype"] == "uint16" + assert info["name"] == "scan" + + def test_files_are_in_slice_order(self, tmp_path): + root = write_stack(tmp_path / "scan", n=12) + names = [p.name for p in tss.tiff_files(root)] + assert names == sorted(names) + assert names[0].endswith("0000.tif") and names[-1].endswith("0011.tif") + + def test_rejects_inconsistent_numbering(self, tmp_path): + # Unpadded names sort lexically as img_1, img_10, img_2 — registering + # that as-is would build a shuffled volume, silently. + root = tmp_path / "scan" + root.mkdir() + for i in (1, 2, 10): + tifffile.imwrite(str(root / f"img_{i}.tif"), np.zeros((4, 4), np.uint16)) + with pytest.raises(HTTPException) as excinfo: + tss.inspect_tiff_stack(str(root)) + assert excinfo.value.status_code == 422 + assert "order" in excinfo.value.detail + + def test_rejects_a_single_image(self, tmp_path): + root = write_stack(tmp_path / "scan", n=1) + with pytest.raises(HTTPException) as excinfo: + tss.inspect_tiff_stack(str(root)) + assert excinfo.value.status_code == 422 + assert "volume" in excinfo.value.detail + + def test_rejects_a_directory_with_no_tiffs(self, tmp_path): + (tmp_path / "empty").mkdir() + with pytest.raises(HTTPException) as excinfo: + tss.inspect_tiff_stack(str(tmp_path / "empty")) + assert excinfo.value.status_code == 422 + + def test_rejects_a_relative_path(self, tmp_path): + with pytest.raises(HTTPException) as excinfo: + tss.inspect_tiff_stack("relative/scan") + assert excinfo.value.status_code == 400 + + def test_reports_no_reads_when_no_pyramid_is_needed(self, tmp_path): + # A small stack needs no generated levels, so registration is instant — + # the UI should not warn about a long job. + info = tss.inspect_tiff_stack(str(write_stack(tmp_path / "scan", n=4, h=8, w=8))) + assert info["pyramid_plan"] == [] + assert info["slices_to_read"] == 0 + + +class TestPyramidStore: + """The generated levels are written as a real on-disk Zarr store. + + Not with Tiled's ``write_array``: its ``/zarr/v2`` chunk route only serves + externally-managed arrays, and the Zarr façade is exactly what the 3-D viewer + reads. See this module's docstring. + """ + + def test_writes_a_readable_zarr_v2_group(self, tmp_path, monkeypatch): + zarr = pytest.importorskip("zarr") + monkeypatch.setenv("VOLUME_CACHE_DIR", str(tmp_path / "cache")) + levels = { + "scale1": np.arange(2 * 4 * 4, dtype=np.float32).reshape(2, 4, 4), + "scale2": np.ones((1, 2, 2), dtype=np.float32), + } + store = tss.write_pyramid_store("sample__volume", levels) + + assert store.is_dir() + group = zarr.open_group(str(store), mode="r") + for name, expected in levels.items(): + assert np.array_equal(group[name][:], expected) + + def test_is_zarr_v2_not_v3(self, tmp_path, monkeypatch): + # Tiled's /zarr/v2 facade and the viewer's reader both speak v2; a v3 + # store would write `zarr.json` and the viewer would find no `.zarray`. + monkeypatch.setenv("VOLUME_CACHE_DIR", str(tmp_path / "cache")) + store = tss.write_pyramid_store( + "sample__volume", {"scale1": np.zeros((1, 2, 2), np.float32)} + ) + assert (store / ".zgroup").exists() + assert (store / "scale1" / ".zarray").exists() + + def test_chunks_one_slice_at_a_time(self, tmp_path, monkeypatch): + zarr = pytest.importorskip("zarr") + monkeypatch.setenv("VOLUME_CACHE_DIR", str(tmp_path / "cache")) + store = tss.write_pyramid_store( + "sample__volume", {"scale1": np.zeros((6, 8, 8), np.float32)} + ) + assert zarr.open_group(str(store), mode="r")["scale1"].chunks == (1, 8, 8) + + def test_replaces_a_previous_build(self, tmp_path, monkeypatch): + # Re-registering must not leave a stale volume behind for the same key. + zarr = pytest.importorskip("zarr") + monkeypatch.setenv("VOLUME_CACHE_DIR", str(tmp_path / "cache")) + tss.write_pyramid_store("sample__volume", {"scale1": np.zeros((4, 4, 4), np.float32)}) + store = tss.write_pyramid_store( + "sample__volume", {"scale1": np.ones((2, 2, 2), np.float32)} + ) + group = zarr.open_group(str(store), mode="r") + assert group["scale1"].shape == (2, 2, 2) + assert list(group.array_keys()) == ["scale1"] + + def test_honours_the_cache_dir_override(self, tmp_path, monkeypatch): + # Beamline reconstruction directories are routinely read-only, so the + # pyramid must never be written next to the source. + monkeypatch.setenv("VOLUME_CACHE_DIR", str(tmp_path / "elsewhere")) + store = tss.write_pyramid_store( + "sample__volume", {"scale1": np.zeros((1, 2, 2), np.float32)} + ) + assert str(store).startswith(str(tmp_path / "elsewhere")) + + +class TestRegisteredKey: + def test_is_a_sidecar_of_the_dataset_key(self, tmp_path): + # The per-slice container Annotate reads already owns the unsuffixed key. + # Colliding with it would make the 3-D view impossible for precisely the + # datasets it exists to serve. + root = write_stack(tmp_path / "rec20230224_sea_shell") + key = tss.registered_key(root) + assert key == f"rec20230224_sea_shell{tss.VOLUME_SUFFIX}" + + def test_matches_the_existing_sidecar_convention(self, tmp_path): + # Same shape as __v_thumbs / __masks elsewhere in the catalog. + assert tss.VOLUME_SUFFIX.startswith("__") + + +class TestBuildLevel: + def test_downsamples_a_real_stack_correctly(self, tmp_path): + # Slice i has value 2i, so a factor-2 level's slice j is the mean of + # source slices 2j and 2j+1 — i.e. 4j+1. Even values keep the expected + # result an exact integer, so this checks the averaging rather than the + # rounding (which has its own test below). + root = tmp_path / "scan" + root.mkdir() + for i in range(8): + tifffile.imwrite(str(root / f"img_{i:04d}.tif"), np.full((8, 8), 2 * i, np.uint16)) + out = tss._build_level(tss.tiff_files(root), [2, 2, 2], [4, 4, 4], np.dtype(np.uint16)) + assert out.shape == (4, 4, 4) + assert out.dtype == np.uint16 + for j in range(4): + assert np.all(out[j] == 4 * j + 1) + + def test_rounds_rather_than_truncating_on_integer_output(self, tmp_path): + # Mean of 0 and 3 is 1.5. Truncation would give 1; the cast must round. + root = tmp_path / "scan" + root.mkdir() + for i, value in enumerate((0, 3)): + tifffile.imwrite(str(root / f"img_{i:04d}.tif"), np.full((4, 4), value, np.uint16)) + out = tss._build_level(tss.tiff_files(root), [2, 1, 1], [1, 4, 4], np.dtype(np.uint16)) + assert np.all(out == 2) + + def test_counts_every_source_slice_once(self, tmp_path): + root = write_stack(tmp_path / "scan", n=8, h=8, w=8) + reads = [] + tss._build_level( + tss.tiff_files(root), [2, 2, 2], [4, 4, 4], np.dtype(np.uint16), lambda: reads.append(1) + ) + assert len(reads) == 8 + + def test_supports_anisotropic_factors(self, tmp_path): + # z left alone, xy halved — the shape a short, wide stack produces. + root = write_stack(tmp_path / "scan", n=4, h=8, w=8, dtype=np.float32) + out = tss._build_level(tss.tiff_files(root), [1, 2, 2], [4, 4, 4], np.dtype(np.float32)) + assert out.shape == (4, 4, 4) + for z in range(4): + assert np.allclose(out[z], z) # no z averaging, so values survive + + def test_preserves_float_input_without_rounding(self, tmp_path): + root = write_stack(tmp_path / "scan", n=4, h=8, w=8, dtype=np.float32) + out = tss._build_level(tss.tiff_files(root), [2, 2, 2], [2, 4, 4], np.dtype(np.float32)) + assert out.dtype == np.float32 + assert np.allclose(out[0], 0.5) # mean of slices valued 0 and 1 + + +class TestCascade: + """Coarse levels are built from the level above, not from the source. + + That is what makes registration read every source slice once instead of once + per level. It is only a valid shortcut because the factors are powers of two, + so averaging an average equals averaging the source — these tests pin that. + """ + + def test_matches_a_direct_downsample(self): + rng = np.random.default_rng(7) + source = rng.random((16, 32, 32), dtype=np.float32) + + direct = tss.downsample_array(source, [4, 4, 4]) + cascaded = tss.downsample_array(tss.downsample_array(source, [2, 2, 2]), [2, 2, 2]) + + assert direct.shape == cascaded.shape + assert np.allclose(direct, cascaded, atol=1e-5) + + def test_matches_a_direct_downsample_with_anisotropic_factors(self): + rng = np.random.default_rng(11) + source = rng.random((4, 32, 32), dtype=np.float32) + + direct = tss.downsample_array(source, [1, 4, 4]) + cascaded = tss.downsample_array(tss.downsample_array(source, [1, 2, 2]), [1, 2, 2]) + + assert np.allclose(direct, cascaded, atol=1e-5) + + def test_relative_factors_between_plan_levels_are_whole_numbers(self): + # The cascade divides each level's factor by the previous level's; a + # non-integer ratio would silently truncate and misalign the level. + plan = tss.pyramid_plan((604, 2560, 2560), target_dim=384, levels=3) + for previous, level in zip(plan, plan[1:]): + for axis in range(3): + assert level["factor"][axis] % previous["factor"][axis] == 0 + + def test_shrinks_on_every_axis_that_still_can(self): + out = tss.downsample_array(np.ones((8, 8, 8), dtype=np.float32), [2, 2, 2]) + assert out.shape == (4, 4, 4) + assert np.allclose(out, 1.0) # a constant volume stays constant + + +class FakeContainer: + """Duck-typed fake Tiled container — enough surface for preflight_tiff_stack's + navigation (`_walk`/`_child_keys`) without touching a real Tiled server.""" + + def __init__(self, children=None, metadata=None): + self._children = dict(children or {}) + self.metadata = metadata or {} + + def __iter__(self): + return iter(self._children) + + def __getitem__(self, key): + return self._children[key] + + def __len__(self): + return len(self._children) + + def keys(self): + return list(self._children.keys()) + + +class TestPreflightTiffStack: + """The registration path itself needs a live Tiled server (see this file's + module docstring), but preflight_tiff_stack never calls it — it only + inspects the local directory and navigates a fake Tiled client, so it is + fully testable without one.""" + + def test_no_collision_when_key_absent(self, tmp_path, monkeypatch): + root = write_stack(tmp_path / "scan") + client = FakeContainer({"browse": FakeContainer({})}) + monkeypatch.setattr(tss, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(tss, "api_key_for_uri", lambda uri: None) + result = tss.preflight_tiff_stack(None, str(root), "browse") + assert result["exists"] is False + assert result["existing"] is None + assert result["key"] == tss.registered_key(root) + + def test_collision_reports_external_registration(self, tmp_path, monkeypatch): + root = write_stack(tmp_path / "scan") + key = tss.registered_key(root) + existing_node = FakeContainer( + {"scale0": object()}, + metadata={"source_format": "tiff-stack-3d", "sample_name": "scan"}, + ) + client = FakeContainer({"browse": FakeContainer({key: existing_node})}) + monkeypatch.setattr(tss, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(tss, "api_key_for_uri", lambda uri: None) + result = tss.preflight_tiff_stack(None, str(root), "browse") + assert result["exists"] is True + assert result["existing"]["external"] is True + assert result["existing"]["child_count"] == 1 + assert result["existing"]["sample_name"] == "scan" + + def test_collision_with_internally_managed_data_is_not_external(self, tmp_path, monkeypatch): + root = write_stack(tmp_path / "scan") + key = tss.registered_key(root) + existing_node = FakeContainer({"img_0000.tif": object()}, metadata={}) + client = FakeContainer({"browse": FakeContainer({key: existing_node})}) + monkeypatch.setattr(tss, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(tss, "api_key_for_uri", lambda uri: None) + result = tss.preflight_tiff_stack(None, str(root), "browse") + assert result["existing"]["external"] is False + + def test_missing_target_container_reports_no_collision(self, tmp_path, monkeypatch): + root = write_stack(tmp_path / "scan") + client = FakeContainer({}) + monkeypatch.setattr(tss, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(tss, "api_key_for_uri", lambda uri: None) + result = tss.preflight_tiff_stack(None, str(root), "browse/missing") + assert result["exists"] is False + + def test_invalid_path_propagates_the_http_exception(self, monkeypatch): + with pytest.raises(HTTPException) as exc: + tss.preflight_tiff_stack(None, "not/absolute") + assert exc.value.status_code == 400 diff --git a/backend/tests/test_tiled_annotation_sync.py b/backend/tests/test_tiled_annotation_sync.py index f207608..86e3ad0 100644 --- a/backend/tests/test_tiled_annotation_sync.py +++ b/backend/tests/test_tiled_annotation_sync.py @@ -2,8 +2,11 @@ from __future__ import annotations +import pytest + +import arrays as arrays_mod from source_keys import parse_source_key -from tiled_annotation_sync import STUDIO_ANNOTATED, annotation_metadata +from tiled_annotation_sync import STUDIO_ANNOTATED, annotation_metadata, sync_annotation_metadata def test_parse_local_source_key() -> None: @@ -25,3 +28,84 @@ def test_annotation_metadata_no_when_empty() -> None: meta = annotation_metadata({"classes": [], "slices": {}}) assert meta[STUDIO_ANNOTATED] == "no" assert meta["studio_shape_count"] == 0 + + +def test_annotation_metadata_counts_classes_and_multiple_slices() -> None: + meta = annotation_metadata({ + "classes": [{"classId": 1}, {"classId": 2}], + "slices": {"0": [{"id": "s1"}, {"id": "s2"}], "1": [{"id": "s3"}]}, + }) + assert meta["studio_shape_count"] == 3 + assert meta["studio_class_count"] == 2 + + +def test_annotation_metadata_non_list_classes_counts_as_zero() -> None: + meta = annotation_metadata({"classes": "not-a-list", "slices": {}}) + assert meta["studio_class_count"] == 0 + + +def test_annotation_metadata_includes_iso_timestamp() -> None: + meta = annotation_metadata({"classes": [], "slices": {}}) + assert "T" in meta["studio_updated_at"] + + +class FakeTiledNode: + def __init__(self): + self.updates: list[dict] = [] + + def update_metadata(self, metadata): + self.updates.append(metadata) + + +class TestSyncAnnotationMetadata: + def test_local_source_is_a_no_op(self, monkeypatch: pytest.MonkeyPatch): + called = [] + monkeypatch.setattr(arrays_mod, "resolve_array", lambda *a, **k: called.append(1)) + sync_annotation_metadata("local:foo.tif", {"classes": [], "slices": {}}) + assert called == [] + + def test_tiled_source_updates_node_metadata(self, monkeypatch: pytest.MonkeyPatch): + node = FakeTiledNode() + monkeypatch.setattr(arrays_mod, "resolve_array", lambda path, kind, server_uri: node) + sync_annotation_metadata( + "tiled::browse/sample", {"classes": [{"classId": 1}], "slices": {"0": [{"id": "s1"}]}}, + ) + assert len(node.updates) == 1 + assert node.updates[0][STUDIO_ANNOTATED] == "yes" + assert node.updates[0]["studio_shape_count"] == 1 + + def test_node_without_update_metadata_is_skipped_not_raised(self, monkeypatch: pytest.MonkeyPatch): + class NoUpdateNode: + pass + + monkeypatch.setattr(arrays_mod, "resolve_array", lambda path, kind, server_uri: NoUpdateNode()) + # Should not raise even though the resolved node can't be updated. + sync_annotation_metadata("tiled::browse/sample", {"classes": [], "slices": {}}) + + def test_empty_path_after_tiled_prefix_is_a_no_op(self, monkeypatch: pytest.MonkeyPatch): + called = [] + monkeypatch.setattr(arrays_mod, "resolve_array", lambda *a, **k: called.append(1)) + sync_annotation_metadata("tiled:", {"classes": [], "slices": {}}) + assert called == [] + + def test_passes_parsed_server_uri_through(self, monkeypatch: pytest.MonkeyPatch): + import source_keys + import tiled_config + + monkeypatch.setattr( + tiled_config, "get_tiled_servers", lambda: {"local": {"uri": "http://x:1"}}, + ) + monkeypatch.setattr(source_keys, "get_tiled_servers", tiled_config.get_tiled_servers) + + seen = [] + node = FakeTiledNode() + + def fake_resolve(path, kind, server_uri): + seen.append((path, kind, server_uri)) + return node + + monkeypatch.setattr(arrays_mod, "resolve_array", fake_resolve) + sync_annotation_metadata( + "tiled:http://x:1:browse/sample", {"classes": [], "slices": {}}, + ) + assert seen[0] == ("browse/sample", "tiled", "http://x:1") diff --git a/backend/tests/test_tiled_clients.py b/backend/tests/test_tiled_clients.py new file mode 100644 index 0000000..3bb7b35 --- /dev/null +++ b/backend/tests/test_tiled_clients.py @@ -0,0 +1,183 @@ +"""Tests for tiled_clients.py using plain nested dicts as fake Tiled nodes +(only __getitem__/__len__ are needed) and a monkeypatched tiled.client.from_uri +so get_tiled_client never makes a real network call. +""" +from __future__ import annotations + +import pytest + +import tiled_clients +import tiled_config + + +@pytest.fixture(autouse=True) +def clear_client_cache(): + tiled_clients._client_cache.clear() + yield + tiled_clients._client_cache.clear() + + +# --------------------------------------------------------------------------- +# api_key_for_uri +# --------------------------------------------------------------------------- + +class TestApiKeyForUri: + def test_none_uri_returns_none(self): + assert tiled_clients.api_key_for_uri(None) is None + + def test_matches_configured_server_ignoring_trailing_slash(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr( + tiled_config, "get_tiled_servers", + lambda: {"local": {"uri": "http://x:1/", "api_key": "secret"}}, + ) + monkeypatch.setattr(tiled_clients, "get_tiled_servers", tiled_config.get_tiled_servers) + assert tiled_clients.api_key_for_uri("http://x:1") == "secret" + + def test_unmatched_uri_returns_none(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(tiled_config, "get_tiled_servers", lambda: {"local": {"uri": "http://x:1"}}) + monkeypatch.setattr(tiled_clients, "get_tiled_servers", tiled_config.get_tiled_servers) + assert tiled_clients.api_key_for_uri("http://other:2") is None + + +# --------------------------------------------------------------------------- +# get_tiled_client +# --------------------------------------------------------------------------- + +class TestGetTiledClient: + def test_caches_by_uri_and_api_key(self, monkeypatch: pytest.MonkeyPatch): + calls = [] + + def fake_from_uri(uri, **kwargs): + calls.append((uri, kwargs)) + return object() + + monkeypatch.setattr("tiled.client.from_uri", fake_from_uri) + monkeypatch.setattr(tiled_clients, "get_tiled_base", lambda: "http://default:1") + monkeypatch.setattr(tiled_clients, "get_tiled_api_key", lambda: None) + monkeypatch.setattr(tiled_clients, "api_key_for_uri", lambda uri: None) + + first = tiled_clients.get_tiled_client("http://x:1") + second = tiled_clients.get_tiled_client("http://x:1") + assert first is second + assert len(calls) == 1 + + def test_different_uri_creates_a_new_client(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr("tiled.client.from_uri", lambda uri, **kwargs: object()) + monkeypatch.setattr(tiled_clients, "get_tiled_base", lambda: "http://default:1") + monkeypatch.setattr(tiled_clients, "get_tiled_api_key", lambda: None) + monkeypatch.setattr(tiled_clients, "api_key_for_uri", lambda uri: None) + + a = tiled_clients.get_tiled_client("http://x:1") + b = tiled_clients.get_tiled_client("http://y:1") + assert a is not b + + def test_falls_back_to_default_base_uri(self, monkeypatch: pytest.MonkeyPatch): + seen_uris = [] + monkeypatch.setattr("tiled.client.from_uri", lambda uri, **kwargs: seen_uris.append(uri) or object()) + monkeypatch.setattr(tiled_clients, "get_tiled_base", lambda: "http://default:1") + monkeypatch.setattr(tiled_clients, "get_tiled_api_key", lambda: None) + monkeypatch.setattr(tiled_clients, "api_key_for_uri", lambda uri: None) + + tiled_clients.get_tiled_client(None) + assert seen_uris == ["http://default:1"] + + def test_explicit_api_key_wins_over_configured_and_default(self, monkeypatch: pytest.MonkeyPatch): + seen_kwargs = [] + monkeypatch.setattr("tiled.client.from_uri", lambda uri, **kwargs: seen_kwargs.append(kwargs) or object()) + monkeypatch.setattr(tiled_clients, "get_tiled_base", lambda: "http://default:1") + monkeypatch.setattr(tiled_clients, "get_tiled_api_key", lambda: "global-key") + monkeypatch.setattr(tiled_clients, "api_key_for_uri", lambda uri: "configured-key") + + tiled_clients.get_tiled_client("http://x:1", "explicit-key") + assert seen_kwargs == [{"api_key": "explicit-key"}] + + def test_no_api_key_omits_the_kwarg(self, monkeypatch: pytest.MonkeyPatch): + seen_kwargs = [] + monkeypatch.setattr("tiled.client.from_uri", lambda uri, **kwargs: seen_kwargs.append(kwargs) or object()) + monkeypatch.setattr(tiled_clients, "get_tiled_base", lambda: "http://default:1") + monkeypatch.setattr(tiled_clients, "get_tiled_api_key", lambda: None) + monkeypatch.setattr(tiled_clients, "api_key_for_uri", lambda uri: None) + + tiled_clients.get_tiled_client("http://x:1") + assert seen_kwargs == [{}] + + +# --------------------------------------------------------------------------- +# get_browse_container / get_browse_container_for +# --------------------------------------------------------------------------- + +class TestGetBrowseContainer: + def test_env_var_path_wins_when_present_and_non_empty(self, monkeypatch: pytest.MonkeyPatch): + client = {"custom": {"path": {"a": 1, "b": 2}}} + monkeypatch.setenv("TILED_BROWSE_PATH", "custom/path") + node, prefix = tiled_clients.get_browse_container(client) + assert node == {"a": 1, "b": 2} + assert prefix == "custom/path" + + def test_env_var_path_ignored_when_empty_container(self, monkeypatch: pytest.MonkeyPatch): + client = {"custom": {"path": {}}, "browse": {"s1": {}}} + monkeypatch.setenv("TILED_BROWSE_PATH", "custom/path") + node, prefix = tiled_clients.get_browse_container(client) + assert prefix == "browse" + + def test_env_var_path_ignored_when_missing(self, monkeypatch: pytest.MonkeyPatch): + client = {"browse": {"s1": {}}} + monkeypatch.setenv("TILED_BROWSE_PATH", "does/not/exist") + node, prefix = tiled_clients.get_browse_container(client) + assert prefix == "browse" + + def test_first_matching_candidate_wins_in_priority_order(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("TILED_BROWSE_PATH", raising=False) + client = { + "beamlines": {"bl733": {"s1": {}}}, + "browse": {"s2": {}}, + } + node, prefix = tiled_clients.get_browse_container(client) + # "beamlines/bl733" is checked before the plain "browse" fallback. + assert prefix == "beamlines/bl733" + + def test_falls_back_to_plain_browse(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("TILED_BROWSE_PATH", raising=False) + client = {"browse": {"s1": {}}} + node, prefix = tiled_clients.get_browse_container(client) + assert prefix == "browse" + assert node == {"s1": {}} + + def test_empty_candidate_containers_are_skipped(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("TILED_BROWSE_PATH", raising=False) + client = {"browse": {"generated_data": {}}, "beamlines": {"bl733": {"s1": {}}}} + node, prefix = tiled_clients.get_browse_container(client) + assert prefix == "beamlines/bl733" + + def test_falls_back_to_client_root_when_nothing_matches(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("TILED_BROWSE_PATH", raising=False) + client = {} + node, prefix = tiled_clients.get_browse_container(client) + assert node is client + assert prefix == "" + + +class TestGetBrowseContainerFor: + def test_explicit_path_navigates_directly(self): + client = {"a": {"b": {"x": 1}}} + node, prefix = tiled_clients.get_browse_container_for(client, "a/b") + assert node == {"x": 1} + assert prefix == "a/b" + + def test_strips_leading_and_trailing_slashes(self): + client = {"a": {"x": 1}} + node, prefix = tiled_clients.get_browse_container_for(client, "/a/") + assert node == {"x": 1} + assert prefix == "a" + + def test_no_path_falls_back_to_heuristic_discovery(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("TILED_BROWSE_PATH", raising=False) + client = {"browse": {"s1": {}}} + node, prefix = tiled_clients.get_browse_container_for(client, None) + assert prefix == "browse" + + def test_empty_string_path_falls_back_to_heuristic_discovery(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("TILED_BROWSE_PATH", raising=False) + client = {"browse": {"s1": {}}} + node, prefix = tiled_clients.get_browse_container_for(client, " ") + assert prefix == "browse" diff --git a/backend/tests/test_tiled_mask_sync.py b/backend/tests/test_tiled_mask_sync.py index 2a6e8bd..c66d4b4 100644 --- a/backend/tests/test_tiled_mask_sync.py +++ b/backend/tests/test_tiled_mask_sync.py @@ -1,7 +1,13 @@ """Unit tests for tiled_mask_sync.build_mask_volumes (pure rasterization → volumes).""" +import io +from typing import Any + import numpy as np +import pytest +from PIL import Image as PILImage -from schemas import AnnotationClass, ExportSourceItem +import tiled_mask_sync +from schemas import AnnotationClass, ExportSourceItem, PredictedSlicePointer from tiled_mask_sync import build_mask_volumes, merge_mask_volumes H = W = 32 @@ -67,6 +73,63 @@ def test_returns_none_without_slices(): assert build_mask_volumes(item, _classes(), {"height": H, "width": W}) is None +def _fake_commit_png(raw_label_map: np.ndarray) -> bytes: + buf = io.BytesIO() + PILImage.fromarray(raw_label_map, mode="L").save(buf, format="PNG") + return buf.getvalue() + + +def test_predicted_slices_reads_straight_from_the_run_commit_png(monkeypatch: pytest.MonkeyPatch): + """Regression test for the lazy-vectorization plan item: a slice with no + real shapes but a predicted-slice pointer must be rasterized directly + from the ipred run's commit.png (fetched server-side), never requiring + the frontend to have vectorized it into polygon shapes first.""" + raw = np.zeros((H, W), dtype=np.uint8) + raw[2:6, 2:6] = 10 # raw frontend classId, matching AnnotationClass(classId=10, "Cell") + calls: list[str] = [] + + def fake_run_commit_png(run_id: str) -> bytes: + calls.append(run_id) + assert run_id == "run-abc" + return _fake_commit_png(raw) + + monkeypatch.setattr(tiled_mask_sync.ipred_client, "run_commit_png", fake_run_commit_png) + + item = ExportSourceItem( + kind="tiled", source="browse/ds/img", slices={}, + predicted_slices={"7": PredictedSlicePointer(run_id="run-abc", class_ids=[10, 20])}, + ) + vols = build_mask_volumes(item, _classes(), {"height": H, "width": W}) + + assert vols is not None + assert calls == ["run-abc"] + assert vols["slice_indices"] == [7] + assert vols["semantic"][0, 3, 3] == 1 # remapped from raw classId 10 -> legend id 1 (Cell) + assert vols["class_vols"]["Cell"][0, 3, 3] == 255 + assert vols["class_vols"]["Wall"][0].sum() == 0 + + +def test_real_shapes_win_over_a_predicted_pointer_for_the_same_slice(monkeypatch: pytest.MonkeyPatch): + """A slice already vectorized/edited into real shapes must never fall + back to its (possibly stale) predicted pointer — matches the frontend's + own precedence in handleCommitVolumeApply/predictedRasterStore.""" + def fake_run_commit_png(run_id: str) -> bytes: + raise AssertionError("must not fetch commit.png when real shapes already cover this slice") + + monkeypatch.setattr(tiled_mask_sync.ipred_client, "run_commit_png", fake_run_commit_png) + + slices = {"7": [{"id": "a", "kind": "rectangle", "classId": 20, "x": 1, "y": 1, "w": 3, "h": 3}]} + item = ExportSourceItem( + kind="tiled", source="browse/ds/img", slices=slices, + predicted_slices={"7": PredictedSlicePointer(run_id="run-abc", class_ids=[10, 20])}, + ) + vols = build_mask_volumes(item, _classes(), {"height": H, "width": W}) + + assert vols is not None + assert vols["slice_indices"] == [7] + assert vols["semantic"][0].max() == 2 # Wall (classId 20), from the real shape — not the pointer + + def _volumes(slices, negatives=None): item = ExportSourceItem( kind="tiled", source="browse/ds/img", slices=slices, negative_slices=negatives or [], @@ -108,3 +171,255 @@ def test_merge_fresh_when_no_existing(): assert merged["slice_indices"] == [3] assert merged["updated_indices"] == [3] assert merged["semantic"].shape == (1, H, W) + + +# --------------------------------------------------------------------------- +# _read_existing_masks / write_masks_to_tiled — fake Tiled container, no live +# server. Only mask_pyramid.register_mask_pyramid actually needs a live Tiled +# server (see mask_pyramid.py's own test file for that established boundary), +# so it's stubbed here; everything else about write_masks_to_tiled is real +# container navigation/merge logic that can run against a fake. +# --------------------------------------------------------------------------- + +class FakeMaskContainer: + def __init__(self, metadata=None, children=None): + self.metadata = metadata or {} + self._children: dict[str, Any] = dict(children or {}) + self.deleted = False + self.written: dict[str, dict] = {} + self.created: list[str] = [] + + def __iter__(self): + return iter(self._children) + + def __getitem__(self, key): + return self._children[key] + + def __len__(self): + return len(self._children) + + def keys(self): + return list(self._children.keys()) + + def delete_contents(self, recursive=True, external_only=False): + self.deleted = True + self._children = {} + + def update_metadata(self, metadata): + self.metadata = metadata + + def write_array(self, arr, key, dims=None, metadata=None): + self.written[key] = {"arr": arr, "dims": dims, "metadata": metadata} + + def create_container(self, key, metadata): + child = FakeMaskContainer(metadata=metadata) + self._children[key] = child + self.created.append(key) + return child + + +@pytest.fixture(autouse=True) +def stub_register_mask_pyramid(monkeypatch: pytest.MonkeyPatch): + calls: list[dict] = [] + + def fake_register(semantic, key, container, cache_key): + calls.append({"semantic": semantic, "key": key, "container": container, "cache_key": cache_key}) + return {"key": key} + + monkeypatch.setattr(tiled_mask_sync.mask_pyramid, "register_mask_pyramid", fake_register) + return calls + + +@pytest.fixture() +def fake_client(monkeypatch: pytest.MonkeyPatch): + ds = FakeMaskContainer() + root = FakeMaskContainer(children={"browse": FakeMaskContainer(children={"ds": ds})}) + monkeypatch.setattr(tiled_mask_sync, "get_tiled_client", lambda uri, key: root) + monkeypatch.setattr(tiled_mask_sync, "api_key_for_uri", lambda uri: None) + return ds # the "browse/ds" container new mask containers get created under + + +class TestReadExistingMasks: + def test_reads_semantic_via_mask_pyramid_and_class_arrays_directly(self, monkeypatch: pytest.MonkeyPatch): + semantic = np.zeros((2, H, W), dtype=np.uint8) + monkeypatch.setattr(tiled_mask_sync.mask_pyramid, "read_mask_scale0", lambda container, key: semantic) + cell_arr = np.full((2, H, W), 255, dtype=np.uint8) + container = FakeMaskContainer( + metadata={ + "legend": [{"id": 1, "name": "Cell", "color": "#f00"}], + "slice_indices": [1, 2], + "slice_updated_at": {"1": "t1"}, + }, + children={"semantic": object(), "Cell": cell_arr}, + ) + result = tiled_mask_sync._read_existing_masks(container) + assert result is not None + assert result["slice_indices"] == [1, 2] + assert np.array_equal(result["semantic"], semantic) + assert np.array_equal(result["class_arrays"]["Cell"], cell_arr) + assert result["slice_updated_at"] == {"1": "t1"} + + def test_unreadable_container_returns_none(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr( + tiled_mask_sync.mask_pyramid, "read_mask_scale0", + lambda container, key: (_ for _ in ()).throw(RuntimeError("boom")), + ) + container = FakeMaskContainer(metadata={"legend": [], "slice_indices": []}, children={"semantic": object()}) + assert tiled_mask_sync._read_existing_masks(container) is None + + def test_unknown_class_key_is_skipped(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(tiled_mask_sync.mask_pyramid, "read_mask_scale0", lambda container, key: np.zeros((1, H, W))) + container = FakeMaskContainer( + metadata={"legend": [{"id": 1, "name": "Cell"}], "slice_indices": [0]}, + children={"semantic": object(), "orphan_key": np.zeros((1, H, W))}, + ) + result = tiled_mask_sync._read_existing_masks(container) + assert result["class_arrays"] == {} + + +class TestWriteMasksToTiled: + def _volumes_dict(self): + return _volumes({"2": [{"id": "a", "kind": "rectangle", "classId": 10, "x": 2, "y": 2, "w": 6, "h": 6}]}) + + def test_creates_a_new_container_when_none_exists(self, fake_client, stub_register_mask_pyramid): + info = tiled_mask_sync.write_masks_to_tiled("browse/ds/img", None, self._volumes_dict(), _classes()) + assert info["path"] == "browse/ds/img__masks" + assert "img__masks" in fake_client.created + new_container = fake_client["img__masks"] + assert new_container.metadata["studio_type"] == "segmentation_masks" + assert "Cell" in new_container.written + + def test_container_suffix_keeps_producers_independent(self, fake_client, stub_register_mask_pyramid): + info = tiled_mask_sync.write_masks_to_tiled( + "browse/ds/img", None, self._volumes_dict(), _classes(), container_suffix="_deep", + ) + assert info["path"] == "browse/ds/img__masks_deep" + assert stub_register_mask_pyramid[0]["cache_key"] == "img__masks_deep" + + def test_merges_into_an_existing_container(self, fake_client, monkeypatch, stub_register_mask_pyramid): + existing_vols = self._volumes_dict() + existing_container = FakeMaskContainer( + metadata={ + "legend": existing_vols["legend"], "slice_indices": existing_vols["slice_indices"], + }, + children={"semantic": object(), "Cell": existing_vols["class_vols"]["Cell"]}, + ) + monkeypatch.setattr( + tiled_mask_sync.mask_pyramid, "read_mask_scale0", + lambda container, key: existing_vols["semantic"], + ) + fake_client._children["img__masks"] = existing_container + + new_vols = _volumes({"5": [{"id": "b", "kind": "rectangle", "classId": 20, "x": 1, "y": 1, "w": 3, "h": 3}]}) + info = tiled_mask_sync.write_masks_to_tiled("browse/ds/img", None, new_vols, _classes()) + + assert existing_container.deleted is True + assert info["n_slices"] == 2 # slice 2 (existing) + slice 5 (new) + assert info["updated"] == 1 + + def test_shape_mismatch_discards_existing_and_replaces(self, fake_client, monkeypatch, stub_register_mask_pyramid): + mismatched_semantic = np.zeros((1, H * 2, W * 2), dtype=np.uint8) + existing_container = FakeMaskContainer( + metadata={"legend": [], "slice_indices": [9]}, + children={"semantic": object()}, + ) + monkeypatch.setattr(tiled_mask_sync.mask_pyramid, "read_mask_scale0", lambda container, key: mismatched_semantic) + fake_client._children["img__masks"] = existing_container + + info = tiled_mask_sync.write_masks_to_tiled("browse/ds/img", None, self._volumes_dict(), _classes()) + # Old slice 9 is gone entirely — replaced, not merged, due to the H/W mismatch. + assert info["n_slices"] == 1 + assert info["updated"] == 1 + + def test_slice_count_mismatch_discards_existing_and_replaces( + self, fake_client, monkeypatch, stub_register_mask_pyramid, + ): + # Metadata claims 2 slices (9 and 10) but the registered semantic array + # only actually holds 1 frame — e.g. stale metadata from a prior + # interrupted/partial write. Merging positionally by slice_indices would + # index past the array (`old_sem[1]` on a size-1 array) and crash with a + # numpy IndexError; this must be treated the same as an H/W mismatch. + short_semantic = np.zeros((1, H, W), dtype=np.uint8) + existing_container = FakeMaskContainer( + metadata={"legend": [], "slice_indices": [9, 10]}, + children={"semantic": object()}, + ) + monkeypatch.setattr(tiled_mask_sync.mask_pyramid, "read_mask_scale0", lambda container, key: short_semantic) + fake_client._children["img__masks"] = existing_container + + info = tiled_mask_sync.write_masks_to_tiled("browse/ds/img", None, self._volumes_dict(), _classes()) + # Old slices 9/10 are gone entirely — replaced, not merged, due to the + # slice-count mismatch. Confirms no IndexError was raised too. + assert info["n_slices"] == 1 + assert info["updated"] == 1 + + def test_returns_correct_summary_counts(self, fake_client, stub_register_mask_pyramid): + info = tiled_mask_sync.write_masks_to_tiled("browse/ds/img", None, self._volumes_dict(), _classes()) + assert info["n_classes"] == 2 # Cell + Wall + assert info["n_slices"] == 1 + assert info["updated"] == 1 + + +# --------------------------------------------------------------------------- +# run_mask_sync_job +# --------------------------------------------------------------------------- + +class _Payload: + def __init__(self, classes): + self.classes = classes + + +class TestRunMaskSyncJob: + def test_no_tiled_sources_reports_done_with_a_note(self): + import export_jobs + + local_item = ExportSourceItem(kind="local", source="foo.tif", slices={"0": []}) + jid = export_jobs.new_job("x") + tiled_mask_sync.run_mask_sync_job(jid, [local_item], _Payload(_classes())) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["written"] == [] + assert "skipped" in job["result"]["note"] + + def test_writes_masks_for_each_tiled_source(self, monkeypatch, fake_client, stub_register_mask_pyramid): + import export_jobs + + slices = {"2": [{"id": "a", "kind": "rectangle", "classId": 10, "x": 2, "y": 2, "w": 6, "h": 6}]} + item = ExportSourceItem(kind="tiled", source="browse/ds/img", slices=slices) + monkeypatch.setattr(tiled_mask_sync.arrays_mod, "resolve_array", lambda source, kind, server_uri: "node") + monkeypatch.setattr(tiled_mask_sync.arrays_mod, "array_shape_meta", lambda node: {"height": H, "width": W}) + + jid = export_jobs.new_job("x") + tiled_mask_sync.run_mask_sync_job(jid, [item], _Payload(_classes())) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert len(job["result"]["written"]) == 1 + assert job["result"]["written"][0]["source"] == "browse/ds/img" + assert job["result"]["written"][0]["container"] == "browse/ds/img__masks" + + def test_item_with_no_annotated_slices_is_skipped_not_errored(self, monkeypatch, fake_client): + import export_jobs + + item = ExportSourceItem(kind="tiled", source="browse/ds/img", slices={}) + monkeypatch.setattr(tiled_mask_sync.arrays_mod, "resolve_array", lambda source, kind, server_uri: "node") + monkeypatch.setattr(tiled_mask_sync.arrays_mod, "array_shape_meta", lambda node: {"height": H, "width": W}) + + jid = export_jobs.new_job("x") + tiled_mask_sync.run_mask_sync_job(jid, [item], _Payload(_classes())) + job = export_jobs.get_job(jid) + assert job["state"] == "done" + assert job["result"]["written"] == [] + + def test_exception_is_reported_as_a_job_error(self, monkeypatch): + import export_jobs + + item = ExportSourceItem(kind="tiled", source="browse/ds/img", slices={"0": [{"id": "a", "kind": "rectangle", "classId": 10, "x": 0, "y": 0, "w": 2, "h": 2}]}) + monkeypatch.setattr( + tiled_mask_sync.arrays_mod, "resolve_array", + lambda source, kind, server_uri: (_ for _ in ()).throw(RuntimeError("resolve failed")), + ) + jid = export_jobs.new_job("x") + tiled_mask_sync.run_mask_sync_job(jid, [item], _Payload(_classes())) + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "resolve failed" in job["error"] diff --git a/backend/tests/test_tiling.py b/backend/tests/test_tiling.py new file mode 100644 index 0000000..ade9620 --- /dev/null +++ b/backend/tests/test_tiling.py @@ -0,0 +1,229 @@ +"""Tests for tiling.py — real (no mocks) since torch+qlty are both installed +in this environment; this module's whole point is qlty geometry, so faking +that away would test nothing meaningful. +""" +from __future__ import annotations + +import numpy as np +import pytest + +torch = pytest.importorskip("torch") +pytest.importorskip("qlty") + +import tiling # noqa: E402 +from train_common import IGNORE_INDEX # noqa: E402 + + +class TestStepAndBorder: + def test_step_for_applies_overlap_fraction(self): + assert tiling.step_for(64) == 48 # 64 - round(64*0.25) + + def test_step_for_never_below_one(self): + assert tiling.step_for(1) >= 1 + + def test_border_for_uses_divisor(self): + assert tiling.border_for(64) == 8 # 64 // 8 + + def test_border_for_never_below_one(self): + assert tiling.border_for(1) >= 1 + + +class TestTileOrigins: + def test_exact_multiple_of_step(self): + # dim=100, window=20, step=20 -> origins 0,20,40,60,80 (last clamps to 80 anyway) + origins = tiling.tile_origins(100, 20, 20) + assert origins[0] == 0 + assert origins[-1] == 80 + assert all(o + 20 <= 100 for o in origins) + + def test_dim_equals_window_gives_single_origin(self): + assert tiling.tile_origins(20, 20, 15) == [0] + + def test_last_window_clamped_inside_image(self): + origins = tiling.tile_origins(50, 20, 15) + assert origins[-1] == 30 # 50 - 20 + assert max(origins) + 20 == 50 + + def test_dim_smaller_than_window_raises(self): + with pytest.raises(ValueError, match="smaller than window"): + tiling.tile_origins(10, 20, 5) + + +class TestPadToMin: + def test_already_big_enough_returns_same_array(self): + arr = np.zeros((32, 32), dtype=np.uint8) + assert tiling.pad_to_min(arr, 32, 32, fill=0) is arr + + def test_pads_bottom_right_only(self): + arr = np.ones((10, 10), dtype=np.uint8) + out = tiling.pad_to_min(arr, 16, 20, fill=0) + assert out.shape == (16, 20) + assert np.all(out[:10, :10] == 1) + assert np.all(out[10:, :] == 0) + assert np.all(out[:, 10:] == 0) + + def test_fill_value_used_for_padding(self): + arr = np.zeros((5, 5), dtype=np.uint8) + out = tiling.pad_to_min(arr, 8, 8, fill=IGNORE_INDEX) + assert out[7, 7] == IGNORE_INDEX + + def test_preserves_trailing_channel_dim(self): + arr = np.ones((10, 10, 3), dtype=np.uint8) + out = tiling.pad_to_min(arr, 16, 16, fill=0) + assert out.shape == (16, 16, 3) + + +class TestQltyAvailable: + def test_true_when_installed(self): + assert tiling.qlty_available() is True + + +class TestTilePair: + def test_produces_window_sized_patches(self): + rgb = np.random.default_rng(0).integers(0, 256, size=(64, 64, 3), dtype=np.uint8) + label = np.zeros((64, 64), dtype=np.uint8) + label[10:54, 10:54] = 1 # broad interior annotation + patches = tiling._tile_pair(rgb, label, window=32) + assert len(patches) > 0 + for patch_rgb, patch_label in patches: + assert patch_rgb.shape == (32, 32, 3) + assert patch_label.shape == (32, 32) + + def test_entirely_unannotated_image_yields_no_patches(self): + rgb = np.zeros((64, 64, 3), dtype=np.uint8) + label = np.full((64, 64), IGNORE_INDEX, dtype=np.uint8) + patches = tiling._tile_pair(rgb, label, window=32) + assert patches == [] + + def test_edge_only_annotation_falls_back_to_unmasked_patches(self): + # Annotate only the outermost couple of pixels — inside every patch's + # down-weighted border ring, so the masked pass would keep nothing; + # the edge-only fallback (mask_borders=False) must still find it. + rgb = np.zeros((64, 64, 3), dtype=np.uint8) + label = np.full((64, 64), IGNORE_INDEX, dtype=np.uint8) + label[0:2, 0:2] = 1 + patches = tiling._tile_pair(rgb, label, window=32) + assert len(patches) > 0 + assert any((p_label != IGNORE_INDEX).any() for _, p_label in patches) + + def test_smaller_than_window_image_is_padded_first(self): + rgb = np.ones((10, 10, 3), dtype=np.uint8) + label = np.ones((10, 10), dtype=np.uint8) + patches = tiling._tile_pair(rgb, label, window=32) + assert len(patches) >= 1 + assert patches[0][0].shape == (32, 32, 3) + + +class TestHoldoutValPatches: + def _pairs(self, n): + return [(np.zeros((4, 4)), np.zeros((4, 4))) for _ in range(n)] + + def test_no_op_when_val_already_present(self): + datasets = {"train": self._pairs(20), "val": self._pairs(2)} + out, n = tiling.holdout_val_patches(datasets, seed=0) + assert n == 0 + assert out is datasets + + def test_no_op_below_min_patches(self): + datasets = {"train": self._pairs(4), "val": []} + out, n = tiling.holdout_val_patches(datasets, seed=0, min_patches=8) + assert n == 0 + assert out["train"] == datasets["train"] + + def test_holds_out_a_seeded_fraction(self): + datasets = {"train": self._pairs(20), "val": []} + out, n = tiling.holdout_val_patches(datasets, seed=0, fraction=0.1) + assert n == 2 # round(20 * 0.1) + assert len(out["train"]) == 18 + assert len(out["val"]) == 2 + + def test_deterministic_given_same_seed(self): + datasets = {"train": self._pairs(20), "val": []} + out1, _ = tiling.holdout_val_patches(datasets, seed=42) + out2, _ = tiling.holdout_val_patches(datasets, seed=42) + assert len(out1["val"]) == len(out2["val"]) + + +class TestTileDatasets: + def test_tiles_every_split(self): + rgb = np.random.default_rng(0).integers(0, 256, size=(64, 64, 3), dtype=np.uint8) + label = np.ones((64, 64), dtype=np.uint8) + datasets = {"train": [(rgb, label)], "val": [(rgb, label)]} + out = tiling.tile_datasets(datasets, window=32) + assert len(out["train"]) > 0 + assert len(out["val"]) > 0 + + def test_cancellation_returns_none(self): + rgb = np.ones((64, 64, 3), dtype=np.uint8) + label = np.ones((64, 64), dtype=np.uint8) + datasets = {"train": [(rgb, label), (rgb, label)]} + out = tiling.tile_datasets(datasets, window=32, cancel_cb=lambda: True) + assert out is None + + +def _tiny_model(): + model = torch.nn.Conv2d(3, 2, kernel_size=3, padding=1) + + def forward_fn(batch): + return model(batch) + + def to_tensor_fn(arr): + t = torch.from_numpy(np.ascontiguousarray(arr)).float() / 255.0 + return t.permute(2, 0, 1) if t.ndim == 3 else t.unsqueeze(0) + + return forward_fn, to_tensor_fn + + +class TestBlendTiledForwardAndFriends: + def test_predict_label_map_tiled_shape_and_dtype(self): + forward_fn, to_tensor_fn = _tiny_model() + rgb = np.random.default_rng(0).integers(0, 256, size=(48, 48, 3), dtype=np.uint8) + label = tiling.predict_label_map_tiled( + rgb, forward_fn=forward_fn, to_tensor_fn=to_tensor_fn, + window=32, min_confidence=0.0, device="cpu", + ) + assert label.shape == (48, 48) + assert label.dtype == np.uint8 + + def test_predict_label_map_tiled_high_confidence_threshold_yields_background(self): + forward_fn, to_tensor_fn = _tiny_model() + rgb = np.zeros((48, 48, 3), dtype=np.uint8) + label = tiling.predict_label_map_tiled( + rgb, forward_fn=forward_fn, to_tensor_fn=to_tensor_fn, + window=32, min_confidence=1.1, device="cpu", # impossible to reach + ) + assert np.all(label == 0) + + def test_predict_label_map_tiled_cancellation_returns_none(self): + forward_fn, to_tensor_fn = _tiny_model() + rgb = np.zeros((48, 48, 3), dtype=np.uint8) + label = tiling.predict_label_map_tiled( + rgb, forward_fn=forward_fn, to_tensor_fn=to_tensor_fn, + window=32, min_confidence=0.5, device="cpu", cancel_cb=lambda: True, + ) + assert label is None + + def test_denoise_image_tiled_single_channel_squeezed(self): + model = torch.nn.Conv2d(1, 1, kernel_size=3, padding=1) + + def forward_fn(batch): + return model(batch) + + def to_tensor_fn(arr): + return torch.from_numpy(np.ascontiguousarray(arr)).float().unsqueeze(0) / 255.0 + + gray = np.random.default_rng(0).integers(0, 256, size=(48, 48), dtype=np.uint8) + out = tiling.denoise_image_tiled( + gray, forward_fn=forward_fn, to_tensor_fn=to_tensor_fn, window=32, device="cpu", + ) + assert out.shape == (48, 48) + assert out.dtype == np.float32 + + def test_image_smaller_than_window_is_handled(self): + forward_fn, to_tensor_fn = _tiny_model() + rgb = np.zeros((10, 10, 3), dtype=np.uint8) + label = tiling.predict_label_map_tiled( + rgb, forward_fn=forward_fn, to_tensor_fn=to_tensor_fn, + window=32, min_confidence=0.0, device="cpu", + ) + assert label.shape == (10, 10) diff --git a/backend/tests/test_train_common.py b/backend/tests/test_train_common.py new file mode 100644 index 0000000..18d62af --- /dev/null +++ b/backend/tests/test_train_common.py @@ -0,0 +1,406 @@ +"""train_common.py + the denoise_bake.py model-denoiser fix it enables. + +Ported alongside train_common/tiling/denoise_train/denoise_runtime/ +autoencoder_runtime (Phase 5, Stage A). The `ml` extra (torch/dlsia/qlty) may +or may not be installed in whatever environment runs this suite — these tests +don't assume either way. The property that matters is: every module here is +import-safe regardless, and denoise_bake.py's `_ModelDenoiser` (which already +called into these modules before they existed) fails with clear, specific +errors instead of a raw `ModuleNotFoundError` when a dependency truly is +missing — verified here by monkeypatching availability rather than by relying +on the dev machine's actual install state. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import numpy as np +import pytest +from fastapi import HTTPException + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import denoise_bake # noqa: E402 +import train_common # noqa: E402 + + +def test_modules_import_cleanly(): + """All Stage A modules must be import-safe whether or not torch/dlsia are installed.""" + import autoencoder_runtime # noqa: F401 + import denoise_runtime # noqa: F401 + import denoise_train # noqa: F401 + import dlsia_runtime # noqa: F401 + import tiling # noqa: F401 + + +def test_capability_never_raises_and_has_no_dinov3_fields(): + """DINOv3 is deferred to Phase 5.5 — capability() must not report it at all, + regardless of what's actually installed.""" + result = train_common.capability() + assert "dinov3" not in result + assert "error" not in result + assert set(result) == { + "torch_available", "torch_version", "device", "dlsia", "tiling", + "denoise", "runs_dir", "busy", + } + + +def test_capability_reports_unavailable_without_torch(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(train_common, "torch_available", lambda: False) + monkeypatch.setattr(train_common, "dlsia_available", lambda: False) + result = train_common.capability() + assert result["torch_available"] is False + assert result["torch_version"] is None + assert result["device"] is None + assert result["dlsia"] == {"available": False} + + +@pytest.fixture() +def runs_dir(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Path: + monkeypatch.setenv("DINO_RUNS_DIR", str(tmp_path / "runs")) + return tmp_path / "runs" + + +def _write_run_config(runs_dir: Path, run_id: str, config: dict) -> None: + d = runs_dir / run_id + d.mkdir(parents=True, exist_ok=True) + (d / "config.json").write_text(json.dumps(config)) + + +def test_model_denoiser_unknown_run_is_a_clean_404(runs_dir: Path) -> None: + """Previously: ModuleNotFoundError before train_common.py existed.""" + with pytest.raises(Exception) as exc_info: + denoise_bake._ModelDenoiser("nonexistent-run", None, {}) + assert "Unknown run" in str(exc_info.value) + + +def test_model_denoiser_refuses_non_denoiser_run(runs_dir: Path) -> None: + _write_run_config(runs_dir, "seg-run", {"model_family": "dlsia_tunet", "task": "segmentation"}) + with pytest.raises(ValueError, match="not a denoiser run"): + denoise_bake._ModelDenoiser("seg-run", None, {}) + + +def test_model_denoiser_reports_missing_dlsia_by_name( + runs_dir: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """A tunet-architecture denoiser run needs dlsia — confirm the error names it + specifically, not just a generic import failure. Forces the "not installed" + branch via monkeypatch so this doesn't depend on the test machine's actual + dlsia install state.""" + monkeypatch.setattr(train_common, "dlsia_available", lambda: False) + _write_run_config( + runs_dir, + "denoiser-run", + { + "model_family": "dlsia_denoiser", + "task": "denoising", + "model_config": {"architecture": "tunet"}, + }, + ) + with pytest.raises(ValueError, match="dlsia is not installed"): + denoise_bake._ModelDenoiser("denoiser-run", None, {}) + + +def test_model_denoiser_reports_missing_torch_by_name( + runs_dir: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """A cnn_ae denoiser needs no dlsia, but still needs torch for a device.""" + import tiling + + monkeypatch.setattr(tiling, "qlty_available", lambda: True) + monkeypatch.setattr(train_common, "pick_device", lambda: None) + _write_run_config( + runs_dir, + "ae-run", + { + "model_family": "dlsia_denoiser", + "task": "denoising", + "model_config": {"architecture": "cnn_ae"}, + "image_size": 256, + }, + ) + with pytest.raises(ValueError, match="torch is not installed"): + denoise_bake._ModelDenoiser("ae-run", None, {}) + + +# --------------------------------------------------------------------------- +# pick_device +# --------------------------------------------------------------------------- + +class TestPickDevice: + def test_no_torch_returns_none(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(train_common, "torch_available", lambda: False) + assert train_common.pick_device() is None + + @pytest.mark.parametrize("override", ["mps", "cuda", "cpu"]) + def test_train_device_env_override_short_circuits_detection(self, monkeypatch: pytest.MonkeyPatch, override): + # A valid override returns before ever importing torch's backends — + # exercised for real (torch IS available in this env) without needing + # the overridden device to actually exist on this machine. + monkeypatch.setenv("TRAIN_DEVICE", override) + assert train_common.pick_device() == override + + def test_invalid_override_falls_through_to_real_detection(self, monkeypatch: pytest.MonkeyPatch): + pytest.importorskip("torch") + monkeypatch.setenv("TRAIN_DEVICE", "not-a-real-device") + assert train_common.pick_device() in ("mps", "cuda", "cpu") + + +# --------------------------------------------------------------------------- +# run_dir / _validate_run_id +# --------------------------------------------------------------------------- + +class TestRunDirValidation: + @pytest.mark.parametrize("bad_id", ["", ".", "..", "a/b", "a\\b", "a\x00b"]) + def test_rejects_unsafe_run_ids(self, runs_dir: Path, bad_id): + with pytest.raises(HTTPException) as exc: + train_common.run_dir(bad_id) + assert exc.value.status_code == 400 + + def test_accepts_a_safe_run_id(self, runs_dir: Path): + d = train_common.run_dir("my-run-1") + assert d.name == "my-run-1" + assert d.parent == train_common.runs_dir() + + +# --------------------------------------------------------------------------- +# Run persistence: list_runs / delete_run / load_run_config +# (config.json written directly via the existing _write_run_config helper — +# these three never touch torch, unlike save_run/load_adapter_state below) +# --------------------------------------------------------------------------- + +class TestListRuns: + def test_empty_when_runs_dir_does_not_exist(self, runs_dir: Path): + assert train_common.list_runs() == [] + + def test_lists_saved_runs_sorted_newest_first(self, runs_dir: Path): + _write_run_config(runs_dir, "old", {"model_family": "dlsia_tunet", "created_at": "2024-01-01T00:00:00"}) + _write_run_config(runs_dir, "new", {"model_family": "dlsia_tunet", "created_at": "2024-06-01T00:00:00"}) + results = train_common.list_runs() + assert [r["created_at"] for r in results] == ["2024-06-01T00:00:00", "2024-01-01T00:00:00"] + + def test_backward_compat_missing_task_defaults_to_segmentation(self, runs_dir: Path): + _write_run_config(runs_dir, "old-run", {"model_family": "dlsia_tunet", "created_at": "x"}) + results = train_common.list_runs() + assert results[0]["task"] == "segmentation" + + def test_non_directory_entries_are_skipped(self, runs_dir: Path): + runs_dir.mkdir(parents=True, exist_ok=True) + (runs_dir / "stray_file.txt").write_text("not a run") + assert train_common.list_runs() == [] + + def test_malformed_run_directory_is_skipped_not_raised(self, runs_dir: Path): + d = runs_dir / "broken" + d.mkdir(parents=True, exist_ok=True) + (d / "config.json").write_text("{ not valid json") + assert train_common.list_runs() == [] + + def test_metrics_defaults_to_empty_dict_when_missing(self, runs_dir: Path): + _write_run_config(runs_dir, "no-metrics", {"model_family": "dlsia_tunet", "created_at": "x"}) + results = train_common.list_runs() + assert results[0]["metrics"] == {} + + +class TestDeleteRun: + def test_unknown_run_is_404(self, runs_dir: Path): + with pytest.raises(HTTPException) as exc: + train_common.delete_run("nonexistent") + assert exc.value.status_code == 404 + + def test_removes_the_run_directory(self, runs_dir: Path): + _write_run_config(runs_dir, "to-delete", {"model_family": "dlsia_tunet"}) + d = train_common.run_dir("to-delete") + assert d.is_dir() + train_common.delete_run("to-delete") + assert not d.exists() + + +class TestLoadRunConfig: + def test_unknown_run_is_404(self, runs_dir: Path): + with pytest.raises(HTTPException) as exc: + train_common.load_run_config("nonexistent") + assert exc.value.status_code == 404 + + def test_corrupt_config_is_500(self, runs_dir: Path): + d = runs_dir / "corrupt" + d.mkdir(parents=True, exist_ok=True) + (d / "config.json").write_text("{ broken") + with pytest.raises(HTTPException) as exc: + train_common.load_run_config("corrupt") + assert exc.value.status_code == 500 + + def test_missing_task_defaults_to_segmentation(self, runs_dir: Path): + _write_run_config(runs_dir, "old", {"model_family": "dlsia_tunet"}) + config = train_common.load_run_config("old") + assert config["task"] == "segmentation" + + +class TestSaveRunAndLoadAdapterState: + def test_round_trips_config_metrics_and_weights(self, runs_dir: Path): + torch = pytest.importorskip("torch") + train_common.save_run( + "run1", + model_family="dlsia_tunet", + model_config={"depth": 2}, + classes=[{"classId": 1, "label": "a"}], + render={}, + image_size=64, + hyperparams={"depth": 2}, + source_keys=["local:x.tif"], + adapter_state={"weight": torch.zeros(2, 2)}, + metrics={"epochs_completed": 1}, + ) + config = train_common.load_run_config("run1") + assert config["model_family"] == "dlsia_tunet" + assert config["classes"] == [{"classId": 1, "label": "a"}] + + adapter = train_common.load_adapter_state("run1") + assert torch.equal(adapter["weight"], torch.zeros(2, 2)) + + def test_load_adapter_state_missing_weights_is_404(self, runs_dir: Path): + pytest.importorskip("torch") + _write_run_config(runs_dir, "no-weights", {"model_family": "dlsia_tunet"}) + with pytest.raises(HTTPException) as exc: + train_common.load_adapter_state("no-weights") + assert exc.value.status_code == 404 + + +# --------------------------------------------------------------------------- +# letterbox / letterbox_params / unletterbox — pure numpy/PIL, no torch +# --------------------------------------------------------------------------- + +class TestLetterbox: + def test_already_square_is_returned_unchanged(self): + img = np.zeros((32, 32, 3), dtype=np.uint8) + lbl = np.ones((32, 32), dtype=np.uint8) + out_img, out_lbl = train_common.letterbox(img, lbl, 32) + assert out_img is img + assert out_lbl is lbl + + def test_pads_a_non_square_image_to_a_square_canvas(self): + img = np.full((20, 40, 3), 255, dtype=np.uint8) + lbl = np.ones((20, 40), dtype=np.uint8) + out_img, out_lbl = train_common.letterbox(img, lbl, 40) + assert out_img.shape == (40, 40, 3) + assert out_lbl.shape == (40, 40) + + def test_padded_label_regions_use_ignore_index(self): + img = np.full((10, 40, 3), 255, dtype=np.uint8) + lbl = np.ones((10, 40), dtype=np.uint8) + _out_img, out_lbl = train_common.letterbox(img, lbl, 40) + # Top/bottom padding bands must be IGNORE_INDEX, not a real class. + assert np.all(out_lbl[0, :] == train_common.IGNORE_INDEX) + assert np.all(out_lbl[-1, :] == train_common.IGNORE_INDEX) + + +class TestLetterboxParamsAndUnletterbox: + def test_letterbox_params_centers_the_scaled_image(self): + params = train_common.letterbox_params(20, 40, 40) + assert params["nh"] == 20 + assert params["nw"] == 40 + assert params["top"] == 10 + assert params["left"] == 0 + + def test_unletterbox_inverts_letterbox_for_an_exact_fit(self): + img = np.zeros((20, 40, 3), dtype=np.uint8) + lbl = np.full((20, 40), 3, dtype=np.uint8) + boxed_img, boxed_lbl = train_common.letterbox(img, lbl, 40) + recovered = train_common.unletterbox(boxed_lbl, 20, 40, 40) + assert recovered.shape == (20, 40) + assert np.all(recovered == 3) + + def test_unletterbox_resizes_back_when_original_was_downscaled(self): + # image_size smaller than the original -> nh/nw != orig, exercising the + # resize-back branch rather than the exact-crop shortcut. + img = np.zeros((100, 200, 3), dtype=np.uint8) + lbl = np.full((100, 200), 5, dtype=np.uint8) + boxed_img, boxed_lbl = train_common.letterbox(img, lbl, 40) + recovered = train_common.unletterbox(boxed_lbl, 100, 200, 40) + assert recovered.shape == (100, 200) + # Nearest-neighbor resize of a constant label stays constant. + assert np.all(recovered == 5) + + +# --------------------------------------------------------------------------- +# denoising_render_slice_fn / _denoise_params +# --------------------------------------------------------------------------- + +class TestDenoiseParams: + def test_none_denoise_returns_none_method(self): + assert train_common._denoise_params(None) == (None, 0.5) + + def test_falsy_dict_returns_none_method(self): + assert train_common._denoise_params({}) == (None, 0.5) + + def test_method_none_string_is_treated_as_no_denoise(self): + method, strength = train_common._denoise_params({"method": "none", "strength": 0.7}) + assert method is None + assert strength == 0.7 + + def test_model_method_is_not_supported_as_a_preprocessor(self): + method, _ = train_common._denoise_params({"method": "model", "strength": 0.5}) + assert method is None + + def test_real_classical_method_passes_through(self): + assert train_common._denoise_params({"method": "median", "strength": 0.3}) == ("median", 0.3) + + def test_object_with_method_attribute(self): + class Opts: + method = "gaussian" + strength = 0.4 + + assert train_common._denoise_params(Opts()) == ("gaussian", 0.4) + + +class TestDenoisingRenderSliceFn: + def test_no_denoise_returns_plain_render_slice(self): + import images as images_mod + + fn = train_common.denoising_render_slice_fn(None) + assert fn is images_mod.render_slice + + def test_denoise_method_filters_grayscale_input_before_rendering(self, monkeypatch): + calls = [] + import denoise as denoise_mod + + def fake_denoise_slice(arr, method, strength): + calls.append((method, strength)) + return arr + + monkeypatch.setattr(denoise_mod, "denoise_slice", fake_denoise_slice) + fn = train_common.denoising_render_slice_fn({"method": "median", "strength": 0.2}) + arr = np.zeros((8, 8), dtype=np.float32) + fn(arr, {"norm": "slice", "scale": "linear", "vmin_pct": 1.0, "vmax_pct": 99.0, "cmap": "gray"}, None) + assert calls == [("median", 0.2)] + + def test_rgb_input_skips_denoising(self, monkeypatch): + calls = [] + import denoise as denoise_mod + + monkeypatch.setattr(denoise_mod, "denoise_slice", lambda *a: calls.append(1)) + fn = train_common.denoising_render_slice_fn({"method": "median", "strength": 0.2}) + arr = np.zeros((8, 8, 3), dtype=np.uint8) + fn(arr, {"norm": "slice", "scale": "linear", "vmin_pct": 1.0, "vmax_pct": 99.0, "cmap": "gray"}, None) + assert calls == [] + + +# --------------------------------------------------------------------------- +# capability()'s exception-swallowing path +# --------------------------------------------------------------------------- + +class TestCapabilityErrorHandling: + def test_internal_failure_is_reported_not_raised(self, monkeypatch: pytest.MonkeyPatch): + def boom(): + raise RuntimeError("tiling import exploded") + + monkeypatch.setattr(train_common, "torch_available", boom) + result = train_common.capability() + assert "error" in result + # The raw exception message must never reach an HTTP response (CodeQL + # py/stack-trace-exposure) — only a generic marker crosses the API + # boundary; the real detail goes to the server log only. + assert "tiling import exploded" not in result["error"] + assert result["torch_available"] is False diff --git a/backend/tests/test_train_common_ml.py b/backend/tests/test_train_common_ml.py new file mode 100644 index 0000000..548182d --- /dev/null +++ b/backend/tests/test_train_common_ml.py @@ -0,0 +1,223 @@ +"""Real (no mocks) torch round trips for train_common.py's model-construction +and training-loop functions — build_family, compute_miou, run_training_loop, +_evaluate. Kept in a separate file from test_train_common.py so that file can +stay import-safe/runnable without torch (its own stated design), while this +one skips outright when the `ml` extra isn't installed. +""" +from __future__ import annotations + +import sys +from pathlib import Path + +import numpy as np +import pytest + +torch = pytest.importorskip("torch") + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import train_common # noqa: E402 +from schemas import DlsiaDenoiserConfig, DlsiaTunetConfig # noqa: E402 + +IMAGE_SIZE = 64 + + +def _tunet_config(**overrides): + hp = {"depth": 2, "base_channels": 4, "growth_rate": 1.2, "image_size": IMAGE_SIZE, "batch_size": 2} + hp.update(overrides) + return DlsiaTunetConfig(hyperparams=hp) + + +def _denoiser_config(architecture="tunet", **overrides): + hp = {"depth": 2, "base_channels": 4, "growth_rate": 1.2, "image_size": IMAGE_SIZE, "batch_size": 2} + hp.update(overrides) + scheme = "ae" if architecture == "cnn_ae" else "n2n" + return DlsiaDenoiserConfig(architecture=architecture, hyperparams=hp, training_scheme=scheme) + + +class TestBuildFamilyTunet: + def test_builds_from_scratch(self): + built = train_common.build_family(_tunet_config(), n_classes=3, device="cpu", log_cb=lambda m: None) + assert built.model_config_snapshot == {"depth": 2, "base_channels": 4, "growth_rate": 1.2} + assert len(built.trainable_params) > 0 + + def test_dlsia_unavailable_raises(self, monkeypatch): + monkeypatch.setattr(train_common, "dlsia_available", lambda: False) + with pytest.raises(RuntimeError, match="dlsia is not installed"): + train_common.build_family(_tunet_config(), n_classes=3, device="cpu", log_cb=lambda m: None) + + def test_resume_uses_saved_topology_not_the_request(self): + built = train_common.build_family(_tunet_config(), n_classes=3, device="cpu", log_cb=lambda m: None) + state = built.adapter_state_fn() + resumed = train_common.build_family( + _tunet_config(depth=6), n_classes=3, device="cpu", log_cb=lambda m: None, init_state=state, + ) + # depth=6 in the new request must be ignored; the saved topology wins. + assert resumed.model_config_snapshot["depth"] == 2 + + def test_forward_produces_correct_class_channels(self): + built = train_common.build_family(_tunet_config(), n_classes=3, device="cpu", log_cb=lambda m: None) + batch = torch.zeros((1, 3, IMAGE_SIZE, IMAGE_SIZE)) + out = built.forward_fn(batch) + assert out.shape == (1, 3, IMAGE_SIZE, IMAGE_SIZE) + + +class TestBuildFamilyDenoiserTunet: + def test_builds_from_scratch(self): + built = train_common.build_family(_denoiser_config("tunet"), n_classes=2, device="cpu", log_cb=lambda m: None) + assert built.model_config_snapshot["architecture"] == "tunet" + + def test_dlsia_unavailable_raises(self, monkeypatch): + monkeypatch.setattr(train_common, "dlsia_available", lambda: False) + with pytest.raises(RuntimeError, match="dlsia is not installed"): + train_common.build_family(_denoiser_config("tunet"), n_classes=2, device="cpu", log_cb=lambda m: None) + + def test_resume_uses_saved_topology(self): + built = train_common.build_family(_denoiser_config("tunet"), n_classes=2, device="cpu", log_cb=lambda m: None) + state = built.adapter_state_fn() + resumed = train_common.build_family( + _denoiser_config("tunet", depth=6), n_classes=2, device="cpu", log_cb=lambda m: None, init_state=state, + ) + assert resumed.model_config_snapshot["depth"] == 2 + + +class TestBuildFamilyDenoiserCnnAe: + def test_builds_from_scratch_no_dlsia_gate(self, monkeypatch): + # cnn_ae is plain torch — must succeed even when dlsia is "unavailable". + monkeypatch.setattr(train_common, "dlsia_available", lambda: False) + built = train_common.build_family(_denoiser_config("cnn_ae"), n_classes=2, device="cpu", log_cb=lambda m: None) + assert built.model_config_snapshot["architecture"] == "cnn_ae" + assert built.model_config_snapshot["latent_channels"] > 0 + + def test_resume_uses_saved_topology(self): + built = train_common.build_family(_denoiser_config("cnn_ae"), n_classes=2, device="cpu", log_cb=lambda m: None) + state = built.adapter_state_fn() + resumed = train_common.build_family( + _denoiser_config("cnn_ae", depth=6), n_classes=2, device="cpu", log_cb=lambda m: None, init_state=state, + ) + assert resumed.model_config_snapshot["depth"] == 2 + + +class TestBuildFamilyUnknownConfig: + def test_unrecognized_config_type_raises_value_error(self): + class NotAModelConfig: + hyperparams = None + model_family = "bogus" + + with pytest.raises(ValueError, match="Unknown model family"): + train_common.build_family(NotAModelConfig(), n_classes=2, device="cpu", log_cb=lambda m: None) + + +class TestComputeMiou: + def test_perfect_match_gives_miou_one(self): + pred = torch.tensor([0, 0, 1, 1]) + target = torch.tensor([0, 0, 1, 1]) + assert train_common.compute_miou(pred, target, n_classes=2) == pytest.approx(1.0) + + def test_no_overlap_gives_miou_zero(self): + pred = torch.tensor([0, 0, 0, 0]) + target = torch.tensor([1, 1, 1, 1]) + assert train_common.compute_miou(pred, target, n_classes=2) == pytest.approx(0.0) + + def test_ignore_index_pixels_are_excluded(self): + pred = torch.tensor([0, 0, 1, 1]) + target = torch.tensor([0, 0, 255, 255]) + # Only the first two pixels count; both match -> class 0 IoU 1.0, class 1 has no valid pixels (skipped). + assert train_common.compute_miou(pred, target, n_classes=2, ignore_index=255) == pytest.approx(1.0) + + def test_all_ignored_returns_zero(self): + pred = torch.tensor([0, 1]) + target = torch.tensor([255, 255]) + assert train_common.compute_miou(pred, target, n_classes=2, ignore_index=255) == 0.0 + + +def _synthetic_pairs(n, n_classes=2, size=IMAGE_SIZE, seed=0): + rng = np.random.default_rng(seed) + return [ + ( + rng.integers(0, 256, size=(size, size, 3), dtype=np.uint8), + rng.integers(0, n_classes, size=(size, size), dtype=np.uint8), + ) + for _ in range(n) + ] + + +class TestRunTrainingLoop: + def _built(self, n_classes=2): + return train_common.build_family(_tunet_config(), n_classes=n_classes, device="cpu", log_cb=lambda m: None) + + def test_no_training_data_raises(self): + built = self._built() + with pytest.raises(ValueError, match="No training data"): + train_common.run_training_loop( + train_pairs=[], val_pairs=[], image_size=IMAGE_SIZE, n_classes=2, epochs=1, + batch_size=2, seed=0, flip_augment=False, to_tensor_fn=built.to_tensor_fn, + forward_fn=built.forward_fn, trainable_params=built.trainable_params, lr=1e-3, device="cpu", + set_train_mode=built.set_train_mode, + ) + + def test_completes_epochs_and_reports_loss(self): + built = self._built() + result = train_common.run_training_loop( + train_pairs=_synthetic_pairs(4), val_pairs=[], image_size=IMAGE_SIZE, n_classes=2, + epochs=2, batch_size=2, seed=0, flip_augment=True, to_tensor_fn=built.to_tensor_fn, + forward_fn=built.forward_fn, trainable_params=built.trainable_params, lr=1e-3, device="cpu", + set_train_mode=built.set_train_mode, + ) + assert result["epochs_completed"] == 2 + assert result["cancelled"] is False + assert result["final_val_loss"] is None + assert result["val_miou"] is None + + def test_validation_pairs_produce_val_loss_and_miou(self): + built = self._built() + result = train_common.run_training_loop( + train_pairs=_synthetic_pairs(4), val_pairs=_synthetic_pairs(2, seed=1), + image_size=IMAGE_SIZE, n_classes=2, epochs=1, batch_size=2, seed=0, flip_augment=False, + to_tensor_fn=built.to_tensor_fn, forward_fn=built.forward_fn, + trainable_params=built.trainable_params, lr=1e-3, device="cpu", + set_train_mode=built.set_train_mode, + ) + assert result["final_val_loss"] is not None + assert result["val_miou"] is not None + + def test_on_batch_cancel_stops_immediately(self): + built = self._built() + result = train_common.run_training_loop( + train_pairs=_synthetic_pairs(8), val_pairs=[], image_size=IMAGE_SIZE, n_classes=2, + epochs=3, batch_size=2, seed=0, flip_augment=False, to_tensor_fn=built.to_tensor_fn, + forward_fn=built.forward_fn, trainable_params=built.trainable_params, lr=1e-3, device="cpu", + on_batch=lambda: True, set_train_mode=built.set_train_mode, + ) + assert result["cancelled"] is True + assert result["epochs_completed"] == 1 + + def test_on_epoch_cancel_stops_after_that_epoch(self): + built = self._built() + epochs_seen = [] + result = train_common.run_training_loop( + train_pairs=_synthetic_pairs(4), val_pairs=[], image_size=IMAGE_SIZE, n_classes=2, + epochs=5, batch_size=2, seed=0, flip_augment=False, to_tensor_fn=built.to_tensor_fn, + forward_fn=built.forward_fn, trainable_params=built.trainable_params, lr=1e-3, device="cpu", + on_epoch=lambda epoch, tl, vl, miou: epochs_seen.append(epoch) or epoch >= 2, + set_train_mode=built.set_train_mode, + ) + assert result["cancelled"] is True + assert result["epochs_completed"] == 2 + assert epochs_seen == [1, 2] + + def test_custom_optimizer_factory_is_used(self): + built = self._built() + seen = [] + + def make_optimizer(params, lr): + seen.append(lr) + return torch.optim.SGD(params, lr=lr) + + train_common.run_training_loop( + train_pairs=_synthetic_pairs(2), val_pairs=[], image_size=IMAGE_SIZE, n_classes=2, + epochs=1, batch_size=2, seed=0, flip_augment=False, to_tensor_fn=built.to_tensor_fn, + forward_fn=built.forward_fn, trainable_params=built.trainable_params, lr=5e-4, device="cpu", + make_optimizer_fn=make_optimizer, set_train_mode=built.set_train_mode, + ) + assert seen == [5e-4] diff --git a/backend/tests/test_train_e2e_real_ml.py b/backend/tests/test_train_e2e_real_ml.py new file mode 100644 index 0000000..23062f0 --- /dev/null +++ b/backend/tests/test_train_e2e_real_ml.py @@ -0,0 +1,174 @@ +"""Real (not mocked) end-to-end check of the core train/save/load contract. + +Skips automatically if the `ml` extra (torch/dlsia) isn't installed. When it +IS installed, this proves the ported train_common/dlsia_runtime/ +autoencoder_runtime machinery actually works — builds a real model, runs a +real training step, saves real weights, reloads them, and confirms inference +on the reloaded model reproduces the trained model's output — not just that +the error paths are clean (see test_train_common.py for those). +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import numpy as np +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +torch = pytest.importorskip("torch") +pytest.importorskip("dlsia") + +import autoencoder_runtime # noqa: E402 +import dlsia_runtime # noqa: E402 +import train_common # noqa: E402 +from schemas import DlsiaDenoiserConfig, DlsiaTunetConfig # noqa: E402 + + +@pytest.fixture() +def runs_dir(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Path: + monkeypatch.setenv("DINO_RUNS_DIR", str(tmp_path / "runs")) + return tmp_path / "runs" + + +def _synthetic_seg_pairs(n: int, size: int, n_classes: int, seed: int = 0): + rng = np.random.default_rng(seed) + pairs = [] + for _ in range(n): + rgb = rng.integers(0, 256, size=(size, size, 3), dtype=np.uint8) + label = rng.integers(0, n_classes, size=(size, size), dtype=np.uint8) + pairs.append((rgb, label)) + return pairs + + +def test_dlsia_tunet_train_save_reload_infer_round_trip(runs_dir: Path) -> None: + n_classes = 2 + image_size = 64 + model_cfg = DlsiaTunetConfig( + hyperparams={ + "epochs": 1, + "depth": 2, + "base_channels": 4, + "growth_rate": 1.2, + "batch_size": 2, + "image_size": image_size, + "tiling": False, + } + ) + + logs: list[str] = [] + built = train_common.build_family(model_cfg, n_classes, "cpu", logs.append) + assert built.set_train_mode is not None + + train_pairs = _synthetic_seg_pairs(4, image_size, n_classes) + metrics = train_common.run_training_loop( + train_pairs=train_pairs, + val_pairs=[], + image_size=image_size, + n_classes=n_classes, + epochs=model_cfg.hyperparams.epochs, + batch_size=model_cfg.hyperparams.batch_size, + seed=model_cfg.hyperparams.seed, + flip_augment=False, + to_tensor_fn=built.to_tensor_fn, + forward_fn=built.forward_fn, + trainable_params=built.trainable_params, + lr=model_cfg.hyperparams.lr, + device="cpu", + set_train_mode=built.set_train_mode, + ) + assert metrics["epochs_completed"] == 1 + assert metrics["cancelled"] is False + assert np.isfinite(metrics["final_train_loss"]) + + run_id = "test-tunet-run" + train_common.save_run( + run_id, + model_family=model_cfg.model_family, + model_config=built.model_config_snapshot, + classes=[{"classId": 1, "label": "a", "color": "#f00"}, {"classId": 2, "label": "b", "color": "#0f0"}], + render={}, + image_size=image_size, + hyperparams=model_cfg.hyperparams.model_dump(), + source_keys=["local:fake.tif"], + adapter_state=built.adapter_state_fn(), + metrics=metrics, + ) + + # list_runs / load_run_config round-trip. + runs = train_common.list_runs() + assert any(r["run_id"] == run_id for r in runs) + loaded_config = train_common.load_run_config(run_id) + assert loaded_config["model_family"] == "dlsia_tunet" + assert loaded_config["task"] == "segmentation" + + # Reload the ACTUAL saved weights into a fresh model and confirm inference + # runs and reproduces the just-trained model's output (not just "doesn't crash"). + loaded_state = train_common.load_adapter_state(run_id) + reloaded_model = dlsia_runtime.load_model(loaded_state, "cpu") + reloaded_model.eval() + reloaded_forward = dlsia_runtime.make_forward_fn(reloaded_model) + + # eval() on both sides: TUNet's BatchNorm gives different output in train + # mode (batch statistics) vs eval mode (running stats) — comparing them in + # mismatched modes would fail even with byte-identical weights. + built.set_train_mode(False) + sample_rgb, _ = train_pairs[0] + batch = built.to_tensor_fn(sample_rgb).unsqueeze(0) + with torch.no_grad(): + original_logits = built.forward_fn(batch) + reloaded_logits = reloaded_forward(batch) + assert original_logits.shape == (1, n_classes, image_size, image_size) + torch.testing.assert_close(original_logits, reloaded_logits) + + train_common.delete_run(run_id) + assert not any(r["run_id"] == run_id for r in train_common.list_runs()) + + +def test_dlsia_denoiser_cnn_ae_build_save_reload_round_trip(runs_dir: Path) -> None: + """Pure-torch denoiser architecture (no dlsia needed for this one) — build, + save, and reload real weights, confirming the autoencoder_runtime six-function + template round-trips correctly through train_common.build_family/save_run.""" + image_size = 64 + model_cfg = DlsiaDenoiserConfig( + architecture="cnn_ae", + training_scheme="ae", + ae_compression=4, + hyperparams={"depth": 2, "base_channels": 4, "image_size": image_size}, + ) + + built = train_common.build_family(model_cfg, 0, "cpu", lambda _msg: None) + + run_id = "test-ae-run" + train_common.save_run( + run_id, + model_family=model_cfg.model_family, + model_config={**built.model_config_snapshot, "training_scheme": "ae", "architecture": "cnn_ae"}, + classes=[], + render={}, + image_size=image_size, + hyperparams=model_cfg.hyperparams.model_dump(), + source_keys=[], + adapter_state=built.adapter_state_fn(), + metrics={"epochs_completed": 0, "cancelled": False}, + task="denoising", + ) + + loaded_config = train_common.load_run_config(run_id) + assert loaded_config["task"] == "denoising" + assert loaded_config["model_config"]["architecture"] == "cnn_ae" + + loaded_state = train_common.load_adapter_state(run_id) + reloaded = autoencoder_runtime.load_model(loaded_state, "cpu") + reloaded.eval() + reloaded_forward = autoencoder_runtime.make_forward_fn(reloaded) + + rng = np.random.default_rng(1) + gray = rng.integers(0, 256, size=(image_size, image_size), dtype=np.uint8) + to_tensor = autoencoder_runtime.make_to_tensor_fn() + batch = to_tensor(gray).unsqueeze(0) + with torch.no_grad(): + out = reloaded_forward(batch) + assert out.shape == (1, 1, image_size, image_size) diff --git a/backend/tests/test_train_jobs.py b/backend/tests/test_train_jobs.py new file mode 100644 index 0000000..b961fc1 --- /dev/null +++ b/backend/tests/test_train_jobs.py @@ -0,0 +1,202 @@ +"""Tests for train_jobs.py — the pure resume/validation helpers directly, and +run_train_job's early guard/error paths (lock contention, task/model +mismatch). The core train/save/load contract itself is proven for real (no +mocks) in test_train_e2e_real_ml.py; a full run_train_job segmentation round +trip would need real annotated sources assembled the way prepare_datasets +expects, which is out of scope for this pass. +""" +from __future__ import annotations + +import pytest + +import export_jobs +import train_common +import train_jobs +from schemas import ( + AnnotationClass, + DlsiaDenoiserConfig, + DlsiaTunetConfig, + ExportSourceItem, + TrainRequest, +) + + +def _classes(): + return [AnnotationClass(classId=1, label="A", color="#f00"), AnnotationClass(classId=2, label="B", color="#0f0")] + + +def _sources(): + return [ExportSourceItem(kind="local", source="fake.tif", slices={"0": []})] + + +def _tunet_request(**overrides): + defaults = dict(sources=_sources(), classes=_classes(), model=DlsiaTunetConfig()) + defaults.update(overrides) + return TrainRequest(**defaults) + + +def _denoiser_request(**overrides): + defaults = dict( + task="denoising", sources=_sources(), classes=[], + model=DlsiaDenoiserConfig(architecture="cnn_ae", training_scheme="ae", ae_compression=4), + ) + defaults.update(overrides) + return TrainRequest(**defaults) + + +class TestNewRunId: + def test_includes_model_family_and_is_unique(self): + a = train_jobs.new_run_id("dlsia_tunet") + b = train_jobs.new_run_id("dlsia_tunet") + assert "dlsia_tunet" in a + assert a != b + + +class TestCheckTaskMatchesModel: + def test_segmentation_task_with_denoiser_family_raises(self): + # classes=_classes() needed to get past the schema's own "segmentation + # needs >=1 class" validator so this reaches check_task_matches_model's + # family-specific check instead of failing at construction for an + # unrelated reason. + request = _denoiser_request(task="segmentation", classes=_classes()) + with pytest.raises(ValueError, match="cannot be trained as a"): + train_jobs.check_task_matches_model(request) + + def test_denoising_task_with_segmentation_family_raises(self): + with pytest.raises(ValueError, match="needs a denoiser model family"): + train_jobs.check_task_matches_model(_tunet_request(task="denoising")) + + def test_matching_pairs_do_not_raise(self): + train_jobs.check_task_matches_model(_tunet_request()) + train_jobs.check_task_matches_model(_denoiser_request()) + + +class TestCheckResumeCompatible: + def test_model_family_mismatch_raises(self): + parent = {"model_family": "dlsia_denoiser", "task": "segmentation", "classes": []} + with pytest.raises(ValueError, match="Cannot continue fine-tuning a dlsia_denoiser run"): + train_jobs.check_resume_compatible(parent, _tunet_request()) + + def test_task_mismatch_raises(self): + parent = {"model_family": "dlsia_tunet", "task": "denoising", "classes": []} + with pytest.raises(ValueError, match="Cannot continue fine-tuning a 'denoising'-task run"): + train_jobs.check_resume_compatible(parent, _tunet_request()) + + def test_denoiser_architecture_mismatch_raises(self): + parent = { + "model_family": "dlsia_denoiser", "task": "denoising", + "model_config": {"architecture": "tunet"}, + } + with pytest.raises(ValueError, match="Cannot continue fine-tuning a 'tunet' denoiser"): + train_jobs.check_resume_compatible(parent, _denoiser_request()) + + def test_denoiser_matching_architecture_skips_class_check(self): + parent = { + "model_family": "dlsia_denoiser", "task": "denoising", + "model_config": {"architecture": "cnn_ae"}, + } + train_jobs.check_resume_compatible(parent, _denoiser_request()) # no raise + + def test_denoiser_defaults_missing_architecture_to_tunet(self): + parent = {"model_family": "dlsia_denoiser", "task": "denoising"} + with pytest.raises(ValueError, match="Cannot continue fine-tuning a 'tunet' denoiser"): + train_jobs.check_resume_compatible(parent, _denoiser_request()) + + def test_class_list_changed_raises(self): + parent = { + "model_family": "dlsia_tunet", "task": "segmentation", + "classes": [{"label": "A"}, {"label": "C"}], + } + with pytest.raises(ValueError, match="classes changed since that run was trained"): + train_jobs.check_resume_compatible(parent, _tunet_request()) + + def test_class_list_case_and_whitespace_insensitive_match(self): + parent = { + "model_family": "dlsia_tunet", "task": "segmentation", + "classes": [{"label": " a "}, {"label": "B"}], + } + train_jobs.check_resume_compatible(parent, _tunet_request()) # no raise + + def test_missing_task_defaults_to_segmentation(self): + parent = {"model_family": "dlsia_tunet", "classes": [{"label": "a"}, {"label": "b"}]} + train_jobs.check_resume_compatible(parent, _tunet_request()) # no raise + + +class TestApplyParentArchitecture: + def test_tunet_inherits_geometry_and_topology(self): + request = _tunet_request() + parent = { + "hyperparams": {"image_size": 128, "tiling": True, "depth": 5, "base_channels": 8, "growth_rate": 1.5}, + } + train_jobs._apply_parent_architecture(parent, request) + hp = request.model.hyperparams + assert hp.image_size == 128 + assert hp.tiling is True + assert hp.depth == 5 + assert hp.base_channels == 8 + assert hp.growth_rate == 1.5 + + def test_tunet_inherits_denoise_settings(self): + request = _tunet_request() + parent = {"hyperparams": {}, "denoise": {"method": "gaussian", "strength": 0.5}} + train_jobs._apply_parent_architecture(parent, request) + assert request.denoise is not None + assert request.denoise.method == "gaussian" + + def test_tunet_clears_denoise_when_parent_had_none(self): + from schemas import DenoiseTrainOpts + request = _tunet_request(denoise=DenoiseTrainOpts(method="gaussian", strength=0.5)) + train_jobs._apply_parent_architecture({"hyperparams": {}}, request) + assert request.denoise is None + + def test_denoiser_cnn_ae_forces_scheme_and_compression(self): + request = _denoiser_request() + request.model.architecture = "tunet" # caller sent something else + parent = { + "hyperparams": {}, + "model_config": {"architecture": "cnn_ae", "ae_compression": 8}, + } + train_jobs._apply_parent_architecture(parent, request) + assert request.model.architecture == "cnn_ae" + assert request.model.training_scheme == "ae" + assert request.model.ae_compression == 8 + + def test_denoiser_defaults_missing_parent_architecture_to_tunet(self): + request = _denoiser_request() + train_jobs._apply_parent_architecture({"hyperparams": {}, "model_config": {}}, request) + assert request.model.architecture == "tunet" + + +class TestSourceKeys: + def test_tiled_and_local_formats(self): + request = _tunet_request( + sources=[ + ExportSourceItem(kind="tiled", source="a/b", server_uri="http://x:1", slices={}), + ExportSourceItem(kind="local", source="c/d", slices={}), + ], + ) + assert train_jobs._source_keys(request) == ["tiled:http://x:1:a/b", "local:c/d"] + + +class TestRunTrainJobGuards: + def test_reports_error_when_ml_lock_already_held(self): + jid = export_jobs.new_job("x") + assert train_common.ML_LOCK.acquire(blocking=False) + try: + train_jobs.run_train_job(jid, _tunet_request(), "run-1") + finally: + train_common.ML_LOCK.release() + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert "already running" in job["error"] + + def test_task_model_mismatch_reports_error_not_a_crash(self): + jid = export_jobs.new_job("x") + # task defaults to "segmentation" but the model is a denoiser -> caught + # by check_task_matches_model inside run_train_job's try block. + request = _tunet_request() + request.model = DlsiaDenoiserConfig(architecture="cnn_ae", training_scheme="ae") + train_jobs.run_train_job(jid, request, "run-2") + job = export_jobs.get_job(jid) + assert job["state"] == "error" + assert not train_common.ML_LOCK.locked() diff --git a/backend/tests/test_train_routes.py b/backend/tests/test_train_routes.py new file mode 100644 index 0000000..c843ee1 --- /dev/null +++ b/backend/tests/test_train_routes.py @@ -0,0 +1,117 @@ +"""HTTP-level tests for the /api/train/* routes. + +Route wiring, validation, and job-status polling for the parts that don't +need a real registered dataset (capability, runs list/delete, batch-size +probe). The core train/save/load contract itself is proven for real (no +mocks) in test_train_e2e_real_ml.py; a full train-via-HTTP round trip would +also need real Tiled/local array registration, which is exercised elsewhere +for export and is out of scope for this file. +""" + +from __future__ import annotations + +import asyncio +import time + +import pytest +from httpx import ASGITransport, AsyncClient + +import export_jobs +import train_common +from annotation_server import app + + +@pytest.fixture() +def runs_dir(tmp_path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("DINO_RUNS_DIR", str(tmp_path / "runs")) + return tmp_path / "runs" + + +async def _await_job(jid: str, timeout: float = 30.0) -> dict: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + job = export_jobs.get_job(jid) + assert job is not None + if job["state"] in ("done", "error"): + return job + await asyncio.sleep(0.05) + raise AssertionError(f"job {jid} did not finish within {timeout}s") + + +@pytest.mark.asyncio +async def test_capability_route(runs_dir) -> None: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/train/capability") + assert response.status_code == 200 + body = response.json() + assert "dinov3" not in body + assert "torch_available" in body + + +@pytest.mark.asyncio +async def test_runs_list_and_delete_unknown(runs_dir) -> None: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.get("/api/train/runs") + assert response.status_code == 200 + assert response.json() == {"runs": []} + + response = await client.delete("/api/train/runs/nonexistent") + assert response.status_code == 404 + + +@pytest.mark.asyncio +async def test_train_start_rejects_segmentation_with_no_classes(runs_dir) -> None: + """Pydantic validation (TrainRequest.validate_train_taxonomy) surfaces as 422.""" + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/train/start", + json={ + "task": "segmentation", + "sources": [{"kind": "local", "source": "fake.tif", "slices": {"0": []}}], + "classes": [], + "model": {"model_family": "dlsia_tunet"}, + }, + ) + assert response.status_code == 422 + + +@pytest.mark.asyncio +async def test_train_infer_unknown_run_is_404(runs_dir) -> None: + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post( + "/api/train/infer", + json={ + "run_id": "nonexistent", + "kind": "local", + "source": "fake.tif", + "slice_indices": [0], + }, + ) + if not train_common.torch_available(): + assert response.status_code == 503 + else: + assert response.status_code == 404 + + +@pytest.mark.asyncio +async def test_estimate_batch_probe_runs_for_real_when_torch_available(runs_dir) -> None: + """A real (not mocked) tiny batch-size probe, when torch is installed — + otherwise just confirms the clean 503 degradation.""" + payload = { + "model": { + "model_family": "dlsia_tunet", + "hyperparams": {"depth": 2, "base_channels": 4, "image_size": 64, "batch_size": 2}, + }, + "n_classes": 2, + } + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + response = await client.post("/api/train/estimate-batch", json=payload) + if not train_common.torch_available(): + assert response.status_code == 503 + return + assert response.status_code == 200 + jid = response.json()["job_id"] + job = await _await_job(jid) + assert job["state"] == "done" + assert job["result"]["suggested_batch_size"] is not None + assert job["result"]["largest_ok"] >= 1 diff --git a/backend/tests/test_volume_build.py b/backend/tests/test_volume_build.py new file mode 100644 index 0000000..740029c --- /dev/null +++ b/backend/tests/test_volume_build.py @@ -0,0 +1,132 @@ +"""Tests for building a 3-D volume from a slice stack already in Tiled. + +The registration half needs a live catalog, so what is pinned here is the part +that decides whether the feature is usable at all: what gets inspected, which +inputs are refused (and with a message that says what to do instead), and that +full resolution is deliberately not duplicated. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import numpy as np +import pytest +from fastapi import HTTPException + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import tiff_stack_source as tss # noqa: E402 +import volume_build # noqa: E402 + + +@pytest.fixture +def fake_stack(monkeypatch): + """Stand in for a Tiled stack: `resolve_array` + `array_shape_meta` + `read_slice`.""" + + # 4096 > TARGET_DIM (2048) by default, so an un-overridden fake_stack still + # represents "a large stack that needs a pyramid" — this used to be true at + # 1024 back when TARGET_DIM was 384. + def install(n_slices=64, height=4096, width=4096, dtype="uint16", is_rgb=False): + meta = { + "n_slices": n_slices, + "height": height, + "width": width, + "dtype": dtype, + "is_rgb": is_rgb, + } + monkeypatch.setattr(volume_build.arrays_mod, "resolve_array", lambda *a, **k: object()) + monkeypatch.setattr(volume_build.arrays_mod, "array_shape_meta", lambda *a, **k: meta) + # Slice i is filled with i, so a downsampled level's values are checkable. + monkeypatch.setattr( + volume_build.arrays_mod, + "read_slice", + lambda node, m, idx: np.full((height, width), idx, dtype=np.dtype(dtype)), + ) + return meta + + return install + + +class TestInspect: + def test_describes_the_pyramid_that_would_be_built(self, fake_stack): + fake_stack(n_slices=64, height=4096, width=4096) + info = volume_build.inspect_volume_build("browse/stack") + assert info["full_shape"] == [64, 4096, 4096] + assert info["pyramid_plan"] + assert info["already_small"] is False + + def test_reads_every_source_slice_once(self, fake_stack): + # The coarse levels cascade in memory, so the cost is one pass. + fake_stack(n_slices=690) + assert volume_build.inspect_volume_build("browse/stack")["slices_to_read"] == 690 + + def test_flags_a_stack_that_needs_no_downsampling(self, fake_stack): + fake_stack(n_slices=8, height=64, width=64) + info = volume_build.inspect_volume_build("browse/stack") + assert info["already_small"] is True + assert info["slices_to_read"] == 0 + + def test_refuses_a_single_image(self, fake_stack): + fake_stack(n_slices=1) + with pytest.raises(HTTPException) as excinfo: + volume_build.inspect_volume_build("browse/one") + assert excinfo.value.status_code == 422 + assert "single image" in excinfo.value.detail + + def test_refuses_colour_images(self, fake_stack): + # The renderer draws one scalar volume; RGB has no meaning there. + fake_stack(is_rgb=True) + with pytest.raises(HTTPException) as excinfo: + volume_build.inspect_volume_build("browse/rgb") + assert excinfo.value.status_code == 422 + assert "Colour" in excinfo.value.detail + + +class TestBuildGuards: + def test_refuses_non_tiled_sources(self): + with pytest.raises(HTTPException) as excinfo: + volume_build.build_volume("foo.tif", kind="local") + assert excinfo.value.status_code == 422 + + def test_refuses_an_already_small_stack_with_advice(self, fake_stack): + # No pyramid is needed, but the data is still per-slice and unstreamable. + # Saying only "nothing to do" would leave the user stuck. + fake_stack(n_slices=8, height=64, width=64) + with pytest.raises(HTTPException) as excinfo: + volume_build.build_volume("browse/small") + assert excinfo.value.status_code == 422 + assert "source images" in excinfo.value.detail + + def test_refuses_to_place_a_volume_at_the_catalog_root(self, fake_stack): + fake_stack() + with pytest.raises(HTTPException) as excinfo: + volume_build.build_volume("stack") + assert excinfo.value.status_code == 422 + assert "root" in excinfo.value.detail + + +class TestMetadataShape: + def test_omits_scale0_when_full_resolution_is_not_copied(self): + # Full resolution stays in the per-slice nodes. Declaring a scale0 that + # does not exist would make the viewer request a missing level. + plan = tss.pyramid_plan((690, 2560, 2560)) + ms = tss.multiscales_metadata("v", plan, include_scale0=False)["attributes"][ + "multiscales" + ][0] + paths = [d["path"] for d in ms["datasets"]] + assert all(p.startswith(f"{tss.PYRAMID_KEY}/") for p in paths) + assert "scale0" not in paths + + def test_scale_transforms_stay_relative_to_full_resolution(self): + # Dropping scale0 must not renumber the others: a level downsampled 16x + # still describes 16 full-res voxels per voxel, or the volume renders at + # the wrong physical size. + plan = tss.pyramid_plan((690, 2560, 2560)) + ms = tss.multiscales_metadata("v", plan, include_scale0=False)["attributes"][ + "multiscales" + ][0] + for dataset, level in zip(ms["datasets"], plan): + scale = dataset["coordinateTransformations"][0]["scale"] + assert scale == [float(f) for f in level["factor"]] diff --git a/backend/tests/test_volume_nodes.py b/backend/tests/test_volume_nodes.py new file mode 100644 index 0000000..f9b017b --- /dev/null +++ b/backend/tests/test_volume_nodes.py @@ -0,0 +1,136 @@ +"""Tests for locating a dataset's renderable 3-D volume. + +This is the module that answers "which node do I point the viewer at?". Getting +it wrong is not a subtle failure — the viewer reports +``openOmeZarr: missing multiscales in root .zattrs``, which tells the user +nothing about the real answer ("look at the sidecar" / "none has been built"). + +A fake catalog stands in for Tiled: the logic under test is metadata inspection +and path walking, both of which a real server would only slow down. +""" + +from __future__ import annotations + +import sys +from pathlib import Path + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import volume_nodes # noqa: E402 +from tiff_stack_source import VOLUME_SUFFIX # noqa: E402 + +MULTISCALES = {"attributes": {"multiscales": [{"datasets": [{"path": "scale0"}]}]}} + + +class FakeNode: + """A catalog node: metadata plus children, indexable like a Tiled container.""" + + def __init__(self, metadata=None, children=None): + self.metadata = metadata or {} + self._children = children or {} + + def __getitem__(self, key): + return self._children[key] + + +@pytest.fixture +def catalog(monkeypatch): + """A catalog with one of each shape the resolver has to tell apart.""" + root = FakeNode(children={ + "browse": FakeNode(children={ + # A registered Zarr volume: multiscale itself, with level children. + "zarr_vol": FakeNode( + metadata=dict(MULTISCALES, zarr_path="/data/zarr_vol.zarr"), + children={"scale0": FakeNode(children={"image": FakeNode()})}, + ), + # A TIFF stack: per-slice 2-D arrays, no multiscales anywhere. + "tiff_stack": FakeNode(metadata={"n_images": 604}), + # ...and the sidecar built for it. + f"tiff_stack{VOLUME_SUFFIX}": FakeNode( + metadata=dict(MULTISCALES, tiff_dir="/data/tiff_stack"), + ), + # A stack nobody has built a volume for. + "lonely_stack": FakeNode(metadata={"n_images": 12}), + }) + }) + monkeypatch.setattr(volume_nodes, "get_tiled_client", lambda *a, **k: root) + monkeypatch.setattr(volume_nodes, "api_key_for_uri", lambda *a, **k: None) + return root + + +class TestHasMultiscales: + def test_true_where_tiled_will_serve_it(self): + # Tiled's .zattrs route returns metadata["attributes"] verbatim, so that + # is the only place the viewer can see multiscales. + assert volume_nodes.has_multiscales(FakeNode(metadata=MULTISCALES)) + + def test_false_when_nested_anywhere_else(self): + # A common near-miss: right key, wrong level. It would never reach .zattrs. + assert not volume_nodes.has_multiscales( + FakeNode(metadata={"multiscales": [{"datasets": []}]}) + ) + + def test_false_for_empty_or_missing(self): + assert not volume_nodes.has_multiscales(FakeNode(metadata={})) + assert not volume_nodes.has_multiscales(FakeNode(metadata={"attributes": {}})) + assert not volume_nodes.has_multiscales( + FakeNode(metadata={"attributes": {"multiscales": []}}) + ) + + +class TestResolveVolume: + def test_open_node_is_already_a_volume(self, catalog): + result = volume_nodes.resolve_volume(None, "browse/zarr_vol") + assert result["mode"] == "self" + assert result["path"] == "browse/zarr_vol" + + def test_walks_up_from_a_pyramid_level(self, catalog): + # Opening `/scale0/image` is ordinary — that array has no + # multiscales of its own, but its grandparent does. + result = volume_nodes.resolve_volume(None, "browse/zarr_vol/scale0/image") + assert result["mode"] == "ancestor" + assert result["path"] == "browse/zarr_vol" + + def test_finds_the_sidecar_for_a_tiff_stack(self, catalog): + # The bug this module exists for: the open node is a container of 2-D + # slices, and the volume is its __volume sibling. + result = volume_nodes.resolve_volume(None, "browse/tiff_stack") + assert result["mode"] == "sidecar" + assert result["path"] == f"browse/tiff_stack{VOLUME_SUFFIX}" + + def test_reports_none_with_an_actionable_message(self, catalog): + result = volume_nodes.resolve_volume(None, "browse/lonely_stack") + assert result["mode"] == "none" + assert result["path"] is None + assert "built" in result["message"] + + def test_reports_none_for_a_missing_node(self, catalog): + result = volume_nodes.resolve_volume(None, "browse/does_not_exist") + assert result["mode"] == "none" + + def test_handles_an_empty_source(self, catalog): + result = volume_nodes.resolve_volume(None, "") + assert result["mode"] == "none" + assert result["message"] + + def test_tolerates_surrounding_slashes(self, catalog): + assert volume_nodes.resolve_volume(None, "/browse/zarr_vol/")["mode"] == "self" + + def test_never_ascends_past_the_catalog_root(self, catalog): + # Walking up must not wander off the top and resolve some unrelated node. + assert volume_nodes.resolve_volume(None, "browse")["mode"] == "none" + + +class TestSourceDir: + def test_surfaces_the_registered_tiff_directory(self, catalog): + result = volume_nodes.resolve_volume(None, "browse/tiff_stack") + assert result["source_dir"] == "/data/tiff_stack" + + def test_surfaces_the_registered_zarr_path(self, catalog): + result = volume_nodes.resolve_volume(None, "browse/zarr_vol") + assert result["source_dir"] == "/data/zarr_vol.zarr" + + def test_is_none_when_nothing_was_recorded(self): + assert volume_nodes.source_dir_of(FakeNode(metadata={})) is None diff --git a/backend/tests/test_zarr_source.py b/backend/tests/test_zarr_source.py new file mode 100644 index 0000000..39d69fe --- /dev/null +++ b/backend/tests/test_zarr_source.py @@ -0,0 +1,497 @@ +"""Tests for on-disk Zarr registration and multiscale slice resolution. + +Fixtures build tiny multiscale stores on the fly, so these run anywhere without +the multi-GB reference data. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import numpy as np +import pytest +from fastapi import HTTPException + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import arrays as arrays_mod # noqa: E402 +import zarr_source # noqa: E402 + + +def _make_multiscale(root: Path, shapes: list[tuple[int, int, int]]) -> Path: + """Write a v2 multiscale store shaped like the reference data. + + Layout mirrors the real volumes: ``scaleN/image`` arrays plus a + ``multiscales`` attribute naming them, so path discovery is exercised rather + than shape guessing. + """ + import zarr + + store = root / "vol.zarr" + group = zarr.open_group(str(store), mode="w") + for i, shape in enumerate(shapes): + sub = group.create_group(f"scale{i}") + arr = sub.create_array("image", shape=shape, dtype="float32", chunks=(2, 8, 8)) + arr[:] = np.linspace(0, 1, int(np.prod(shape)), dtype="float32").reshape(shape) + + base = shapes[0] + datasets = [] + for i, shape in enumerate(shapes): + factor = base[1] / shape[1] + datasets.append( + { + "path": f"scale{i}/image", + "coordinateTransformations": [ + {"type": "scale", "scale": [0.5 * factor, 0.5 * factor, 0.5 * factor]} + ], + } + ) + (store / ".zattrs").write_text( + json.dumps( + { + "multiscales": [ + { + "axes": [ + {"name": "z", "type": "space", "unit": "micrometer"}, + {"name": "y", "type": "space", "unit": "micrometer"}, + {"name": "x", "type": "space", "unit": "micrometer"}, + ], + "datasets": datasets, + } + ] + } + ) + ) + return store + + +@pytest.fixture +def pyramid(tmp_path: Path) -> Path: + # Deliberately NOT a clean power of two in z (9 -> 5 -> 3), matching the real + # data where 690 -> 172 gives a factor of 4.0116. + return _make_multiscale(tmp_path, [(9, 32, 32), (5, 16, 16), (3, 8, 8)]) + + +class TestInspect: + def test_lists_levels_finest_first(self, pyramid: Path) -> None: + info = zarr_source.inspect_zarr(str(pyramid)) + assert [lv["path"] for lv in info["levels"]] == [ + "scale0/image", + "scale1/image", + "scale2/image", + ] + assert info["full_shape"] == [9, 32, 32] + assert info["dtype"] == "float32" + + def test_reports_downsample_and_voxel_size(self, pyramid: Path) -> None: + info = zarr_source.inspect_zarr(str(pyramid)) + assert info["levels"][0]["downsample"] == [1.0, 1.0, 1.0] + assert info["levels"][1]["downsample"][1] == 2.0 + assert info["voxel_size"] == [0.5, 0.5, 0.5] + assert info["voxel_unit"] == "micrometer" + + def test_rejects_missing_path(self) -> None: + with pytest.raises(HTTPException) as exc: + zarr_source.inspect_zarr("/nope/missing.zarr") + assert exc.value.status_code == 404 + + def test_rejects_relative_path(self) -> None: + with pytest.raises(HTTPException) as exc: + zarr_source.inspect_zarr("relative/vol.zarr") + assert exc.value.status_code == 400 + + def test_rejects_non_zarr_directory(self, tmp_path: Path) -> None: + plain = tmp_path / "images" + plain.mkdir() + with pytest.raises(HTTPException) as exc: + zarr_source.inspect_zarr(str(plain)) + assert exc.value.status_code == 422 + assert "not a Zarr store" in exc.value.detail + + def test_rejects_zip_with_actionable_message(self, tmp_path: Path) -> None: + archive = tmp_path / "vol.zarr.zip" + archive.write_bytes(b"PK\x03\x04") + with pytest.raises(HTTPException) as exc: + zarr_source.inspect_zarr(str(archive)) + assert exc.value.status_code == 422 + assert "unzip" in exc.value.detail.lower() + + def test_rejects_group_with_no_arrays(self, tmp_path: Path) -> None: + import zarr + + empty = tmp_path / "empty.zarr" + zarr.open_group(str(empty), mode="w") + with pytest.raises(HTTPException) as exc: + zarr_source.inspect_zarr(str(empty)) + assert exc.value.status_code == 422 + assert "no 3-D arrays" in exc.value.detail + + def test_discovers_arrays_without_multiscales_metadata(self, tmp_path: Path) -> None: + import zarr + + store = tmp_path / "bare.zarr" + group = zarr.open_group(str(store), mode="w") + group.create_array("volume", shape=(4, 8, 8), dtype="float32", chunks=(2, 4, 4)) + info = zarr_source.inspect_zarr(str(store)) + assert [lv["path"] for lv in info["levels"]] == ["volume"] + assert info["full_shape"] == [4, 8, 8] + + def test_handles_a_bare_array_store_with_no_group_wrapper(self, tmp_path: Path) -> None: + """A store that's just a 3-D array at its own root (.zarray, no + .zgroup) — e.g. written via `zarr.open(mode='w', ...)` / `dask.array + .to_zarr()` with no OME-NGFF multiscale wrapper. `_is_zarr_dir` + already recognizes this shape; `inspect_zarr` must not reject it.""" + import zarr + + store = tmp_path / "plain.zarr" + arr = zarr.open(str(store), mode="w", shape=(6, 10, 12), dtype="uint16", zarr_format=2) + arr[:] = np.arange(6 * 10 * 12, dtype="uint16").reshape(6, 10, 12) + assert not (store / ".zgroup").exists() # confirm this really is group-less + + info = zarr_source.inspect_zarr(str(store)) + assert len(info["levels"]) == 1 + level = info["levels"][0] + assert level["path"] == "" + assert level["shape"] == [6, 10, 12] + assert level["downsample"] == [1.0, 1.0, 1.0] + assert info["full_shape"] == [6, 10, 12] + assert info["dtype"] == "uint16" + assert info["voxel_size"] is None + + def test_rejects_a_bare_array_that_is_not_3d(self, tmp_path: Path) -> None: + import zarr + + store = tmp_path / "plain2d.zarr" + zarr.open(str(store), mode="w", shape=(10, 12), dtype="uint16", zarr_format=2) + with pytest.raises(HTTPException) as exc: + zarr_source.inspect_zarr(str(store)) + assert exc.value.status_code == 422 + assert "not 3-D" in exc.value.detail + + +class TestScanAndRegisterZarrs: + """scan_and_register_zarrs's own logic (candidate discovery, shadow + detection, aggregation) — register_zarr itself needs a live Tiled server + (see TestPreflightZarr's own docstring), so it's mocked here rather than + re-proven. The shadow pre-check navigates a fake Tiled client (see + FakeContainer, defined below in this file) the same way preflight_zarr's + own tests do.""" + + def _fake_client(self, monkeypatch, browse_children=None): + client = FakeContainer({"browse": FakeContainer(browse_children or {})}) + monkeypatch.setattr(zarr_source, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(zarr_source, "api_key_for_uri", lambda uri: None) + return client + + def test_finds_only_top_level_zarr_dirs_and_registers_each(self, tmp_path: Path, monkeypatch) -> None: + import zarr + + for name in ("a.zarr", "b.zarr"): + store = tmp_path / name + zarr.open(str(store), mode="w", shape=(2, 4, 4), dtype="uint8", zarr_format=2) + (tmp_path / "not_a_store").mkdir() + (tmp_path / "readme.txt").write_text("hi") + # A Zarr store's OWN internals must never be treated as a second + # candidate — only immediate children of scan_root are considered. + (tmp_path / "a.zarr" / "nested.zarr").mkdir() + + self._fake_client(monkeypatch) + calls: list[str] = [] + + def fake_register_zarr(server_uri, path, container_path, description="", on_conflict="fail"): + calls.append(Path(path).name) + return {"key": Path(path).stem, "tiled_path": f"browse/{Path(path).stem}", "skipped": False} + + monkeypatch.setattr(zarr_source, "register_zarr", fake_register_zarr) + + result = zarr_source.scan_and_register_zarrs(None, str(tmp_path), "browse") + assert result["scanned"] == 2 + assert sorted(calls) == ["a.zarr", "b.zarr"] + assert sorted(r["name"] for r in result["registered"]) == ["a.zarr", "b.zarr"] + assert result["skipped"] == [] + assert result["shadowed"] == [] + assert result["errors"] == [] + + def test_reports_skipped_entries_separately_from_newly_registered(self, tmp_path: Path, monkeypatch) -> None: + import zarr + + zarr.open(str(tmp_path / "existing.zarr"), mode="w", shape=(2, 4, 4), dtype="uint8", zarr_format=2) + # A same-kind ("zarr") existing node — a legitimate re-scan match, not + # a shadow — so register_zarr's own on_conflict=skip path is what + # reports it, exactly as before. + self._fake_client(monkeypatch, {"existing": FakeContainer({}, metadata={"source_format": "zarr"})}) + + def fake_register_zarr(server_uri, path, container_path, description="", on_conflict="fail"): + return {"key": "existing", "tiled_path": "browse/existing", "skipped": True} + + monkeypatch.setattr(zarr_source, "register_zarr", fake_register_zarr) + + result = zarr_source.scan_and_register_zarrs(None, str(tmp_path), "browse") + assert result["registered"] == [] + assert result["skipped"] == ["existing"] + assert result["shadowed"] == [] + + def test_a_different_kind_collision_is_reported_as_shadowed_not_skipped( + self, tmp_path: Path, monkeypatch + ) -> None: + import zarr + + zarr.open(str(tmp_path / "existing.zarr"), mode="w", shape=(2, 4, 4), dtype="uint8", zarr_format=2) + self._fake_client( + monkeypatch, {"existing": FakeContainer({}, metadata={"source_format": "image-stack"})} + ) + register_calls: list[str] = [] + monkeypatch.setattr( + zarr_source, "register_zarr", + lambda *a, **kw: register_calls.append(a) or {"key": "existing", "tiled_path": "x", "skipped": False}, + ) + + result = zarr_source.scan_and_register_zarrs(None, str(tmp_path), "browse") + assert result["registered"] == [] + assert result["skipped"] == [] + assert result["shadowed"] == [ + {"name": "existing.zarr", "key": "existing", "existing_kind": "image-stack", "suggested_key": "existing_zarr"} + ] + # register_zarr must never be called for a shadowed candidate — no + # blind replace, no misleading "skip" of someone else's data. + assert register_calls == [] + + def test_renames_lets_a_shadowed_candidate_register_under_an_alternate_key( + self, tmp_path: Path, monkeypatch + ) -> None: + import zarr + + zarr.open(str(tmp_path / "existing.zarr"), mode="w", shape=(2, 4, 4), dtype="uint8", zarr_format=2) + self._fake_client( + monkeypatch, {"existing": FakeContainer({}, metadata={"source_format": "image-stack"})} + ) + calls = [] + + def fake_register_zarr(server_uri, path, container_path, on_conflict="fail"): + calls.append(container_path) + return {"key": "existing", "tiled_path": f"{container_path}/existing", "skipped": False} + + monkeypatch.setattr(zarr_source, "register_zarr", fake_register_zarr) + + result = zarr_source.scan_and_register_zarrs( + None, str(tmp_path), "browse", renames={"existing.zarr": "existing_zarr"} + ) + assert result["shadowed"] == [] + assert [r["key"] for r in result["registered"]] == ["existing_zarr"] + # Registered one level deeper, under a container named for the + # chosen alternate key — the only way to make Tiled's own + # filename-derived key land under a different name. + assert calls == ["browse/existing_zarr"] + assert result["registered"][0]["tiled_path"] == "browse/existing_zarr/existing" + + def test_one_bad_store_does_not_abort_the_rest(self, tmp_path: Path, monkeypatch) -> None: + import zarr + + for name in ("good.zarr", "bad.zarr"): + zarr.open(str(tmp_path / name), mode="w", shape=(2, 4, 4), dtype="uint8", zarr_format=2) + + self._fake_client(monkeypatch) + + def fake_register_zarr(server_uri, path, container_path, description="", on_conflict="fail"): + if Path(path).name == "bad.zarr": + raise HTTPException(502, "boom") + return {"key": "good", "tiled_path": "browse/good", "skipped": False} + + monkeypatch.setattr(zarr_source, "register_zarr", fake_register_zarr) + + result = zarr_source.scan_and_register_zarrs(None, str(tmp_path), "browse") + assert [r["name"] for r in result["registered"]] == ["good.zarr"] + assert result["errors"] == [{"name": "bad.zarr", "error": "boom"}] + + def test_rejects_relative_scan_root(self) -> None: + with pytest.raises(HTTPException) as exc: + zarr_source.scan_and_register_zarrs(None, "relative/dir", "browse") + assert exc.value.status_code == 400 + + def test_rejects_missing_scan_root(self, tmp_path: Path) -> None: + with pytest.raises(HTTPException) as exc: + zarr_source.scan_and_register_zarrs(None, str(tmp_path / "nope"), "browse") + assert exc.value.status_code == 404 + + def test_empty_directory_scans_cleanly_with_nothing_found(self, tmp_path: Path, monkeypatch) -> None: + self._fake_client(monkeypatch) + result = zarr_source.scan_and_register_zarrs(None, str(tmp_path), "browse") + assert result == {"scanned": 0, "registered": [], "skipped": [], "shadowed": [], "errors": []} + + def test_invalid_on_conflict_falls_back_to_skip(self, tmp_path: Path, monkeypatch) -> None: + import zarr + + zarr.open(str(tmp_path / "a.zarr"), mode="w", shape=(2, 4, 4), dtype="uint8", zarr_format=2) + self._fake_client(monkeypatch) + seen_on_conflict = [] + + def fake_register_zarr(server_uri, path, container_path, description="", on_conflict="fail"): + seen_on_conflict.append(on_conflict) + return {"key": "a", "tiled_path": "browse/a", "skipped": False} + + monkeypatch.setattr(zarr_source, "register_zarr", fake_register_zarr) + zarr_source.scan_and_register_zarrs(None, str(tmp_path), "browse", on_conflict="fail") + assert seen_on_conflict == ["skip"] + + +class TestRegisteredKey: + def test_strips_the_zarr_extension(self, pyramid: Path) -> None: + # Tiled keys on the stem; the collision check must use the same key or it + # silently misses an existing node and fails deep inside registration. + assert zarr_source.registered_key(pyramid) == "vol" + + +class _FakeArray: + def __init__(self, shape: tuple[int, ...]) -> None: + self.shape = shape + self.dtype = np.dtype("float32") + self.structure_family = "array" + + def __getitem__(self, idx: int) -> np.ndarray: + assert 0 <= idx < self.shape[0], f"slice {idx} out of range for {self.shape}" + return np.zeros(self.shape[1:], dtype="float32") + + +class _FakeContainer(dict): + structure_family = "container" + + def __init__(self, children: dict) -> None: + super().__init__(children) + + +def _fake_pyramid() -> _FakeContainer: + """Container mimicking a registered volume: scaleN -> {image: array}.""" + return _FakeContainer( + { + "scale0": _FakeContainer({"image": _FakeArray((9, 32, 32))}), + "scale1": _FakeContainer({"image": _FakeArray((5, 16, 16))}), + "scale2": _FakeContainer({"image": _FakeArray((3, 8, 8))}), + } + ) + + +class TestMultiscaleResolution: + def test_detects_pyramid_levels_in_order(self) -> None: + assert arrays_mod.multiscale_levels(_fake_pyramid()) == ["scale0", "scale1", "scale2"] + + def test_ignores_non_pyramid_containers(self) -> None: + stack = _FakeContainer({"a": _FakeArray((4, 4)), "b": _FakeArray((4, 4))}) + assert arrays_mod.multiscale_levels(stack) is None + + def test_selecting_the_volume_opens_the_finest_array(self) -> None: + # The generic walk would return `scale0` as a one-element stack, showing a + # whole 3-D volume as a single slice. + node = arrays_mod._descend_to_stack(_fake_pyramid()) + assert isinstance(node, _FakeArray) + assert node.shape == (9, 32, 32) + + +class TestFullResolutionCoordinates: + def _meta(self, level_shape: tuple[int, int, int], full: list[int]) -> dict: + pyramid = { + "level_key": "scale2", + "level_index": 2, + "level_count": 3, + "full_shape": full, + "z_downsample": full[0] / level_shape[0], + } + return arrays_mod.array_shape_meta(_FakeArray(level_shape), pyramid) + + def test_reports_finest_geometry_for_a_coarse_level(self) -> None: + meta = self._meta((3, 8, 8), [9, 32, 32]) + # Annotation coordinates must be full-resolution whichever level is shown. + assert (meta["n_slices"], meta["height"], meta["width"]) == (9, 32, 32) + assert (meta["level_n_slices"], meta["level_height"]) == (3, 8) + + def test_maps_full_resolution_index_onto_the_level(self) -> None: + level = _FakeArray((3, 8, 8)) + meta = self._meta((3, 8, 8), [9, 32, 32]) + # Every full-res index must land in range — _FakeArray asserts otherwise. + for idx in range(9): + assert arrays_mod.read_slice(level, meta, idx).shape == (8, 8) + + def test_handles_non_power_of_two_z_ratios(self) -> None: + # 690 -> 172 is 4.0116; a fixed integer factor would run off the end. + level = _FakeArray((172, 160, 160)) + meta = self._meta((172, 160, 160), [690, 2560, 2560]) + assert arrays_mod.read_slice(level, meta, 689).shape == (160, 160) + assert meta["width"] == 2560 + + def test_leaves_plain_volumes_untouched(self) -> None: + meta = arrays_mod.array_shape_meta(_FakeArray((7, 16, 16))) + assert (meta["n_slices"], meta["height"], meta["width"]) == (7, 16, 16) + assert "z_downsample" not in meta + + +class FakeContainer: + """Duck-typed fake Tiled container — enough surface for preflight_zarr's + navigation (`_walk`/`_child_keys`) without touching a real Tiled server.""" + + def __init__(self, children=None, metadata=None): + self._children = dict(children or {}) + self.metadata = metadata or {} + + def __iter__(self): + return iter(self._children) + + def __getitem__(self, key): + return self._children[key] + + def __len__(self): + return len(self._children) + + def keys(self): + return list(self._children.keys()) + + +class TestPreflightZarr: + """register_zarr itself needs a live Tiled server; preflight_zarr never + calls it — it only inspects the local store and navigates a fake Tiled + client, so it is fully testable without one.""" + + def test_no_collision_when_key_absent(self, pyramid: Path, monkeypatch) -> None: + client = FakeContainer({"browse": FakeContainer({})}) + monkeypatch.setattr(zarr_source, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(zarr_source, "api_key_for_uri", lambda uri: None) + result = zarr_source.preflight_zarr(None, str(pyramid), "browse") + assert result["exists"] is False + assert result["existing"] is None + assert result["key"] == zarr_source.registered_key(pyramid) + + def test_collision_reports_external_registration(self, pyramid: Path, monkeypatch) -> None: + key = zarr_source.registered_key(pyramid) + existing_node = FakeContainer( + {"scale0": object()}, + metadata={"source_format": "zarr", "sample_name": "vol", "n_images": 9}, + ) + client = FakeContainer({"browse": FakeContainer({key: existing_node})}) + monkeypatch.setattr(zarr_source, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(zarr_source, "api_key_for_uri", lambda uri: None) + result = zarr_source.preflight_zarr(None, str(pyramid), "browse") + assert result["exists"] is True + assert result["existing"]["external"] is True + assert result["existing"]["sample_name"] == "vol" + assert result["existing"]["n_images"] == 9 + + def test_collision_with_internally_managed_data_is_not_external(self, pyramid: Path, monkeypatch) -> None: + key = zarr_source.registered_key(pyramid) + existing_node = FakeContainer({"img_0000.tif": object()}, metadata={}) + client = FakeContainer({"browse": FakeContainer({key: existing_node})}) + monkeypatch.setattr(zarr_source, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(zarr_source, "api_key_for_uri", lambda uri: None) + result = zarr_source.preflight_zarr(None, str(pyramid), "browse") + assert result["existing"]["external"] is False + + def test_missing_target_container_reports_no_collision(self, pyramid: Path, monkeypatch) -> None: + client = FakeContainer({}) + monkeypatch.setattr(zarr_source, "get_tiled_client", lambda uri, key: client) + monkeypatch.setattr(zarr_source, "api_key_for_uri", lambda uri: None) + result = zarr_source.preflight_zarr(None, str(pyramid), "browse/missing") + assert result["exists"] is False + + def test_invalid_path_propagates_the_http_exception(self) -> None: + with pytest.raises(HTTPException) as exc: + zarr_source.preflight_zarr(None, "relative/path") + assert exc.value.status_code == 400 diff --git a/backend/tiff_stack_source.py b/backend/tiff_stack_source.py new file mode 100644 index 0000000..53b83ee --- /dev/null +++ b/backend/tiff_stack_source.py @@ -0,0 +1,702 @@ +"""Expose a directory of TIFF slices as a streamable 3-D Zarr volume. + +Why this exists +--------------- +Tiled 0.2.12 mounts a Zarr v2 router at ``/zarr/v2``: every array node is +already readable as a Zarr store, with ``.zarray`` synthesized from the Tiled +structure and ``/{i.j.k}`` serving a chunk. The 3-D viewer streams from that +directly — no export, no second endpoint. + +Two things stop a TIFF stack from working that way today: + +1. :mod:`ingest` writes one 2-D array **per file**, so ``/zarr/v2/`` is a + *group of N 2-D arrays* — not a volume. That layout is load-bearing for the + 2-D canvas (``arrays._stack_keys`` / ``read_slice``) and for per-frame Browse + metadata, so it stays exactly as it is. This module registers a 3-D **view + alongside** it; nothing here changes how slices are read. +2. Tiled's ``.zattrs`` returns ``metadata["attributes"]`` verbatim, so OME-NGFF + ``multiscales`` appears only if something writes it. This module writes it. + +Layout produced +--------------- +The same OME-NGFF shape :mod:`zarr_source` already documents, so both paths look +identical to the viewer. The volume is a **sidecar** of the dataset, keyed with a +``__volume`` suffix in the style of the existing ``__v_thumbs`` and ``__masks`` +siblings, because the per-slice container occupies the unsuffixed key:: + + / # untouched: the per-slice 2-D nodes Annotate reads + __volume/ # container; .zattrs carries "multiscales" + scale0 # (N, H, W) — the TIFF files, registered IN PLACE + pyramid/scale1 # downsampled, in an on-disk Zarr sidecar + pyramid/scale2 # ... + +Why the coarse levels go to disk rather than into Tiled +------------------------------------------------------- +Writing them with ``write_array`` would be less machinery, but Tiled 0.2.12's +``/zarr/v2`` chunk route only works for **externally-managed** arrays: for one it +writes itself, ``entry.read(slice=)`` reaches ``NDSlice.__getitem__`` with +a tuple and raises ``TypeError: tuple indices must be integers or slices``. The +``/api/v1`` block route serves the same array fine, so this is specific to the +Zarr façade — and the Zarr façade is the whole point here. Measured on a real +604 x 2560 x 2560 stack: ``scale0`` (external TIFF sequence) served a 19 MB chunk +in 0.2 s while every ``write_array`` level returned HTTP 500. + +Writing a normal Zarr store and registering it in place sidesteps that, and has +the same shape as :mod:`zarr_source`'s path — one fewer special case, not one +more. Revisit if the upstream route is fixed. + +Why a pyramid is not optional +----------------------------- +The renderer only lists levels that fit ``maxTextureDimension3D`` (commonly +2048) and its voxel budget. A single-level ``multiscales`` over a +2000x3232x3232 stack therefore offers *nothing renderable*, which looks +identical to a broken viewer. The coarse levels are the ones actually drawn; +``scale0`` is never fully read for the 3-D view, and earns its place by making +full-resolution ROI work possible later. +""" + +from __future__ import annotations + +import asyncio +import logging +import os +import shutil +from pathlib import Path +from typing import Any, Callable + +import numpy as np +from fastapi import HTTPException + +import ingest as ingest_mod +from tiled_clients import api_key_for_uri, get_tiled_client + +logger = logging.getLogger("tiff_stack_source") + +TIFF_SUFFIXES: tuple[str, ...] = (".tif", ".tiff") + +# Longest edge (voxels) the finest GENERATED level may have. 2048 matches +# WebGPU's guaranteed-minimum `maxTextureDimension3D` — the viewer picks +# whichever registered level is the coarsest that still fits the actual +# device's limit, so a level above what a given GPU can take is simply +# skipped, never a hard failure. Raising this raises fidelity but also +# memory/time to build it: this module assembles each generated level fully +# in RAM before writing it, so cost scales with actual voxel count +# (nx * ny * nz, not nx**3 — z is normally far smaller than the in-plane +# edges for a tomography stack). +TARGET_DIM = 2048 + +# How many levels to generate below scale0. Three gives the renderer a choice of +# detail without the cost growing: each is 1/8 the voxels of the one above. +GENERATED_LEVELS = 3 + +#: Key of the sub-group holding the generated levels, inside the volume node. +PYRAMID_KEY = "pyramid" + + +def pyramid_cache_root() -> Path: + """Directory the generated Zarr pyramids are written to. + + Deliberately *not* next to the source TIFFs: beamline reconstruction + directories are routinely read-only or on shared storage, and failing to + register a volume because the source's parent cannot be written to would be a + confusing way to discover that. Override with ``VOLUME_CACHE_DIR``. + """ + default = Path(__file__).resolve().parent.parent / ".tiled" / "volumes" + return Path(os.getenv("VOLUME_CACHE_DIR", str(default))).expanduser().resolve() + + +def _resolve_dir(raw: str) -> Path: + """Expand and validate a user-supplied absolute path to a TIFF directory. + + Raises: + HTTPException: 4xx with a user-facing message, so the UI never has to + surface a traceback. + """ + if not (raw or "").strip(): + raise HTTPException(400, "Enter the path to a directory of TIFF slices.") + path = Path(raw).expanduser() + if not path.is_absolute(): + raise HTTPException(400, f"Path must be absolute: {raw!r}") + if not path.exists(): + raise HTTPException(404, f"No such path: {path}") + if not path.is_dir(): + raise HTTPException(422, f"Not a directory: {path}") + return path.resolve() + + +def tiff_files(path: Path) -> list[Path]: + """Sorted TIFF files directly inside *path*. + + Sorted lexically, which is slice order for the zero-padded names tomography + reconstruction writes. Non-padded names (``img_2`` before ``img_10``) would + sort wrong — :func:`inspect_tiff_stack` rejects those rather than silently + building a shuffled volume. + """ + return sorted( + p for p in path.iterdir() if p.is_file() and p.suffix.lower() in TIFF_SUFFIXES + ) + + +def _zero_padding_is_consistent(files: list[Path]) -> bool: + """True when lexical order is numeric order. + + Only meaningful when the names carry numbers at all; a set of names with no + digits has nothing to get wrong, so it passes. + """ + import re + + numbers: list[int] = [] + for f in files: + match = re.search(r"(\d+)(?!.*\d)", f.stem) + if not match: + return True # no numbering scheme to violate + numbers.append(int(match.group(1))) + return numbers == sorted(numbers) + + +def pyramid_plan( + shape: tuple[int, int, int], + target_dim: int = TARGET_DIM, + levels: int = GENERATED_LEVELS, +) -> list[dict[str, Any]]: + """Downsample factors and shapes for the levels to generate below ``scale0``. + + Factors are **per axis**, and each is the smallest power of two bringing that + axis to ``target_dim`` or below. Powers of two keep the block mean exact and + the coordinate transforms simple. + + Per-axis rather than one shared factor because tomography stacks are often + strongly anisotropic — 4 slices of 4096x4096 is an ordinary shape. A single + factor chosen for the wide axes would ask for 4//16 = 0 slices, so *every* + level would be dropped as degenerate and the volume would end up with no + renderable level at all: indistinguishable, from the outside, from a broken + viewer. An axis therefore stops halving once it reaches 1. + + Args: + shape: ``(n_slices, height, width)`` of the full-resolution stack. + target_dim: Longest edge allowed for the finest generated level. + levels: How many levels to generate. + + Returns: + One dict per level with ``path``, ``factor`` (a 3-list, z/y/x) and + ``shape``. Empty when the source is already at or below ``target_dim`` on + every axis — there is then nothing to generate and ``scale0`` alone is + renderable. + """ + if any(s <= 0 for s in shape): + return [] + if max(shape) <= target_dim: + return [] + + def axis_factor(size: int, cap: int) -> int: + f = 1 + # `size // (f * 2) >= 1` is the guard that keeps a short axis alive. + while size / f > cap and size // (f * 2) >= 1: + f *= 2 + return f + + factors = [axis_factor(s, target_dim) for s in shape] + + plan: list[dict[str, Any]] = [] + for index in range(levels): + level_factors = [ + # Halve again only where the axis can still spare it; an axis already + # at 1 stays at 1 rather than dragging the level into degeneracy. + f * 2 if index and shape[axis] // (f * 2) >= 1 else f + for axis, f in enumerate(factors) + ] + factors = level_factors + level_shape = [shape[axis] // f for axis, f in enumerate(factors)] + if min(level_shape) < 1: + break + # A level identical to the one above adds bytes and no detail. + if plan and level_shape == plan[-1]["shape"]: + break + plan.append( + { + "path": f"scale{len(plan) + 1}", + "factor": list(factors), + "shape": level_shape, + } + ) + return plan + + +def multiscales_metadata( + name: str, + plan: list[dict[str, Any]], + voxel_size: tuple[float, float, float] = (1.0, 1.0, 1.0), + include_scale0: bool = True, +) -> dict[str, Any]: + """OME-NGFF ``multiscales`` for the generated levels, and optionally ``scale0``. + + Returned under an ``attributes`` key because Tiled's ``.zattrs`` route + returns ``metadata["attributes"]`` verbatim — put it anywhere else and the + viewer sees an empty attribute set and reports "missing multiscales". + + Args: + include_scale0: True when the full-resolution data is registered in place + as a ``scale0`` sibling (the TIFF-directory path). False when the + volume was built from data already in Tiled, where full resolution + stays in the per-slice nodes and is not duplicated — the viewer never + loads a level that large anyway. + """ + datasets = [] + if include_scale0: + datasets.append( + { + "path": "scale0", + "coordinateTransformations": [ + {"type": "scale", "scale": [float(v) for v in voxel_size]} + ], + } + ) + for level in plan: + factors = level["factor"] + datasets.append( + { + # Generated levels live in the sidecar sub-group, so the path the + # viewer follows is nested. Tiled serves nested groups fine — the + # Zarr volumes `zarr_source` registers use `scale0/image`. + "path": f"{PYRAMID_KEY}/{level['path']}", + "coordinateTransformations": [ + { + "type": "scale", + "scale": [float(v) * float(f) for v, f in zip(voxel_size, factors)], + } + ], + } + ) + return { + "attributes": { + "multiscales": [ + { + "version": "0.4", + "name": name, + "axes": [ + {"name": "z", "type": "space"}, + {"name": "y", "type": "space"}, + {"name": "x", "type": "space"}, + ], + "datasets": datasets, + } + ] + } + } + + +def block_mean(block: np.ndarray, fy: int, fx: int) -> np.ndarray: + """Mean-downsample a ``(k, H, W)`` block to ``(H//fy, W//fx)``. + + Averages across the whole block in z as well as over each ``fy x fx`` tile, + so one output voxel is the mean of every input voxel it covers. Trailing + rows/columns that do not fill a tile are dropped rather than partially + averaged — a partial tile would be brighter or darker than its neighbours + purely because of where the edge fell. + + Computed in float32 regardless of input dtype (uint16 sums overflow fast), + then cast back by the caller. + """ + if fy < 1 or fx < 1: + raise ValueError("downsample factors must be >= 1") + h = (block.shape[1] // fy) * fy + w = (block.shape[2] // fx) * fx + trimmed = block[:, :h, :w].astype(np.float32, copy=False) + return trimmed.reshape(trimmed.shape[0], h // fy, fy, w // fx, fx).mean(axis=(0, 2, 4)) + + +def inspect_tiff_stack(raw_path: str) -> dict[str, Any]: + """Describe a TIFF directory and the pyramid that would be built for it. + + Reads only the first file, so this is cheap enough to call on every keystroke + in a path field. + + Raises: + HTTPException: 4xx with a user-facing message for anything unusable. + """ + import tifffile + + path = _resolve_dir(raw_path) + files = tiff_files(path) + if not files: + raise HTTPException(422, f"No .tif/.tiff files directly inside {path.name!r}.") + if len(files) < 2: + raise HTTPException( + 422, + f"{path.name!r} holds a single image, not a volume. The 3-D view needs a stack.", + ) + if not _zero_padding_is_consistent(files): + raise HTTPException( + 422, + f"{path.name!r} has inconsistently numbered files (e.g. 'img_2' next to " + "'img_10'), so filename order is not slice order. Zero-pad the numbers " + "and try again — registering as-is would build a shuffled volume.", + ) + + try: + first = tifffile.imread(str(files[0])) + except Exception as exc: # noqa: BLE001 — surface as a clean 422 + raise HTTPException(422, f"Could not read {files[0].name!r}: {exc}") from exc + if first.ndim != 2: + raise HTTPException( + 422, + f"{files[0].name!r} is {first.ndim}-D; this path expects one 2-D slice per file.", + ) + + shape = (len(files), int(first.shape[0]), int(first.shape[1])) + plan = pyramid_plan(shape) + return { + "name": path.name, + "path": str(path), + "n_slices": shape[0], + "height": shape[1], + "width": shape[2], + "dtype": str(first.dtype), + "full_shape": list(shape), + "pyramid_plan": plan, + # Surfaced so the UI can say "this will take a while" honestly: every + # source slice is read once to build the levels. + "slices_to_read": shape[0] if plan else 0, + } + + +def _build_level( + files: list[Path], + factor: list[int], + out_shape: list[int], + dtype: np.dtype, + on_slice: Callable[[], None] | None = None, +) -> np.ndarray: + """Read the source stack and mean-downsample it by the per-axis *factor*. + + Reads at most ``factor[0]`` slices at a time, so peak memory is set by the + output level plus a handful of input slices — never by the whole stack, + which can be tens of GB. + """ + import tifffile + + fz, fy, fx = (int(f) for f in factor) + out = np.empty(out_shape, dtype=np.float32) + for z in range(out_shape[0]): + block = [] + for k in range(fz): + index = z * fz + k + if index >= len(files): + break + block.append(tifffile.imread(str(files[index]))) + if on_slice: + on_slice() + out[z] = block_mean(np.stack(block), fy, fx) + return _cast_like(out, dtype) + + +def downsample_array(volume: np.ndarray, factor: list[int]) -> np.ndarray: + """Mean-downsample an in-memory volume by the per-axis *factor*. + + Used to build each coarse level from the level above rather than from the + source. Because the factors are powers of two, averaging an average over + matching block sizes gives the same result as averaging the source directly, + while reading every source slice **once** instead of once per level. On a + 604 x 2560 x 2560 float32 stack that is the difference between ~16 GB of + reads and ~48 GB. + + (The two agree exactly only where the divisions are exact; a trailing partial + tile is dropped at each step, so the very last row/column of a level may + differ from a direct downsample. That is a sub-voxel edge effect on a preview + level, not something a viewer can show.) + """ + fz, fy, fx = (int(f) for f in factor) + z = (volume.shape[0] // fz) * fz + y = (volume.shape[1] // fy) * fy + x = (volume.shape[2] // fx) * fx + trimmed = volume[:z, :y, :x].astype(np.float32, copy=False) + return trimmed.reshape(z // fz, fz, y // fy, fy, x // fx, fx).mean(axis=(1, 3, 5)) + + +def write_pyramid_store(key: str, levels: dict[str, np.ndarray]) -> Path: + """Write the generated levels as a Zarr v2 group and return its path. + + One store per dataset, under :func:`pyramid_cache_root`, replacing any + previous build for the same key — a re-registration should not accumulate + stale copies of a volume. + + Written as Zarr **v2** to match what Tiled's ``/zarr/v2`` façade and the + viewer's OME-NGFF reader both speak. + """ + import zarr + + root = pyramid_cache_root() / key + root.mkdir(parents=True, exist_ok=True) + store_path = root / f"{PYRAMID_KEY}.zarr" + if store_path.exists(): + shutil.rmtree(store_path) + + group = zarr.open_group(str(store_path), mode="w", zarr_format=2) + for name, array in levels.items(): + group.create_array( + name, + shape=array.shape, + dtype=array.dtype, + # One chunk per output slice: matches how the viewer streams, and + # keeps any single request small. + chunks=(1, *array.shape[1:]), + )[:] = array + return store_path + + +def _cast_like(values: np.ndarray, dtype: np.dtype) -> np.ndarray: + """Cast a float working array back to the source dtype, rounding integers. + + Keeps the rendered volume in the same intensity units as the 2-D canvas + instead of silently becoming float — and rounds rather than truncating, so a + mean of 1.5 does not become 1. + """ + if np.issubdtype(dtype, np.integer): + info = np.iinfo(dtype) + return np.clip(np.rint(values), info.min, info.max).astype(dtype) + return values.astype(dtype) + + +#: Suffix marking the 3-D volume node as a sidecar of the dataset it describes. +#: Matches the existing ``__v_thumbs`` / ``__masks`` convention. +VOLUME_SUFFIX = "__volume" + + +def registered_key(path: Path) -> str: + """The node key for *path*'s 3-D volume. + + Suffixed, because the 2-D per-slice container that Annotate reads is keyed on + the very same directory name. Without the suffix the volume would land on top + of it — and since that container holds ingested data, the collision check + below would (correctly) refuse, making the 3-D view impossible for exactly + the datasets it is for. A sidecar key lets the two coexist, which is the + whole design: this is additive, and the 2-D read path is untouched. + + The stem comes from Tiled's own helper rather than assuming ``path.name`` — + the same care :mod:`zarr_source` takes, and for the same reason: a key + mismatch makes the collision check inspect a node that is not the one about + to be created, and Tiled then resolves the real collision itself, deep inside + registration where its only remedy is deletion. + """ + from tiled.client.register import Settings + + return f"{Settings.init().key_from_filename(path.name)}{VOLUME_SUFFIX}" + + +def preflight_tiff_stack( + server_uri: str | None, raw_path: str, container_path: str = "browse" +) -> dict[str, Any]: + """Report whether registering *raw_path* would collide, changing nothing.""" + path = _resolve_dir(raw_path) + key = registered_key(path) + client = get_tiled_client(server_uri, api_key_for_uri(server_uri)) + parts = [p for p in container_path.strip("/").split("/") if p] + target = ingest_mod._walk(client, parts) + existing = None + if target is not None and key in ingest_mod._child_keys(target): + node = target[key] + meta: dict[str, Any] = {} + try: + meta = dict(getattr(node, "metadata", {}) or {}) + except Exception: # noqa: BLE001 — best-effort description + pass + children: list[str] = [] + try: + children = list(node) + except Exception: # noqa: BLE001 + pass + existing = { + "child_count": len(children), + # Only a previous registration by THIS module is safe to replace: + # dropping it removes catalog rows and the generated levels, never + # the source TIFFs. Anything else may be internally-managed data, + # where deleting the node deletes the files. + "external": meta.get("source_format") == "tiff-stack-3d", + "sample_name": meta.get("sample_name") or "", + } + return {"key": key, "exists": existing is not None, "existing": existing} + + +def register_tiff_stack( + server_uri: str | None, + raw_path: str, + container_path: str = "browse", + description: str = "", + on_conflict: str = "fail", + progress: Callable[[str, int, int], None] | None = None, +) -> dict[str, Any]: + """Register a TIFF directory as a 3-D multiscale volume, copying no slices. + + ``scale0`` points at the TIFF files where they already are. Only the + generated (small) levels are written into Tiled. + + Args: + server_uri: Connected Tiled server URI. + raw_path: Absolute path to the directory of TIFF slices. + container_path: Slash-separated target container (e.g. ``browse``). + description: Optional keyword(s), as ingest stores them, so the volume is + filterable in Browse. + on_conflict: ``"fail"``, ``"replace"`` or ``"skip"`` when the key exists. + progress: Optional ``(message, done, total)`` callback for the job UI. + + Returns: + The inspection result plus the registered ``key`` and ``tiled_path``. + """ + info = inspect_tiff_stack(raw_path) + path = Path(info["path"]) + files = tiff_files(path) + plan: list[dict[str, Any]] = info["pyramid_plan"] + description = (description or "").strip() + keywords = ingest_mod.parse_keywords(description) + if on_conflict not in ingest_mod.ON_CONFLICT_MODES: + on_conflict = "fail" + + client = get_tiled_client(server_uri, api_key_for_uri(server_uri)) + parts = [p for p in container_path.strip("/").split("/") if p] + target = ingest_mod._ensure_container(client, parts) + + key = registered_key(path) + if key in ingest_mod._child_keys(target): + existing = preflight_tiff_stack(server_uri, raw_path, container_path)["existing"] + if on_conflict == "skip": + return {**info, "key": key, "tiled_path": "/".join([*parts, key]), "skipped": True} + if on_conflict == "replace": + if not (existing or {}).get("external"): + raise HTTPException( + 409, + f"{key!r} already exists in {container_path!r} and does not hold a " + "registered TIFF volume. Replacing it could delete uploaded data — " + "load into a different container, or remove that dataset yourself first.", + ) + target.delete_contents(key, recursive=True, external_only=False) + else: + raise HTTPException( + 409, + f"{key!r} already exists in {container_path!r}. Choose Replace or a " + "different destination.", + ) + + volume = target.create_container(key=key, metadata={}) + + # scale0: the files themselves, registered in place as ONE 3-D array. + # `register_image_sequence` builds a TiffSequenceAdapter over the sorted list + # and stores external Assets — shape (N, H, W), one chunk per slice, which is + # exactly the granularity a streaming volume viewer wants. + from tiled.client.register import Settings, register_image_sequence + + if progress: + progress("Registering full-resolution slices", 0, info["slices_to_read"]) + try: + asyncio.run(register_image_sequence(volume, "scale0", files, Settings.init())) + except Exception as exc: # noqa: BLE001 — classified for the UI + logger.warning("tiff sequence registration failed for %s: %s", path, exc) + raise HTTPException(502, ingest_mod._classify_error(exc)["message"]) from exc + + if ingest_mod._walk(volume, ["scale0"]) is None: + # register_image_sequence logs and swallows adapter errors, so a missing + # node here is the signal that registration did not actually happen. + raise HTTPException( + 502, + f"Tiled did not register the slices of {key!r}. Check the backend log for " + "the adapter error, and that the Tiled server can read this path.", + ) + + # Generated levels, built as a cascade: the finest from the TIFF files, each + # coarser one from the level above (already in memory, and tiny). Reading the + # source once instead of once per level is the difference between ~16 GB and + # ~48 GB of I/O on a 604 x 2560 x 2560 float32 stack. + dtype = np.dtype(info["dtype"]) + total_reads = info["slices_to_read"] + done = 0 + previous: np.ndarray | None = None + previous_factor: list[int] | None = None + built: dict[str, np.ndarray] = {} + + for level in plan: + label = f"Building {level['path']}" + if progress: + progress(label, done, total_reads) + + if previous is None or previous_factor is None: + def _tick() -> None: + nonlocal done + done += 1 + # Throttled: one update per 25 slices is smooth enough for a + # progress bar and keeps the job lock uncontended. + if progress and done % 25 == 0: + progress(label, done, total_reads) + + working = _build_level(files, level["factor"], level["shape"], dtype, _tick).astype( + np.float32, copy=False + ) + else: + relative = [ + int(level["factor"][axis] // previous_factor[axis]) for axis in range(3) + ] + working = downsample_array(previous, relative) + + built[level["path"]] = _cast_like(working, dtype) + previous, previous_factor = working, list(level["factor"]) + + # Write the generated levels as a plain on-disk Zarr store and register it in + # place, exactly as an externally-supplied Zarr volume would be — see this + # module's docstring for why `write_array` cannot be used here. + if built: + if progress: + progress("Writing pyramid", total_reads, total_reads) + sidecar = write_pyramid_store(key, built) + from tiled.client.register import register_single_item + + try: + asyncio.run( + register_single_item( + volume, sidecar, is_directory=True, settings=Settings.init() + ) + ) + except Exception as exc: # noqa: BLE001 — classified for the UI + logger.warning("pyramid registration failed for %s: %s", sidecar, exc) + raise HTTPException(502, ingest_mod._classify_error(exc)["message"]) from exc + + # A node existing is not enough. When the store sits outside Tiled's + # `readable_storage`, registration still creates the node — it just has + # no children, and every chunk request then 500s at read time, far from + # the cause. Check for contents, and name the likely fix. + pyramid_node = ingest_mod._walk(volume, [PYRAMID_KEY]) + if pyramid_node is None or not list(pyramid_node): + raise HTTPException( + 502, + f"Tiled registered no levels for {key!r} from {sidecar}. The usual " + "cause is that this path is not in the Tiled server's " + "`readable_storage` (see tiled/config.yml) — add it, or point " + "VOLUME_CACHE_DIR somewhere already readable.", + ) + + meta: dict[str, Any] = { + "sample_name": key, + "n_images": info["n_slices"], + "source_format": "tiff-stack-3d", + "tiff_dir": str(path), + "full_shape": info["full_shape"], + "pyramid_plan": plan, + **multiscales_metadata(key, plan), + } + if description: + meta["description"] = description + if keywords: + meta["keywords"] = keywords + try: + volume.update_metadata(metadata=meta) + except Exception as exc: # noqa: BLE001 — best-effort; the data is registered + logger.warning("could not set metadata on %s: %s", key, exc) + + if progress: + progress("Done", total_reads, total_reads) + + # Spread `info` FIRST so its filesystem "path" cannot shadow the Tiled path + # the caller needs to open the dataset. + return { + **info, + "key": key, + "tiled_path": "/".join([*parts, key]), + "skipped": False, + } diff --git a/backend/tiled_mask_sync.py b/backend/tiled_mask_sync.py index 35c71b2..be1ff78 100644 --- a/backend/tiled_mask_sync.py +++ b/backend/tiled_mask_sync.py @@ -25,12 +25,30 @@ import arrays as arrays_mod import export_jobs +import ipred_client +import mask_pyramid from coco_export import _safe_name, shape_to_mask from tiled_clients import api_key_for_uri, get_tiled_client logger = logging.getLogger(__name__) +def _read_predicted_label_map(run_id: str) -> np.ndarray: + """Fetch an ipred run's commit.png and decode it into the same ``(H, W)`` + uint8 array shape a rasterized shape would produce — pixel value is the + raw frontend classId directly (ipred's own convention, matching the + frontend's ``labelMapToPolygonShapes``), NOT yet remapped to this + export's legend ids; the caller does that remap the same way it already + does for real shapes. + """ + import io + + from PIL import Image + + png_bytes = ipred_client.run_commit_png(run_id) + return np.asarray(Image.open(io.BytesIO(png_bytes))) + + def _category_maps(classes: list[Any]) -> tuple[dict[int, int], dict[int, str], list[dict[str, Any]]]: """1-based COCO ids + legend, matching ``build_export_plan``'s convention.""" cat_id_map: dict[int, int] = {} @@ -56,15 +74,26 @@ def build_mask_volumes( there are no slices to write. ``semantic`` is ``(n,H,W)`` uint8 (class index, 0=bg); ``class_vols`` maps class label → ``(n,H,W)`` uint8 (0/255). Slices are every annotated key plus any ``negative_slices`` (emitted as all-zero frames - for hard negatives), in sorted numeric order. + for hard negatives) plus any ``predicted_slices`` key not already covered by + real shapes, in sorted numeric order. + + ``predicted_slices`` (see :class:`schemas.PredictedSlicePointer`) lets an + un-vectorized iPred volume-apply result go straight into the mask volume — + fetching its commit.png from the ipred service and rasterizing THAT, + instead of requiring the frontend to first trace it into polygon shapes + just so this function can immediately rasterize them back into a mask. + A slice with real shapes always wins over its predicted pointer (matches + the frontend's own precedence in ``handleCommitVolumeApply``/ + `usePredictedRasterStore`). """ h, w = int(meta["height"]), int(meta["width"]) cat_id_map, cat_id_to_name, legend = _category_maps(classes) slices: dict[str, list[dict[str, Any]]] = item.slices or {} neg = {str(k) for k in (item.negative_slices or [])} + predicted: dict[str, Any] = getattr(item, "predicted_slices", None) or {} keys = sorted( - {k for k, shapes in slices.items() if shapes} | neg, + {k for k, shapes in slices.items() if shapes} | neg | {k for k, p in predicted.items() if p}, key=lambda k: int(k), ) if not keys: @@ -77,14 +106,27 @@ def build_mask_volumes( for key in keys: label = np.zeros((h, w), dtype=np.uint8) acc: dict[str, np.ndarray] = {name: np.zeros((h, w), dtype=bool) for name in class_names} - for shape in slices.get(key, []): - shape_dict = shape if isinstance(shape, dict) else shape.model_dump() - mask = shape_to_mask(shape_dict, h, w) - if float(mask.sum()) < 1: - continue - cat_id = cat_id_map.get(int(shape_dict.get("classId", 1)), 1) - label[mask] = cat_id - acc[cat_id_to_name.get(cat_id, "")] |= mask + shapes_here = slices.get(key, []) + if shapes_here: + for shape in shapes_here: + shape_dict = shape if isinstance(shape, dict) else shape.model_dump() + mask = shape_to_mask(shape_dict, h, w) + if float(mask.sum()) < 1: + continue + cat_id = cat_id_map.get(int(shape_dict.get("classId", 1)), 1) + label[mask] = cat_id + acc[cat_id_to_name.get(cat_id, "")] |= mask + elif key in predicted and predicted[key]: + pointer = predicted[key] + run_id = pointer["run_id"] if isinstance(pointer, dict) else pointer.run_id + raw = _read_predicted_label_map(run_id) + for raw_id in np.unique(raw): + if raw_id == 0: + continue + cat_id = cat_id_map.get(int(raw_id), 1) + mask = raw == raw_id + label[mask] = cat_id + acc[cat_id_to_name.get(cat_id, "")] |= mask sem_list.append(label) for name in class_names: class_lists[name].append((acc[name] * 255).astype(np.uint8)) @@ -188,7 +230,10 @@ def _read_existing_masks(container: Any) -> dict[str, Any] | None: meta = dict(container.metadata) legend = meta.get("legend") or meta.get("classes") or [] indices = [int(i) for i in (meta.get("slice_indices") or [])] - semantic = np.asarray(container["semantic"][...]) + # "semantic" is now a registered multiscale node (mask_pyramid.py), not + # a flat array — the merge logic always operates on native resolution, + # never the downsampled viewer-only preview levels. + semantic = mask_pyramid.read_mask_scale0(container, "semantic") safe_to_name = {_safe_name(e["name"]): e["name"] for e in legend} class_arrays: dict[str, np.ndarray] = {} for key in list(container): @@ -214,13 +259,26 @@ def write_masks_to_tiled( server_uri: str | None, volumes: dict[str, Any], classes: list[Any], + container_suffix: str = "", ) -> dict[str, Any]: - """Merge stacked mask volumes into a ``__masks`` sibling container. + """Merge stacked mask volumes into a ``__masks`` + sibling container. Slices in this push overwrite the same index; previously-pushed slices are kept (merge). Metadata records ``updated_at`` and a per-slice ``slice_updated_at`` map plus ``last_updated_slices`` so the latest version of each slice is explicit. Returns ``{path, n_slices, updated, n_classes}``. + + ``container_suffix`` keeps independent producers of masks for the same + source from silently merging into one blob: the manual "sync masks to + Tiled" action and a dlsia run's "write to Tiled" (``infer_jobs.py``) both + call this function, and without a suffix they'd write the exact same + ``__masks`` container — a later push from one would merge onto + (and, on overlapping slices, overwrite) the other's, making it impossible + to keep both results around to compare, e.g. side-by-side in the 3-D + viewer's two independent mask layers. Empty by default (the manual sync + action's own container, unsuffixed, for backward compatibility with + anything already pointing at ``__masks``). """ api_key = api_key_for_uri(server_uri) client = get_tiled_client(server_uri, api_key) @@ -231,7 +289,7 @@ def write_masks_to_tiled( for part in parts[:-1]: parent = parent[part] - container_key = f"{stem}__masks" + container_key = f"{stem}__masks{container_suffix}" try: container: Any = parent[container_key] except KeyError: @@ -242,6 +300,19 @@ def write_masks_to_tiled( if existing is not None and existing["semantic"].shape[1:] != volumes["semantic"].shape[1:]: logger.warning("mask merge: shape changed for %s — replacing existing masks", source) existing = None + # slice_indices (container metadata) and the registered semantic array's own + # slice count can disagree — e.g. a prior interrupted/partial write, or stale + # metadata left over from before a fix landed — and merge_mask_volumes indexes + # the array positionally by `enumerate(slice_indices)`, so a longer metadata + # list than the array actually holds raises "index N is out of bounds for + # axis 0" deep inside the merge. Same "can't trust it, replace" contract as + # the H/W-mismatch guard above, rather than crashing the whole write. + if existing is not None and existing["semantic"].shape[0] != len(existing["slice_indices"]): + logger.warning( + "mask merge: slice_indices (%d) != stored semantic slices (%d) for %s — replacing existing masks", + len(existing["slice_indices"]), existing["semantic"].shape[0], source, + ) + existing = None merged = merge_mask_volumes(existing, volumes) @@ -270,9 +341,13 @@ def write_masks_to_tiled( container = parent.create_container(key=container_key, metadata=container_meta) dims = ["slice", "y", "x"] - container.write_array( - merged["semantic"], key="semantic", dims=dims, - metadata={"studio_type": "segmentation_semantic"}, + # Registered as a real OME-NGFF multiscale node (mask_pyramid.py), not a + # bare write_array — that's what lets the volume viewer's loadMask() + # actually open this as a Zarr store instead of rejecting it for + # "missing multiscales". Only `semantic` needs this: it's the one array + # the viewer's single combined class-id mask texture reads. + mask_pyramid.register_mask_pyramid( + merged["semantic"], key="semantic", container=container, cache_key=container_key, ) for name, vol in merged["class_vols"].items(): container.write_array( diff --git a/backend/tiling.py b/backend/tiling.py new file mode 100644 index 0000000..34ca144 --- /dev/null +++ b/backend/tiling.py @@ -0,0 +1,432 @@ +"""Patch-based ("tiled") training and inference for images larger than the model input. + +Without tiling, a slice is letterbox-rescaled to a single ``image_size`` square +(see :func:`train_common.letterbox`), so a 4096px slice at the default 512 is +downsampled 8× before the model ever sees it. Fine structure is unresolvable. + +Tiling instead cuts ``window``-sized windows at **native resolution**, runs the +model per window, and recombines. Training patches come out exactly +``window``x``window``, which :func:`train_common.letterbox` passes through +untouched, so the existing training loop, flip augmentation and validation work +unchanged. Inference blends overlapping windows into a full-resolution +prediction. + +Geometry comes from :mod:`qlty` (``NCYXQuilt``), which is already a dlsia +dependency and written by dlsia's author. Two of its properties matter here: + +* Windows overlap, and a window's outermost ring is the least reliable part of + its prediction (least surrounding context), so blending down-weights that + ring and training masks it out of the loss entirely. +* The last window in each axis is *clamped* to sit inside the image rather than + the image being zero-padded, so edge windows simply overlap more. + +Softmax is applied **after** recombining, never per window: averaging softmaxed +patches is not the softmax of averaged logits (qlty's own docs call this out). +""" + +from __future__ import annotations + +import logging +from typing import Any, Callable + +import numpy as np + +from train_common import IGNORE_INDEX + +logger = logging.getLogger(__name__) + +# Fraction of the window shared by consecutive windows. At 25% every interior +# pixel gets several predictions to average, which hides seams, while the window +# count only grows ~1/(1-f)² — 4/3× per axis here. +OVERLAP_FRACTION = 0.25 +# The outermost ``window // BORDER_DIVISOR`` ring of each window is the +# down-weighted/ignored border. At 25% overlap this ring is exactly one half +# of the overlap (`border_for(w) == (w - step_for(w)) // 2`), so a +# down-weighted edge is always covered by a neighbour's full-weight interior, +# which extends the other half of the way across the overlap. +BORDER_DIVISOR = 8 +# Weight given to border pixels when blending (interior is 1.0). Small but +# non-zero, so a border pixel still contributes where nothing else covers it. +BORDER_WEIGHT = 0.1 +# Windows per forward pass during inference. Keeps device memory bounded +# independently of how many windows the image yields. +INFER_TILE_BATCH = 4 + + +def step_for(window: int) -> int: + """Stride between consecutive windows.""" + return max(1, window - int(round(window * OVERLAP_FRACTION))) + + +def border_for(window: int) -> int: + """Width of the down-weighted/ignored ring at each window edge.""" + return max(1, window // BORDER_DIVISOR) + + +def tile_origins(dim: int, window: int, step: int) -> list[int]: + """Start offsets of every window along one axis. + + Mirrors qlty's own grid, including its clamped final window + (``min(i * step, dim - window)``) — so a streamed traversal lands on exactly + the windows ``NCYXQuilt.unstitch`` would produce. Requires ``dim >= window`` + (pad first, see :func:`pad_to_min`). + """ + if dim < window: + raise ValueError(f"dim {dim} smaller than window {window}; pad first") + full_steps = (dim - window) // step + count = full_steps + 2 if dim > full_steps * step + window else full_steps + 1 + return [min(i * step, dim - window) for i in range(count)] + + +def pad_to_min(arr: np.ndarray, min_h: int, min_w: int, fill: int) -> np.ndarray: + """Bottom/right-pad *arr* so it is at least ``min_h`` x ``min_w``. + + An image smaller than one window can't be tiled at all; padding to the + window (rather than upscaling) keeps every real pixel at native scale. The + padding is cropped back off after stitching, and for labels *fill* should be + :data:`train_common.IGNORE_INDEX` so it never reaches the loss. + """ + h, w = arr.shape[:2] + pad_h, pad_w = max(0, min_h - h), max(0, min_w - w) + if pad_h == 0 and pad_w == 0: + return arr + pad_width = [(0, pad_h), (0, pad_w)] + [(0, 0)] * (arr.ndim - 2) + return np.pad(arr, pad_width, mode="constant", constant_values=fill) + + +def _quilt(height: int, width: int, window: int) -> Any: + """``NCYXQuilt`` for one image's geometry, with this module's overlap/border.""" + from qlty.qlty2D import NCYXQuilt # noqa: PLC0415 — optional dependency + + step = step_for(window) + return NCYXQuilt( + Y=height, + X=width, + window=(window, window), + step=(step, step), + border=border_for(window), + border_weight=BORDER_WEIGHT, + ) + + +def qlty_available() -> bool: + """True if :mod:`qlty` can be imported (ships with dlsia; see pyproject's ml extra).""" + import importlib.util # noqa: PLC0415 + + return importlib.util.find_spec("qlty") is not None + + +# --------------------------------------------------------------------------- +# Training +# --------------------------------------------------------------------------- + + +def _tile_pair(rgb: np.ndarray, label: np.ndarray, window: int) -> list[tuple[np.ndarray, np.ndarray]]: + """Cut one full-resolution ``(rgb, label)`` pair into window-sized patches. + + Patches whose labels are entirely unannotated are dropped — with sparse + annotations most of an image is :data:`train_common.IGNORE_INDEX` and such + patches contribute nothing to the loss. + """ + import torch # noqa: PLC0415 + from qlty.cleanup import weed_sparse_classification_training_pairs_2D # noqa: PLC0415 + + img = pad_to_min(rgb, window, window, fill=0) + lbl = pad_to_min(label, window, window, fill=IGNORE_INDEX) + quilt = _quilt(img.shape[0], img.shape[1], window) + + # uint8 throughout: 4× less memory than float32 for the patch stack, and + # IGNORE_INDEX (255) is representable. Conversion to float happens per batch + # inside the training loop's to_tensor_fn. + img_t = torch.from_numpy(np.ascontiguousarray(img.transpose(2, 0, 1))).unsqueeze(0) + lbl_t = torch.from_numpy(np.ascontiguousarray(lbl)).unsqueeze(0) + + def _cut_and_weed(*, mask_borders: bool) -> tuple[Any, Any]: + """Cut patches, then drop those with no labelled pixel left. + + With *mask_borders*, each patch's border ring is blanked in the target so + the loss only supervises interiors — where the model has full context — + and only interior pixels count towards keeping a patch (``weed`` masks + its validity check by the same border tensor). + """ + patches_in, patches_out = quilt.unstitch_data_pair( + img_t, lbl_t, missing_label=IGNORE_INDEX if mask_borders else None + ) + border = quilt.border_tensor() if mask_borders else torch.ones(quilt.window) + return weed_sparse_classification_training_pairs_2D( + patches_in, patches_out, missing_label=IGNORE_INDEX, border_tensor=border + )[:2] + + kept_in, kept_out = _cut_and_weed(mask_borders=True) + + # An annotation lying entirely within the image's outermost ring falls in no + # patch's interior, so border-masking would discard the slice's only labels. + # Keep those patches whole rather than silently training on nothing. + if len(kept_in) == 0 and bool((lbl != IGNORE_INDEX).any()): + kept_in, kept_out = _cut_and_weed(mask_borders=False) + logger.info("Tiling: slice annotated only near its edge — kept %d unmasked patch(es)", len(kept_in)) + + return [ + ( + np.ascontiguousarray(kept_in[i].numpy().transpose(1, 2, 0)), + np.ascontiguousarray(kept_out[i].numpy()), + ) + for i in range(len(kept_in)) + ] + + +# Share of training patches held out for validation when there are no validation +# slices — matches the 0.1 "valid" ratio of the default auto-split. +VAL_HOLDOUT_FRACTION = 0.1 +# Below this many patches, holding any out costs more training signal than the +# metric is worth. +VAL_HOLDOUT_MIN_PATCHES = 8 + + +def holdout_val_patches( + datasets: dict[str, list[tuple[np.ndarray, np.ndarray]]], + *, + seed: int, + fraction: float = VAL_HOLDOUT_FRACTION, + min_patches: int = VAL_HOLDOUT_MIN_PATCHES, +) -> tuple[dict[str, list[tuple[np.ndarray, np.ndarray]]], int]: + """Move a deterministic share of training patches into an empty val split. + + Splits are assigned per *slice* before tiling, so annotating one or two + slices puts them all in train and the run reports no ``val_loss``/``val_miou`` + at all. Tiling turns those slices into many patches, which makes a holdout + possible — so a small, seeded share becomes validation data. + + Only fires when validation is otherwise **empty**: an explicit or auto-assigned + validation slice is always left alone. + + Caveat worth surfacing to the user: held-out patches come from the same + image(s) as the ones trained on, so the resulting mIoU is optimistic compared + with validating on a genuinely unseen slice. + + Returns: + ``(datasets, n_held_out)`` — the input unchanged with ``0`` when the + holdout doesn't apply. + """ + train = datasets.get("train") or [] + if datasets.get("val") or len(train) < min_patches: + return datasets, 0 + + n_val = max(1, int(round(len(train) * fraction))) + if n_val >= len(train): + return datasets, 0 + + order = np.random.default_rng(seed).permutation(len(train)) + val_idx = {int(i) for i in order[:n_val]} + return ( + { + **datasets, + "train": [p for i, p in enumerate(train) if i not in val_idx], + "val": [train[i] for i in sorted(val_idx)], + }, + n_val, + ) + + +def tile_datasets( + datasets: dict[str, list[tuple[np.ndarray, np.ndarray]]], + window: int, + progress_cb: Callable[[str], None] | None = None, + cancel_cb: Callable[[], bool] | None = None, +) -> dict[str, list[tuple[np.ndarray, np.ndarray]]] | None: + """Replace each full-resolution pair with its window-sized patches. + + Takes and returns :func:`train_common.prepare_datasets`' shape, so it drops + straight into the training pipeline. + + Checks *cancel_cb* (if given) between slices — tiling a large annotated set + can itself take a while, and previously ran to completion uncancellably with + no progress shown beyond one summary line per split at the very end. Returns + ``None`` if cancelled partway, mirroring :func:`predict_label_map_tiled`'s + cancellation contract so callers check for it the same way. + """ + out: dict[str, list[tuple[np.ndarray, np.ndarray]]] = {} + for split, pairs in datasets.items(): + patches: list[tuple[np.ndarray, np.ndarray]] = [] + for i, (rgb, label) in enumerate(pairs): + if cancel_cb is not None and cancel_cb(): + return None + before = len(patches) + patches.extend(_tile_pair(rgb, label, window)) + if progress_cb is not None: + progress_cb(f"{split} slice {i + 1}/{len(pairs)}: {len(patches) - before} patch(es)") + out[split] = patches + if progress_cb is not None: + progress_cb(f"{split}: {len(pairs)} slice(s) → {len(patches)} patch(es) of {window}px") + return out + + +# --------------------------------------------------------------------------- +# Inference +# --------------------------------------------------------------------------- + + +def _blend_tiled_forward( + rgb: np.ndarray, + *, + forward_fn: Callable[[Any], Any], + to_tensor_fn: Callable[[np.ndarray], Any], + window: int, + device: str, + cancel_cb: Callable[[], bool] | None = None, + progress_cb: Callable[[str], None] | None = None, +) -> Any | None: + """Blend per-window ``forward_fn`` output into a full-resolution canvas. + + This is the shared core behind :func:`predict_label_map_tiled` + (classification: softmax/argmax/confidence-thresholding layered on top) and + :func:`denoise_image_tiled` (regression: the blended canvas IS the result, + no further post-processing) — anything that needs qlty's tiled-window + blending but differs only in what it does with the blended output. + *forward_fn* is fully generic here: nothing classification- or + regression-specific lives in this function, only the tiling/accumulation + machinery. + + Windows are forwarded in small batches, so peak DEVICE memory depends on + :data:`INFER_TILE_BATCH` rather than on the image size — only one batch of + windows is ever on the GPU/MPS device at once. The weighted output canvas + they accumulate into lives on the host instead, precisely so it can scale + with image size x channel count without competing for device memory; for a + very large image or channel count that host allocation is still the actual + memory ceiling here, just not a device one. The accumulation is + arithmetically the same weighted mean ``NCYXQuilt.stitch`` computes. + + Returns a ``(C, orig_h, orig_w)`` float32 CPU tensor — *C* is whatever + ``forward_fn`` emits per window (n_classes for segmentation logits, 1 for a + single-channel denoiser) — or ``None`` if *cancel_cb* asked to stop partway. + """ + import torch # noqa: PLC0415 + + orig_h, orig_w = rgb.shape[:2] + padded = pad_to_min(rgb, window, window, fill=0) + height, width = padded.shape[:2] + + quilt = _quilt(height, width, window) + weight = quilt.weight # (window, window): 1.0 interior, BORDER_WEIGHT ring + step = step_for(window) + origins = [(y, x) for y in tile_origins(height, window, step) for x in tile_origins(width, window, step)] + + image = to_tensor_fn(padded) # (3, H, W) float, CPU + canvas: Any = None # (channels, H, W), allocated once the channel count is known + norm = torch.zeros((height, width), dtype=torch.float32) + + if progress_cb is not None: + progress_cb(f"{len(origins)} window(s) of {window}px @ {step}px step") + + for batch_start in range(0, len(origins), INFER_TILE_BATCH): + if cancel_cb is not None and cancel_cb(): + return None + batch_origins = origins[batch_start : batch_start + INFER_TILE_BATCH] # noqa: E203 + batch = torch.stack([image[:, y : y + window, x : x + window] for y, x in batch_origins]) # noqa: E203 + + out = forward_fn(batch.to(device)).detach().to("cpu", torch.float32) + if canvas is None: + canvas = torch.zeros((out.shape[1], height, width), dtype=torch.float32) + + for i, (y, x) in enumerate(batch_origins): + canvas[:, y : y + window, x : x + window] += out[i] * weight # noqa: E203 + norm[y : y + window, x : x + window] += weight # noqa: E203 + + if canvas is None: # unreachable for a real image (always ≥1 window) + raise RuntimeError("Tiled inference produced no windows") + + # Every pixel is covered by ≥1 window, but clamp anyway so a zero can never + # turn into inf/NaN and poison whatever runs on top of this. + blended = canvas / norm.clamp(min=1e-8) + return blended[:, :orig_h, :orig_w] # drop any padding + + +def predict_label_map_tiled( + rgb: np.ndarray, + *, + forward_fn: Callable[[Any], Any], + to_tensor_fn: Callable[[np.ndarray], Any], + window: int, + min_confidence: float, + device: str, + cancel_cb: Callable[[], bool] | None = None, + progress_cb: Callable[[str], None] | None = None, +) -> np.ndarray | None: + """Predict a full-resolution label map by blending per-window predictions. + + The blending itself — streamed accumulation into a weighted canvas, + arithmetically identical to ``NCYXQuilt.stitch`` — lives in + :func:`_blend_tiled_forward`, shared with :func:`denoise_image_tiled`. This + function only adds what's specific to classification: softmax is applied + **after** recombining, never per window (averaging softmaxed patches is + not the softmax of averaged logits — qlty's own docs call this out), then + argmax and confidence-thresholding turn it into a label map. + + Returns a ``(H, W)`` uint8 map using the pipeline's convention — + ``0`` = below ``min_confidence`` (background), ``1..n`` = class index + 1 — + or ``None`` if *cancel_cb* asked to stop partway. + """ + import torch # noqa: PLC0415 + import torch.nn.functional as F # noqa: PLC0415 + + blended = _blend_tiled_forward( + rgb, + forward_fn=forward_fn, + to_tensor_fn=to_tensor_fn, + window=window, + device=device, + cancel_cb=cancel_cb, + progress_cb=progress_cb, + ) + if blended is None: + return None + + probs = F.softmax(blended, dim=0) # after stitching — never per window + confidence, pred_class = probs.max(dim=0) + label = torch.where( + confidence >= min_confidence, + (pred_class + 1).to(torch.uint8), + torch.zeros_like(pred_class, dtype=torch.uint8), + ) + return label.numpy() + + +def denoise_image_tiled( + rgb: np.ndarray, + *, + forward_fn: Callable[[Any], Any], + to_tensor_fn: Callable[[np.ndarray], Any], + window: int, + device: str, + cancel_cb: Callable[[], bool] | None = None, + progress_cb: Callable[[str], None] | None = None, +) -> np.ndarray | None: + """Denoise a full-resolution image by blending per-window regression output. + + Same tiled-window blending as :func:`predict_label_map_tiled` — see + :func:`_blend_tiled_forward`, which both share — but *forward_fn* here is a + regression model (continuous-valued output, e.g. a trained denoiser) + rather than a classifier, so blending is the END of the pipeline: no + softmax, argmax or confidence-thresholding, all of which are meaningless + on continuous output. The blended canvas IS the denoised image. + + Used both for a live single-slice preview and (by a caller driving it once + per slice of a volume) a whole-volume "apply trained denoiser" bake job. + + Returns a float32 array at *rgb*'s original resolution — ``(H, W)`` when + *forward_fn* emits a single channel (the expected case for a denoiser), + else ``(C, H, W)`` — or ``None`` if *cancel_cb* asked to stop partway. + """ + blended = _blend_tiled_forward( + rgb, + forward_fn=forward_fn, + to_tensor_fn=to_tensor_fn, + window=window, + device=device, + cancel_cb=cancel_cb, + progress_cb=progress_cb, + ) + if blended is None: + return None + out = blended.numpy() + return out[0] if out.shape[0] == 1 else out diff --git a/backend/train_common.py b/backend/train_common.py new file mode 100644 index 0000000..264b628 --- /dev/null +++ b/backend/train_common.py @@ -0,0 +1,870 @@ +"""Shared plumbing for the Train tab. + +Two model families share this module's data prep, generic training loop, and +run persistence: + +* ``"dlsia_tunet"`` — dlsia's tunable U-Net trained from scratch, no + pretrained checkpoint (see :mod:`dlsia_runtime`). Needs ``torch`` + the + optional ``dlsia`` package. +* ``"dlsia_denoiser"`` — a self-supervised single-channel denoiser, either a + dlsia TUNet (:mod:`denoise_runtime`) or a plain convolutional autoencoder + (:mod:`autoencoder_runtime`). + +Both are optional dependencies (``backend/pyproject.toml``'s ``ml`` extra) — +this module is import-safe without torch installed; only the functions that +actually need it import it lazily and are guarded by :func:`torch_available`. + +DINOv3 LoRA fine-tuning is deliberately NOT supported here — see Phase 5.5 in +the integration plan. It needs vendored, Meta-licensed (research/non-commercial) +model code and checkpoint provisioning this repo doesn't have a story for yet. +""" + +from __future__ import annotations + +import importlib.util +import io +import json +import logging +import os +import shutil +from pathlib import Path +from typing import Any, Callable, NamedTuple + +import numpy as np +from fastapi import HTTPException +from PIL import Image as PILImage + +logger = logging.getLogger(__name__) + +# One fine-tune (or inference) job at a time — they contend for the same GPU +# memory, and neither this hand-rolled loop nor dlsia's TUNet is written to be +# safely reentrant across concurrent callers. +import threading # noqa: E402 + +ML_LOCK = threading.Lock() + +# Separate from ML_LOCK: ML_LOCK gates whole JOBS against each other (a +# training run, an inference run, a denoise bake, a batch probe — all four +# hold it for their entire duration so no two heavy ML jobs contend for the +# GPU at once). GPU_FORWARD_LOCK is finer-grained — it serializes only the +# actual model forward call, so a per-job worker pool (see infer_jobs.py's +# per-slice pool) can run I/O, preprocessing, and vectorization for multiple +# slices concurrently while still guaranteeing only one thread ever calls +# into the model at a given instant. Do not use this in place of ML_LOCK — +# it does not protect against two different JOBS running at once, only +# against two threads calling the model at the same time within one job. +GPU_FORWARD_LOCK = threading.Lock() + +IGNORE_INDEX = 255 # unannotated pixels — see coco_export.build_export_plan(lightly=True) + + +def torch_available() -> bool: + """True if ``torch`` is importable, without paying the import cost.""" + return importlib.util.find_spec("torch") is not None + + +def dlsia_available() -> bool: + """True if ``dlsia`` is importable, without paying the import cost.""" + return importlib.util.find_spec("dlsia") is not None + + +def pick_device() -> str | None: + """Return ``"mps"|"cuda"|"cpu"``, or ``None`` if torch is unavailable. + + ``TRAIN_DEVICE`` env var overrides auto-detection (e.g. to force ``cpu`` + on a machine where MPS is flaky for a particular op). + """ + if not torch_available(): + return None + override = (os.getenv("TRAIN_DEVICE") or "").strip().lower() + if override in {"mps", "cuda", "cpu"}: + return override + import torch # noqa: PLC0415 — optional dependency, imported lazily + + if torch.backends.mps.is_available(): + return "mps" + if torch.cuda.is_available(): + return "cuda" + return "cpu" + + +def runs_dir() -> Path: + """Server-owned directory holding saved fine-tune runs (both families).""" + configured = (os.getenv("DINO_RUNS_DIR") or "").strip() + if configured: + return Path(configured).expanduser().resolve() + local_root = Path(os.getenv("LOCAL_DATA_ROOT", "~/data")).expanduser().resolve() + return (local_root / "models" / "runs").resolve() + + +def _validate_run_id(run_id: str) -> str: + """Reject anything that isn't a single safe path component.""" + if not run_id or run_id in {".", ".."} or "/" in run_id or "\\" in run_id or "\x00" in run_id: + raise HTTPException(400, "Invalid run_id") + return run_id + + +def run_dir(run_id: str) -> Path: + """Resolved, containment-checked directory for one run.""" + base = runs_dir() + candidate = (base / _validate_run_id(run_id)).resolve() + if not candidate.is_relative_to(base): + raise HTTPException(400, "Invalid run_id") + return candidate + + +# --------------------------------------------------------------------------- +# Run persistence (shared shape across both families) +# --------------------------------------------------------------------------- + + +def save_run( + run_id: str, + *, + model_family: str, + model_config: dict[str, Any], + classes: list[dict[str, Any]], + render: dict[str, Any], + image_size: int, + hyperparams: dict[str, Any], + source_keys: list[str], + adapter_state: dict[str, Any], + metrics: dict[str, Any], + resumed_from: str | None = None, + task: str = "segmentation", + denoise: dict[str, Any] | None = None, +) -> None: + """Persist one run's config + adapter/weights + metrics to ``runs_dir()``. + + ``resumed_from`` records the run this one continued fine-tuning from, so a + chain of successive refinements stays traceable (a resume always writes a + NEW run — the parent is never modified). + + ``task`` distinguishes a segmentation run from a self-supervised denoiser + run (see ``schemas.TrainRequest.task``); defaults to ``"segmentation"`` + for callers that predate the denoiser family. + + ``denoise`` records any denoising applied to the model's INPUT pixels, so + inference can reapply exactly the same preprocessing off the run instead of + trusting the caller to remember it (see :func:`denoising_render_slice_fn`). + """ + import torch # noqa: PLC0415 + + d = run_dir(run_id) + d.mkdir(parents=True, exist_ok=True) + config = { + "run_id": run_id, + "model_family": model_family, + "task": task, + "model_config": model_config, + "classes": classes, + "render": render, + "denoise": denoise, + "image_size": image_size, + "hyperparams": hyperparams, + "source_keys": source_keys, + "resumed_from": resumed_from, + "created_at": _now_iso(), + } + (d / "config.json").write_text(json.dumps(config, indent=2)) + (d / "metrics.json").write_text(json.dumps(metrics, indent=2)) + torch.save(adapter_state, d / "adapter.pt") + + +def _now_iso() -> str: + from datetime import datetime, timezone + + return datetime.now(timezone.utc).isoformat() + + +def list_runs() -> list[dict[str, Any]]: + """List saved runs, newest first. Malformed run directories are skipped.""" + base = runs_dir() + if not base.exists(): + return [] + results: list[dict[str, Any]] = [] + for d in base.iterdir(): + if not d.is_dir(): + continue + try: + config = json.loads((d / "config.json").read_text()) + # Runs saved before `task` existed have no such key on disk — treat + # that as "segmentation", the only thing this app trained before. + config.setdefault("task", "segmentation") + metrics = json.loads((d / "metrics.json").read_text()) if (d / "metrics.json").exists() else {} + results.append({**config, "metrics": metrics}) + except Exception as exc: # noqa: BLE001 — one bad run dir must not break the list + logger.warning("Skipping malformed run directory %s: %s", d, exc) + continue + results.sort(key=lambda r: r.get("created_at", ""), reverse=True) + return results + + +def delete_run(run_id: str) -> None: + """Permanently remove a saved run's directory (config, metrics, weights).""" + d = run_dir(run_id) + if not d.is_dir(): + raise HTTPException(404, f"Unknown run: {run_id!r}") + shutil.rmtree(d) + + +def load_run_config(run_id: str) -> dict[str, Any]: + """Return the parsed ``config.json`` for *run_id*, or raise 404. + + ``task`` is a field added after this app already had saved runs on disk — + an absent key means the run predates the field, and every run trained + before it existed was a segmentation run, so that is the default filled + in here rather than leaving callers (e.g. ``check_resume_compatible``) to + each re-derive the same backward-compat fallback. + """ + d = run_dir(run_id) + config_path = d / "config.json" + if not config_path.exists(): + raise HTTPException(404, f"Unknown run: {run_id!r}") + try: + config = json.loads(config_path.read_text()) + except json.JSONDecodeError as exc: + raise HTTPException(500, "Run config is corrupt") from exc + config.setdefault("task", "segmentation") + return config + + +def load_adapter_state(run_id: str) -> dict[str, Any]: + """Load the saved ``adapter.pt`` state dict for *run_id* onto the CPU.""" + import torch # noqa: PLC0415 + + d = run_dir(run_id) + adapter_path = d / "adapter.pt" + if not adapter_path.exists(): + raise HTTPException(404, f"No saved weights for run: {run_id!r}") + return torch.load(adapter_path, map_location="cpu", weights_only=False) + + +# --------------------------------------------------------------------------- +# Training-data preparation (shared: reuses the Lightly export path) +# --------------------------------------------------------------------------- + + +def denoising_render_slice_fn(denoise: Any) -> Callable[..., np.ndarray]: + """A ``render_slice``-compatible callable that denoises the RAW slice first. + + The single place model-input denoising is applied, used by BOTH training + (:func:`prepare_datasets`) and inference (``infer_jobs``). Keeping one + implementation is the point: a model must see the same pixel distribution at + predict time that it saw during training, and two copies of this would be + free to drift apart silently — the failure mode is not an error, just + quietly worse predictions. + + Denoising runs on raw intensity units, before ``render_slice`` normalises, + matching where the Annotate preview applies it. + + Args: + denoise: Anything with ``.method``/``.strength`` (a + :class:`schemas.DenoiseTrainOpts`) or the equivalent dict as + reloaded from a run's ``config.json``. Falsy, or a ``"none"`` + method, yields plain ``images.render_slice``. + """ + import images as images_mod # noqa: PLC0415 + + method, strength = _denoise_params(denoise) + if method is None: + return images_mod.render_slice + + import denoise as denoise_mod # noqa: PLC0415 + + def _render(arr: np.ndarray, opts: dict[str, Any], global_range: Any = None) -> np.ndarray: + # Colour sources are left alone: the classical filters are 2-D + # grayscale, and render_slice early-returns for RGB anyway. + if arr.ndim == 2: + arr = denoise_mod.denoise_slice(arr, method, strength) + return images_mod.render_slice(arr, opts, global_range) + + return _render + + +def _denoise_params(denoise: Any) -> tuple[str | None, float]: + """Normalise a DenoiseTrainOpts / dict / None into ``(method, strength)``. + + Returns ``(None, ...)`` when no denoising should be applied — including for + ``"model"``, which would need its own run and a GPU pass per slice and is + deliberately not supported as a preprocessor for another model. + """ + if not denoise: + return None, 0.5 + if hasattr(denoise, "method"): + method, strength = denoise.method, float(denoise.strength) + else: + method, strength = denoise.get("method"), float(denoise.get("strength", 0.5)) + if not method or method in ("none", "model"): + return None, strength + return method, strength + + +def prepare_datasets( + sources: list[Any], + classes: list[Any], + render: Any, + auto_split: dict[str, Any], + progress_cb: Callable[[str], None] | None = None, + denoise: Any = None, +) -> dict[str, list[tuple[np.ndarray, np.ndarray]]]: + """Render + rasterize every annotated source into in-memory train/val pairs. + + Reuses ``coco_export.build_export_plan(lightly=True)`` per source — the + same in-memory rendering/rasterization the Lightly export already does — + so no disk export is needed just to assemble training tensors. Valid + + test folds into val (same convention as the Lightly export). + + Returns: + ``{"train": [(rgb_uint8_hwc, label_uint8_hw), ...], "val": [...]}``. + """ + import arrays as arrays_mod + import images as images_mod + from coco_export import build_export_plan + from schemas import ExportRequest + + train_pairs: list[tuple[np.ndarray, np.ndarray]] = [] + val_pairs: list[tuple[np.ndarray, np.ndarray]] = [] + + for item in sources: + node = arrays_mod.resolve_array(item.source, item.kind, item.server_uri) + shim = ExportRequest( + kind=item.kind, + source=item.source, + server_uri=item.server_uri, + slices=item.slices, + split_by_slice=item.split_by_slice, + negative_slices=item.negative_slices, + classes=classes, + render=render, + auto_split=auto_split, + ) + plan = build_export_plan( + node, + shim, + render_slice_fn=denoising_render_slice_fn(denoise), + array_shape_meta_fn=arrays_mod.array_shape_meta, + read_slice_fn=arrays_mod.read_slice, + sample_global_stats_fn=images_mod._sample_global_stats, + progress_cb=progress_cb, + lightly=True, + ) + for split_name, split_data in plan["splits"].items(): + bucket = train_pairs if split_name == "train" else val_pairs + for img in split_data["images"]: + rgb = np.asarray(PILImage.open(io.BytesIO(img["png_bytes"])).convert("RGB")) + label = np.asarray(PILImage.open(io.BytesIO(img["label_png_bytes"]))) + bucket.append((rgb, label)) + + return {"train": train_pairs, "val": val_pairs} + + +def letterbox(image: np.ndarray, label: np.ndarray, size: int) -> tuple[np.ndarray, np.ndarray]: + """Resize-keep-aspect + pad *image* (uint8 HWC) and *label* (uint8 HW) to + a ``size`` x ``size`` square. Label padding uses :data:`IGNORE_INDEX` so + padded pixels never contribute to the loss. + + Also returns enough to invert the transform (see :func:`unletterbox`), + packed into the returned label's dtype-preserving companion — callers that + need to invert should use :func:`letterbox_params` directly instead of + re-deriving it, since floating-point resize choices must match exactly. + """ + h, w = image.shape[:2] + if (h, w) == (size, size) and label.shape[:2] == (size, size): + # Already exactly the target square — every tiled-training patch hits + # this. The full computation below is a no-op here anyway (scale=1, + # zero padding), just paid for with 4 array copies per patch per epoch; + # skip straight to returning the inputs. Safe to hand back unchanged + # (not a defensive copy): callers only ever read these, or derive a NEW + # array via reversal/dtype-cast, never mutate in place. + return image, label + scale = min(size / h, size / w) + nh, nw = max(1, round(h * scale)), max(1, round(w * scale)) + + img_pil = PILImage.fromarray(image).resize((nw, nh), PILImage.Resampling.BILINEAR) + lbl_pil = PILImage.fromarray(label).resize((nw, nh), PILImage.Resampling.NEAREST) + + out_img = np.zeros((size, size, image.shape[2]), dtype=np.uint8) + out_lbl = np.full((size, size), IGNORE_INDEX, dtype=np.uint8) + top, left = (size - nh) // 2, (size - nw) // 2 + out_img[top : top + nh, left : left + nw] = np.asarray(img_pil) # noqa: E203 + out_lbl[top : top + nh, left : left + nw] = np.asarray(lbl_pil) # noqa: E203 + return out_img, out_lbl + + +def letterbox_params(h: int, w: int, size: int) -> dict[str, int]: + """Return the placement used by :func:`letterbox` for an ``h``x``w`` image, + so a prediction on the ``size``x``size`` canvas can be cropped/resized back + to the original resolution (see :func:`unletterbox`).""" + scale = min(size / h, size / w) + nh, nw = max(1, round(h * scale)), max(1, round(w * scale)) + top, left = (size - nh) // 2, (size - nw) // 2 + return {"top": top, "left": left, "nh": nh, "nw": nw} + + +def unletterbox(label: np.ndarray, orig_h: int, orig_w: int, size: int) -> np.ndarray: + """Invert :func:`letterbox` on a predicted label map: crop the padding, + then nearest-resize back to ``(orig_h, orig_w)``.""" + p = letterbox_params(orig_h, orig_w, size) + cropped = label[p["top"] : p["top"] + p["nh"], p["left"] : p["left"] + p["nw"]] # noqa: E203 + if (p["nh"], p["nw"]) == (orig_h, orig_w): + return cropped + return np.asarray(PILImage.fromarray(cropped).resize((orig_w, orig_h), PILImage.Resampling.NEAREST)) + + +# --------------------------------------------------------------------------- +# Model construction (shared: the one place a model gets built from a +# schemas.ModelConfig, used by both train_jobs.py and batch_probe.py) +# --------------------------------------------------------------------------- + + +class BuiltFamily(NamedTuple): + """Everything :func:`run_training_loop` (or the batch-size probe) needs + from a constructed model, plus enough to save it afterward. + + ``adapter_state_fn`` defers building the saved-weights dict (the dlsia + network dict) until it's actually called: only :mod:`train_jobs` calls + it, after training completes; :mod:`batch_probe` never saves anything and + simply ignores the field. + """ + + forward_fn: Callable[["Any"], "Any"] + to_tensor_fn: Callable[["Any"], "Any"] + trainable_params: list["Any"] + set_train_mode: Callable[[bool], None] | None + model_config_snapshot: dict[str, Any] + adapter_state_fn: Callable[[], dict[str, Any]] + + +def denoiser_runtime_for(config: dict[str, Any]) -> "Any": + """The runtime module that can load *config*'s saved denoiser. + + Both inference paths (the single-slice canvas preview in + ``annotation_server`` and the whole-volume bake in ``denoise_bake``) need + this, and both need the same backward-compatible default — hence one + implementation rather than two that could drift. + + A run's ``model_config["architecture"]`` names the network. Runs saved + before that field existed have no such key and are dlsia TUNets by + definition, so an absent value MUST default to ``"tunet"`` or every existing + run on disk becomes unloadable. + + Raises: + ValueError: the run names an architecture this build doesn't know (e.g. + a run produced by a newer version). + """ + architecture = (config.get("model_config") or {}).get("architecture", "tunet") + if architecture == "cnn_ae": + import autoencoder_runtime # noqa: PLC0415 + + return autoencoder_runtime + if architecture == "tunet": + import denoise_runtime # noqa: PLC0415 + + return denoise_runtime + raise ValueError(f"Unknown denoiser architecture on this run: {architecture!r}") + + +def denoiser_needs_dlsia(config: dict[str, Any]) -> bool: + """Whether *config*'s denoiser architecture requires the dlsia dependency. + + Only the TUNet architecture does; the convolutional autoencoder is plain + torch. Lets the inference gates refuse for the right reason instead of + demanding dlsia for a network that never touches it. + """ + return (config.get("model_config") or {}).get("architecture", "tunet") == "tunet" + + +def build_family( + model_cfg: "Any", + n_classes: int, + device: str, + log_cb: Callable[[str], None], + init_state: dict[str, Any] | None = None, +) -> BuiltFamily: + """Build the model, forward pass, and tensor conversion for *model_cfg*'s + family. + + This is the single place a model gets constructed from hyperparameters — + :mod:`train_jobs` and :mod:`batch_probe` both call it, so the probe's + "builds exactly like training does" claim is enforced by sharing code + rather than by keeping two copies in sync by hand. + + ``init_state`` is a saved run's ``adapter.pt`` to warm-start from (continue + fine-tuning) instead of starting from scratch. Optimizer/scheduler state is + not saved and so is not restored — a resume gets a fresh AdamW and a cosine + schedule starting again at ``lr``, which is the normal fine-tune-again + behaviour. + + Every recognized ``schemas.ModelConfig`` member has its own explicit + branch below; anything else raises :class:`ValueError` rather than being + silently treated as a dlsia TUNet — that used to be this function's + implicit ``else``, which would happily (and wrongly) build a segmentation + TUNet for a config type nobody had actually written a branch for yet. + """ + from schemas import DlsiaDenoiserConfig, DlsiaTunetConfig # noqa: PLC0415 — avoid a hard import-time cycle + + hp = model_cfg.hyperparams + + if isinstance(model_cfg, DlsiaTunetConfig): + import dlsia_runtime as fam # noqa: PLC0415 + + if not dlsia_available(): + raise RuntimeError("dlsia is not installed on this server") + if init_state is not None: + # load_model rebuilds the net from the saved topo_dict, so the resumed + # topology is the run's own — depth/base_channels/growth_rate/image_size + # from the request are deliberately NOT used here (they can't be: a + # different topology cannot load these weights). run_train_job forces + # them to match the saved run before getting here. + log_cb(f"Resuming dlsia TUNet from saved weights on {device}…") + model = fam.load_model(init_state, device) + topo = init_state.get("topo_dict", {}) + snapshot = { + "depth": topo.get("depth", hp.depth), + "base_channels": topo.get("base_channels", hp.base_channels), + "growth_rate": topo.get("growth_rate", hp.growth_rate), + } + else: + log_cb(f"Building dlsia TUNet (depth={hp.depth}, base_channels={hp.base_channels}) on {device}…") + model = fam.build_model(n_classes, hp.image_size, hp.depth, hp.base_channels, hp.growth_rate, device) + snapshot = {"depth": hp.depth, "base_channels": hp.base_channels, "growth_rate": hp.growth_rate} + return BuiltFamily( + forward_fn=fam.make_forward_fn(model), + to_tensor_fn=fam.make_to_tensor_fn(), + trainable_params=list(model.parameters()), + set_train_mode=fam.make_set_train_mode_fn(model), + model_config_snapshot=snapshot, + adapter_state_fn=lambda: fam.network_dict(model), + ) + + if isinstance(model_cfg, DlsiaDenoiserConfig): + # Two architectures share this family; both expose the identical + # six-function runtime template, so everything below the module choice + # is common. See DlsiaDenoiserConfig for why they aren't separate + # model_family values. + if model_cfg.architecture == "cnn_ae": + import autoencoder_runtime as fam # noqa: PLC0415 + + # No dlsia_available() gate here, unlike the TUNet branch below: + # this network is plain torch and needs no optional dependency. + if init_state is not None: + log_cb(f"Resuming autoencoder denoiser from saved weights on {device}…") + model = fam.load_model(init_state, device) + topo = init_state.get("topo_dict", {}) + snapshot = { + "architecture": "cnn_ae", + "depth": topo.get("depth", hp.depth), + "base_channels": topo.get("base_channels", hp.base_channels), + "ae_compression": topo.get("compression", model_cfg.ae_compression), + "latent_channels": topo.get("latent_channels"), + } + else: + latent = fam.latent_channels_for(hp.depth, model_cfg.ae_compression) + log_cb( + f"Building autoencoder denoiser (depth={hp.depth}, " + f"base_channels={hp.base_channels}, {model_cfg.ae_compression}x compression " + f"-> {latent} latent channels) on {device}…" + ) + model = fam.build_model( + hp.image_size, hp.depth, hp.base_channels, model_cfg.ae_compression, device + ) + snapshot = { + "architecture": "cnn_ae", + "depth": hp.depth, + "base_channels": hp.base_channels, + "ae_compression": model_cfg.ae_compression, + "latent_channels": latent, + } + else: + import denoise_runtime as fam # noqa: PLC0415 + + if not dlsia_available(): + raise RuntimeError("dlsia is not installed on this server") + if init_state is not None: + # Same warm-start contract as the segmentation TUNet branch above: + # the saved topo_dict wins, the request's topology knobs don't. + log_cb(f"Resuming dlsia denoiser TUNet from saved weights on {device}…") + model = fam.load_model(init_state, device) + topo = init_state.get("topo_dict", {}) + snapshot = { + "architecture": "tunet", + "depth": topo.get("depth", hp.depth), + "base_channels": topo.get("base_channels", hp.base_channels), + "growth_rate": topo.get("growth_rate", hp.growth_rate), + } + else: + log_cb( + f"Building dlsia denoiser TUNet (depth={hp.depth}, base_channels={hp.base_channels}) " + f"on {device}…" + ) + # No n_classes: denoise_runtime.build_model fixes in/out channels to + # 1 (single-channel regression) — there is nothing to parameterize. + model = fam.build_model(hp.image_size, hp.depth, hp.base_channels, hp.growth_rate, device) + snapshot = { + "architecture": "tunet", + "depth": hp.depth, + "base_channels": hp.base_channels, + "growth_rate": hp.growth_rate, + } + return BuiltFamily( + forward_fn=fam.make_forward_fn(model), + to_tensor_fn=fam.make_to_tensor_fn(), + trainable_params=list(model.parameters()), + set_train_mode=fam.make_set_train_mode_fn(model), + model_config_snapshot=snapshot, + adapter_state_fn=lambda: fam.network_dict(model), + ) + + raise ValueError( + f"Unknown model family/config type for build_family: {type(model_cfg).__name__!r} " + f"(model_family={getattr(model_cfg, 'model_family', None)!r})" + ) + + +# --------------------------------------------------------------------------- +# Generic training loop (family-agnostic: takes a forward callable + params) +# --------------------------------------------------------------------------- + + +def compute_miou(pred: "Any", target: "Any", n_classes: int, ignore_index: int = IGNORE_INDEX) -> float: + """Mean IoU over ``n_classes`` (0-indexed), ignoring ``ignore_index`` pixels. + + ``pred``/``target`` are torch tensors of matching shape (after argmax); + a class absent from both prediction and target on this batch is skipped + (rather than counted as a perfect or zero score) so small batches don't + bias the running average toward classes that simply didn't appear. + """ + valid = target != ignore_index + ious = [] + for c in range(n_classes): + p = (pred == c) & valid + t = (target == c) & valid + union = (p | t).sum().item() + if union == 0: + continue + intersection = (p & t).sum().item() + ious.append(intersection / union) + return float(sum(ious) / len(ious)) if ious else 0.0 + + +def run_training_loop( + *, + train_pairs: list[tuple[np.ndarray, np.ndarray]], + val_pairs: list[tuple[np.ndarray, np.ndarray]], + image_size: int, + n_classes: int, + epochs: int, + batch_size: int, + seed: int, + flip_augment: bool, + to_tensor_fn: Callable[["Any"], "Any"], + forward_fn: Callable[["Any"], "Any"], + trainable_params: list["Any"], + lr: float, + device: str, + make_optimizer_fn: Callable[[list["Any"], float], "Any"] | None = None, + on_batch: Callable[[], bool] | None = None, + on_epoch: Callable[[int, float, float | None, float | None], bool] | None = None, + set_train_mode: Callable[[bool], None] | None = None, +) -> dict[str, Any]: + """Shared epoch/batch loop used by both model families. + + ``to_tensor_fn(rgb_uint8_hwc) -> float tensor (C,H,W)`` applies whatever + normalisation the family wants (plain /255 for dlsia). ``forward_fn(batch_images) + -> logits (B,n_classes,H,W)`` runs the family's model. Both + ``on_batch``/``on_epoch`` return ``True`` to request a cooperative stop + (checkpoints are still saved by the caller either way). + + ``set_train_mode(is_training)`` toggles ``nn.Module.train()``/``eval()`` + around validation — matters for the dlsia family (TUNet defaults to + BatchNorm2d, whose running stats should stay fixed during validation + rather than being computed from that batch). Omit for a model with no such + distinction. + + Returns ``{"epochs_completed", "final_train_loss", "final_val_loss", + "val_miou", "cancelled"}``. + """ + import torch + import torch.nn as nn + + if not train_pairs: + raise ValueError("No training data: at least one annotated slice is required") + + rng = np.random.default_rng(seed) + optimizer = (make_optimizer_fn or (lambda params, lr_: torch.optim.AdamW(params, lr=lr_, weight_decay=0.01)))( + trainable_params, lr + ) + scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max(1, epochs)) + criterion = nn.CrossEntropyLoss(ignore_index=IGNORE_INDEX) + + def _prep(rgb: np.ndarray, label: np.ndarray, flip: bool) -> tuple["Any", "Any"]: + img_l, lbl_l = letterbox(rgb, label, image_size) + if flip: + img_l = np.ascontiguousarray(img_l[:, ::-1]) + lbl_l = np.ascontiguousarray(lbl_l[:, ::-1]) + return to_tensor_fn(img_l), torch.from_numpy(lbl_l.astype(np.int64)) + + cancelled = False + final_train_loss = 0.0 + final_val_loss: float | None = None + val_miou: float | None = None + epochs_completed = 0 + + for epoch in range(epochs): + order = rng.permutation(len(train_pairs)) + epoch_loss = 0.0 + n_batches = 0 + for start in range(0, len(order), batch_size): + batch_idx = order[start : start + batch_size] # noqa: E203 + imgs, lbls = [], [] + for i in batch_idx: + rgb, label = train_pairs[int(i)] + flip = bool(flip_augment and rng.random() < 0.5) + img_t, lbl_t = _prep(rgb, label, flip) + imgs.append(img_t) + lbls.append(lbl_t) + batch_img = torch.stack(imgs).to(device) + batch_lbl = torch.stack(lbls).to(device) + + logits = forward_fn(batch_img) + loss = criterion(logits, batch_lbl) + optimizer.zero_grad() + loss.backward() + optimizer.step() + + epoch_loss += float(loss.detach().item()) + n_batches += 1 + if on_batch is not None and on_batch(): + cancelled = True + break + scheduler.step() + final_train_loss = epoch_loss / max(1, n_batches) + epochs_completed = epoch + 1 + + if val_pairs and not cancelled: + if set_train_mode is not None: + set_train_mode(False) + final_val_loss, val_miou = _evaluate( + val_pairs, image_size, n_classes, to_tensor_fn, forward_fn, device, criterion, batch_size + ) + if set_train_mode is not None: + set_train_mode(True) + + if on_epoch is not None and on_epoch(epochs_completed, final_train_loss, final_val_loss, val_miou): + cancelled = True + if cancelled: + break + + return { + "epochs_completed": epochs_completed, + "final_train_loss": final_train_loss, + "final_val_loss": final_val_loss, + "val_miou": val_miou, + "cancelled": cancelled, + } + + +def _evaluate( + val_pairs: list[tuple[np.ndarray, np.ndarray]], + image_size: int, + n_classes: int, + to_tensor_fn: Callable[["Any"], "Any"], + forward_fn: Callable[["Any"], "Any"], + device: str, + criterion: "Any", + batch_size: int = 1, +) -> tuple[float, float]: + """Average validation loss + per-sample mIoU over ``val_pairs``. + + Batches the forward pass by ``batch_size`` (mirrors the training loop) + instead of one sample at a time — tiling's holdout can leave dozens of + validation patches, and running under ``torch.no_grad()`` means the + activation memory a bigger batch needs stays bounded regardless. + + mIoU is still scored per sample, not per batch: :func:`compute_miou` + pools every matching pixel in whatever tensor it's given, so handing it a + whole batch at once would silently compute one pooled-pixel score across + samples instead of the average of each sample's own score. That per-sample + split is cheap CPU-side tensor slicing after the one batched forward, not + a second forward pass, so it doesn't undo the batching's benefit. + """ + import torch + + total_loss = 0.0 + total_miou = 0.0 + n = len(val_pairs) + with torch.no_grad(): + for start in range(0, n, batch_size): + chunk = val_pairs[start : start + batch_size] # noqa: E203 + imgs, lbls = [], [] + for rgb, label in chunk: + img_l, lbl_l = letterbox(rgb, label, image_size) + imgs.append(to_tensor_fn(img_l)) + lbls.append(torch.from_numpy(lbl_l.astype(np.int64))) + batch_img = torch.stack(imgs).to(device) + batch_lbl = torch.stack(lbls).to(device) + + logits = forward_fn(batch_img) + # criterion's mean reduction is over the whole batch, so weight by + # sample count to keep this a per-sample average overall — same + # convention the training loop already uses for its own loss. + total_loss += float(criterion(logits, batch_lbl).item()) * len(chunk) + preds = logits.argmax(dim=1) + for i in range(len(chunk)): + total_miou += compute_miou(preds[i], batch_lbl[i], n_classes) + return total_loss / n, total_miou / n + + +# --------------------------------------------------------------------------- +# Capability probe (never raises — feeds GET /api/train/capability) +# --------------------------------------------------------------------------- + + +def capability() -> dict[str, Any]: + """Best-effort snapshot of Train-tab readiness. Never raises. + + No DINOv3 fields — that model family is deferred to Phase 5.5 and isn't + scaffolded here at all. + """ + result: dict[str, Any] = { + "torch_available": False, + "torch_version": None, + "device": None, + "dlsia": {"available": False}, + "tiling": {"available": False}, + # Classical denoising is pure CPU (scipy/skimage) and never touches + # ML_LOCK, so it stays usable while a training job runs. + "denoise": {"available": False, "methods": []}, + "runs_dir": str(runs_dir()), + "busy": ML_LOCK.locked(), + } + try: + result["torch_available"] = torch_available() + if result["torch_available"]: + import torch # noqa: PLC0415 + + result["torch_version"] = torch.__version__ + result["device"] = pick_device() + + result["dlsia"] = {"available": result["torch_available"] and dlsia_available()} + + import tiling # noqa: PLC0415 — avoid a hard import-time cycle (tiling imports train_common) + + result["tiling"] = {"available": tiling.qlty_available()} + + import denoise # noqa: PLC0415 — cheap, but keep the probe self-contained + + # `available: True` unconditionally: scipy/skimage are hard dependencies + # here (unlike torch/dlsia), so the classical filters always work. The + # per-method `available` flags carry the real gating — wavelet needs + # PyWavelets, which skimage imports lazily and which is not installed. + result["denoise"] = {"available": True, "methods": denoise.describe_methods()} + except Exception as exc: # noqa: BLE001 — a capability probe must never 500 + # Full exception detail goes to the server log only — the raw message + # (module paths, internal state) must never reach an HTTP response + # (CodeQL py/stack-trace-exposure). The client only needs to know the + # probe failed; every field above already defaults to unavailable. + logger.warning("Train capability probe failed: %s", exc) + result["error"] = "capability probe failed" + return result diff --git a/backend/train_jobs.py b/backend/train_jobs.py new file mode 100644 index 0000000..37bf240 --- /dev/null +++ b/backend/train_jobs.py @@ -0,0 +1,575 @@ +"""Background job orchestration for the Train tab. + +Dispatches a :class:`schemas.TrainRequest` to the requested model family +(dlsia TUNet segmentation, or the dlsia/autoencoder denoiser family), sharing +data preparation, the generic training loop, and run persistence via +:mod:`train_common`. Progress/cancellation are reported through the existing +:mod:`export_jobs` registry — the frontend polls the same +``GET /api/export/status/{job_id}`` route already used by exports and +mask-sync jobs. + +DINOv3 LoRA fine-tuning is out of scope here (see Phase 5.5). + +Two *tasks* are dispatched from here, distinguished by ``TrainRequest.task``: + +* ``"segmentation"`` — annotation-driven, via + :func:`train_common.prepare_datasets` and + :func:`train_common.run_training_loop`. +* ``"denoising"`` — self-supervised (Noise2Noise / Noise2Void / autoencoder), + via :mod:`denoise_train`'s raw-slice samplers and its own regression loop. + +Everything either task has in common — the ML lock, device selection, resume +resolution, the qlty guard, ``export_jobs`` progress/cancel reporting, and +persistence through :func:`train_common.save_run` — is shared in +:func:`run_train_job`; only data preparation, the loss, and the validation +metric differ, and those live behind the split at the end of its preamble. +""" + +from __future__ import annotations + +import logging +import uuid +from datetime import datetime, timezone + +import export_jobs +import train_common +from schemas import ( + DenoiseTrainOpts, + DlsiaDenoiserConfig, + DlsiaTunetConfig, + TrainRequest, +) + +logger = logging.getLogger(__name__) + + +def new_run_id(model_family: str) -> str: + """``__`` — sorts newest-last + lexically within a day, globally unique enough for a single-user tool.""" + ts = datetime.now(timezone.utc).strftime("%Y%m%d_%H%M%S") + return f"{ts}_{model_family}_{uuid.uuid4().hex[:4]}" + + +def check_resume_compatible(parent_config: dict, request: TrainRequest) -> None: + """Raise if *request* cannot continue fine-tuning from *parent_config*'s run. + + The saved head has one output channel per class, in the parent's class + ORDER (``infer_jobs`` maps channel ``c`` to ``classes[c]["classId"]``), so + the class list has to line up positionally for a resume to mean anything. + A count mismatch would fail loudly inside ``load_state_dict`` anyway, but a + same-count-different-labels resume would train happily against the wrong + semantics and never say so — which is the case this exists to catch. + + A self-supervised denoiser has no class taxonomy at all, so the class-list + comparison is skipped when BOTH the parent run and *request* are + ``task == "denoising"``. That is not the same as skipping the check + whenever either side is denoising: resuming a denoiser as a segmentation + run (or vice versa) is a genuine incompatibility — the saved weights are + single-channel regression output, not per-class logits, or the reverse — + and must still be reported as such rather than silently allowed through. + ``parent_config`` may predate the ``task`` field entirely (see + ``train_common.load_run_config``), so an absent key defaults to + ``"segmentation"`` here too, not just at load time. + """ + if parent_config.get("model_family") != request.model.model_family: + raise ValueError( + f"Cannot continue fine-tuning a {parent_config.get('model_family')} run " + f"as {request.model.model_family} — pick the matching model family, or train a new run." + ) + + parent_task = parent_config.get("task", "segmentation") + request_task = request.task + if parent_task != request_task: + raise ValueError( + f"Cannot continue fine-tuning a {parent_task!r}-task run as a {request_task!r}-task " + "request — a segmentation model and a denoiser are not interchangeable. Pick the " + "matching task, or train a new run." + ) + if parent_task == "denoising": + # Both sides are confirmed "denoising" above — there is no class + # taxonomy to compare for a self-supervised denoiser. But the two + # architectures in this family share a model_family, so the check above + # passes for a TUNet-vs-autoencoder mismatch; without this, the resume + # would reach load_state_dict and fail on an opaque key/shape mismatch + # instead of saying what is actually wrong. Runs saved before + # `architecture` existed are TUNets by definition. + parent_arch = (parent_config.get("model_config") or {}).get("architecture", "tunet") + request_arch = getattr(request.model, "architecture", "tunet") + if parent_arch != request_arch: + raise ValueError( + f"Cannot continue fine-tuning a {parent_arch!r} denoiser as {request_arch!r} — " + "the two architectures have different weights entirely. Pick the matching " + "architecture, or train a new run." + ) + return + + parent_labels = [str(c.get("label", "")).strip().lower() for c in parent_config.get("classes", [])] + current_labels = [c.label.strip().lower() for c in request.classes] + if parent_labels != current_labels: + raise ValueError( + "The classes changed since that run was trained, so its saved weights no longer apply " + f"(run has {len(parent_labels)}: {', '.join(parent_labels) or '—'}; " + f"now {len(current_labels)}: {', '.join(current_labels) or '—'}). " + "Restore the original classes to continue fine-tuning, or train a new run instead." + ) + + +def _apply_parent_architecture(parent_config: dict, request: TrainRequest) -> None: + """Force the architecture-defining settings to the parent run's values. + + Anything that changes tensor shapes cannot differ across a resume — the + saved weights simply would not load. Rather than trusting the client to + echo these back correctly (or 400-ing on every mismatch), the server just + overrides them, so a resume is always weight-compatible by construction. + Genuinely re-tunable knobs — epochs, lr, batch_size, seed, flip_augment — + are left as the caller sent them; that is the point of resuming. + + Every recognized ``schemas.ModelConfig`` member has its own explicit + branch below; anything else raises :class:`ValueError` — mirrors + ``train_common.build_family``'s dispatch. + """ + model_cfg = request.model + hp = model_cfg.hyperparams + parent_model = parent_config.get("model_config", {}) + parent_hp = parent_config.get("hyperparams", {}) + + # image_size + tiling define the geometry the weights were fit to. + for field in ("image_size", "tiling"): + if field in parent_hp: + setattr(hp, field, parent_hp[field]) + + # Input denoising is inherited for the same reason as the geometry above: + # the parent's weights were fit to that pixel distribution, so continuing to + # train them on differently-preprocessed pixels degrades them silently. It + # is NOT a re-tunable knob like epochs or lr — leaving it to the caller + # meant a fine-tune of a denoise-trained run quietly reverted to raw pixels + # whenever the client forgot to echo it back. + parent_denoise = parent_config.get("denoise") + request.denoise = ( + DenoiseTrainOpts(**parent_denoise) if parent_denoise else None + ) + + if isinstance(model_cfg, DlsiaTunetConfig): + # TUNet's topology comes back from the saved topo_dict regardless; mirror + # it onto the request so the saved config records what actually ran. + for field in ("depth", "base_channels", "growth_rate"): + if field in parent_hp: + setattr(hp, field, parent_hp[field]) + elif isinstance(model_cfg, DlsiaDenoiserConfig): + # Same TUNet topology knobs as the segmentation family above — dlsia's + # TUNet is the shared architecture underneath both; only the fixed + # in/out channel counts differ, and those aren't user-configurable. + for field in ("depth", "base_channels", "growth_rate"): + if field in parent_hp: + setattr(hp, field, parent_hp[field]) + # The architecture itself, and the bottleneck width that follows from it, + # are inherited for the same reason as the geometry above: they define + # the tensor shapes, so a resume that changed them could not load the + # saved weights. A run saved before `architecture` existed has no such + # key and is a TUNet by definition. + parent_arch = parent_model.get("architecture", "tunet") + model_cfg.architecture = parent_arch + if parent_arch == "cnn_ae": + # `ae` is the only scheme the schema permits with cnn_ae, so a + # resume must land on it or the request would be self-inconsistent. + model_cfg.training_scheme = "ae" + if parent_model.get("ae_compression") is not None: + model_cfg.ae_compression = int(parent_model["ae_compression"]) + else: + raise ValueError( + f"Unknown model config type for resume: {type(model_cfg).__name__!r} " + f"(model_family={getattr(model_cfg, 'model_family', None)!r})" + ) + + +def check_task_matches_model(request: TrainRequest) -> None: + """Raise unless ``request.task`` and the model family agree. + + ``task`` and ``model.model_family`` are independent fields on the schema, + and ``task`` defaults to ``"segmentation"`` — so a client that sends a + ``dlsia_denoiser`` config and forgets ``task`` produces a request that + validates cleanly and then means something incoherent. Left unchecked it + would route into the annotation-driven path and train a single-channel + regression network against ``CrossEntropyLoss`` over zero classes. The + mirror case (``task="denoising"`` with a segmentation family) would ask + :mod:`denoise_train` to feed grayscale into a 3-channel model. Both are + caught here, before any data is read. + """ + is_denoiser_family = isinstance(request.model, DlsiaDenoiserConfig) + if request.task == "denoising" and not is_denoiser_family: + raise ValueError( + "task='denoising' needs a denoiser model family, but got " + f"{request.model.model_family!r}. Use model_family='dlsia_denoiser', or set task='segmentation'." + ) + if request.task != "denoising" and is_denoiser_family: + raise ValueError( + "model_family='dlsia_denoiser' is a self-supervised denoiser and cannot be trained as a " + f"{request.task!r} task — send task='denoising' (and classes: []) instead." + ) + + +def _source_keys(request: TrainRequest) -> list[str]: + """Stable per-source identifiers recorded on the saved run.""" + return [ + (f"tiled:{item.server_uri or ''}:{item.source}" if item.kind == "tiled" else f"local:{item.source}") + for item in request.sources + ] + + +def _run_denoise_training( + jid: str, + request: TrainRequest, + run_id: str, + *, + device: str, + model_cfg: DlsiaDenoiserConfig, + init_state: dict | None, + resume_id: str | None, + progress_cb, +) -> None: + """Self-supervised denoiser branch of :func:`run_train_job`. + + Called with :data:`train_common.ML_LOCK` already held and the resume + already resolved, and reports through ``export_jobs`` with the same phase + names (``preparing`` → ``tiling`` → ``training`` → ``saving`` → ``done``) + and the same cancellation contract as the segmentation path, so the + frontend's existing job polling needs no denoiser-specific handling. + + Weights and run metadata go through the ordinary + :func:`train_common.save_run`, with ``task="denoising"`` and an empty class + list, so the run appears in ``list_runs`` and the Learned Denoiser panel's + ``model_family == "dlsia_denoiser"`` filter finds it. + """ + import denoise_train + import tiling + + hp = model_cfg.hyperparams + scheme = model_cfg.training_scheme + + # phase is already "preparing" — set by run_train_job before the split. + scheme_label = { + "n2n": "Noise2Noise", + "n2v": "Noise2Void", + "ae": "autoencoder", + }.get(scheme, scheme) + export_jobs.log( + jid, + f"Preparing self-supervised {scheme_label} data from raw slices (no annotations needed)…", + ) + if scheme == "n2n": + datasets = denoise_train.prepare_noise2noise_datasets( + request.sources, request.render, progress_cb=progress_cb + ) + else: + # n2v and dae are both single-slice schemes: each item is a slice paired + # with itself, and the objective differs only in how the INPUT is + # perturbed (blind-spot masking vs. added synthetic noise). So they share + # this sampler rather than needing a third one. + datasets = denoise_train.prepare_noise2void_datasets( + request.sources, request.render, progress_cb=progress_cb + ) + + if hp.tiling: + export_jobs.update(jid, phase="tiling") + export_jobs.log(jid, f"Cutting {hp.image_size}px windows…") + datasets = denoise_train.tile_denoise_datasets( + datasets, + hp.image_size, + progress_cb=progress_cb, + cancel_cb=lambda: export_jobs.cancel_requested(jid), + ) + if datasets is None: + # Same "cancelled before a model existed" shape the segmentation + # path uses — deliberately not the partial-run result, which implies + # save_run produced something. + export_jobs.update(jid, state="done", phase="done", result={"cancelled": True}) + export_jobs.log(jid, "Training cancelled while tiling; nothing was trained yet.") + return + + # The denoiser's samplers put everything in "train" (there are no + # annotation-driven splits to inherit), so the seeded patch holdout is what + # produces a validation set at all. Applied whether or not tiling ran: with + # tiling off the items are whole slices, and the function no-ops below its + # minimum count rather than starving a small run of training data. + datasets, n_held = tiling.holdout_val_patches(datasets, seed=hp.seed) + if n_held: + export_jobs.log( + jid, + f"Held back {n_held} training patch(es) for validation. They come from the same " + "slice(s) as the training patches, and both schemes' targets are themselves noisy, " + "so the reported correlation is a convergence signal — not an image-quality score.", + ) + + n_train = len(datasets["train"]) + if n_train == 0: + raise ValueError("No slices to train the denoiser on") + + # n_classes is ignored by build_family's denoiser branch (in/out channels + # are fixed at 1); 0 is passed to make that explicit rather than incidental. + built = train_common.build_family(model_cfg, 0, device, progress_cb, init_state=init_state) + + batches_per_epoch = max(1, -(-n_train // hp.batch_size)) + export_jobs.set_total(jid, hp.epochs * batches_per_epoch) + export_jobs.update(jid, phase="training") + + def _on_batch() -> bool: + export_jobs.bump(jid, 1) + return export_jobs.cancel_requested(jid) + + def _on_epoch(epoch: int, train_loss: float, val_loss: float | None, val_metric: float | None) -> bool: + msg = f"epoch {epoch}/{hp.epochs} — train loss {train_loss:.5f}" + if val_loss is not None: + # Labelled as correlation with the NOISY target, never as quality. + msg += f", val loss {val_loss:.5f}, noisy-target r {val_metric:.3f}" + export_jobs.log(jid, msg) + return export_jobs.cancel_requested(jid) + + metrics = denoise_train.run_denoise_training_loop( + train_pairs=datasets["train"], + val_pairs=datasets["val"], + image_size=hp.image_size, + training_scheme=scheme, + epochs=hp.epochs, + batch_size=hp.batch_size, + seed=hp.seed, + flip_augment=hp.flip_augment, + to_tensor_fn=built.to_tensor_fn, + forward_fn=built.forward_fn, + trainable_params=built.trainable_params, + lr=hp.lr, + device=device, + on_batch=_on_batch, + on_epoch=_on_epoch, + set_train_mode=built.set_train_mode, + ) + + export_jobs.update(jid, phase="saving") + train_common.save_run( + run_id, + model_family=model_cfg.model_family, + # training_scheme belongs on the run: n2n and n2v produce different + # models from the same topology, and the runs list surfaces which. + model_config={ + **built.model_config_snapshot, + "training_scheme": scheme, + # Recorded only where it means something, so a run's config doesn't + # imply a knob that had no effect on how it was trained. + # `architecture` is always recorded: both inference sites dispatch + # on it, defaulting to "tunet" for runs saved before it existed. + "architecture": model_cfg.architecture, + }, + classes=[], + render=request.render.model_dump(), + image_size=hp.image_size, + hyperparams=hp.model_dump(), + source_keys=_source_keys(request), + adapter_state=built.adapter_state_fn(), + metrics=metrics, + resumed_from=resume_id, + task="denoising", + ) + + result = {"run_id": run_id, **metrics} + export_jobs.update(jid, state="done", phase="done", result=result) + export_jobs.log( + jid, + "Training cancelled; partial run saved." if metrics["cancelled"] else "Denoiser training complete.", + ) + + +def run_train_job(jid: str, request: TrainRequest, run_id: str) -> None: + """Background worker: prepare data, build the requested model family, run + the shared training loop, and persist the resulting run. + + Holds :data:`train_common.ML_LOCK` for the whole job — training and + inference contend for the same device memory, so only one ML job (of + either kind) runs at a time. + """ + if not train_common.ML_LOCK.acquire(blocking=False): + export_jobs.update( + jid, + state="error", + phase="error", + error="Another training or inference job is already running", + ) + return + try: + export_jobs.update(jid, state="running", phase="preparing") + device = train_common.pick_device() + if device is None: + raise RuntimeError("torch is not installed on this server") + + def _progress(msg: str) -> None: + export_jobs.log(jid, msg) + + model_cfg = request.model + check_task_matches_model(request) + n_classes = len(request.classes) + + # Resolve a resume BEFORE prepare_datasets: an incompatible one must fail + # in seconds, not after minutes of rendering. Loading the weights here + # too means a corrupt/missing adapter.pt is caught just as early. + init_state = None + resume_id = request.resume_from_run_id + if resume_id: + parent_config = train_common.load_run_config(resume_id) + check_resume_compatible(parent_config, request) + _apply_parent_architecture(parent_config, request) + init_state = train_common.load_adapter_state(resume_id) + export_jobs.log(jid, f"Continuing fine-tuning from run {resume_id}.") + + hp = model_cfg.hyperparams + + # Check qlty BEFORE the render/rasterize pass below (prepare_datasets can + # take minutes on a large source list) so a torch-only install fails fast + # instead of after paying for it. + if hp.tiling: + import tiling + + if not tiling.qlty_available(): + raise RuntimeError("Tiling requires the 'qlty' package, which is not installed on this server") + + # Everything above is task-agnostic (lock, device, resume, qlty guard). + # From here the two tasks diverge: a denoiser reads raw slices instead + # of rendering annotations, and optimises a regression loss. + if request.task == "denoising": + _run_denoise_training( + jid, + request, + run_id, + device=device, + model_cfg=model_cfg, + init_state=init_state, + resume_id=resume_id, + progress_cb=_progress, + ) + return + + datasets = train_common.prepare_datasets( + request.sources, + request.classes, + request.render, + request.auto_split, + progress_cb=_progress, + denoise=request.denoise, + ) + if request.denoise is not None: + _progress( + f"Training on {request.denoise.method}-denoised input " + f"({request.denoise.strength:.0%}); inference will reapply it automatically." + ) + + # Tiling: keep native resolution by cutting `image_size` windows out of each + # slice, instead of letting the training loop rescale whole slices down to + # `image_size`. Patches come out exactly window-sized, which letterbox() + # passes through, so the loop below is unchanged either way. + if model_cfg.hyperparams.tiling: + import tiling # already confirmed available above; re-import is a cheap sys.modules hit + + export_jobs.update(jid, phase="tiling") + export_jobs.log(jid, f"Tiling slices into {model_cfg.hyperparams.image_size}px windows…") + + datasets = tiling.tile_datasets( + datasets, + model_cfg.hyperparams.image_size, + progress_cb=_progress, + cancel_cb=lambda: export_jobs.cancel_requested(jid), + ) + if datasets is None: + # Cancelled before any model existed to save — a bare "cancelled" + # result, not the usual partial-run shape from a training-loop + # cancel (train_common.run_training_loop's own cancel path, below). + export_jobs.update(jid, state="done", phase="done", result={"cancelled": True}) + export_jobs.log(jid, "Training cancelled while tiling; nothing was trained yet.") + return + # Splits are per-slice, so a one- or two-slice dataset leaves validation + # empty and the run reports no metrics. Tiling yields enough patches to + # hold a few back instead. + datasets, n_held = tiling.holdout_val_patches( + datasets, seed=model_cfg.hyperparams.seed + ) + if n_held: + export_jobs.log( + jid, + f"No validation slices — held back {n_held} training patch(es) for validation. " + "They come from the same image(s) as the training patches, so mIoU reads " + "optimistically compared with an unseen slice.", + ) + + n_train = len(datasets["train"]) + if n_train == 0: + raise ValueError("No annotated slices to train on") + + built = train_common.build_family(model_cfg, n_classes, device, _progress, init_state=init_state) + forward_fn = built.forward_fn + to_tensor_fn = built.to_tensor_fn + trainable = built.trainable_params + set_train_mode = built.set_train_mode + model_config_snapshot = built.model_config_snapshot + + batches_per_epoch = max(1, -(-n_train // hp.batch_size)) + export_jobs.set_total(jid, hp.epochs * batches_per_epoch) + export_jobs.update(jid, phase="training") + + def _on_batch() -> bool: + export_jobs.bump(jid, 1) + return export_jobs.cancel_requested(jid) + + def _on_epoch(epoch: int, train_loss: float, val_loss: float | None, val_miou: float | None) -> bool: + msg = f"epoch {epoch}/{hp.epochs} — train loss {train_loss:.4f}" + if val_loss is not None: + msg += f", val loss {val_loss:.4f}, mIoU {val_miou:.3f}" + export_jobs.log(jid, msg) + return export_jobs.cancel_requested(jid) + + metrics = train_common.run_training_loop( + train_pairs=datasets["train"], + val_pairs=datasets["val"], + image_size=hp.image_size, + n_classes=n_classes, + epochs=hp.epochs, + batch_size=hp.batch_size, + seed=hp.seed, + flip_augment=hp.flip_augment, + to_tensor_fn=to_tensor_fn, + forward_fn=forward_fn, + trainable_params=trainable, + lr=hp.lr, + device=device, + on_batch=_on_batch, + on_epoch=_on_epoch, + set_train_mode=set_train_mode, + ) + + export_jobs.update(jid, phase="saving") + adapter_state = built.adapter_state_fn() + + train_common.save_run( + run_id, + model_family=model_cfg.model_family, + model_config=model_config_snapshot, + classes=[c.model_dump() for c in request.classes], + render=request.render.model_dump(), + image_size=hp.image_size, + hyperparams=hp.model_dump(), + source_keys=_source_keys(request), + adapter_state=adapter_state, + metrics=metrics, + resumed_from=resume_id, + task=request.task, + # Recorded so inference reapplies the same input preprocessing + # without the caller having to remember it. + denoise=request.denoise.model_dump() if request.denoise else None, + ) + + result = {"run_id": run_id, **metrics} + export_jobs.update(jid, state="done", phase="done", result=result) + export_jobs.log( + jid, + "Training cancelled; partial run saved." if metrics["cancelled"] else "Training complete.", + ) + except Exception as exc: # noqa: BLE001 — reported as a job error, never a crash + logger.error("Training job %s failed: %s", jid, exc) + export_jobs.update(jid, state="error", phase="error", error=str(exc)) + finally: + train_common.ML_LOCK.release() diff --git a/backend/volume_build.py b/backend/volume_build.py new file mode 100644 index 0000000..b84c336 --- /dev/null +++ b/backend/volume_build.py @@ -0,0 +1,203 @@ +"""Build a renderable 3-D volume from a slice stack already in Tiled. + +:mod:`tiff_stack_source` registers a *directory of TIFFs* and additionally +streams the full-resolution slices in place. That is the right path when the +source images are still on the server — but most datasets here arrived through +drag-and-drop ingest, which copies each slice into Tiled's own storage as a 2-D +array. For those there is no source directory to point at, and asking the user +to find one would be asking them to re-supply data the app already has. + +So this module builds the pyramid from whatever :mod:`arrays` can already read. +It needs no arguments beyond the dataset that is open, which is what makes +"Build 3-D volume" a single button rather than a form. + +What it does *not* do is copy the full-resolution data. Only the downsampled +levels are written, because those are the only ones the renderer ever uploads — +it picks a level that fits ``maxTextureDimension3D``, and full resolution never +does. Full resolution stays exactly where it is, read by the 2-D canvas as +always. +""" + +from __future__ import annotations + +import asyncio +import logging +from typing import Any, Callable + +import numpy as np +from fastapi import HTTPException + +import arrays as arrays_mod +import ingest as ingest_mod +import tiff_stack_source as tss +from tiled_clients import api_key_for_uri, get_tiled_client + +logger = logging.getLogger("volume_build") + + +def _stack_shape(source: str, kind: str, server_uri: str | None) -> tuple[Any, dict, tuple[int, int, int]]: + """Resolve *source* to a readable stack and report its ``(z, y, x)`` shape.""" + node = arrays_mod.resolve_array(source, kind, server_uri) + meta = arrays_mod.array_shape_meta(node) + n_slices = int(meta.get("n_slices") or 0) + height = int(meta.get("height") or 0) + width = int(meta.get("width") or 0) + if n_slices < 2: + raise HTTPException( + 422, + "This dataset is a single image, not a stack — there is no volume to build.", + ) + if meta.get("is_rgb"): + raise HTTPException( + 422, + "Colour images are not supported by the 3-D view, which renders a single " + "scalar volume.", + ) + return node, meta, (n_slices, height, width) + + +def inspect_volume_build( + source: str, kind: str = "tiled", server_uri: str | None = None +) -> dict[str, Any]: + """Describe the volume that would be built, without building it.""" + _node, meta, shape = _stack_shape(source, kind, server_uri) + plan = tss.pyramid_plan(shape) + return { + "full_shape": list(shape), + "dtype": str(meta.get("dtype") or "float32"), + "pyramid_plan": plan, + # Every source slice is read once; the coarser levels cascade in memory. + "slices_to_read": shape[0] if plan else 0, + "already_small": not plan, + } + + +def build_volume( + source: str, + kind: str = "tiled", + server_uri: str | None = None, + container_path: str | None = None, + progress: Callable[[str, int, int], None] | None = None, +) -> dict[str, Any]: + """Build and register a 3-D volume sidecar for the stack open at *source*. + + Args: + source: Tiled path of the per-slice dataset. + kind: ``"tiled"`` (the only kind with a catalog to register into). + server_uri: Connected Tiled server URI. + container_path: Where to put the sidecar; defaults to *source*'s parent, + so the volume lands next to the dataset it describes. + progress: Optional ``(message, done, total)`` callback for the job UI. + + Returns: + The registered ``key`` and ``tiled_path``, plus the build description. + """ + if kind != "tiled": + raise HTTPException(422, "Only datasets in the Tiled catalog can be built into volumes.") + + node, meta, shape = _stack_shape(source, kind, server_uri) + plan = tss.pyramid_plan(shape) + if not plan: + raise HTTPException( + 422, + f"This stack is already small enough ({shape[0]}x{shape[1]}x{shape[2]}) that " + "no downsampling is needed — but it is stored as separate 2-D slices, which " + "cannot be streamed as a volume. Re-register it from its source images.", + ) + + parts = [p for p in source.strip("/").split("/") if p] + stem = parts[-1] + key = f"{stem}{tss.VOLUME_SUFFIX}" + target_parts = ( + [p for p in container_path.strip("/").split("/") if p] + if container_path + else parts[:-1] + ) + if not target_parts: + raise HTTPException(422, "Cannot place a volume at the catalog root.") + + dtype = np.dtype(meta.get("dtype") or "float32") + total = shape[0] + built: dict[str, np.ndarray] = {} + previous: np.ndarray | None = None + previous_factor: list[int] | None = None + + for level in plan: + label = f"Building {level['path']}" + if previous is None or previous_factor is None: + # Finest generated level: read the stack once, a block of slices at a + # time, so peak memory is the output level plus a few input slices. + fz, fy, fx = (int(f) for f in level["factor"]) + out = np.empty(level["shape"], dtype=np.float32) + for z in range(level["shape"][0]): + block = [] + for k in range(fz): + index = z * fz + k + if index >= shape[0]: + break + block.append(arrays_mod.read_slice(node, meta, index)) + if progress and (z * fz + k) % 25 == 0: + progress(label, z * fz + k, total) + out[z] = tss.block_mean(np.stack(block), fy, fx) + previous_working = out + else: + relative = [int(level["factor"][a] // previous_factor[a]) for a in range(3)] + previous_working = tss.downsample_array(previous, relative) + + built[level["path"]] = tss._cast_like(previous_working, dtype) + previous, previous_factor = previous_working, list(level["factor"]) + + if progress: + progress("Writing pyramid", total, total) + sidecar = tss.write_pyramid_store(key, built) + + client = get_tiled_client(server_uri, api_key_for_uri(server_uri)) + target = ingest_mod._ensure_container(client, target_parts) + + # Replace any previous build for this dataset: re-running must refresh the + # volume, not fail or accumulate duplicates. + if key in ingest_mod._child_keys(target): + target.delete_contents(key, recursive=True, external_only=False) + volume = target.create_container(key=key, metadata={}) + + from tiled.client.register import Settings, register_single_item + + try: + asyncio.run( + register_single_item(volume, sidecar, is_directory=True, settings=Settings.init()) + ) + except Exception as exc: # noqa: BLE001 — classified for the UI + logger.warning("volume registration failed for %s: %s", sidecar, exc) + raise HTTPException(502, ingest_mod._classify_error(exc)["message"]) from exc + + pyramid_node = ingest_mod._walk(volume, [tss.PYRAMID_KEY]) + if pyramid_node is None or not list(pyramid_node): + raise HTTPException( + 502, + f"Tiled registered no levels for {key!r} from {sidecar}. The usual cause is " + "that this path is not in the Tiled server's `readable_storage` (see " + "tiled/config.yml).", + ) + + volume.update_metadata( + metadata={ + "sample_name": key, + "n_images": shape[0], + "source_format": "tiled-stack-3d", + "built_from": source, + "full_shape": list(shape), + "pyramid_plan": plan, + # No scale0: full resolution stays in the per-slice nodes rather than + # being duplicated. The renderer never uploads a level that large. + **tss.multiscales_metadata(key, plan, include_scale0=False), + } + ) + + if progress: + progress("Done", total, total) + return { + "key": key, + "tiled_path": "/".join([*target_parts, key]), + "full_shape": list(shape), + "pyramid_plan": plan, + } diff --git a/backend/volume_nodes.py b/backend/volume_nodes.py new file mode 100644 index 0000000..be000c5 --- /dev/null +++ b/backend/volume_nodes.py @@ -0,0 +1,128 @@ +"""Find the Tiled node that holds a dataset's renderable 3-D volume. + +The 3-D viewer needs an OME-NGFF multiscale group. Which node that is depends on +how the dataset got into the catalog, and the answer is not something the +frontend can work out from a path: + +* **Registered Zarr volumes** already are one (:mod:`zarr_source`), so the open + node — or an ancestor of it, when a specific pyramid level is open — is the + volume. +* **TIFF stacks** are stored as a container of per-slice 2-D arrays, which is not + a volume at all. Their 3-D view lives in the ``__volume`` sidecar built by + :mod:`tiff_stack_source`, and it only exists once someone has built it. + +Without this the viewer is pointed straight at whatever is open and fails with +``openOmeZarr: missing multiscales in root .zattrs`` — technically accurate and +useless to the person reading it, since the real answer is either "look at the +sidecar" or "no volume has been built for this dataset yet". + +Detection is on ``metadata["attributes"]["multiscales"]``: Tiled's ``.zattrs`` +route returns ``metadata["attributes"]`` verbatim, so that key is exactly what +the viewer will see. +""" + +from __future__ import annotations + +import logging +from typing import Any + +from tiff_stack_source import VOLUME_SUFFIX +from tiled_clients import api_key_for_uri, get_tiled_client + +logger = logging.getLogger("volume_nodes") + +#: How far to walk up from an open node looking for the multiscale root. A +#: pyramid level sits at most two levels down (``/scale0/image``). +_MAX_ASCENT = 3 + + +def has_multiscales(node: Any) -> bool: + """True if *node* carries OME-NGFF ``multiscales`` where Tiled will serve it.""" + try: + attributes = (dict(getattr(node, "metadata", {}) or {})).get("attributes") or {} + except Exception: # noqa: BLE001 — a node we cannot describe is not a volume + return False + multiscales = attributes.get("multiscales") + return isinstance(multiscales, list) and bool(multiscales) + + +def _node_at(client: Any, parts: list[str]) -> Any | None: + node = client + for part in parts: + try: + node = node[part] + except Exception: # noqa: BLE001 — missing child, or not a container + return None + return node + + +def source_dir_of(node: Any) -> str | None: + """Filesystem path a volume was registered from, if it records one.""" + try: + meta = dict(getattr(node, "metadata", {}) or {}) + except Exception: # noqa: BLE001 + return None + for key in ("tiff_dir", "zarr_path"): + value = meta.get(key) + if isinstance(value, str) and value: + return value + return None + + +def resolve_volume(server_uri: str | None, source: str) -> dict[str, Any]: + """Locate the renderable volume for the dataset open at *source*. + + Args: + server_uri: Connected Tiled server URI; ``None`` uses the default. + source: Tiled path of the open dataset. + + Returns: + ``mode`` is one of: + + * ``"self"`` — the open node is the volume. + * ``"ancestor"`` — a specific level was open; ``path`` is its root. + * ``"sidecar"`` — the ``__volume`` node built for a TIFF stack. + * ``"none"`` — no volume exists yet; ``message`` says so in plain terms + and ``source_dir`` carries a build candidate when one is known. + """ + parts = [p for p in (source or "").strip("/").split("/") if p] + if not parts: + return {"path": None, "mode": "none", "source_dir": None, + "message": "Open a dataset from Browse to view it in 3D."} + + client = get_tiled_client(server_uri, api_key_for_uri(server_uri)) + + # 1. The open node itself. + node = _node_at(client, parts) + if node is not None and has_multiscales(node): + return {"path": "/".join(parts), "mode": "self", + "source_dir": source_dir_of(node), "message": ""} + + # 2. An ancestor — the user may have opened one pyramid level directly, e.g. + # `/scale0/image`, which is an array and has no multiscales itself. + for depth in range(1, min(_MAX_ASCENT, len(parts)) + 1): + ancestor_parts = parts[:-depth] + if not ancestor_parts: + break + ancestor = _node_at(client, ancestor_parts) + if ancestor is not None and has_multiscales(ancestor): + return {"path": "/".join(ancestor_parts), "mode": "ancestor", + "source_dir": source_dir_of(ancestor), "message": ""} + + # 3. The sidecar built for a per-slice TIFF stack. + sidecar_parts = [*parts[:-1], f"{parts[-1]}{VOLUME_SUFFIX}"] + sidecar = _node_at(client, sidecar_parts) + if sidecar is not None and has_multiscales(sidecar): + return {"path": "/".join(sidecar_parts), "mode": "sidecar", + "source_dir": source_dir_of(sidecar), "message": ""} + + return { + "path": None, + "mode": "none", + "source_dir": None, + "message": ( + "No 3-D volume has been built for this dataset yet. It is stored as " + "individual 2-D slices, which the 3-D view cannot stream — build one " + "from the source images to enable it." + ), + } diff --git a/backend/zarr_source.py b/backend/zarr_source.py new file mode 100644 index 0000000..8003553 --- /dev/null +++ b/backend/zarr_source.py @@ -0,0 +1,571 @@ +"""Register on-disk Zarr volumes with Tiled, without copying any data. + +The drag-and-drop ingest path (:mod:`ingest`) streams image files through the +browser and writes each one into Tiled. That cannot express a tomography volume: +these datasets are 7-56 GB, so the bytes must stay where they are. + +Tiled can serve a Zarr store in place — ``.zarr`` maps to ``application/x-zarr`` +and ``ZarrGroupAdapter`` reads chunks lazily off disk — so "loading" a volume is +really just *registering* a path. A full-resolution 2560x2560 slice of a +(690, 2560, 2560) float32 store comes back over Tiled's HTTP layer in under +0.2 s, which is comfortably interactive. + +Layout +------ +The supported stores are OME-NGFF-style multiscale groups:: + + .zarr/ + .zattrs # "multiscales" -> datasets[].path + coordinateTransformations + scale0/image # (z, y, x) float32, full resolution + scale1/image # downsampled + ... + +After registration the Tiled path of one level is +``//scale0/image``, which :mod:`arrays` resolves as an +``(N, H, W)`` slice stack with no further plumbing. +""" + +from __future__ import annotations + +import asyncio +import json +import logging +from pathlib import Path +from typing import Any + +from fastapi import HTTPException + +import ingest as ingest_mod +from tiled_clients import api_key_for_uri, get_tiled_client + +logger = logging.getLogger("zarr_source") + +# Group-level metadata files that mark a directory as a Zarr store. v2 uses +# `.zgroup`/`.zarray`; v3 uses a single `zarr.json`. +_ZARR_MARKERS: tuple[str, ...] = (".zgroup", ".zarray", "zarr.json") + + +def _is_zarr_dir(path: Path) -> bool: + """True if *path* looks like the root of a Zarr store.""" + return path.is_dir() and any((path / marker).exists() for marker in _ZARR_MARKERS) + + +def _resolve_path(raw: str) -> Path: + """Expand and validate a user-supplied absolute path to a Zarr store. + + Raises: + HTTPException: 400/404/422 with a message aimed at the user, so the UI + never has to surface a traceback. + """ + if not (raw or "").strip(): + raise HTTPException(400, "Enter the path to a .zarr directory.") + path = Path(raw).expanduser() + if not path.is_absolute(): + raise HTTPException(400, f"Path must be absolute: {raw!r}") + if not path.exists(): + raise HTTPException(404, f"No such path: {path}") + if path.is_file(): + # The common near-miss: a zipped store. Zarr can read these via ZipStore, + # but Tiled registers a directory asset, so say so plainly. + if path.suffix.lower() == ".zip": + raise HTTPException( + 422, + "Zipped Zarr archives are not supported — unzip it first and " + "point at the resulting .zarr directory.", + ) + raise HTTPException(422, f"Not a directory: {path}") + if not _is_zarr_dir(path): + raise HTTPException( + 422, + f"{path.name!r} is not a Zarr store (no .zgroup/.zarray/zarr.json " + "at its root).", + ) + return path.resolve() + + +def _multiscale_datasets(path: Path) -> list[dict[str, Any]] | None: + """Read OME-NGFF ``multiscales`` from ``.zattrs``, if present. + + Returns the ``datasets`` list (each with ``path`` and + ``coordinateTransformations``), or None when the store is not multiscale. + """ + attrs_file = path / ".zattrs" + if not attrs_file.exists(): + return None + try: + attrs = json.loads(attrs_file.read_text()) + except (OSError, json.JSONDecodeError): + return None + multiscales = attrs.get("multiscales") + if not isinstance(multiscales, list) or not multiscales: + return None + datasets = multiscales[0].get("datasets") + return datasets if isinstance(datasets, list) and datasets else None + + +def _scale_vector(dataset: dict[str, Any]) -> list[float] | None: + """Pull the ``scale`` transform (voxel size per axis) out of a dataset entry.""" + for transform in dataset.get("coordinateTransformations") or []: + if transform.get("type") == "scale": + scale = transform.get("scale") + if isinstance(scale, list) and len(scale) == 3: + return [float(v) for v in scale] + return None + + +def _walk_arrays(group: Any, prefix: str = "", depth: int = 0) -> list[tuple[str, Any]]: + """Depth-first list of ``(path, array)`` for every 3-D array in *group*.""" + import zarr + + found: list[tuple[str, Any]] = [] + if depth > 3: + return found + for key in group.keys(): + try: + child = group[key] + except Exception: # noqa: BLE001 — a broken child shouldn't sink the scan + continue + child_path = f"{prefix}/{key}" if prefix else key + if isinstance(child, zarr.Array): + if child.ndim == 3: + found.append((child_path, child)) + else: + found.extend(_walk_arrays(child, child_path, depth + 1)) + return found + + +def inspect_zarr(raw_path: str) -> dict[str, Any]: + """Describe a Zarr store's resolution pyramid without registering anything. + + Args: + raw_path: Absolute path to a ``.zarr`` directory on the server. + + Returns: + Dict with ``name``, ``levels`` (finest first) and voxel metadata. Each + level carries the array's Tiled sub-path, shape, dtype, and its + downsample factor relative to the finest level — the factor the UI needs + to explain that a coarse level addresses only every f-th slice. + + Raises: + HTTPException: 4xx with a user-facing message for anything unusable. + """ + import zarr + + path = _resolve_path(raw_path) + try: + group = zarr.open_group(str(path), mode="r") + except Exception as group_exc: # noqa: BLE001 — try a bare array next + # Not a group — a store with a bare 3-D array at its root (no OME-NGFF + # multiscales wrapper) is an equally common, simpler way to save a + # volume (e.g. `zarr.save`/`dask.array.to_zarr` with no group). `_is_ + # zarr_dir` already accepts this shape (its own `.zarray`/`zarr.json` + # marker check) — rejecting it here would contradict that. + try: + arr = zarr.open_array(str(path), mode="r") + except Exception as arr_exc: # noqa: BLE001 — genuinely neither + raise HTTPException( + 422, f"Could not open {path.name!r} as Zarr: {group_exc}" + ) from arr_exc + if arr.ndim != 3: + raise HTTPException( + 422, + f"{path.name!r} is a Zarr array but not 3-D (shape {arr.shape}) " + "— expected (z, y, x).", + ) + shape = [int(v) for v in arr.shape] + return { + "name": path.name, + "path": str(path), + # Empty: the store registers as a single leaf array node directly + # at the container key, with no sub-path to descend into (see the + # frontend's own guard against appending a trailing empty segment). + "levels": [ + { + "path": "", + "shape": shape, + "dtype": str(arr.dtype), + "n_slices": shape[0], + "height": shape[1], + "width": shape[2], + "downsample": [1.0, 1.0, 1.0], + } + ], + "full_shape": shape, + "dtype": str(arr.dtype), + "voxel_size": None, + "voxel_unit": None, + "pixel_size": None, + } + + # Prefer the declared multiscales order; fall back to discovering 3-D arrays. + datasets = _multiscale_datasets(path) + entries: list[tuple[str, Any, list[float] | None]] = [] + if datasets: + for dataset in datasets: + key = str(dataset.get("path") or "").strip("/") + if not key: + continue + try: + arr = group[key] + except Exception: # noqa: BLE001 — declared but missing; skip it + logger.warning("multiscales lists %r but it is not readable", key) + continue + if getattr(arr, "ndim", 0) == 3: + entries.append((key, arr, _scale_vector(dataset))) + if not entries: + entries = [(key, arr, None) for key, arr in _walk_arrays(group)] + + if not entries: + raise HTTPException( + 422, + f"{path.name!r} contains no 3-D arrays to annotate. (An empty Zarr " + "group has nothing to load.)", + ) + + # Finest first, by voxel count. + entries.sort(key=lambda e: -(e[1].shape[0] * e[1].shape[1] * e[1].shape[2])) + base_shape = entries[0][1].shape + base_scale = entries[0][2] + + levels: list[dict[str, Any]] = [] + for key, arr, scale in entries: + shape = [int(v) for v in arr.shape] + levels.append( + { + "path": key, + "shape": shape, + "dtype": str(arr.dtype), + "n_slices": shape[0], + "height": shape[1], + "width": shape[2], + # How many finest-level voxels one voxel here spans, per axis. + "downsample": [ + round(base_shape[i] / shape[i], 4) if shape[i] else 1.0 for i in range(3) + ], + } + ) + + voxel_size = base_scale[0] if base_scale else None + return { + "name": path.name, + "path": str(path), + "levels": levels, + "full_shape": [int(v) for v in base_shape], + "dtype": str(entries[0][1].dtype), + # Finest-level voxel size (z, y, x) and its unit, for the Measurement panel. + "voxel_size": base_scale, + "voxel_unit": _voxel_unit(path), + "pixel_size": voxel_size, + } + + +def _voxel_unit(path: Path) -> str | None: + """Axis unit declared in the store's ``multiscales`` axes (e.g. micrometer).""" + attrs_file = path / ".zattrs" + if not attrs_file.exists(): + return None + try: + axes = json.loads(attrs_file.read_text())["multiscales"][0]["axes"] + except (OSError, json.JSONDecodeError, KeyError, IndexError, TypeError): + return None + for axis in axes: + unit = axis.get("unit") + if unit: + return str(unit) + return None + + +def registered_key(path: Path) -> str: + """The node key Tiled will use for *path*. + + Tiled strips the extension, so ``foo.zarr`` registers as ``foo``. Deriving + this with Tiled's own helper (rather than assuming ``path.name``) keeps the + collision check below looking at the key that will actually be created — a + mismatch here means the check passes and Tiled then hits the collision + itself, deep inside registration. + """ + from tiled.client.register import Settings + + return Settings.init().key_from_filename(path.name) + + +def _existing_node_info(node: Any) -> dict[str, Any]: + """Describe what already occupies a key, for the conflict prompt. + + ``external`` distinguishes a previous Zarr registration (safe to drop — it + only removes catalog rows) from internally-managed data such as an uploaded + image stack, where deleting the node also deletes the files. + """ + children: list[str] = [] + try: + children = list(node) + except Exception: # noqa: BLE001 — best-effort description + pass + meta = {} + try: + meta = dict(getattr(node, "metadata", {}) or {}) + except Exception: # noqa: BLE001 + pass + return { + "child_count": len(children), + "external": meta.get("source_format") == "zarr", + "sample_name": meta.get("sample_name") or "", + "n_images": meta.get("n_images"), + } + + +def preflight_zarr( + server_uri: str | None, raw_path: str, container_path: str = "browse" +) -> dict[str, Any]: + """Report whether registering *raw_path* would collide, without changing anything. + + Returns the derived ``key``, whether it ``exists``, and — when it does — what + is there, so the UI can warn before offering Replace. + """ + path = _resolve_path(raw_path) + key = registered_key(path) + client = get_tiled_client(server_uri, api_key_for_uri(server_uri)) + parts = [p for p in container_path.strip("/").split("/") if p] + target = ingest_mod._walk(client, parts) + existing = None + if target is not None and key in ingest_mod._child_keys(target): + existing = _existing_node_info(target[key]) + return {"key": key, "exists": existing is not None, "existing": existing} + + +def register_zarr( + server_uri: str | None, + raw_path: str, + container_path: str, + description: str = "", + on_conflict: str = "fail", +) -> dict[str, Any]: + """Register a Zarr store into a Tiled container, copying no data. + + Args: + server_uri: Connected Tiled server URI. + raw_path: Absolute path to the ``.zarr`` directory. + container_path: Slash-separated target container (e.g. ``browse``). + description: Optional keyword(s) stored on the node, exactly as ingest + does, so the volume is filterable in Browse. + on_conflict: ``"fail"``, ``"replace"`` or ``"skip"`` when the key exists. + + Returns: + Dict with the registered ``key``, its Tiled ``path``, and the inspection + result (so the caller can offer the level picker without a second call). + """ + info = inspect_zarr(raw_path) + path = Path(info["path"]) + description = (description or "").strip() + keywords = ingest_mod.parse_keywords(description) + if on_conflict not in ingest_mod.ON_CONFLICT_MODES: + on_conflict = "fail" + + client = get_tiled_client(server_uri, api_key_for_uri(server_uri)) + parts = [p for p in container_path.strip("/").split("/") if p] + target = ingest_mod._ensure_container(client, parts) + + # Tiled derives the key from the filename (dropping ".zarr"), so resolve the + # collision against THAT key. Getting this wrong lets registration proceed + # until Tiled hits the collision internally, where it can only offer to + # delete external assets and aborts on anything else. + key = registered_key(path) + if key in ingest_mod._child_keys(target): + existing = _existing_node_info(target[key]) + if on_conflict == "skip": + return {**info, "key": key, "tiled_path": "/".join([*parts, key]), "skipped": True} + if on_conflict == "replace": + if not existing["external"]: + # The occupant is internally-managed data (e.g. an uploaded image + # stack that happens to share this name). Deleting it would delete + # the files themselves, which is not what "replace this Zarr" + # should ever mean — make the user rename or remove it explicitly. + raise HTTPException( + 409, + f"{key!r} already exists in {container_path!r} and holds uploaded " + f"data ({existing['child_count']} items), not a registered Zarr. " + "Replacing it would delete that data — load into a different " + "container, or delete that dataset yourself first.", + ) + # A previous Zarr registration: dropping it removes catalog rows only, + # never the store on disk. + target.delete_contents(key, recursive=True, external_only=False) + else: + raise HTTPException( + 409, + f"{key!r} already exists in {container_path!r}. Choose Replace or " + "a different destination.", + ) + + # Tiled's own single-item registration: it resolves .zarr -> application/x-zarr + # and stores an external Asset pointing at the directory. Nothing is copied. + from tiled.client.register import Settings, register_single_item + + try: + asyncio.run( + register_single_item(target, path, is_directory=True, settings=Settings.init()) + ) + except Exception as exc: # noqa: BLE001 — classified for the UI below + logger.warning("zarr registration failed for %s: %s", path, exc) + raise HTTPException(502, ingest_mod._classify_error(exc)["message"]) from exc + + node = ingest_mod._walk(target, [key]) + if node is None: + # register_single_item logs and swallows adapter errors, returning None — + # so a missing node here is the signal that registration did not happen. + raise HTTPException( + 502, + f"Tiled did not register {key!r}. Check the backend log for the " + "adapter error, and that the Tiled server can read this path.", + ) + + # Same metadata shape the dropzone writes, so Browse treats this like any + # other sample, plus the pyramid description the level picker needs. + meta: dict[str, Any] = { + "sample_name": key, + "n_images": info["levels"][0]["n_slices"], + "source_format": "zarr", + "zarr_path": str(path), + "zarr_levels": info["levels"], + "full_shape": info["full_shape"], + } + if info.get("voxel_size"): + meta["voxel_size"] = info["voxel_size"] + if info.get("voxel_unit"): + meta["voxel_unit"] = info["voxel_unit"] + if description: + meta["description"] = description + if keywords: + meta["keywords"] = keywords + try: + node.update_metadata(metadata=meta) + except Exception as exc: # noqa: BLE001 — best-effort; the data is registered + logger.warning("could not set metadata on %s: %s", key, exc) + + # Spread `info` FIRST: it carries the filesystem path under "path", which + # must not shadow the Tiled path the caller needs to open the dataset. + return { + **info, + "key": key, + "tiled_path": "/".join([*parts, key]), + "skipped": False, + } + + +def scan_and_register_zarrs( + server_uri: str | None, + scan_root: str, + container_path: str = "browse", + on_conflict: str = "skip", + renames: dict[str, str] | None = None, +) -> dict[str, Any]: + """Walk *scan_root* for Zarr stores and register each one not already present. + + For a directory of already-reconstructed volumes (e.g. a bind-mounted host + folder) that should all show up in Browse without registering each one + individually through the UI. Non-recursive by design: only immediate + subdirectories of *scan_root* that look like a Zarr store (`_is_zarr_dir`) + are candidates — a store's own internal structure (`scale0/`, chunk files) + must never be treated as separate stores to register. + + A Zarr store and an unrelated dataset (e.g. a raw image folder of the same + acquisition) commonly share the same stem name — in which case they'd + derive the identical Tiled key. Rather than silently treating that as + "already registered" (misleading — this store was never actually + registered) or blindly replacing someone else's data, a same-key collision + with a DIFFERENT kind of registration (per the ``source_format`` tag; see + :func:`ingest.node_source_kind`) is reported as **shadowed**, distinctly + from a same-kind ``skipped`` match from a previous run of this same scan. + + Args: + server_uri: Connected Tiled server URI. + scan_root: Absolute directory to scan. + container_path: Target container every discovered store registers into. + on_conflict: Passed through to :func:`register_zarr` for each same-kind + match — ``"skip"`` (default) leaves already-registered entries + alone, so re-running the scan after adding new datasets is always + safe. ``"fail"`` is rejected: one conflicting store shouldn't be + able to abort a bulk scan the way it correctly can for a single + register. + renames: Optional ``{folder_name: alternate_key}`` override, so a + shadowed candidate can be retried under a different key without + re-scanning everything else. + + Returns: + Dict with ``scanned`` (candidate count), ``registered`` (newly + registered, each with ``name``/``key``/``tiled_path``), ``skipped`` + (names already present), ``shadowed`` (``name``/``key``/ + ``existing_kind``/``suggested_key`` — a different-kind collision, + nothing registered), and ``errors`` (``name``/``error`` pairs for + anything that failed to register — a bad store never aborts the rest). + + Raises: + HTTPException: 400/404 if *scan_root* itself is unusable. + """ + if on_conflict not in ("skip", "replace"): + on_conflict = "skip" + renames = renames or {} + + root = Path(scan_root).expanduser() + if not root.is_absolute(): + raise HTTPException(400, f"Path must be absolute: {scan_root!r}") + if not root.is_dir(): + raise HTTPException(404, f"No such directory: {root}") + + candidates = sorted((p for p in root.iterdir() if _is_zarr_dir(p)), key=lambda p: p.name) + + parts = [p for p in container_path.strip("/").split("/") if p] + target = ingest_mod._walk(get_tiled_client(server_uri, api_key_for_uri(server_uri)), parts) + + registered: list[dict[str, Any]] = [] + skipped: list[str] = [] + shadowed: list[dict[str, str]] = [] + errors: list[dict[str, str]] = [] + for candidate in candidates: + default_key = registered_key(candidate) + key = renames.get(candidate.name, default_key) + + if target is not None and key in ingest_mod._child_keys(target): + existing_kind = ingest_mod.node_source_kind(target[key]) + if existing_kind != "zarr": + shadowed.append( + { + "name": candidate.name, + "key": key, + "existing_kind": existing_kind or "unknown", + "suggested_key": f"{key}_zarr", + } + ) + continue + + try: + if key != default_key: + # register_zarr always derives the key from the path's own + # filename — the only way to register under a chosen + # alternate key is to nest one level deeper, in a container + # named for it. Only affects this explicit, rare rename- + # recovery path, not registration in general. + nested_container = f"{container_path}/{key}".strip("/") + result = register_zarr(server_uri, str(candidate), nested_container, on_conflict=on_conflict) + if not result.get("skipped"): + result = {**result, "tiled_path": f"{nested_container}/{result['key']}"} + else: + result = register_zarr(server_uri, str(candidate), container_path, on_conflict=on_conflict) + except HTTPException as exc: + errors.append({"name": candidate.name, "error": str(exc.detail)}) + continue + except Exception as exc: # noqa: BLE001 — one bad store must not sink the scan + errors.append({"name": candidate.name, "error": str(exc)}) + continue + if result.get("skipped"): + skipped.append(key) + else: + registered.append({"name": candidate.name, "key": key, "tiled_path": result["tiled_path"]}) + + return { + "scanned": len(candidates), + "registered": registered, + "skipped": skipped, + "shadowed": shadowed, + "errors": errors, + } diff --git a/docker-compose.als-prod.yml b/docker-compose.als-prod.yml new file mode 100644 index 0000000..03085a2 --- /dev/null +++ b/docker-compose.als-prod.yml @@ -0,0 +1,45 @@ +# The ":als-prod" shape: backend (with ml/dlsia) + ipred bundled, pointed at +# ALS's real PRODUCTION Tiled server by default — Dockerfile's app-ml stage, +# base path baked in at /bl832/seg_studio/ (hub.als.lbl.gov's own path +# prefix). This is what .github/workflows/publish-image.yml publishes under +# the `:als-prod` tag; this compose file is for running that published image +# (or building the same thing locally to test before it ships). +# +# Tiled itself stays external — ALS already runs its own production Tiled; +# bundling a second, empty one here would be actively wrong. Only +# TILED_API_KEY needs setting per deployment (a per-user/service credential — +# see docs/reference/deployment.md's open question on per-proposal Tiled auth +# before this goes live with real data). +# +# docker compose -f docker-compose.als-prod.yml up --build +# +# To run the already-published ghcr.io image instead of building locally, +# skip `build:` and use `image: ghcr.io//:als-prod`. +# +name: segmentation_annotation_studio +services: + app-als-prod: + build: + context: . + target: app-ml + args: + VITE_BASE_PATH: "/bl832/seg_studio/" + ports: + - "8002:8002" # the app itself (SPA + API) — the only port most setups need + - "8003:8003" # ipred, direct access (optional) + environment: + # ALS's real production Tiled. Override via TILED_URI in your own .env + # only if this ever needs to point somewhere else temporarily. + TILED_URI: "${TILED_URI:-https://tiled.als.lbl.gov}" + TILED_API_KEY: "${TILED_API_KEY:-}" + # Real institutional path into ALS's production catalog — confirm with + # whoever operates that Tiled server before assuming this is exactly + # right for a different beamline. + TILED_BROWSE_PATH: "${TILED_BROWSE_PATH:-beamlines/bl832/processed}" + LOCAL_DATA_ROOT: "/data" + BROWSE_ALLOWED_ORIGINS: "${BROWSE_ALLOWED_ORIGINS:-}" + volumes: + - annotation-data-als-prod:/data + +volumes: + annotation-data-als-prod: diff --git a/docker-compose.als-staging.yml b/docker-compose.als-staging.yml new file mode 100644 index 0000000..b2c303c --- /dev/null +++ b/docker-compose.als-staging.yml @@ -0,0 +1,33 @@ +# The ":als-staging" shape — identical to docker-compose.als-prod.yml except +# it defaults to ALS's STAGING Tiled server instead of production. See that +# file for the full explanation; this one exists so pointing at staging vs. +# production is a matter of which compose file you run, not remembering to +# override TILED_URI correctly every time. +# +# docker compose -f docker-compose.als-staging.yml up --build +# +# To run the already-published ghcr.io image instead of building locally, +# skip `build:` and use `image: ghcr.io//:als-staging`. +# +name: segmentation_annotation_studio +services: + app-als-staging: + build: + context: . + target: app-ml + args: + VITE_BASE_PATH: "/bl832/seg_studio/" + ports: + - "8002:8002" + - "8003:8003" + environment: + TILED_URI: "${TILED_URI:-https://tiled-staging.als.lbl.gov}" + TILED_API_KEY: "${TILED_API_KEY:-}" + TILED_BROWSE_PATH: "${TILED_BROWSE_PATH:-beamlines/bl832/processed}" + LOCAL_DATA_ROOT: "/data" + BROWSE_ALLOWED_ORIGINS: "${BROWSE_ALLOWED_ORIGINS:-}" + volumes: + - annotation-data-als-staging:/data + +volumes: + annotation-data-als-staging: diff --git a/docker-compose.full.yml b/docker-compose.full.yml new file mode 100644 index 0000000..abae5ce --- /dev/null +++ b/docker-compose.full.yml @@ -0,0 +1,40 @@ +# Batteries-included stack: Tiled + backend (with the ml/dlsia extra) + ipred, +# all bundled into ONE container (Dockerfile's app-full stage) — nothing +# external required, unlike docker-compose.yml's lean `app` service, which +# needs a separately-run Tiled. Use this for a single-command full-stack +# try-it (iPred + dlsia both work), or docker-compose.yml for a lighter +# deployment that already has its own Tiled. +# +# docker compose -f docker-compose.full.yml up --build +# +name: segmentation_annotation_studio +services: + app-full: + build: + context: . + target: app-full + ports: + - "8002:8002" # the app itself (SPA + API) — the only port most setups need + - "8003:8003" # ipred, direct access (optional) + - "8010:8010" # Tiled, direct access (optional) + environment: + LOCAL_DATA_ROOT: "/data" + BROWSE_ALLOWED_ORIGINS: "${BROWSE_ALLOWED_ORIGINS:-}" + # Only relevant if you point this stack's bundled Tiled at pre-existing + # data rather than this app's own ingest — leave unset otherwise: + TILED_BROWSE_PATH: "${TILED_BROWSE_PATH:-}" + # Set this to keep the same Tiled API key across container restarts — + # otherwise docker-entrypoint-full.sh generates a fresh one each time, + # which is fine (Tiled's catalog itself persists in the volume below) + # but means any client that cached the old key needs to reconnect. + TILED_API_KEY: "${TILED_API_KEY:-}" + volumes: + - annotation-data-full:/data + # Bind-mount your own source datasets to /data/processed (matches + # tiled/config.docker.yml's readable_storage). Set LOCAL_SOURCE_DIR in + # your own .env (see .env.example at the repo root) to a real host + # directory; defaults to an empty placeholder so this works unset. + - "${LOCAL_SOURCE_DIR:-./.local-source}:/data/processed" + +volumes: + annotation-data-full: diff --git a/docker-compose.local.yml b/docker-compose.local.yml new file mode 100644 index 0000000..8c221f8 --- /dev/null +++ b/docker-compose.local.yml @@ -0,0 +1,61 @@ +# The ":local" shape: fully bundled (Tiled + backend/ml + ipred, same as +# docker-compose.full.yml) but with VITE_BASE_PATH baked in at +# /seg_studio/ so it's reachable at the familiar dev-server-shaped URL +# http://localhost:5173/seg_studio/ by default — exercising the same +# subpath-hosting behavior ALS's own :als-prod/:als-staging deployments use, +# rather than leaving that path only tested at ALS. Host port is configurable +# via HOST_PORT for anyone who wants a different one. Published by +# .github/workflows/publish-image.yml (mirroring als-computing/ +# view_tomography_recon_app's own publish-image.yml — same tag names, same +# one-image-many-tags shape). +# +# The app container itself has NO idea it's hosted under a subpath — same as +# the real ALS deployment, it depends entirely on a reverse proxy stripping +# the prefix before the request arrives (see docs/reference/architecture.md's +# "Subpath hosting" note). A tiny nginx service (docker/nginx-local.conf) +# provides that here, so this compose file actually reproduces the real +# hosting behavior instead of only baking in URLs nothing then serves. +# +# docker compose -f docker-compose.local.yml up --build +# +# To run the already-published ghcr.io image instead of building locally, +# skip `build:` on app-local and use `image: ghcr.io//:local`. +# +name: segmentation_annotation_studio +services: + app-local: + build: + context: . + target: app-full + args: + VITE_BASE_PATH: "/seg_studio/" + expose: + - "8002" + ports: + - "8003:8003" # ipred, direct access (optional) + - "8010:8010" # Tiled, direct access (optional) + environment: + LOCAL_DATA_ROOT: "/data" + BROWSE_ALLOWED_ORIGINS: "${BROWSE_ALLOWED_ORIGINS:-}" + TILED_API_KEY: "${TILED_API_KEY:-}" + volumes: + - annotation-data-local:/data + # Bind-mount your own source datasets to /data/processed — Tiled's + # ingest/register-in-place routes (Browse's Zarr loader, the Connect + # page's ingest flow) read directly from here. Must match tiled/config + # .docker.yml's readable_storage exactly. Set LOCAL_SOURCE_DIR in your + # own .env (see .env.example at the repo root) to a real host + # directory; defaults to an empty placeholder so this works unset. + - "${LOCAL_SOURCE_DIR:-./.local-source}:/data/processed" + + proxy: + image: nginx:stable-alpine + depends_on: + - app-local + ports: + - "${HOST_PORT:-5173}:80" + volumes: + - ./docker/nginx-local.conf:/etc/nginx/conf.d/default.conf:ro + +volumes: + annotation-data-local: diff --git a/docker-compose.ml.yml b/docker-compose.ml.yml new file mode 100644 index 0000000..a54afb1 --- /dev/null +++ b/docker-compose.ml.yml @@ -0,0 +1,50 @@ +# Backend (with the ml/dlsia extra) + ipred bundled; Tiled stays external — +# point TILED_URI/TILED_API_KEY at your own Tiled server (Dockerfile's app-ml +# stage — the same shape ghcr.io publishes under the ALS-specific `:als-prod`/ +# `:als-staging` tags, see docker-compose.als-prod.yml/docker-compose.als- +# staging.yml for those with real defaults baked in). This file is the +# generic, no-institution-specific-defaults version — use it when you already +# have a Tiled server running (ALS's or anyone else's) and want Train/iPred +# to work without a second, separately-maintained ipred deployment. For a +# fully bundled stack with no external services at all, use docker- +# compose.full.yml instead; for the lightest option (no ipred/ML), use +# docker-compose.yml. +# +# docker compose -f docker-compose.ml.yml up --build +# +name: segmentation_annotation_studio +services: + app-ml: + build: + context: . + target: app-ml + args: + # Bare path only, never a scheme/host — see .env.example and + # docs/reference/deployment.md's "Hosting under a URL prefix". Leave + # unset for root-hosted. + VITE_BASE_PATH: "${VITE_BASE_PATH:-}" + ports: + - "8002:8002" # the app itself (SPA + API) — the only port most setups need + - "8003:8003" # ipred, direct access (optional) + environment: + # External Tiled server (required) — set these in your own .env for a + # staging/production deployment (see .env.example at the repo root; + # docker compose loads it automatically): + TILED_URI: "${TILED_URI:-http://host.docker.internal:8010}" + TILED_API_KEY: "${TILED_API_KEY:-}" + # Path into the Tiled tree Browse treats as its root — set this to the + # real institutional path (e.g. beamlines/bl832/processed) when + # TILED_URI points at an existing catalog rather than a from-scratch one: + TILED_BROWSE_PATH: "${TILED_BROWSE_PATH:-}" + # Persisted annotation drafts/versions/exports live under this path: + LOCAL_DATA_ROOT: "/data" + # Same-origin SPA → CORS can be empty; set if you split origins: + BROWSE_ALLOWED_ORIGINS: "${BROWSE_ALLOWED_ORIGINS:-}" + volumes: + - annotation-data-ml:/data + # host.docker.internal lets the container reach a Tiled on the host (Linux): + extra_hosts: + - "host.docker.internal:host-gateway" + +volumes: + annotation-data-ml: diff --git a/docker-compose.yml b/docker-compose.yml index c1fa1b7..16f6ae3 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -3,6 +3,7 @@ # # docker compose up --build # +name: segmentation_annotation_studio services: app: build: . @@ -12,6 +13,10 @@ services: # External Tiled server (required): TILED_URI: "${TILED_URI:-http://host.docker.internal:8010}" TILED_API_KEY: "${TILED_API_KEY:-}" + # Path into the Tiled tree Browse treats as its root — set this for a + # deployment pointed at an existing institutional Tiled catalog rather + # than a from-scratch one (see .env.example): + TILED_BROWSE_PATH: "${TILED_BROWSE_PATH:-}" # Persisted annotation drafts/versions/exports live under this path: LOCAL_DATA_ROOT: "/data" # Same-origin SPA → CORS can be empty; set if you split origins: diff --git a/docker-entrypoint-full.sh b/docker-entrypoint-full.sh new file mode 100755 index 0000000..e71cb9c --- /dev/null +++ b/docker-entrypoint-full.sh @@ -0,0 +1,102 @@ +#!/bin/sh +# Entrypoint for the app-full Docker image (Dockerfile's app-full stage): +# starts Tiled and ipred as background processes, then execs the backend in +# the foreground as the container's main process. All three talk over +# 127.0.0.1 since they share one container — no Docker networking needed. +# +# Known limitation, deliberate for this "batteries included, single command" +# image: Tiled/ipred are plain backgrounded processes, not supervised — if +# either crashes after startup, the container keeps running (its main +# process is the backend) but that service silently stays down until the +# whole container is restarted. Fine for local/demo use; a production +# deployment that needs real service supervision should run Tiled and ipred +# as separate containers instead (see docker-compose.yml + an external Tiled, +# or split ipred out the same way if this ever needs to be hardened further). +set -e + +# Generate a Tiled API key if one wasn't provided, mirroring start_all.sh's +# own local-dev behavior — without one, Tiled's `allow_anonymous_access: true` +# only permits reads (see tiled/config.docker.yml), so writes (ingest, mask +# sync, volume registration) would fail with no explanation on first run. +if [ -z "${TILED_API_KEY:-}" ]; then + export TILED_API_KEY="$(python3 -c 'import secrets, string; print("".join(secrets.choice(string.ascii_letters + string.digits) for _ in range(32)))')" + echo "Generated a Tiled API key for this container (not persisted — set TILED_API_KEY yourself to keep one across restarts)." +fi + +mkdir -p /data/.tiled/data /data/.tiled/volumes /data/processed + +tiled serve config /app/tiled/config.docker.yml --host 0.0.0.0 --port 8010 --api-key "$TILED_API_KEY" & +uvicorn ipred.api:app --host 0.0.0.0 --port 8003 & + +# Wait for Tiled to actually accept connections before starting the backend — +# without this, the backend's first Tiled-dependent request can race a still- +# initializing catalog and fail. python3 (not curl, which python:3.12-slim +# doesn't ship) is guaranteed present in this image already. +python3 -c " +import socket, time +for _ in range(60): + try: + socket.create_connection(('127.0.0.1', 8010), timeout=1).close() + break + except OSError: + time.sleep(0.5) +else: + print('Tiled did not become reachable on port 8010 in time — starting the backend anyway.') +" + +export TILED_URI="http://127.0.0.1:8010" +export IPRED_URL="http://127.0.0.1:8003" +# Keep the mask/volume pyramid cache under the same persisted /data volume, +# in the exact path tiled/config.docker.yml's readable_storage expects it — +# same "keep in step" rule tiled/config.yml documents for local dev. +export VOLUME_CACHE_DIR="/data/.tiled/volumes" + +# Auto-register any Zarr stores AND image-slice folders already sitting under +# the bind-mounted LOCAL_SOURCE_DIR (see docker-compose.full.yml/docker- +# compose.local.yml) — without this, a container brought up against a folder +# of pre-existing data shows nothing in Browse until someone manually scans or +# registers each one through the UI. Safe on every restart: both scans skip +# anything already registered, so this never re-registers or duplicates +# existing entries. The image-stack scan deliberately does NOT build a 3-D +# pyramid here — that's comparatively expensive (reads every slice) and is +# left to the on-demand "Build 3D volume" button on the 3D page, so startup +# stays fast regardless of how large a folder of TIFFs/PNGs is. Sequential, +# not concurrent: both scans call _ensure_container against the same target +# container, and running them in parallel would race on its creation. +# Non-fatal on failure — a scan problem must never block the app from +# starting; it just leaves auto-discovery for that run to be retried +# manually via the "Scan folder for datasets" button in the Zarr loader +# (POST /api/scan-datasets). +python3 -c " +import sys +sys.path.insert(0, '/app') +import ingest +import zarr_source + +def report(label, result): + print( + f\"{label}: {len(result['registered'])} new, \" + f\"{len(result['skipped'])} already present, \" + f\"{len(result['shadowed'])} shadowed, {len(result['errors'])} failed.\" + ) + for s in result['shadowed']: + print( + f\" - {s['name']}: same name already registered as {s['existing_kind']!r} — \" + f\"retry via the Zarr loader's Scan button with a different key \" + f\"(e.g. {s['suggested_key']!r}) if you want both.\" + ) + for err in result['errors']: + print(f\" - {err['name']}: {err['error']}\") + +try: + report('Zarr auto-registration', zarr_source.scan_and_register_zarrs('http://127.0.0.1:8010', '/data/processed', 'browse')) +except Exception as exc: + print(f'Zarr auto-registration scan failed (non-fatal): {exc}') + +try: + report('Image-stack auto-registration', ingest.scan_and_register_image_stacks('http://127.0.0.1:8010', '/data/processed', 'browse')) +except Exception as exc: + print(f'Image-stack auto-registration scan failed (non-fatal): {exc}') +" || true + +exec uvicorn annotation_server:app --host 0.0.0.0 --port 8002 diff --git a/docker-entrypoint-ml.sh b/docker-entrypoint-ml.sh new file mode 100755 index 0000000..422b0f9 --- /dev/null +++ b/docker-entrypoint-ml.sh @@ -0,0 +1,21 @@ +#!/bin/sh +# Entrypoint for the app-ml Docker image (Dockerfile's app-ml stage): starts +# ipred as a background process, then execs the backend in the foreground as +# the container's main process. Unlike app-full's entrypoint, this does NOT +# start (or need) a Tiled process at all — TILED_URI must be set at `docker +# run` time to point at the deployment's own external, already-running Tiled +# (e.g. ALS's production Tiled for the `:als` tag). +# +# Known limitation, same as app-full's entrypoint: ipred is a plain +# backgrounded process, not supervised — if it crashes after startup, the +# container keeps running (its main process is the backend) but iPred/Train +# silently stay down until the whole container is restarted. +set -e + +mkdir -p /data + +uvicorn ipred.api:app --host 0.0.0.0 --port 8003 & + +export IPRED_URL="http://127.0.0.1:8003" + +exec uvicorn annotation_server:app --host 0.0.0.0 --port 8002 diff --git a/docker/nginx-local.conf b/docker/nginx-local.conf new file mode 100644 index 0000000..c419ee3 --- /dev/null +++ b/docker/nginx-local.conf @@ -0,0 +1,50 @@ +# Minimal stripping reverse proxy for docker-compose.local.yml — simulates +# the hub.als.lbl.gov-style proxy the `:als`/`:local` base-path design +# actually depends on (see the "Base-path (subpath) deployment support" +# section of the deployment plan). The app container itself has NO idea it's +# hosted under /seg_studio/ — its static file serving and SPA catch-all only +# ever see plain root-relative paths. This proxy is what turns +# "localhost:5173/seg_studio/..." into "app-local:8002/..." before it ever +# reaches the container, exactly like the real reverse proxy would strip +# "hub.als.lbl.gov/bl832/seg_studio/..." down to "/...". +server { + listen 80; + + # nginx's default (1MB) is nowhere near enough for the Ingest dropzone's + # multi-file uploads (a folder of TIFFs easily reaches hundreds of MB) — + # without this, an upload past 1MB fails with a raw nginx 413 before the + # request ever reaches the backend. 0 = unlimited; the backend's own + # ingest route is the real, size-aware boundary already (job-based, + # streamed to disk), so there's no reason to cap here too. A real + # ALS-hub-style production proxy needs this exact same setting — it is + # NOT specific to this local test proxy. + client_max_body_size 0; + + location = / { + return 302 /seg_studio/; + } + + location /seg_studio/ { + # Resolve at request time (Docker's embedded DNS), not once at nginx + # startup — a plain `proxy_pass http://app-local:8002` fails hard at + # boot if app-local isn't up yet, and keeps a stale IP across an + # app-local restart/recreate. `set` must come BEFORE the `rewrite... + # break`, since `break` halts further rewrite-module directives + # (including a later `set`) in this block. + resolver 127.0.0.11 valid=10s; + set $upstream app-local:8002; + rewrite ^/seg_studio/(.*)$ /$1 break; + proxy_pass http://$upstream; + proxy_set_header Host $host; + proxy_set_header X-Forwarded-Proto $scheme; + proxy_http_version 1.1; + proxy_set_header Upgrade $http_upgrade; + proxy_set_header Connection "upgrade"; + # Large multi-file uploads over a slow/local link can legitimately + # take a while; don't let nginx's default 60s read/send timeouts cut + # them off mid-transfer. + proxy_read_timeout 600s; + proxy_send_timeout 600s; + client_body_timeout 600s; + } +} diff --git a/docs/getting-started/installation.md b/docs/getting-started/installation.md index 1f0470d..5223c78 100644 --- a/docs/getting-started/installation.md +++ b/docs/getting-started/installation.md @@ -3,7 +3,11 @@ There are two ways to run Segmentation Annotation Studio: - **Local development** — one command starts everything (recommended for annotators and evaluation). -- **Docker** — a single production container that serves the app against an external Tiled server. +- **Docker** — a prebuilt or locally-built container image; pick the shape that matches + whether you already have a Tiled server and whether you want iPred/Train available. + +For hosting a shared, ALS-style deployment behind a reverse proxy (rather than running +it yourself), see [Production deployment](../reference/deployment.md) instead. --- @@ -67,30 +71,82 @@ FRONTEND_PORT=5200 BACKEND_PORT=8100 TILED_PORT=8110 ./start_all.sh --- -## Option 2 — Docker (production) +## Option 2 — Docker + +Three Dockerfile targets/compose files cover different needs — pick the one that +matches what you already have running and whether you want iPred/Train available. +None replaces the others; `app-ml` and `app-full` each build on the previous stage +rather than duplicating install steps. -Docker runs a **single container** that serves the built frontend and the API -together on port **8002**. Tiled is **not** included — you point the container -at an existing Tiled server. +| Compose file | Image target | Tiled | iPred / Train | Use when | +| --- | --- | --- | --- | --- | +| `docker-compose.yml` | `app` | External (you provide one) | Not included | You already have a Tiled server and only need Connect/Browse/Annotate/Export. | +| `docker-compose.ml.yml` | `app-ml` | External (you provide one) | **Bundled** | You already have a Tiled server but also want Train/iPred to work, without standing up a separate iPred service. | +| `docker-compose.full.yml` | `app-full` | **Bundled** | **Bundled** | You want the whole stack (Tiled + backend + iPred) with nothing external to set up — the simplest way to try everything. | + +### Lean (`app`) — bring your own Tiled ```bash +TILED_URI=https://tiled.example.com \ +TILED_API_KEY=your-key \ docker compose up --build ``` Then open . -Configure the connection to your external Tiled through environment variables -(see [Environment variables](#environment-variables)): +### With iPred/Train, external Tiled (`app-ml`) ```bash TILED_URI=https://tiled.example.com \ TILED_API_KEY=your-key \ -docker compose up --build +docker compose -f docker-compose.ml.yml up --build +``` + +Then open . iPred is also reachable directly on `:8003` if needed. + +### Fully bundled (`app-full`) — nothing external required + +```bash +docker compose -f docker-compose.full.yml up --build ``` +Then open . Tiled (`:8010`) and iPred (`:8003`) are also +reachable directly if you want to hit them outside the app. + !!! warning "Persisting your data" - Mount `LOCAL_DATA_ROOT` as a volume so annotation drafts, versions, and - exports survive container restarts. + Every compose file above mounts `LOCAL_DATA_ROOT` (`/data`) as a named volume, + so annotation drafts, versions, and exports already survive a container + restart. + +!!! tip "Pointing at your own datasets" + For `app-full` (or the published `:local` image, which builds from the same + target — see [Deployment](../reference/deployment.md)), set `LOCAL_SOURCE_DIR` in your own `.env` (copy + the repo root's `.env.example`) to an absolute host directory — it's + bind-mounted to `/data/processed` inside the container, so the bundled + Tiled server (and the Zarr loader's "Browse…" directory picker on the + Connect page) can read your real data directly. This matters specifically + in Docker: a path from your own machine (e.g. one you'd type into the Zarr + loader by hand) means nothing to the container unless it's actually + mounted like this — `LOCAL_SOURCE_DIR` is what makes it visible. + ```bash + LOCAL_SOURCE_DIR=/absolute/path/to/your/data docker compose -f docker-compose.full.yml up --build + ``` + + **Everything under that directory is registered into Tiled and shows up in + Browse automatically** — on every container start, and again whenever you + click **"Scan folder for datasets"** in the Zarr loader's directory + browser (useful after adding new files without restarting). This covers + both `.zarr` stores and plain folders of TIFF/PNG/JPG slices — the latter + register as fast, per-slice ingest with no 3-D pyramid built yet (build + one on demand from the 3D tab when you actually need it, so a large + dataset doesn't delay startup). + + If a raw image folder and a `.zarr` reconstruction of the same acquisition + share a name, only one can occupy that key — the scan reports the second + one as **shadowed** rather than silently skipping it, and offers a + "Register as…" action (in the UI, or `POST /api/scan-datasets` / + `/api/ingest/scan` with a `renames` field) to register it under a + different key so both show up side by side. --- @@ -137,7 +193,14 @@ npm test # run the Vitest unit tests ## Environment variables Backend configuration lives in `backend/.env` (created from -`backend/.env.example` on first launch). The most relevant settings: +`backend/.env.example` on first launch) when running via `start_all.sh` / +directly with `uvicorn`. **Running via `docker compose` instead**, copy the +repo root's own `.env.example` to `.env` — `docker compose` loads that +automatically for every `docker-compose*.yml` file, so this is where +`TILED_URI`/`TILED_API_KEY`/`TILED_BROWSE_PATH` (and, for `app-ml`, +`VITE_BASE_PATH`) actually get substituted in. See +[Production deployment](../reference/deployment.md#setting-these-for-docker-compose) +for details. The most relevant settings either way: | Variable | Purpose | Default | | --- | --- | --- | @@ -147,6 +210,12 @@ Backend configuration lives in `backend/.env` (created from | `EXPORT_ROOT` | Override output folder for exports | `~/data/exports` | | `BROWSE_CACHE_TTL_SECONDS` | Cache lifetime for Tiled browse listings | `300` | | `BROWSE_ALLOWED_ORIGINS` | CORS origins (only needed for split frontend/backend hosting) | *(empty)* | +| `TILED_BROWSE_PATH` | Path into the Tiled tree Browse treats as its root | *(unset — this repo's own ingest root)* | + +`VITE_BASE_PATH` is a **frontend build-time** setting (a Docker build-arg, not a +runtime env var — see [Production deployment](../reference/deployment.md)) for +hosting under a URL prefix rather than at the domain root; leave it unset for +everything on this page. !!! danger "Never commit secrets" `backend/.env` is git-ignored. Never commit it, and never expose diff --git a/docs/guide/index.md b/docs/guide/index.md index 1c86815..b8c43ad 100644 --- a/docs/guide/index.md +++ b/docs/guide/index.md @@ -12,7 +12,7 @@ Every screen shares the same layout: - **Header** — the ALS logo and the app title, *"Segmentation Annotation Tool"*. - **Main area** — the current tab's content. -The sidebar contains four tabs by default: +The sidebar contains six tabs by default: | Tab | Icon | What it's for | | --- | --- | --- | @@ -20,6 +20,12 @@ The sidebar contains four tabs by default: | **Browse** | magnifier | Filter and pick a sample to annotate. | | **Reference** | book | Write class descriptions (the annotation guide). | | **Annotate** | pencil | Draw masks, manage versions, and export. | +| **3D** | cube | View the reconstruction and mask layers in 3D. | +| **Train** | brain | Train a dlsia deep-learning segmentation model and run it on a whole volume. | + +A small indicator in the header shows live Tiled connection status (green +"Tiled connected" / red "Tiled disconnected" — click it to jump back to +Connect if it drops). !!! tip "Customize which tabs you see" Click the floating **Customize Layout** button (top-right) to open @@ -33,3 +39,5 @@ The sidebar contains four tabs by default: 3. [Annotate](annotate.md) — the core drawing workflow. 4. [Annotation guide](reference-guide.md) — keep multi-annotator projects consistent. 5. [Export & download](export.md) — produce a COCO dataset. +6. [Train a deep model](train.md) — optional, for a volume-wide model beyond the fast in-tab classifier. +7. [3D volume view](volume.md) — optional, for viewing the reconstruction and mask layers in 3D. diff --git a/docs/guide/train.md b/docs/guide/train.md new file mode 100644 index 0000000..a342ccc --- /dev/null +++ b/docs/guide/train.md @@ -0,0 +1,130 @@ +# 6. Train a deep model + +The **Train** tab fine-tunes a real deep-learning segmentation model (a +**dlsia TUNet**) on this session's annotated samples, then runs it across a +whole volume — a different, slower, more accurate path than the Annotate +tab's own fast in-browser pixel classifier (the **Predict** stage under +Assist/Predict). Use Train when you want a model that generalizes well beyond +the slices you've personally annotated, or that you plan to reuse across +sessions. + +!!! note "This is optional" + Nothing else in the tool requires the Train tab. If the fast pixel + classifier already gives you good enough results, you never need to open + this page. + +## Is training available? + +At the top of the tab, a status banner reports readiness: + +- **Ready**: *"torch {version} · {device}"* (e.g. `mps`, `cuda`, or `cpu`) and + *"dlsia (TUNet): available"*. +- **Unavailable**: *"Training is unavailable on this server."* — the server + wasn't started with ML dependencies installed (the Docker image + intentionally omits them to stay lightweight; run via `start_all.sh` with + `INSTALL_ML=1` — the default on Apple Silicon — on a machine that can + install them). +- If a job is already running, the banner adds *"A training/inference job is + currently running."* — only one training or inference job runs at a time. + +--- + +## Training data + +The **Training data** panel lists every sample you've annotated **this +session**, with its shape count, as checkboxes: + +- If nothing shows up: *"No annotated samples yet this session. Annotate a + few slices in the Annotate tab, then come back here."* +- Check the samples to include; the header shows *"N of M selected."* + +!!! tip + More annotated slices across more samples generally makes for a better + model — this is real deep-learning training, not the fast classifier's + per-sample fit. + +### Train on denoised input (optional) + +If you've tuned a denoise filter in the Annotate tab's Display panel, a +**"Train on denoised input"** checkbox appears, naming the exact filter (e.g. +*"(median, 40%)"*) rather than offering a second, separate copy of the +controls. This is **off by default** and changes what the model actually +*learns* — unlike every other denoise control in the app, which is +display-only. The setting is recorded on the saved run, and inference +automatically reapplies the same filter, so training and prediction can never +disagree about what the model is looking at. + +If no denoise filter is set (or the current one can't be used for training), +the checkbox is disabled with an explanation instead of silently doing +nothing. + +--- + +## Hyperparameters (advanced) + +Collapsed by default — sensible defaults are supplied, so most users never +need to open this. When you do: + +| Field | What it controls | +| --- | --- | +| **Run name** | Optional label; auto-generated if left blank. | +| **Epochs**, **Learning rate** | Standard training controls. | +| **Batch size** | How many patches process at once. Click **Estimate max** to have the server run real training steps at increasing batch sizes and find the largest that fits your GPU/memory — this needs the device to itself, so it's disabled while another job is running. | +| **Patch/Image size (px)** | The window size the model trains on. | +| **Random flip augmentation** | Cheap data augmentation. | +| **Tile large images** | **On** (default): cuts native-resolution patches with 25% overlap and blends predictions back together — keeps fine detail on images larger than the patch size. **Off**: shrinks each whole slice to the patch size before training — faster, but loses detail on large images. | +| **Depth**, **Base channels**, **Growth rate** | TUNet architecture parameters. | + +--- + +## Start training + +Click **Start training**. A progress bar tracks epochs and (once available) +mIoU; cancel is cooperative — a slice already inside a GPU forward pass +finishes before honoring cancel. + +## Saved runs + +Every completed (or cancelled-but-partial) run appears under **Saved runs**: +model family, timestamp, mIoU, epochs completed, whether it was trained on +denoised input, and its class list. Select a run's radio button to use it for +inference below. Click the trash icon to permanently delete a run's saved +weights (asks for confirmation, naming the run precisely since two runs +trained close together can otherwise look identical). + +--- + +## Inference + +With a saved run selected and a sample open (in Browse/Annotate), pick a +scope: + +- **Current slice** — just the slice open in the viewer. +- **Slice range** — a start/end slice. +- **All slices** — the whole volume. + +Click **Run inference**. Progress shows *"{done}/{total} slices"* with a live +log of region counts per slice, and **the preview slider becomes usable as +soon as the first slice finishes** — it keeps extending as more slices +complete while the job is still running, not just after it finishes. +**Cancel** stops a running job cooperatively. + +Once done (or, for slices already predicted, while still running), you get: + +- **Import as annotations** — vectorizes the predicted regions into real, + editable shapes back in the Annotate tab. +- **Write masks to Tiled** *(Tiled sources only)* — writes the prediction + directly into a `__masks_deep` container in Tiled, independent of + whatever the Annotate tab's own "Push masks to Tiled" wrote — so the 3D + view's **Deep** mask layer and **Fast** mask layer can show two genuinely + different results side by side. See [3D volume view](volume.md#mask-layers). + +!!! tip "Large inference jobs can take a while to write" + A big "Write masks to Tiled" write (hundreds of slices) can legitimately + take a while over the network. If the 3D view's mask panel shows *"Still + loading…"* rather than an error, it's still working — a **Check again** + button appears if it's still not done after a very generous wait. + +--- + +Next: [3D volume view →](volume.md) diff --git a/docs/guide/volume.md b/docs/guide/volume.md new file mode 100644 index 0000000..fdeac59 --- /dev/null +++ b/docs/guide/volume.md @@ -0,0 +1,116 @@ +# 7. 3D volume view + +The **3D** tab renders the open dataset's reconstruction and mask layers +directly in 3D, streamed straight from Tiled — there is no export step. It +uses WebGPU, so it needs a modern Chromium-based browser (Chrome or Edge +113+) served over a secure context (`https://`, or `localhost`/`127.0.0.1`). + +!!! note "This is optional" + Nothing else in the tool requires the 3D view — it's there for inspecting + results in context, not for annotating. + +## No volume yet? + +A dataset ingested as individual 2D slices has no 3D pyramid until one is +built. If you see **"No 3D volume for this dataset yet"**, the panel shows the +source shape/dtype and the pyramid it will build, then a **Build 3D volume** +button. This reads every slice once; full resolution is **not** copied — it +stays where it is, and the 3D view never loads a level too large for the GPU. + +If a volume already exists but needs refreshing (e.g. after a fidelity setting +changed), a small **Rebuild volume** button sits in the bottom-left corner of +the 3D view itself. + +--- + +## The render HUD + +A docked sidebar on the right controls how the volume looks: + +| Panel | Controls | +| --- | --- | +| **Data** | Resolution/LOD picker, voxel size, high-res ROI streaming. | +| **Transfer Function** | Colormap, opacity curve (drag/add/remove points on the histogram), color range (with **Auto**/**Equalize**), clip limit. Also has a **Bands** mode for multiple independent intensity ranges. | +| **Slices** | Axis-aligned slice planes through the volume. | +| **Crop** | An ROI crop box. | +| **Measure** | On-canvas measurement between points in the volume. | +| **Annotations** | The built-in mask-layer controls (this app also has its own, described below — either produces the identical visual result). | +| **Presets** | Save/apply/delete named transfer-function presets. | + +Pan with space+drag, shift/middle/right-wheel zoom to the cursor; ++p++/++ctrl+click++ +picks a point, ++l++ cycles LOD, ++o++ toggles open/collapsed. + +--- + +## Mask layers + +A small panel in the top-left corner controls two independent, fixed mask +layers: + +- **Fast (iPred)** — the Annotate tab's quick pixel classifier / manual + annotations. +- **Deep (dlsia)** — a [Train tab](train.md) model's output. + +Each has its own **Load** button, and (once loaded) an opacity slider and +per-class visibility toggles (click a class row to show/hide it). + +### Live vs. Tiled (Fast slot only) + +The Fast slot defaults to **Live** mode — it rasterizes your **current** +annotation shapes directly in the browser and loads instantly, with **no +Tiled sync required first**. Switch to **Tiled** to load the precise, +backend-rasterized result from `Push masks to Tiled` in the Annotate tab +instead. The Deep slot has no Live equivalent — it's always the saved output +of a from-scratch-trained model, loaded from the `__masks_deep` +container [Train's inference panel](train.md#inference) writes. + +!!! tip "Loading a large Deep mask can take a while" + A real Tiled-backed mask (hundreds of slices, freshly written) can + legitimately take a while to fetch over the network. The **Load** button + shows *"Loading…"* then *"Still loading…"* rather than failing outright — + if it's genuinely still not done after a very generous wait, a **Check + again** button appears rather than making you re-fetch from scratch. + +### Getting here from elsewhere + +- Annotate's iPred panel has its own **Push to Tiled** (stays on the tab) and + **View in 3D** (navigates here, loading the Fast layer) — split into two + independent actions so pushing doesn't force you into the 3D view, and + viewing doesn't force a fresh push. +- Train's inference panel's **Write masks to Tiled** writes the Deep layer + independently, so Fast and Deep can show two genuinely different results + side by side for comparison. + +--- + +## Isolating a feature by intensity (Sampler → 3D bridge) + +If you've used the Annotate tab's Sampler/Threshold-lasso tool (see +[Annotate → The Magic tool](annotate.md#the-magic-tool-smart-ai-classic) for +the sibling Magic tool, or the Threshold Brush's own **Set band from a +region**) to fit an intensity band around a small, materially-distinct +feature, its readout gains a **View band in 3D** button (for a plain, +non-projected fit only). Clicking it isolates that intensity band here: +opaque inside the band, transparent outside, via the same **Transfer +Function** the HUD's own panel controls. + +!!! note "This isolates by intensity, not by traced region" + It shows *everything* in that density range across the whole volume, not + only the specific spot you traced — genuinely useful when the feature's + density is distinct from its surroundings (the same property the 2D + threshold tool already relies on), less useful if it overlaps + similar-density material elsewhere. + +--- + +## Connection issues + +If the dataset can't be resolved because Tiled dropped (rather than because +no volume has been built yet), the view shows **"Couldn't reach the Tiled +server"** with a **Go to Connect** button, instead of leaving you on an +indefinite spinner. + +--- + +Back to [Train a deep model ←](train.md), or return to the +[recommended reading order](index.md#recommended-reading-order). diff --git a/docs/reference/architecture.md b/docs/reference/architecture.md index e637e47..25e6f20 100644 --- a/docs/reference/architecture.md +++ b/docs/reference/architecture.md @@ -10,9 +10,11 @@ debug, or deploy it. ## System context -The application is three cooperating processes. The browser only ever talks to -the **backend**; all catalog and array access is proxied server-side so that -Tiled credentials never reach the client. +The application is **four** cooperating processes (grew from three as dlsia +training/inference and the 3D volume viewer became core, always-on parts of +the stack rather than optional extras). The browser only ever talks to the +**backend**; all catalog/array access and all ML work are proxied server-side +so that Tiled credentials and model internals never reach the client. ```mermaid graph LR @@ -26,6 +28,10 @@ graph LR API["annotation_server.py"] end + subgraph Ipred["ipred · FastAPI :8003"] + IpredAPI["ipred.api:app
feature banks · dlsia train/infer
(torch, dlsia, qlty)"] + end + subgraph Data["Data services"] Tiled["Tiled server :8010
SQLite catalog + storage"] Disk[("LOCAL_DATA_ROOT
drafts · versions · exports")] @@ -33,17 +39,31 @@ graph LR User --> FE FE -->|"/api/* (fetch)"| API + FE -.->|"3D volume: Zarr chunks
fetched directly, read-only"| Tiled API -->|"tiled.client (HTTP)"| Tiled + API -->|"httpx (HTTP)"| IpredAPI + IpredAPI -->|"tiled.client (HTTP)"| Tiled API --> Disk - User -. never direct .-> Tiled + User -. never direct for writes .-> Tiled ``` | Component | Role | Default address | | --- | --- | --- | | **Frontend** | React app the annotator interacts with | | -| **Backend** | FastAPI: renders images, rasterizes masks, builds exports | | +| **Backend** | FastAPI: renders images, rasterizes masks, builds exports, proxies ipred | | +| **ipred** | FastAPI: feature banks, iPred pixel-classifier training, dlsia TUNet train/infer (GPU-capable) | | | **Tiled** | Data catalog for source images and mask write-back | | +!!! note "The 3D viewer is the one place the browser talks to Tiled directly" + Every other feature proxies through the backend (see + [Security boundaries](#security-boundaries)). The 3D **Volume** view is a + deliberate, narrow exception: the vendored WebGPU renderer streams Zarr + chunks straight from Tiled's `/zarr/v2` router (`lib/zarrUrl.ts`) for + performance — there is no export step or downsampled-volume endpoint in + between. This works because anonymous **read** access is enabled on Tiled; + the write-scoped API key never leaves the backend, and nothing here lets + the browser write to Tiled. + ## Technology stack === "Frontend" @@ -73,6 +93,18 @@ graph LR | Export | **pycocotools** | RLE encoding, bbox/area for COCO | | Files | **tifffile**, **imagecodecs** | Scientific TIFF reads | | Config | **python-dotenv** | Loading `backend/.env` | + | ML (optional extra) | **torch**, **dlsia**, **qlty** | dlsia TUNet train/infer, tiling geometry — gated behind the `ml` install extra; the backend runs fine without it, just without the Train tab's real functionality | + +=== "ipred (dlsia + pixel classifier)" + + | Concern | Package | Used for | + | --- | --- | --- | + | API | **fastapi** + **uvicorn** | Its own separate ASGI service, port 8003 | + | ML | **torch**, **dlsia** | TUNet segmentation model, denoiser autoencoder | + | Tiling | **qlty** | Patch tiling/stitching for training and inference on full-resolution slices | + | Classical ML | **scikit-learn** | The Annotate tab's "fast" pixel classifier (iPred) | + | Feature banks | **onnxruntime** (optional) | tomojepa/SAM embeddings for manifold-coverage sampling | + | Device | CUDA, MPS (Apple Silicon), or CPU | `train_common.pick_device()` auto-detects, `TRAIN_DEVICE` env overrides | ## Deployment topology @@ -86,25 +118,38 @@ The dev and production layouts differ mainly in **who serves the SPA** and Browser["Browser"] Vite["Vite :5173"] Backend["FastAPI :8002"] + Ipred["ipred :8003"] Tiled["Tiled :8010"] Disk[("~/data")] Browser -->|localhost:5173| Vite Vite -->|"/api/* proxy"| Backend Backend -->|tiled.client| Tiled + Backend -->|httpx| Ipred + Ipred -->|tiled.client| Tiled Backend --> Disk ``` - `start_all.sh` launches all three: Tiled (8010), backend (8002), and the - Vite dev server (5173). Vite proxies every `/api` request to the backend, - so the frontend uses an empty `API_BASE` and stays same-origin. + `start_all.sh` launches all four: Tiled (`$TILED_PORT`, default 8010), + backend (`$BACKEND_PORT`, default 8002), ipred (`$IPRED_PORT`, default + 8003), and the Vite dev server (`$FRONTEND_PORT`, default 5173) — each env + var overridable, and each port auto-bumped to the next free one if taken. + Vite proxies every `/api` request to the backend, so the frontend uses an + empty `API_BASE` and stays same-origin. -=== "Production (Docker)" +=== "Production (Docker) — three images, pick one" + + Three Dockerfile targets/compose files, layered `app` → `app-ml` → `app-full` + (each `FROM` the previous, so the install steps are never duplicated) — + each covers a different deployment need, none replaces the others. + + **`app` (lean)** — frontend + backend only, for a deployment with its own + external Tiled and no need for iPred/dlsia: ```mermaid flowchart LR Browser["Browser"] - Container["Single container :8002
API + static SPA"] + Container["app image :8002
API + static SPA"] TiledExt["External Tiled
(TILED_URI env)"] Volume[("/data volume")] @@ -113,17 +158,66 @@ The dev and production layouts differ mainly in **who serves the SPA** and Container --> Volume ``` - The build copies the compiled SPA into `backend/static/`, and FastAPI serves - both the API and the static files from one origin on port 8002. Tiled is an - external service referenced by `TILED_URI`; only `tiled[client]` ships in the - image. + **`app-ml`** — `app` plus the `ml` extra (torch/dlsia) and ipred bundled, + Tiled still external — for a deployment that already has its own + production Tiled (bundling a second, empty one would be wrong) but still + wants Train/iPred to work without standing up a separate ipred service. + This is what the `:als-prod`/`:als-staging` ghcr.io tags publish, base path + baked in at `/bl832/seg_studio/`: + + ```mermaid + flowchart LR + Browser["Browser"] + Container["app-ml image
backend :8002 (foreground)
+ ipred :8003 (background)"] + TiledExt["External Tiled
(TILED_URI env)"] + Volume[("/data volume")] + + Browser -->|same origin| Container + Container --> TiledExt + Container --> Volume + ``` + + **`app-full` (batteries-included)** — `app-ml` plus a bundled Tiled server + too, one container running all three backend-side services (Tiled and + ipred as backgrounded processes, backend as the container's foreground/main + process — no Docker networking needed). This is what the `:local` ghcr.io + tag publishes, base path baked in at `/seg_studio/`: + + ```mermaid + flowchart LR + Browser["Browser"] + Container["app-full image
backend :8002 (foreground)
+ ipred :8003 (background)
+ Tiled :8010 (background)"] + Volume[("/data volume
/data/processed bind mount")] + + Browser -->|":8002 only"| Container + Container --> Volume + ``` + + All three copy the compiled SPA into `backend/static/`; FastAPI serves the + API and static files from one origin. Known limitation (documented in + `docker-entrypoint-full.sh`'s/`docker-entrypoint-ml.sh`'s own comments): + Tiled/ipred aren't supervised inside `app-full`/`app-ml` — if one crashes + post-startup the container keeps running (backend is its main process) but + that service stays down until a restart. Fine for local/demo use; a real + deployment needing that resilience should use `app` against a + properly-supervised external Tiled instead. + + **Subpath hosting** (e.g. `hub.als.lbl.gov/bl832/seg_studio/`, behind a + reverse proxy that strips the prefix before forwarding) is a build-time + choice, not a runtime one — Vite bakes `base`/`import.meta.env.BASE_URL` + into the compiled JS, so a root-hosted image and a subpath-hosted image are + two different build artifacts from the same source. Set the `VITE_BASE_PATH` + build-arg (`vite.config.ts`'s `base`, threaded through to `main.tsx`'s + `BrowserRouter basename` and `config.ts`'s `API_BASE`) at `docker build` + time; leave it unset for root-hosted (local dev, and the generic + `:latest`-tagged builds of all three images above are unaffected). ## Frontend architecture ### Component shell The UI follows the ALS **Finch** hub pattern: a fixed icon sidebar, a header, -and a routed main area. Four tabs map to four page components. +and a routed main area. Six tabs map to six page components. ```mermaid graph TB @@ -134,13 +228,15 @@ graph TB APP --> HUB["HubAppLayout"] HUB --> SIDEBAR["HubSidebar"] - HUB --> HEADER["HubHeader"] + HUB --> HEADER["HubHeader
+ connection-status indicator"] HUB --> MAINC["HubMainContent"] MAINC --> CONNECT["ConnectPage
/connect"] MAINC --> BROWSE["BrowsePage
/browse"] MAINC --> REF["ReferencePage
/reference"] MAINC --> ANNOT["AnnotatePage
/annotate"] + MAINC --> VOL["VolumePage
/volume"] + MAINC --> TRAIN["TrainPage
/train"] ``` | Tab | Route | Page component | Key children | @@ -148,17 +244,26 @@ graph TB | **Connect** | `/connect` | `ConnectPage` | `IngestDropzone`, server/folder pickers | | **Browse** | `/browse` | `BrowsePage` | `ColumnBrowser`, `LocalSampleBrowser` | | **Reference** | `/reference` | `ReferencePage` | inline guide editors | -| **Annotate** | `/annotate` | `AnnotatePage` | `AnnotationCanvas`, `Toolbar`, `ClassManager` | +| **Annotate** | `/annotate` | `AnnotatePage` | `AnnotationCanvas`, `Toolbar`, `ClassManager`, stage tabs (Draw/Assist/Predict) | +| **3D** | `/volume` | `VolumePage` | `VolumeViewer` (vendored WebGPU renderer), `MaskLayersPanel`, `BuildVolumePanel`/`RebuildVolumeControl` | +| **Train** | `/train` | `TrainPage` | `TrainingDataPanel`, `HyperparamsPanel`, `RunsPanel`, `InferencePanel` | !!! note "Export is a modal, not a tab" Dataset export (COCO or DINOv3/Lightly) lives in `DownloadModal`, opened from the Annotate sidebar — there is no dedicated Export tab in the current navigation. +`HubAppLayout` also mounts `useConnectionHealth()` once at the app-shell level +(not per-page) — it periodically checks Tiled reachability +(`GET /api/tiled/list`, backoff on failure) and drives `connectionStore.status`, +which `HubHeader` renders as a small persistent indicator (green "Tiled +connected" / red "Tiled disconnected" linking back to Connect). + ### State management -State is split across nine **Zustand** stores. Only the annotation store carries -undo/redo history (via **zundo**), and a couple of stores persist to -`localStorage`. +State is split across a dozen **Zustand** stores. Only the annotation store +carries undo/redo history (via **zundo**), and a handful persist to +`localStorage` (display prefs are plain values inside `AnnotatePage`, not a +store — see `lib/displayPrefs.ts`). ```mermaid flowchart TD @@ -188,13 +293,16 @@ flowchart TD | --- | --- | --- | --- | | `annotationStore` | `byImage[sourceKey][slice] → Shape[]`, splits, negatives | draft autosave to backend | **zundo** | | `toolStore` | active tool, brush size, fill/threshold, selection | memory | — | -| `datasetStore` | active sample, `meta`, `currentSlice`, `renderOpts` | memory | — | +| `datasetStore` | active sample, `meta` (incl. `globalValueRange`), `currentSlice`, `renderOpts` | memory | — | | `classStore` | `AnnotationClass[]` (id, label, color, visibility) | in draft/save payloads | — | -| `connectionStore` | tiled/local URIs, paths, sample count | memory | — | +| `connectionStore` | tiled/local URIs, paths, sample count, **live Tiled `status`** | memory | — | | `referenceGuideStore` | guide entries, notes, `loadedFor` | backend via `useGuideSync` | — | | `clipboardStore` | copied shapes | memory | — | | `settingsStore` | `annotatorName`, `colorblindMode`, anonymous `sessionId` | `localStorage` | — | | `ratingStore` | per-sample star ratings | `localStorage` | — | +| `ipredStore` | Annotate → Assist/Predict stage state: feature bank job, trained classifier, manifold sampling | memory | — | +| `layerVisibilityStore` | per-class layer show/hide (Annotate canvas overlay) | memory | — | +| `predictedRasterStore` | pointers (`{runId, classIds}`) to committed-but-not-yet-vectorized predicted slices — see [Lazy shape vectorization](#lazy-predicted-shape-vectorization) | memory | — | `toolStore` also carries the eraser/select scope (`eraseAllClasses`, `selectScope`), the `clipToOtherClasses` (default **on**) and `mergeOverlappingSameClass` toggles, @@ -202,6 +310,20 @@ and `panReturnTool` (so the brush/eraser cursor stays visible while hold-Space panning). `settingsStore.sessionId` is an anonymous per-install id stamped into the "Feedback" bug-report context. +### Lazy predicted-shape vectorization + +A full-volume iPred "Apply across volume" commit does **not** eagerly +vectorize every predicted slice into real, editable `Shape[]` — that produced +a 297MB draft from one commit (139,004 shapes, 98.8% machine-predicted) before +this was built. Instead, `handleCommitVolumeApply` writes a lightweight +pointer (`predictedRasterStore`: `{runId, classIds}`, tens of bytes) for any +slice that doesn't already have real shapes; the existing Predictions raster +overlay (PNG-driven, not `Shape[]`-driven) keeps showing it from the run's own +`commit.png`. A slice is vectorized into real `Shape[]` — and only that slice — +the moment the user actually interacts with it ("Make this slice editable"). +Export and "Push to Tiled" read straight from the pointer for any slice never +made editable, falling back to real shapes for ones that were. + ### Canvas rendering `AnnotationCanvas` stacks several **react-konva** layers. The in-progress brush @@ -255,6 +377,76 @@ round-trip for editing: | `clahe.ts` / `sharpen.ts` / `stretch.ts` / `colormaps.ts` | Display-only preprocessors and LUTs | | `geometry.ts` / `measure.ts` / `datasetStats.ts` | Hit-testing/util, measurement, and the Insights QA metrics | | `sam/samClient.ts` · `sam/samWorker.ts` · `sam/adjust.ts` | SAM main-thread singleton, the Web Worker, and the display-bake used by tools | +| `displayPrefs.ts` | Persists the cosmetic display sliders (not `denoise`) to `localStorage`, global viewer preference | +| `bandTransferFunction.ts` | Converts a Sampler-fitted 2D intensity band into the 3D viewer's opacity-curve domain — see [3D volume view](#3d-volume-view) | +| `zarrUrl.ts` | Builds the Zarr URL the 3D viewer streams directly from Tiled, and the `__masks[_deep]` mask URL | +| `volumeMaskPreview.ts` | Rasterizes the CURRENT (possibly unsynced) annotation shapes into a coarse class-id volume — the Fast mask slot's zero-network "Live" mode | + +## 3D volume view + +The `/volume` tab renders straight off Tiled, with no export step: the vendored +WebGPU renderer (`frontend/vendor/view_tomography_recon_app`, a pinned git +submodule) streams a Zarr store's chunks directly from Tiled's `/zarr/v2` +router. `VolumeViewer.tsx` owns only the React mount/dispose lifecycle around +its imperative `run(canvas, options) → WebGpuViewerInstance` API; everything +about *how* the volume looks is the vendored renderer's own concern. + +### Mask layers + +Two independent, fixed mask/annotation slots — "Fast (iPred)" and "Deep +(dlsia)" — backed by the viewer's own slot-neutral `loadMask`/ +`loadMaskFromArray`/`setMaskClassColor`/etc. `MaskLayersPanel.tsx` is where the +*meaning* of each slot lives (kept out of the vendored viewer entirely): + +- **Deep** is always Tiled-backed, pointed at the `__masks_deep` + container `tiled_mask_sync.write_masks_to_tiled` writes (`container_suffix` + keeps it independent of the Fast slot's own container). +- **Fast** defaults to a zero-network "Live" mode: `volumeMaskPreview.ts` + rasterizes the sample's CURRENT shapes client-side into a coarse class-id + array, loaded via `loadMaskFromArray` — no "Push masks to Tiled" step + required first. It can also point at the Tiled-backed `__masks` + container (the precise backend-rasterized result) via its own toggle. + +Because `loadMask`/`loadMaskFromArray` are fire-and-forget on the viewer's +public interface (no promise, no error signal — `getMaskClasses(slot)` +returning `undefined` covers both "still loading" and "failed" +indistinguishably), `MaskLayersPanel.tsx` polls for up to 5 minutes rather than +assuming synchronous completion — a real Tiled-backed "Deep" mask (hundreds of +slices, freshly written) can legitimately take a while over the network, and a +short timeout only makes the *panel* look broken without stopping the +underlying load. + +### Transfer function: the threshold-fit → 3D bridge + +The Annotate tab's Sampler/Threshold-lasso tool fits an intensity band from a +traced example (`lib/thresholdFit.ts`) — "View band in 3D" sends that band +(native 0–255 canvas-byte space) to `/volume?bandLo=&bandHi=`, which isolates +it in the viewer's opacity curve (`setRendering({ opacityPoints })`): opaque +inside the band, transparent outside, with **zero upstream viewer changes** — +purely driving an already-existing transfer-function API. This isolates by +*intensity value* across the whole volume, not the traced *spatial region* +specifically — good enough when the traced feature's density is genuinely +distinct from its surroundings (the same property the 2D tool already +exploits), not true region-clipping. + +The one real subtlety: the 2D canvas's 0–255 bytes are normalized against a +per-dataset **percentile** range (`images._sample_global_stats`, exposed to +the frontend as `ImageMeta.globalValueRange`), while the 3D viewer separately +normalizes raw voxels against its own **min/max**-based estimate +(`WebGpuViewerInstance.getValueRange()`) — different statistics from different +data, so a fitted band must be converted byte → raw physical value → the +viewer's own domain (`lib/bandTransferFunction.ts`'s +`mapByteBandToViewerDomain`), not simply divided by 255. Assuming the two +normalizations were the same was tried first and produced a visibly wrong +band (confirmed live) before this was understood. + +### Volume build & rebuild + +A TIFF-stack source has no pyramid to stream until one is built +(`BuildVolumePanel.tsx` → `/api/volume/build`, backed by `volume_build.py`). +`RebuildVolumeControl.tsx` re-triggers a build for a dataset that already has +one (e.g. after a fidelity setting changes) — `build_volume` safely replaces +any prior build for the same key. ## Backend architecture @@ -272,6 +464,9 @@ flowchart TB Annot["/api/annotations/* · /api/guide* · /api/measure"] Export["/api/export/* · /api/masks/* · /api/import/*"] Ingest["/api/ingest/*"] + Volume["/api/volume/*"] + Denoise["/api/denoise/*"] + Train["/api/train/* (proxies to ipred for feature/classifier work)"] end subgraph Modules["Helper modules"] @@ -283,11 +478,15 @@ flowchart TB CE["coco_export · coco_import"] TS["tiled_annotation_sync
tiled_mask_sync"] IG["ingest"] + VB["volume_build · mask_pyramid
tiff_stack_source · zarr_source"] + ML["train_jobs · infer_jobs
denoise_bake · denoise_train
train_common · tiling · batch_probe"] + IC["ipred_client · ipred_routes
(thin proxy)"] end subgraph Ext["External"] Tiled["Tiled :8010"] Disk[("LOCAL_DATA_ROOT")] + Ipred["ipred :8003"] end Browse --> BH --> Tiled @@ -301,6 +500,10 @@ flowchart TB Export --> CE --> Disk Export --> TS --> Tiled Ingest --> IG --> Tiled + Volume --> VB --> Tiled + Denoise --> ML + Train --> ML --> Tiled + Train --> IC --> Ipred ``` ### Module responsibilities @@ -311,11 +514,15 @@ flowchart TB | `tiled_config.py` / `tiled_clients.py` | Server config, cached clients, `api_key_for_uri`, browse-root resolution | | `browse_helpers.py` | Metadata facets and filtered search (`tiled.queries.Key`, `distinct()`) | | `arrays.py` / `local_fs.py` | Resolve and slice arrays from Tiled or the local filesystem | -| `images.py` / `thumbnails.py` | Slice → PNG rendering (normalize, colormap, scale) | +| `images.py` / `thumbnails.py` | Slice → PNG rendering (normalize, colormap, scale) — `images._sample_global_stats` is also what `ImageMeta.globalValueRange` exposes to the frontend | | `drafts.py` / `guides.py` | Autosave drafts, immutable version history, annotation guides | | `coco_export.py` / `coco_import.py` | Shape rasterization, COCO build/write, dataset import | -| `export_jobs.py` / `ingest.py` | In-memory background-job registries | -| `tiled_annotation_sync.py` / `tiled_mask_sync.py` | Write `studio_*` metadata and rasterized mask volumes back to Tiled | +| `export_jobs.py` / `ingest.py` | In-memory background-job registries (shared by export, dlsia train/infer, denoise, ingest — all long-running work polls the same `/status/{id}` shape) | +| `tiled_annotation_sync.py` / `tiled_mask_sync.py` | Write `studio_*` metadata and rasterized mask volumes back to Tiled — `container_suffix` keeps independent mask producers (manual sync vs. dlsia's own write) from colliding | +| `volume_build.py` / `mask_pyramid.py` / `tiff_stack_source.py` / `zarr_source.py` | Build/register the Zarr pyramid the 3D viewer streams, and the mask pyramids it loads into its two mask slots | +| `train_common.py` / `train_jobs.py` / `infer_jobs.py` / `tiling.py` / `batch_probe.py` | dlsia TUNet training/inference: device selection, run persistence, qlty patch tiling, GPU batch-size probing. Two-lock design (`ML_LOCK` for cross-job exclusivity, `GPU_FORWARD_LOCK` scoped to just the GPU forward call) lets I/O/preprocessing/vectorization run concurrently across a per-slice worker pool during inference while still serializing actual GPU calls | +| `denoise.py` / `denoise_bake.py` / `denoise_train.py` / `denoise_runtime.py` / `dlsia_runtime.py` / `autoencoder_runtime.py` | Classical + learned (Noise2Noise/Noise2Void/autoencoder) denoising, and baking a denoiser's output permanently into a new Tiled node | +| `ipred_client.py` / `ipred_routes.py` | Thin `httpx` proxy from the backend's `/api/train/*`/iPred routes to the separate ipred service — the backend never runs feature-extraction/classifier training itself | ### Data model @@ -464,6 +671,60 @@ sequenceDiagram API-->>User: dataset.zip ``` +### Train a dlsia model, run inference, write masks to Tiled + +The Train tab's whole loop stays in the backend's `train_jobs.py`/`infer_jobs.py` +— the ipred service is only involved for the Annotate tab's own fast pixel +classifier and feature banks, not dlsia. `run_infer_job` runs a per-slice +worker pool (I/O/preprocessing/vectorization concurrent, GPU forward calls +serialized via `GPU_FORWARD_LOCK`) and publishes partial results incrementally +so the frontend's preview slider extends live while the job is still running, +not just after it finishes. + +```mermaid +sequenceDiagram + actor User + participant TP as TrainPage + participant Job as useExportJob + participant API as Backend + participant TJ as train_jobs / infer_jobs + participant TMS as tiled_mask_sync + participant T as Tiled + + User->>TP: Start training + TP->>Job: start(hyperparams) + Job->>API: POST /api/train/start + API->>TJ: run_train_job (background thread, ML_LOCK) + loop poll + Job->>API: GET /api/export/status/{job_id} + API-->>Job: epoch progress, mIoU + end + + User->>TP: Run inference (all slices) + TP->>API: POST /api/train/infer + API->>TJ: run_infer_job (worker pool, GPU_FORWARD_LOCK) + loop poll (result grows incrementally) + API-->>TP: preview_slices, done/total + end + + User->>TP: Write masks to Tiled + TP->>API: POST /api/train/infer/write-tiled/{job_id} + API->>TJ: run_write_tiled_job (reads its OWN cached job entry —
not the frontend's live state, so a stale "Tiled?" prop can't hide/misfire this) + TJ->>TMS: write_masks_to_tiled(container_suffix="_deep") + TMS->>T: register_mask_pyramid + class_vols +``` + +!!! note "Why the write-tiled route trusts itself, not the frontend" + `run_write_tiled_job` derives `source`/`kind` from its own cached job entry, + never from the request. This was a deliberate fix (#31 in the project's + working notes): a frontend prop computed from the *live* dataset store can + go stale if the user navigates away during a long job and back, which once + silently hid the "Write masks to Tiled" button for a job that had, in + fact, run against a real Tiled source. Trusting the job's own record — and + letting the backend's already-correct rejection of a genuinely non-Tiled + source surface as a normal job error — is more robust than re-deriving the + same fact twice. + ## Key design decisions The choices below explain *why* the code looks the way it does — useful before @@ -495,9 +756,21 @@ extending it. (`/api/annotations/draft`, disk only); an explicit **save** writes an immutable version + thumbnail and syncs summary metadata to Tiled. Everything is keyed by the canonical **source key** so drafts, versions, guide, and exports line up per sample. -- **Long work is a background job + poll.** Export, mask write-back, and ingest use an - in-memory job registry with a `/status/{id}` poller and (for export) a streamed - `.zip`, rather than blocking the request. +- **Long work is a background job + poll.** Export, mask write-back, dlsia train/infer, + denoising, and ingest all use the same in-memory job registry (`export_jobs.py`) + with a `/status/{id}` poller, rather than blocking the request. +- **Fill flood-time barriers reuse the same rasterizer as post-commit clipping.** + The Fill tool's flood (`magicwand.ts`) can now stop at a pixel already + claimed by a different class (an optional `blocked` grid, OR'd into the + existing gradient-edge wall check) — built via `rasterizeUnion`, the exact + helper the brush-clip-detection path already used, so no second + rasterization concept was introduced for the same problem. +- **Every mask sync/pyramid write is keyed by its own real cache path, not a + shared literal.** `mask_pyramid.register_mask_pyramid` takes a `cache_key` + (the actual on-disk Zarr path) separate from `key` (the Tiled sub-node + name, always `"semantic"`) — a real production incident (two independent + mask producers silently overwriting each other's pixels) came from a + version that conflated the two. ## Security boundaries diff --git a/docs/reference/deployment.md b/docs/reference/deployment.md new file mode 100644 index 0000000..0d47c43 --- /dev/null +++ b/docs/reference/deployment.md @@ -0,0 +1,185 @@ +# Production deployment + +This page is for standing up a shared, long-lived deployment (e.g. behind +`hub.als.lbl.gov`), as opposed to [Installation](../getting-started/installation.md)'s +"run it yourself" Docker options. It covers the published container images, hosting +under a URL prefix, and what's still an open question before a real production rollout. + +--- + +## Published images + +Every merge to `main` builds and publishes images to GitHub Container Registry +(`ghcr.io`) via its own dedicated workflow, `.github/workflows/publish-image.yml` — +this never fires on a pull request, only on `main` (PRs get a build-only check +instead, in `ci.yml`'s `docker-build` job — same Dockerfile targets, no registry +push). The structure mirrors +[als-computing/view_tomography_recon_app](https://github.com/als-computing/view_tomography_recon_app)'s +own `publish-image.yml` — same tag names, same one-image-many-tags shape, so anyone +already familiar with that repo's deployment finds this one working the same way. + +**One image name, three tags** — the Dockerfile target differs per tag (invisible +from the registry, but real underneath): + +| Tag | Dockerfile target | Tiled | iPred / Train | Base path baked in | Compose file | +| --- | --- | --- | --- | --- | --- | +| `:local` | `app-full` | Bundled | Bundled | `/seg_studio/` | `docker-compose.local.yml` | +| `:als-prod` | `app-ml` | External (ALS's production Tiled) | Bundled | `/bl832/seg_studio/` | `docker-compose.als-prod.yml` | +| `:als-staging` | `app-ml` | External (ALS's staging Tiled) | Bundled | `/bl832/seg_studio/` | `docker-compose.als-staging.yml` | + +Pull with `ghcr.io//:local` / `:als-prod` / `:als-staging`. + +`:als-prod` and `:als-staging` build identically today — this app reads `TILED_URI` +at container-*run* time (an ordinary env var), not at build time, so the +prod-vs-staging Tiled distinction lives entirely in which compose file's default you +use, not in the image itself. They're still published as two separate, explicitly +named steps in `publish-image.yml` (not one build tagged twice) so they can diverge +later — e.g. if a build-time-only setting ever needs to differ between the two — +without restructuring the workflow. + +`:local` is fully self-contained (Tiled + backend + iPred, nothing external +required) for anyone standing this up for themselves who still wants to exercise the +same subpath-hosting path `:als-prod`/`:als-staging` use, rather than leaving that +path tested only at ALS — see [Installation](../getting-started/installation.md) for +running it locally. `:als-prod`/`:als-staging` bundle iPred/Train (so they work +without a second, separately-maintained service) but keep Tiled external, since ALS +already runs its own Tiled and a second, empty in-container one would be wrong. + +!!! note "No `:latest`/`-sha` tags — a deliberate simplification" + Earlier in this repo's history, every image also published generic `:latest` + + immutable `:sha-` tags alongside the purpose-named ones. That's been + dropped in favor of exactly the three tags above, matching the reference repo's + own simpler scheme — one clearly-named tag per real deployment target, nothing + else. The trade-off: there's no immutable per-commit tag to roll back to + anymore; rolling back today means re-running `publish-image.yml` against an + older commit (`git push -f` a maintenance branch to `main`, or a manual + workflow dispatch pinned to a specific ref) rather than just pointing at a + `-sha` tag that's already sitting in the registry. + +--- + +## Hosting under a URL prefix + +A shared deployment is typically reached at a path prefix off a shared hub domain +(e.g. `hub.als.lbl.gov/bl832/seg_studio/`) rather than its own subdomain, so it stays +manageable from an ops/ticketing perspective. This requires two things working together: + +1. **A reverse proxy in front of the container that strips the prefix** before + forwarding — the app container itself has no idea it's hosted under a subpath; it + only ever serves and expects plain root-relative paths (`/api/...`, `/assets/...`). + A request to `hub.als.lbl.gov/bl832/seg_studio/api/...` must reach the container as + plain `/api/...`. This is standard nginx/Traefik path-based routing — the backend + itself needs **no code changes** for this. +2. **The frontend build baked for that prefix**, via the `VITE_BASE_PATH` Docker + build-arg (e.g. `--build-arg VITE_BASE_PATH=/bl832/seg_studio/`). Vite bakes this + into the compiled JS at build time (`base` in `vite.config.ts`, consumed by + `main.tsx`'s router `basename` and `config.ts`'s `API_BASE`) — **root-hosted and + subpath-hosted are two different build artifacts from the same source**, not one + image that adapts at runtime. The published `:als-prod`/`:als-staging`/`:local` + tags already have this baked in — building your own image (e.g. `docker- + compose.yml`/`docker-compose.ml.yml`/`docker-compose.full.yml`'s generic, + unpublished targets) with `VITE_BASE_PATH` left unset stays root-hosted. + +If you're standing up your own subpath deployment (a different institution, a +different path), build your own image with the matching `VITE_BASE_PATH` rather than +reusing the `:als-prod`/`:als-staging`/`:local` tags, which are baked for this repo's specific paths. + +!!! warning "`VITE_BASE_PATH` is a path, never a hostname" + Set it to the path segment only (`/bl832/seg_studio/`) — **never** a full URL + with a scheme or host (`https://hub.als.lbl.gov/bl832/seg_studio/` would be + wrong). The image has no business knowing what domain fronts it; only the + reverse proxy does, and it can change (or differ between environments) + without ever touching this build-arg. This is what makes the *same* `:als-prod`/`:als-staging` + image usable unmodified across staging and production — e.g. + `hub-staging.als.lbl.gov/bl832/seg_studio/` and + `hub.als.lbl.gov/bl832/seg_studio/` sharing the identical `/bl832/seg_studio/` + path, routed to whichever environment's container by the proxy's own + hostname-based rule, not by anything baked into the image. Baking a real + FQDN in would force a separate image build per environment for no reason, + and would silently break if that hostname ever changed. + +!!! tip "Testing this locally before it matters" + `docker-compose.local.yml` (see [Installation](../getting-started/installation.md)) + includes a small nginx service that reproduces exactly this proxy-stripping + behavior, so `/seg_studio/` hosting can be verified on a laptop before it's ever + tried against a real hub deployment. Serving the built `dist/` directly through a + plain static file server does **not** reproduce this — the stripping proxy is + load-bearing, not optional. + +!!! warning "A real reverse proxy needs a body-size limit raised too" + Discovered via the local test proxy: nginx's own default `client_max_body_size` + (1MB) rejects the Ingest tab's multi-file uploads (a folder of TIFFs easily + reaches hundreds of MB) with a raw `413` *before the request ever reaches this + app* — this app's own upload handling is never even consulted. `docker/nginx- + local.conf` sets `client_max_body_size 0;` (unlimited; the backend's own + job-based, streamed-to-disk ingest route is the real, size-aware boundary) — + **`hub.als.lbl.gov`'s real reverse proxy needs the equivalent setting**, or + every real-world upload through this app will hit the same wall in production. + +--- + +## Environment variables specific to a shared deployment + +Beyond what's already covered in [Installation](../getting-started/installation.md#environment-variables): + +| Variable | Purpose | +| --- | --- | +| `TILED_BROWSE_PATH` | The real path into an existing institutional Tiled catalog Browse should treat as its root (e.g. `beamlines/bl832/processed`). Confirm the exact value with whoever operates that Tiled server — don't assume it matches another deployment's beamline. | +| `TILED_URI` / `TILED_API_KEY` | Point at the shared production Tiled server. See the open authentication question below — a single shared key is a stopgap, not the final design. | +| `LOCAL_SOURCE_DIR` | `docker-compose.full.yml`/`docker-compose.local.yml` (bundled-Tiled shapes) only — bind-mounts a real host directory to `/data/processed` so the bundled Tiled server can read it directly. Everything under it (Zarr stores and plain image-slice folders) auto-registers into Tiled on every container start (see [Installation](../getting-started/installation.md#option-2-docker)) — no manual ingest step needed for pre-existing data. Not applicable to `:als-prod`/`:als-staging` (`app-ml`), which point at an already-existing external Tiled instead of a bundled one. | + +Persistent storage (`LOCAL_DATA_ROOT`, defaulting to the `/data` volume already +declared in the image) and CORS (`BROWSE_ALLOWED_ORIGINS`) need no special +production-specific handling beyond what's already documented for local Docker use — +mount a real volume, and leave CORS empty since the SPA and API are served same-origin. + +### Setting these for `docker compose` + +Copy `.env.example` (repo root) to `.env` and fill it in — `docker compose` loads a +`.env` file in the same directory automatically, for every compose file: + +```bash +cp .env.example .env +# edit .env: set TILED_API_KEY at minimum — TILED_URI/TILED_BROWSE_PATH/ +# VITE_BASE_PATH already default correctly per environment (see below) +docker compose -f docker-compose.als-prod.yml up --build +``` + +For ALS specifically, use `docker-compose.als-prod.yml`/`docker-compose.als- +staging.yml` — these already default `TILED_URI` to the real ALS production/staging +Tiled hostnames and `VITE_BASE_PATH` to `/bl832/seg_studio/`, so only +`TILED_API_KEY` needs setting per deployment. `docker-compose.ml.yml` is the +generic, no-institution-specific-defaults version of the same shape (any other +Tiled server, any base path) — set `TILED_URI`/`TILED_BROWSE_PATH`/`VITE_BASE_PATH` +yourself there. + +`.env` is git-ignored — never commit a real `TILED_API_KEY` into it. This is a +different file from `backend/.env.example`/`frontend/.env.example`, which configure +`start_all.sh`/plain `npm run dev`/`npm run build` outside a container — the root +`.env` is specifically what `docker compose` itself substitutes into the compose +files' `${VAR}` references. If you're running the already-published `ghcr.io` image +directly (`docker run`/your own orchestration) rather than through one of these +compose files, pass the equivalent `-e`/env vars there instead — `.env` only affects +`docker compose` invocations in this directory. + +--- + +## Open questions before a real production rollout + +These are real, unresolved gaps — not yet implemented, flagged here rather than +glossed over: + +- **Per-user Tiled authentication.** Tiled's authorization is per-proposal + (ESAF) read/write tags, not a blanket grant — a single shared `TILED_API_KEY` + either over- or under-privileges every user. The real design needs per-user + ORCID/OIDC login (for both the browser's direct-to-Tiled 3D viewer calls and the + backend's own proxied calls), which hasn't been built yet. Until then, a shared + `TILED_API_KEY` works but does not correctly represent per-proposal permissions — + fine for evaluation, not for a real multi-user production rollout with real data. +- **Hub-level SSO.** Whether `hub.als.lbl.gov`'s own reverse proxy fully gates access + before a request reaches this app, or whether the app is expected to participate in + auth itself, is a separate open question from Tiled's own data-access control above + — confirm with hub admins before relying on either assumption. +- **GPU availability for `:als-prod`/`:als-staging`.** Bundling iPred/ML there only delivers a + practically usable Train tab if the host running that container provides GPU + passthrough — CPU-only training/inference works but is dramatically slower. diff --git a/frontend/.env.example b/frontend/.env.example index d3c4d89..5d53a81 100644 --- a/frontend/.env.example +++ b/frontend/.env.example @@ -4,6 +4,18 @@ # Backend API base. Leave empty for same-origin (dev uses the /api proxy). # VITE_API_BASE= +# Subpath the app is hosted under, behind a reverse proxy that strips the +# prefix before forwarding (e.g. hub.als.lbl.gov/bl832/seg_studio/ -> container +# sees plain /). Read at BUILD time (Vite bakes it into the compiled asset +# URLs/router basename) — leave unset for root-hosted (local dev, the lean +# :local/:latest images). Must end in a trailing slash. +# +# PATH ONLY — never include the scheme/hostname (not +# "https://hub.als.lbl.gov/bl832/seg_studio/"). The reverse proxy owns the +# hostname and can differ per environment (staging vs. production) without +# ever touching this value or requiring a rebuild. +# VITE_BASE_PATH=/bl832/seg_studio/ + # Documentation site URL (the sidebar "Docs" button). # VITE_DOCS_URL=http://127.0.0.1:8000 diff --git a/frontend/index.html b/frontend/index.html index ef45032..7649d8d 100644 --- a/frontend/index.html +++ b/frontend/index.html @@ -2,7 +2,7 @@ - + =14.0.0" - } - }, "node_modules/@rolldown/pluginutils": { "version": "1.0.0-beta.27", "resolved": "https://registry.npmjs.org/@rolldown/pluginutils/-/pluginutils-1.0.0-beta.27.tgz", @@ -2857,9 +2848,9 @@ } }, "node_modules/@storybook/core/node_modules/semver": { - "version": "7.8.4", - "resolved": "https://registry.npmjs.org/semver/-/semver-7.8.4.tgz", - "integrity": "sha512-rUCObTnP32Q08R2uuIrt7r9PlEonuTmtuXYcW6s5kjdlj3xbnwe+21yXptAUYcMAABLkYYTtnmzb3w3EDZfueA==", + "version": "7.8.5", + "resolved": "https://registry.npmjs.org/semver/-/semver-7.8.5.tgz", + "integrity": "sha512-Y7/KDsb8LjooZpwaqGyulO6DQlksgCncchHGk+sZIY4SBvUocMBEFH5Ur1fI4dV+Jvl0w6cjvucaIi40puRioA==", "license": "ISC", "peer": true, "bin": { @@ -3312,9 +3303,9 @@ } }, "node_modules/@types/express-serve-static-core": { - "version": "4.19.8", - "resolved": "https://registry.npmjs.org/@types/express-serve-static-core/-/express-serve-static-core-4.19.8.tgz", - "integrity": "sha512-02S5fmqeoKzVZCHPZid4b8JH2eM5HzQLZWN2FohQEy/0eXTq8VXZfSN6Pcr3F6N9R/vNrj7cpgbhjie6m/1tCA==", + "version": "4.19.9", + "resolved": "https://registry.npmjs.org/@types/express-serve-static-core/-/express-serve-static-core-4.19.9.tgz", + "integrity": "sha512-QP2ESEe/ImWY0HDwNAnK9PvEffUyhLTnWkk7KXzHfyeWAnlrDe1fN77bXl6ia8KT3wPlmA7t9/VPRpnf4Ex9sg==", "license": "MIT", "dependencies": { "@types/node": "*", @@ -3336,9 +3327,9 @@ "license": "MIT" }, "node_modules/@types/lodash": { - "version": "4.17.24", - "resolved": "https://registry.npmjs.org/@types/lodash/-/lodash-4.17.24.tgz", - "integrity": "sha512-gIW7lQLZbue7lRSWEFql49QJJWThrTFFeIMJdp3eH4tKoxm1OvEPg02rm4wCCSHS0cL3/Fizimb35b7k8atwsQ==", + "version": "4.17.25", + "resolved": "https://registry.npmjs.org/@types/lodash/-/lodash-4.17.25.tgz", + "integrity": "sha512-+K1NIO8I+F9/wNulfVvu23QYd0Pe9/OCqRrim4NoYIf1VoEDL90Ve4ClzpyqBLc7NpGGWRvYNCKZ1BE/Jpf8dQ==", "license": "MIT" }, "node_modules/@types/mime": { @@ -3840,15 +3831,15 @@ } }, "node_modules/@vitest/expect": { - "version": "3.2.6", - "resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-3.2.6.tgz", - "integrity": "sha512-1+7q9BtaKzEmO+fmNT3kYvoNn5Y71XWAx2Q5HRim4tTVRQVRv4uJFAQ5FbK0OPUeNP/WmVCpxYxoJdvuHVjzBQ==", + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/expect/-/expect-3.2.7.tgz", + "integrity": "sha512-E8eBXaKibuvH2pSZErOjdVb5vF4PbKYcrnluBTYxEk1l/VhhwZg1kZQsdtjq+CsF5CFydf2Rdkz7jDHKSisi3w==", "dev": true, "license": "MIT", "dependencies": { "@types/chai": "^5.2.2", - "@vitest/spy": "3.2.6", - "@vitest/utils": "3.2.6", + "@vitest/spy": "3.2.7", + "@vitest/utils": "3.2.7", "chai": "^5.2.0", "tinyrainbow": "^2.0.0" }, @@ -3857,13 +3848,13 @@ } }, "node_modules/@vitest/mocker": { - "version": "3.2.6", - "resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-3.2.6.tgz", - "integrity": "sha512-EZOrpDbkKotFAP7wPAQV1UIyoGOk4oX7ynWhBhLB7v+meMHbQhU16oPpIYGTTe4oFlhpryGpgpcZP/sin3hYuw==", + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/mocker/-/mocker-3.2.7.tgz", + "integrity": "sha512-Trr0hYO9CM3Wj6ksWHRhK9IZpIY6wTMO5u/MqXurMxT57sWBaOPEtP3Oq60ihZuh5JsiagKfz95OcxdEP6dBrA==", "dev": true, "license": "MIT", "dependencies": { - "@vitest/spy": "3.2.6", + "@vitest/spy": "3.2.7", "estree-walker": "^3.0.3", "magic-string": "^0.30.17" }, @@ -3884,9 +3875,9 @@ } }, "node_modules/@vitest/pretty-format": { - "version": "3.2.6", - "resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-3.2.6.tgz", - "integrity": "sha512-lb7XXXzmm2h2ASzFnRvQpDo6onT1NmMJA3tkGTWiBFtRJ9lxGY3d3mm/Apt36gej2bkkOVLL/yTOtufDaFa/jA==", + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/pretty-format/-/pretty-format-3.2.7.tgz", + "integrity": "sha512-KUHlwqVu0sRlhCdyPdQ/wBoTfRahjUky1MubOmYw9fWfIZy1gNoHpuaaQBPAaMaVYdQYHJLurzj8ECCj5OwTqA==", "dev": true, "license": "MIT", "dependencies": { @@ -3897,13 +3888,13 @@ } }, "node_modules/@vitest/runner": { - "version": "3.2.6", - "resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-3.2.6.tgz", - "integrity": "sha512-HYcoSj1w5tcgUnzoF0HcyaAQjpA1gj9ftUJ7iSJSuipc02jW9gKkigwZbjFldAfYHA1fa8UZVRftdMY5msWM9Q==", + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/runner/-/runner-3.2.7.tgz", + "integrity": "sha512-sB9y4ovltoQP+WaUPwmSxO9WIg9Ig694Di5PalVPsYHklAdE027mehpWF2SQSVq+k6sFgaivbTjTJwZLSHbedA==", "dev": true, "license": "MIT", "dependencies": { - "@vitest/utils": "3.2.6", + "@vitest/utils": "3.2.7", "pathe": "^2.0.3", "strip-literal": "^3.0.0" }, @@ -3912,13 +3903,13 @@ } }, "node_modules/@vitest/snapshot": { - "version": "3.2.6", - "resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-3.2.6.tgz", - "integrity": "sha512-H+ZjNTWGpObenh0YnlBctAPnJSI20P81PL8BPzWpx54YXLLTm8hEsWawtcYLMrwvpK48hGxLLbCS+1KRXhsKhw==", + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/snapshot/-/snapshot-3.2.7.tgz", + "integrity": "sha512-7C+MwShwtBSI5Buwoyg3s/iY1eHL9PKAf+O1wVh/TdnjXUtkoL/9YQtre90i4MtNXM6edP1wJ2zOBpfCyhIS7g==", "dev": true, "license": "MIT", "dependencies": { - "@vitest/pretty-format": "3.2.6", + "@vitest/pretty-format": "3.2.7", "magic-string": "^0.30.17", "pathe": "^2.0.3" }, @@ -3927,9 +3918,9 @@ } }, "node_modules/@vitest/spy": { - "version": "3.2.6", - "resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-3.2.6.tgz", - "integrity": "sha512-oq6BbH68WzcWmwtBrU9nqLeaXTR4XwJF7FSLkKEZo4i6eoXcrxjcwSuTvWBIRUTC6VC72nXYunzqgZA+IKdtxg==", + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/spy/-/spy-3.2.7.tgz", + "integrity": "sha512-Q2eQGI6d2L/hBtZ0qNuKcAGid68XK6cv1xsoaIma6PaJhHPoqcEJhYpXZ/5myCMqkNgtP6UKuBhbc0nHKnrkuQ==", "dev": true, "license": "MIT", "dependencies": { @@ -3940,13 +3931,13 @@ } }, "node_modules/@vitest/utils": { - "version": "3.2.6", - "resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-3.2.6.tgz", - "integrity": "sha512-lI23nIs4bnT3T8NIoh+vFaz5s2/DdP0Jgt2jxwgWljvwn82cLJtyi/If+fjFyoLMGIOz0U/fKvWE0d4jsNQEfg==", + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/@vitest/utils/-/utils-3.2.7.tgz", + "integrity": "sha512-x6BDOd7dyo3PFLY3I9/HJ25X/6OurhGXk2/B9gOZNPF7XDVjeBK4k01lQE5uvDpbuheErh91qYuE1E2OEjK3Rw==", "dev": true, "license": "MIT", "dependencies": { - "@vitest/pretty-format": "3.2.6", + "@vitest/pretty-format": "3.2.7", "loupe": "^3.1.4", "tinyrainbow": "^2.0.0" }, @@ -3954,10 +3945,17 @@ "url": "https://opencollective.com/vitest" } }, + "node_modules/@webgpu/types": { + "version": "0.1.72", + "resolved": "https://registry.npmjs.org/@webgpu/types/-/types-0.1.72.tgz", + "integrity": "sha512-0cF7RFM2edNoiIS1ODJp0/Gzv4/xSXhwoR0YCza+OWpJWtn4wmo9DvK91aLlH9+uUnwIriP7ZiC3WitmyhuzBw==", + "dev": true, + "license": "BSD-3-Clause" + }, "node_modules/adm-zip": { - "version": "0.5.17", - "resolved": "https://registry.npmjs.org/adm-zip/-/adm-zip-0.5.17.tgz", - "integrity": "sha512-+Ut8d9LLqwEvHHJl1+PIHqoyDxFgVN847JTVM3Izi3xHDWPE4UtzzXysMZQs64DMcrJfBeS/uoEP4AD3HQHnQQ==", + "version": "0.5.18", + "resolved": "https://registry.npmjs.org/adm-zip/-/adm-zip-0.5.18.tgz", + "integrity": "sha512-ufJnssQGbxzLNS1Ho9bCtX4rQKCCvoVuDLHoJyc3F9dOGDB4BkWs2Ci0kv53lqocAEQ/Cbi+I2XCsNYGqVYqng==", "license": "MIT", "engines": { "node": ">=12.0" @@ -4056,9 +4054,9 @@ } }, "node_modules/ast-types": { - "version": "0.16.1", - "resolved": "https://registry.npmjs.org/ast-types/-/ast-types-0.16.1.tgz", - "integrity": "sha512-6t10qk83GOG8p0vKmaCr8eiilZwO171AvbROMtvvNiwrTly62t+7XkA8RdIIVbpMhCASAsxgAzdRSwh6nw/5Dg==", + "version": "0.16.3", + "resolved": "https://registry.npmjs.org/ast-types/-/ast-types-0.16.3.tgz", + "integrity": "sha512-FvWoWYfSCM6kRxCSH+MGLHIKKGRL6A6AW7Zek2O32REPQRdg131428uRTKMBYAeRd3XXAaHDS60Wpri7CdKDrA==", "license": "MIT", "peer": true, "dependencies": { @@ -4128,9 +4126,9 @@ "license": "MIT" }, "node_modules/baseline-browser-mapping": { - "version": "2.10.36", - "resolved": "https://registry.npmjs.org/baseline-browser-mapping/-/baseline-browser-mapping-2.10.36.tgz", - "integrity": "sha512-lVq/Df7LXlO79MVaaUHztSwWiG9oXoWHlgvNS51v8Dpd4+G4/VIy6qYePTw31nAVls33nUtnfezYeLkYAak9dg==", + "version": "2.11.22", + "resolved": "https://registry.npmjs.org/baseline-browser-mapping/-/baseline-browser-mapping-2.11.22.tgz", + "integrity": "sha512-pWc4w51fBFd7mav43/zKRC+RI6f4yfzQoVlfvE8dECePyfkn1bzLp01Fj0QACcyCZyFhiEMyD2qScfKRWgWibA==", "dev": true, "license": "Apache-2.0", "bin": { @@ -4191,9 +4189,9 @@ "peer": true }, "node_modules/browserslist": { - "version": "4.28.2", - "resolved": "https://registry.npmjs.org/browserslist/-/browserslist-4.28.2.tgz", - "integrity": "sha512-48xSriZYYg+8qXna9kwqjIVzuQxi+KYWp2+5nCYnYKPTr0LvD89Jqk2Or5ogxz0NUMfIjhh2lIUX/LyX9B4oIg==", + "version": "4.28.9", + "resolved": "https://registry.npmjs.org/browserslist/-/browserslist-4.28.9.tgz", + "integrity": "sha512-EWazOblFYUvlGZcfGhPUPmYh3nikUxBVb+y9MJun5f3hBi812X+8MSQTujLBtgK3cf51fJWbWfOjyeO954d+Eg==", "dev": true, "funding": [ { @@ -4211,11 +4209,11 @@ ], "license": "MIT", "dependencies": { - "baseline-browser-mapping": "^2.10.12", - "caniuse-lite": "^1.0.30001782", - "electron-to-chromium": "^1.5.328", - "node-releases": "^2.0.36", - "update-browserslist-db": "^1.2.3" + "baseline-browser-mapping": "^2.11.20", + "caniuse-lite": "^1.0.30001810", + "electron-to-chromium": "^1.5.420", + "node-releases": "^2.0.54", + "update-browserslist-db": "^1.3.2" }, "bin": { "browserslist": "cli.js" @@ -4292,9 +4290,9 @@ } }, "node_modules/caniuse-lite": { - "version": "1.0.30001799", - "resolved": "https://registry.npmjs.org/caniuse-lite/-/caniuse-lite-1.0.30001799.tgz", - "integrity": "sha512-hG1bReV+OUU+MOqK4t/ZWI0tZOyz3rqS9XuhOUz1cIcbwBKjOyJEJuw9ER5JuNyqxNk8u/JUVbGibBOL1yrjFw==", + "version": "1.0.30001810", + "resolved": "https://registry.npmjs.org/caniuse-lite/-/caniuse-lite-1.0.30001810.tgz", + "integrity": "sha512-TITQPUkaz+aVk5GL6NhOdwk1aEaNTSDPsGFWrTuhKGtjTF70jL/Oht2W4c6rXUe5fu7Ie19VIahAXHIIiWWNeg==", "dev": true, "funding": [ { @@ -4636,9 +4634,9 @@ } }, "node_modules/dayjs": { - "version": "1.11.21", - "resolved": "https://registry.npmjs.org/dayjs/-/dayjs-1.11.21.tgz", - "integrity": "sha512-98IT+HOahAisibz/yjKbzuOBwYcjJ7BCLPzARyHiyEBmRz4fatF+KPJszEHXsGYjUG234aH/cOjW1wwTbKUZlA==", + "version": "1.11.23", + "resolved": "https://registry.npmjs.org/dayjs/-/dayjs-1.11.23.tgz", + "integrity": "sha512-QDTCU0M0MxR3hQfnlDJfwekQiaanm1ubOD231u73WBckQ/fsamwRLiE2GBz6D3a/xF1NgfiDLJjXBa1hYOYTtQ==", "license": "MIT" }, "node_modules/debug": { @@ -4793,9 +4791,9 @@ } }, "node_modules/electron-to-chromium": { - "version": "1.5.371", - "resolved": "https://registry.npmjs.org/electron-to-chromium/-/electron-to-chromium-1.5.371.tgz", - "integrity": "sha512-e9htk9mAYL6AzmkEhSvVVw7IWGSBJ/Bqdn2eRyRLrj1g6sncN4WbFt5qnILYoCktktr45pyjIrOiRvBThQ808w==", + "version": "1.5.420", + "resolved": "https://registry.npmjs.org/electron-to-chromium/-/electron-to-chromium-1.5.420.tgz", + "integrity": "sha512-2yD6XreGusOfNV+dUcvipJEXc3n/n7fgr7996aszTG+YY5E4mqM4tOq/3uhP129cazL9YHbVWSpc79ePotWtPA==", "dev": true, "license": "ISC" }, @@ -4956,9 +4954,9 @@ } }, "node_modules/expect-type": { - "version": "1.3.0", - "resolved": "https://registry.npmjs.org/expect-type/-/expect-type-1.3.0.tgz", - "integrity": "sha512-knvyeauYhqjOYvQ66MznSMs83wmHrCycNEN6Ao+2AeYEfxUIkuiVxdEa1qlGEPK+We3n0THiDciYSsCcgW/DoA==", + "version": "1.4.0", + "resolved": "https://registry.npmjs.org/expect-type/-/expect-type-1.4.0.tgz", + "integrity": "sha512-KfYbmpRm0VbLjEvVa9yGwCi9GI34xvi7A/HXYWQO65CSD2u3MczUJSuwXKFIxlGsgBQizV9q5J9NHj4VG0n+pA==", "dev": true, "license": "Apache-2.0", "engines": { @@ -5889,9 +5887,9 @@ } }, "node_modules/nanoid": { - "version": "3.3.12", - "resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.12.tgz", - "integrity": "sha512-ZB9RH/39qpq5Vu6Y+NmUaFhQR6pp+M2Xt76XBnEwDaGcVAqhlvxrl3B2bKS5D3NH3QR76v3aSrKaF/Kiy7lEtQ==", + "version": "3.3.19", + "resolved": "https://registry.npmjs.org/nanoid/-/nanoid-3.3.19.tgz", + "integrity": "sha512-Y2tUNy4ouw6tq5oDSKeQYGOyhkUBhNOcGV/02KC+6kd9eDGqdZd++mjMiIDilrBYvjEnCYvVtsuHCuP+okSfug==", "funding": [ { "type": "github", @@ -5907,9 +5905,9 @@ } }, "node_modules/node-releases": { - "version": "2.0.47", - "resolved": "https://registry.npmjs.org/node-releases/-/node-releases-2.0.47.tgz", - "integrity": "sha512-Uzmd6LXpouKo8EUK68IjH4+E01w/hXyV3R3g/geCJo+rXLNfh1xucB+LOzYEOQPSiUK3h/xZf0cQGcSsmyL2Og==", + "version": "2.0.54", + "resolved": "https://registry.npmjs.org/node-releases/-/node-releases-2.0.54.tgz", + "integrity": "sha512-YHs7BmmcsdAI5Ozuf8JZo6PT0mv2GIWC9vMfvUC3dp65M8hn7Ux8CPL+2oBI7juNuj9d0ndhTcznq2ODBps9cQ==", "dev": true, "license": "MIT", "engines": { @@ -6130,9 +6128,9 @@ } }, "node_modules/postcss": { - "version": "8.5.15", - "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.15.tgz", - "integrity": "sha512-FfR8sjd4em2T6fb3I2MwAJU7HWVMr9zba+enmQeeWFfCbm+UOC/0X4DS8XtpUTMwWMGbjKYP7xjfNekzyGmB3A==", + "version": "8.5.28", + "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.28.tgz", + "integrity": "sha512-RRuzqDtt5Y9h3quz5hWhK+TPnsmVs6WwSU6LkJMeY4HstUEDuYTG8UJSdawMRzmzAtV+KEoG8N3Qg2qLy5vM/A==", "funding": [ { "type": "opencollective", @@ -6149,7 +6147,7 @@ ], "license": "MIT", "dependencies": { - "nanoid": "^3.3.12", + "nanoid": "^3.3.18", "picocolors": "^1.1.1", "source-map-js": "^1.2.1" }, @@ -6329,9 +6327,9 @@ "license": "MIT" }, "node_modules/protobufjs": { - "version": "7.6.4", - "resolved": "https://registry.npmjs.org/protobufjs/-/protobufjs-7.6.4.tgz", - "integrity": "sha512-RJJPTTpvFfHcWLkIa2JFWK4XvtSzS0yEWDmunqHXli1h3JlkbcQZXDZdcWxv+JK3Xsl5/UFDPZ0iGm7DAengYw==", + "version": "7.6.6", + "resolved": "https://registry.npmjs.org/protobufjs/-/protobufjs-7.6.6.tgz", + "integrity": "sha512-dYDWdjSl5RNb7SgPxGQcRU+GtvP7s2fpkrY0r432PcOIaZ0/rBcxEZnQN67iJhFuQiVw754JDoPruPCNdGsbjg==", "hasInstallScript": true, "license": "BSD-3-Clause", "dependencies": { @@ -6362,12 +6360,13 @@ } }, "node_modules/qs": { - "version": "6.15.2", - "resolved": "https://registry.npmjs.org/qs/-/qs-6.15.2.tgz", - "integrity": "sha512-Rzq0KEyX/w/tEybncDgdkZrJgVUsUMk3xjh3t5bv3S1HTAtg+uOYt72+ZfwiQwKdysThkTBdL/rTi6HDmX9Ddw==", + "version": "6.16.0", + "resolved": "https://registry.npmjs.org/qs/-/qs-6.16.0.tgz", + "integrity": "sha512-h6fhOIaRrID2CbEY2fqs+7t+UXZo+MLAnU5gRIq85uFtdiUPCdsApMlHhXogKVM4HM2DVbIjGNTTYH2OcmP1vA==", "license": "BSD-3-Clause", "dependencies": { - "side-channel": "^1.1.0" + "es-define-property": "^1.0.1", + "side-channel": "^1.1.1" }, "engines": { "node": ">=0.6" @@ -6544,9 +6543,9 @@ } }, "node_modules/react-router": { - "version": "7.17.0", - "resolved": "https://registry.npmjs.org/react-router/-/react-router-7.17.0.tgz", - "integrity": "sha512-FDELK7rTMlCHO5+reyXsPlmfr7N1F91lPHsWYfMEGQm/KQ+F4JFM8jGoeQDmDvdTs93Fw9aSilH+uKRb4/jXvQ==", + "version": "7.18.3", + "resolved": "https://registry.npmjs.org/react-router/-/react-router-7.18.3.tgz", + "integrity": "sha512-gyXgtdr5uACJ5b1Q4udzjVV+tb/rlHIMJKuJ0e89R4Kzgz47z/rgP0dIKxktqIEUhDHluGTPJJH/wRha7CyqsA==", "license": "MIT", "dependencies": { "cookie": "^1.0.1", @@ -6565,38 +6564,6 @@ } } }, - "node_modules/react-router-dom": { - "version": "6.30.4", - "resolved": "https://registry.npmjs.org/react-router-dom/-/react-router-dom-6.30.4.tgz", - "integrity": "sha512-q4HvNl+mmDdkS0g+MqiBZNteQJCuimWoOyHMy4T/RQLAn9Z29+E91QXRaxOujeMl2HTzRSS0KFPd7lxX3PjV0Q==", - "license": "MIT", - "dependencies": { - "@remix-run/router": "1.23.3", - "react-router": "6.30.4" - }, - "engines": { - "node": ">=14.0.0" - }, - "peerDependencies": { - "react": ">=16.8", - "react-dom": ">=16.8" - } - }, - "node_modules/react-router-dom/node_modules/react-router": { - "version": "6.30.4", - "resolved": "https://registry.npmjs.org/react-router/-/react-router-6.30.4.tgz", - "integrity": "sha512-SVUsDe+DybHM/WmYKIVYhZh1o5Dcuf16yM6WjG02Q9XVFMZIJyHYhwrr6bFBXZkVP6z69kNkMyBCujt8FaFLJA==", - "license": "MIT", - "dependencies": { - "@remix-run/router": "1.23.3" - }, - "engines": { - "node": ">=14.0.0" - }, - "peerDependencies": { - "react": ">=16.8" - } - }, "node_modules/react-style-singleton": { "version": "2.2.3", "resolved": "https://registry.npmjs.org/react-style-singleton/-/react-style-singleton-2.2.3.tgz", @@ -6656,9 +6623,9 @@ } }, "node_modules/recast": { - "version": "0.23.11", - "resolved": "https://registry.npmjs.org/recast/-/recast-0.23.11.tgz", - "integrity": "sha512-YTUo+Flmw4ZXiWfQKGcwwc11KnoRAYgzAE2E7mXKCjSviTKShtxBsN6YUUBB2gtaBzKzeKunxhUwNHQuRryhWA==", + "version": "0.23.21", + "resolved": "https://registry.npmjs.org/recast/-/recast-0.23.21.tgz", + "integrity": "sha512-mFAyJq9vUbSTARLZUvAEf1z3YxlvAwswbmxMx2mPA/MSm4KmpwvwvhsH/NIrZhyOuwD60Lzyw2qh83uCbgTPYw==", "license": "MIT", "peer": true, "dependencies": { @@ -7245,9 +7212,9 @@ "license": "MIT" }, "node_modules/synchronous-promise": { - "version": "2.0.17", - "resolved": "https://registry.npmjs.org/synchronous-promise/-/synchronous-promise-2.0.17.tgz", - "integrity": "sha512-AsS729u2RHUfEra9xJrE39peJcc2stq2+poBXX8bcM08Y6g9j/i/PUzwNQqkaJde7Ntg1TO7bSREbR5sdosQ+g==", + "version": "2.0.18", + "resolved": "https://registry.npmjs.org/synchronous-promise/-/synchronous-promise-2.0.18.tgz", + "integrity": "sha512-4EEtGWYLkSoy/DjlKpHR6LT2AjmORt8HM+CUCSdKiKyuhL3PFBksgcTibSOj7JIoCvuqRVnZ7P9QFBMXR5kW6Q==", "license": "BSD-3-Clause" }, "node_modules/tailwind-merge": { @@ -7428,9 +7395,9 @@ } }, "node_modules/tinyspy": { - "version": "4.0.4", - "resolved": "https://registry.npmjs.org/tinyspy/-/tinyspy-4.0.4.tgz", - "integrity": "sha512-azl+t0z7pw/z958Gy9svOTuzqIk6xq+NSheJzn5MMWtWTFywIacg2wUlzKFGtt3cthx0r2SxMK0yzJOR0IES7Q==", + "version": "4.0.6", + "resolved": "https://registry.npmjs.org/tinyspy/-/tinyspy-4.0.6.tgz", + "integrity": "sha512-u8KszXvGfU68hVcZpRHKG28T0krMuv2G5nDhiHaMLen/gIuFEgIJhaJuO69qjnXg5paSrbPMFfx3brNuN8eVSg==", "dev": true, "license": "MIT", "engines": { @@ -7579,9 +7546,9 @@ } }, "node_modules/update-browserslist-db": { - "version": "1.2.3", - "resolved": "https://registry.npmjs.org/update-browserslist-db/-/update-browserslist-db-1.2.3.tgz", - "integrity": "sha512-Js0m9cx+qOgDxo0eMiFGEueWztz+d4+M3rGlmKPT+T4IS/jP4ylw3Nwpu6cpTTP8R1MAC1kF4VbdLt3ARf209w==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/update-browserslist-db/-/update-browserslist-db-1.3.2.tgz", + "integrity": "sha512-UQ+MSxlhRm1bzjhU+DcuXfjFO1FzNtqhK5+9Yvlp90ItDLk5vT932A0rFu619nf7RVS+Y/VeaUW1jaRDqZ8VJw==", "dev": true, "funding": [ { @@ -7673,17 +7640,16 @@ "license": "MIT" }, "node_modules/uuid": { - "version": "10.0.0", - "resolved": "https://registry.npmjs.org/uuid/-/uuid-10.0.0.tgz", - "integrity": "sha512-8XkAphELsDnEGrDxUOHB3RGvXz6TeuYSGEZBOjtTtPm2lwhGBjLgOzLHB63IUWfBpNucQjND6d3AOudO+H3RWQ==", - "deprecated": "uuid@10 and below is no longer supported. For ESM codebases, update to uuid@latest. For CommonJS codebases, use uuid@11 (but be aware this version will likely be deprecated in 2028).", + "version": "14.0.2", + "resolved": "https://registry.npmjs.org/uuid/-/uuid-14.0.2.tgz", + "integrity": "sha512-xZe/16rV4aa+HGSOCiY2YeLT1OybRLrrkL/Rqaq7p7GMVXjFh+6wN4oMYgjFmnSnhY8t6Xpdl2l9qmnHYuMHwQ==", "funding": [ "https://github.com/sponsors/broofa", "https://github.com/sponsors/ctavan" ], "license": "MIT", "bin": { - "uuid": "dist/bin/uuid" + "uuid": "dist-node/bin/uuid" } }, "node_modules/vite": { @@ -7836,20 +7802,20 @@ } }, "node_modules/vitest": { - "version": "3.2.6", - "resolved": "https://registry.npmjs.org/vitest/-/vitest-3.2.6.tgz", - "integrity": "sha512-xejya+bT/j/+R/AGa1XOfRxLmNUlLtlwjRsFUILF+xHfzElmGcmFydy2gqqIrd62ptIEfwVMofd19uNWD9L7Nw==", + "version": "3.2.7", + "resolved": "https://registry.npmjs.org/vitest/-/vitest-3.2.7.tgz", + "integrity": "sha512-KrxIJ62Fd89gfysR4WotlgZABiz2dqFPgqGzX7s+CwsqLFomRH7777ZcrOD6+WVAh7khPQP41A+BKbpcJFrdEg==", "dev": true, "license": "MIT", "dependencies": { "@types/chai": "^5.2.2", - "@vitest/expect": "3.2.6", - "@vitest/mocker": "3.2.6", - "@vitest/pretty-format": "^3.2.6", - "@vitest/runner": "3.2.6", - "@vitest/snapshot": "3.2.6", - "@vitest/spy": "3.2.6", - "@vitest/utils": "3.2.6", + "@vitest/expect": "3.2.7", + "@vitest/mocker": "3.2.7", + "@vitest/pretty-format": "^3.2.7", + "@vitest/runner": "3.2.7", + "@vitest/snapshot": "3.2.7", + "@vitest/spy": "3.2.7", + "@vitest/utils": "3.2.7", "chai": "^5.2.0", "debug": "^4.4.1", "expect-type": "^1.2.1", @@ -7879,8 +7845,8 @@ "@edge-runtime/vm": "*", "@types/debug": "^4.1.12", "@types/node": "^18.0.0 || ^20.0.0 || >=22.0.0", - "@vitest/browser": "3.2.6", - "@vitest/ui": "3.2.6", + "@vitest/browser": "3.2.7", + "@vitest/ui": "3.2.7", "happy-dom": "*", "jsdom": "*" }, @@ -7909,9 +7875,9 @@ } }, "node_modules/vitest/node_modules/picomatch": { - "version": "4.0.4", - "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-4.0.4.tgz", - "integrity": "sha512-QP88BAKvMam/3NxH6vj2o21R6MjxZUAd6nlwAS/pnGvN9IVLocLHxGYIzFhg6fUQ+5th6P4dv4eW9jX3DSIj7A==", + "version": "4.0.7", + "resolved": "https://registry.npmjs.org/picomatch/-/picomatch-4.0.7.tgz", + "integrity": "sha512-qcJu88Q2IWqJsDD529JKMdwGm/dvInW4HvQnRwiH9JtihJvzGOscDtHE3x1pBKeUOTysQ8kVmLnJ2kJu7yhcGA==", "dev": true, "license": "MIT", "engines": { diff --git a/frontend/package.json b/frontend/package.json index 2750e66..10b98fb 100644 --- a/frontend/package.json +++ b/frontend/package.json @@ -24,7 +24,7 @@ "react-konva": "^18.2.10", "react-router": "^7.14.1", "tailwind-merge": "^3.2.0", - "uuid": "^10.0.0", + "uuid": "^14.0.2", "zundo": "^2.3.0", "zustand": "^5.0.0" }, @@ -37,6 +37,7 @@ "@types/react-dom": "^18.3.5", "@types/uuid": "^10.0.0", "@vitejs/plugin-react": "^4.3.4", + "@webgpu/types": "^0.1.72", "autoprefixer": "^10.4.20", "jsdom": "^26.1.0", "postcss": "^8.4.49", @@ -45,5 +46,8 @@ "vite": "^6.0.3", "vite-tsconfig-paths": "^5.1.4", "vitest": "^3.2.3" + }, + "overrides": { + "@blueskyproject/tiled": "^0.0.34" } } diff --git a/frontend/scripts/fetch-sam-model.mjs b/frontend/scripts/fetch-sam-model.mjs index c9af71e..5a2c64c 100644 --- a/frontend/scripts/fetch-sam-model.mjs +++ b/frontend/scripts/fetch-sam-model.mjs @@ -15,6 +15,10 @@ * node scripts/fetch-sam-model.mjs * * start_all.sh runs this automatically (best-effort, backgrounded) on startup. + * + * iPred's `slimsam` feature module (ipred/src/ipred/feature_setups.py) also + * resolves its ONNX encoder from this exact directory — vendoring here covers + * both the Magic tool and iPred, no separate download needed. */ import { spawnSync } from 'node:child_process'; import { existsSync } from 'node:fs'; diff --git a/frontend/src/app/App.tsx b/frontend/src/app/App.tsx index f2968eb..dd1fcb9 100644 --- a/frontend/src/app/App.tsx +++ b/frontend/src/app/App.tsx @@ -6,13 +6,19 @@ import HubAppLayout from '@/components/HubAppLayout'; import { useHubSelectedTabs } from '@/hooks/useHubSelectedTabs'; import { DOCS_URL, FEEDBACK_FORM_URL, FEEDBACK_ENTRY_ID } from '@/config'; import { buildFeedbackContext, buildFeedbackUrl } from '@/lib/feedbackContext'; -import { PlugsConnected, PencilSimple, MagnifyingGlass, BookOpen } from '@phosphor-icons/react'; +import { PlugsConnected, PencilSimple, MagnifyingGlass, BookOpen, Cube, Brain } from '@phosphor-icons/react'; // Lazy-loaded pages: keeps the heavy Annotate stack (konva, polygon-clipping, // magicwand, canvas) out of the initial /connect bundle — each page is its own chunk. const ConnectPage = lazy(() => import('./pages/ConnectPage')); const AnnotatePage = lazy(() => import('./pages/AnnotatePage')); const BrowsePage = lazy(() => import('./pages/BrowsePage')); const ReferencePage = lazy(() => import('./pages/ReferencePage')); +// Lazy for the same reason as Annotate, and more so: this chunk carries the +// whole vendored WebGPU renderer, which must never load on /connect. +const VolumePage = lazy(() => import('./pages/VolumePage')); +// Lazy for the same reason: the Train tab's own bundle, kept out of every +// other page's initial load. +const TrainPage = lazy(() => import('./pages/TrainPage')); import CustomizePages from '@/components/CustomizePages'; import IframeModal from '@/components/IframeModal'; @@ -43,6 +49,20 @@ const allRoutes: RouteItem[] = [ element: , isBackgroundTransparent: true, }, + { + path: '/volume', + label: '3D', + icon: , + element: , + isBackgroundTransparent: true, + }, + { + path: '/train', + label: 'Train', + icon: , + element: , + isBackgroundTransparent: true, + }, ]; const DEFAULT_PATHS = allRoutes.map((r) => r.path); diff --git a/frontend/src/app/pages/AnnotatePage.test.tsx b/frontend/src/app/pages/AnnotatePage.test.tsx new file mode 100644 index 0000000..9dac7e9 --- /dev/null +++ b/frontend/src/app/pages/AnnotatePage.test.tsx @@ -0,0 +1,770 @@ +/** + * AnnotatePage — orchestration tests. This page wires together nearly every + * annotate hook/component; children with their own dedicated test files are + * mocked here (AnnotationCanvas, Toolbar, ClassManager, LayersPanel, SaveModal, + * DownloadModal, VersionHistoryModal, DenoiseBakeModal, InsightsModal, + * FeatureChannelsPanel, PixelClassifierPanel), as are the data-hooks the page + * calls directly. Also mocked: DisplayControls, SliceNavigator, MaskToolsPanel, + * MeasurementPanel — none of AnnotatePage's own logic depends on their + * internals, and leaving them real would pull in fetch/useQuery/useMaskOps + * plumbing unrelated to this page's orchestration. Real: the Zustand stores, + * DebouncedSlider (opacity), VersionPreviewBar, StageSwitcher (inline). + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { MemoryRouter, Routes, Route } from 'react-router'; +import AnnotatePage from './AnnotatePage'; +import { useDatasetStore, type ImageMeta } from '@/stores/datasetStore'; +import { useAnnotationStore, type Shape } from '@/stores/annotationStore'; +import { useToolStore } from '@/stores/toolStore'; +import { useClassStore } from '@/stores/classStore'; +import { usePredictedRasterStore } from '@/stores/predictedRasterStore'; + +// ---- Hooks AnnotatePage calls directly ---- + +const useDraftSync = vi.fn(); +vi.mock('@/hooks/useDraftSync', () => ({ useDraftSync: (...args: unknown[]) => useDraftSync(...args) })); + +const useGuideLoad = vi.fn(); +vi.mock('@/hooks/useGuideSync', () => ({ useGuideLoad: (...args: unknown[]) => useGuideLoad(...args) })); + +const useKeybinds = vi.fn(); +vi.mock('@/hooks/useKeybinds', () => ({ useKeybinds: (...args: unknown[]) => useKeybinds(...args) })); + +const save = vi.fn(async () => true); +const buildSavePayload = vi.fn(() => null as null | Record); +const fetchVersionPayload = vi.fn(async () => null as unknown); +const restoreVersion = vi.fn(); +let useSaveReturn: Record; +vi.mock('@/hooks/useSave', () => ({ + useSave: (...args: unknown[]) => useSaveHook(...args), +})); +function useSaveHook(..._args: unknown[]) { + return useSaveReturn; +} + +const featuresCompute = vi.fn(); +let useFeatureChannelsReturn: Record; +vi.mock('@/hooks/useFeatureChannels', () => ({ + useFeatureChannels: (...args: unknown[]) => useFeatureChannelsHook(...args), +})); +function useFeatureChannelsHook(..._args: unknown[]) { + return useFeatureChannelsReturn; +} + +const clfTrain = vi.fn(); +const clfTrainAcrossSlices = vi.fn(); +const clfPredict = vi.fn(); +const clfDismiss = vi.fn(); +const clfApplyAcrossVolume = vi.fn(); +const clfResetVolumeApplyJob = vi.fn(); +let usePixelClassifierReturn: Record; +vi.mock('@/hooks/usePixelClassifier', () => ({ + usePixelClassifier: (...args: unknown[]) => usePixelClassifierHook(...args), +})); +function usePixelClassifierHook(..._args: unknown[]) { + return usePixelClassifierReturn; +} + +let useFeatureManifoldReturn: Record; +vi.mock('@/hooks/useFeatureManifold', () => ({ + useFeatureManifold: (...args: unknown[]) => useFeatureManifoldHook(...args), +})); +function useFeatureManifoldHook(..._args: unknown[]) { + return useFeatureManifoldReturn; +} + +const startMaskSync = vi.fn(); +let useExportJobReturn: Record; +vi.mock('@/hooks/useExportJob', () => ({ + useExportJob: (...args: unknown[]) => useExportJobHook(...args), +})); +function useExportJobHook(..._args: unknown[]) { + return useExportJobReturn; +} + +const loadLabelPng = vi.fn(); +const labelMapToPolygonShapes = vi.fn(); +vi.mock('@/lib/pixelClf', async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + loadLabelPng: (...args: unknown[]) => loadLabelPng(...args), + labelMapToPolygonShapes: (...args: unknown[]) => labelMapToPolygonShapes(...args), + }; +}); + +// ---- Heavy / separately-tested child components ---- + +vi.mock('@/components/annotate/AnnotationCanvas', () => ({ + default: (props: Record) => ( +
+ +
+ ), +})); + +vi.mock('@/components/annotate/Toolbar', () => ({ + default: () =>
, +})); + +vi.mock('@/components/annotate/ClassManager', () => ({ + default: (props: { activeClassId: number | null; onActivate: (id: number) => void; onClassDeleted: (id: number) => void }) => ( +
+ + +
+ ), +})); + +vi.mock('@/components/annotate/LayersPanel', () => ({ + default: () =>
, +})); + +vi.mock('@/components/annotate/DisplayControls', () => ({ + default: () =>
, +})); + +vi.mock('@/components/annotate/SliceNavigator', () => ({ + default: () =>
, +})); + +vi.mock('@/components/annotate/MaskToolsPanel', () => ({ + default: () =>
, +})); + +vi.mock('@/components/annotate/MeasurementPanel', () => ({ + default: () =>
, +})); + +vi.mock('@/components/annotate/FeatureChannelsPanel', () => ({ + default: () =>
, +})); + +vi.mock('@/components/annotate/PixelClassifierPanel', () => ({ + default: (props: Record) => ( +
+ + + + + + + + + + +
+ ), +})); + +vi.mock('@/components/annotate/SaveModal', () => ({ + default: (props: { + shapeCount: number; + classCount: number; + isSaving: boolean; + onSave: (opts: { annotatedBy: string; notes: string }) => void; + onClose: () => void; + }) => ( +
+ + +
+ ), +})); + +vi.mock('@/components/annotate/DownloadModal', () => ({ + default: (props: { onClose: () => void }) => ( +
+ +
+ ), +})); + +vi.mock('@/components/annotate/InsightsModal', () => ({ + default: (props: { onClose: () => void; onFocus: (slice: number, bbox?: { x: number; y: number; w: number; h: number }) => void }) => ( +
+ + +
+ ), +})); + +vi.mock('@/components/annotate/VersionHistoryModal', () => ({ + default: (props: { + versions: { version: number }[]; + onPreview: (v: number) => void; + onRestore: (v: number) => void; + onClose: () => void; + }) => ( +
+ + + +
+ ), +})); + +vi.mock('@/components/annotate/DenoiseBakeModal', () => ({ + default: (props: { open: boolean }) => ( +
+ ), +})); + +const initialDatasetState = useDatasetStore.getState(); +const initialAnnotationState = useAnnotationStore.getState(); +const initialToolState = useToolStore.getState(); +const initialClassState = useClassStore.getState(); +const initialPredictedRasterState = usePredictedRasterStore.getState(); + +const META: ImageMeta = { + nSlices: 5, + height: 100, + width: 100, + dtype: 'uint8', + isRgb: false, + valueRange: [0, 255], +}; + +function shape(id: string, classId: number): Shape { + return { id, classId, kind: 'rectangle', x: 0, y: 0, w: 10, h: 10 }; +} + +/** Loads a dataset (meta present) so AnnotatePage renders the full workspace. */ +function loadDataset(overrides: Partial<{ source: string; kind: string; serverUri: string | null }> = {}) { + useDatasetStore.setState({ + ...useDatasetStore.getState(), + kind: overrides.kind ?? 'tiled', + source: overrides.source ?? 'sample.zarr', + serverUri: overrides.serverUri ?? 'http://localhost:8000/api', + meta: META, + currentSlice: 0, + }); +} + +function renderPage() { + return render( + + + } /> + BROWSE PAGE
} /> + VOLUME PAGE
} /> + TRAIN PAGE
} /> + + , + ); +} + +describe('AnnotatePage', () => { + beforeEach(() => { + useDatasetStore.setState(initialDatasetState, true); + useAnnotationStore.setState(initialAnnotationState, true); + useToolStore.setState(initialToolState, true); + useClassStore.setState(initialClassState, true); + usePredictedRasterStore.setState(initialPredictedRasterState, true); + + useSaveReturn = { + isDirty: false, + isSaving: false, + lastSavedAt: null, + buildSavePayload, + saveSummary: { shapeCount: 0, classCount: 0 }, + save, + versions: [], + fetchVersionPayload, + restoreVersion, + markClean: vi.fn(), + }; + useFeatureChannelsReturn = { + job: null, + channelIndex: null, + channelUrl: null, + computing: false, + error: null, + invalidateJob: vi.fn(), + adoptFeatureBank: vi.fn(), + compute: featuresCompute, + selectChannel: vi.fn(), + cycleChannel: vi.fn(), + clearSelection: vi.fn(), + }; + usePixelClassifierReturn = { + params: { iterations: 200, depth: 6, learningRate: 0.1, alpha: 0.05 }, + setParams: vi.fn(), + model: null, + commitUrl: null, + statusUrl: null, + probaUrl: null, + predictCounts: null, + probaClassIndex: 0, + activeProbaClassId: null, + activeProbaThreshold: 0.5, + cycleProbaClass: vi.fn(), + setProbaThreshold: vi.fn(), + error: null, + training: false, + predicting: false, + train: clfTrain, + trainAcrossSlices: clfTrainAcrossSlices, + predict: clfPredict, + dismiss: clfDismiss, + commitUrl_: null, + compositionId: 'comp-1', + trainerId: 'catboost', + canTrainWithoutJob: true, + multiTrainJob: { done: 0, total: 0 }, + multiTraining: false, + resetMultiTrainJob: vi.fn(), + applyAcrossVolume: clfApplyAcrossVolume, + volumeApplyJob: { done: 0, total: 0, result: null, jobId: null }, + volumeApplying: false, + resetVolumeApplyJob: clfResetVolumeApplyJob, + }; + useFeatureManifoldReturn = { + params: { k: 24, boxSize: 64 }, + setParams: vi.fn(), + points: [], + heatmapUrl: null, + showHeatmap: false, + setShowHeatmap: vi.fn(), + showMarkers: false, + setShowMarkers: vi.fn(), + heatmapOpacity: 0.45, + setHeatmapOpacity: vi.fn(), + sampling: false, + error: null, + meta: null, + sample: vi.fn(), + dismiss: vi.fn(), + placementMask: null, + setPlacementMaskFromShapes: vi.fn(), + clearPlacementMask: vi.fn(), + hasSample: false, + }; + useExportJobReturn = { + state: { status: 'idle', phase: '', done: 0, total: 0, log: [], result: null, error: null, jobId: null }, + start: vi.fn(), + startMaskSync, + startIpredBatchTrain: vi.fn(), + startIpredBatchApply: vi.fn(), + startJob: vi.fn(), + reset: vi.fn(), + downloadUrl: null, + }; + }); + + afterEach(() => { + cleanup(); + vi.clearAllMocks(); + }); + + it('shows "no sample loaded" when no dataset is open, and Go to Browse navigates', async () => { + const user = userEvent.setup(); + renderPage(); + expect(screen.getByText(/No sample loaded/)).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: /Go to Browse/ })); + expect(await screen.findByText('BROWSE PAGE')).toBeInTheDocument(); + }); + + it('renders the draw-stage panels by default and switches stages via the tab bar', async () => { + const user = userEvent.setup(); + loadDataset(); + renderPage(); + + expect(screen.getByTestId('toolbar')).toBeInTheDocument(); + expect(screen.getByTestId('display-controls')).toBeInTheDocument(); + expect(screen.getByTestId('mask-tools-panel')).toBeInTheDocument(); + expect(screen.getByTestId('measurement-panel')).toBeInTheDocument(); + expect(screen.queryByTestId('feature-channels-panel')).not.toBeInTheDocument(); + expect(screen.queryByTestId('pixel-classifier-panel')).not.toBeInTheDocument(); + + await user.click(screen.getByRole('tab', { name: 'Assist' })); + expect(screen.getByTestId('feature-channels-panel')).toBeInTheDocument(); + expect(screen.queryByTestId('toolbar')).not.toBeInTheDocument(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + expect(screen.getByTestId('pixel-classifier-panel')).toBeInTheDocument(); + expect(screen.queryByTestId('feature-channels-panel')).not.toBeInTheDocument(); + }); + + it('auto-activates the first class once classes load, and clicking a class updates the canvas + clears the in-progress brush', async () => { + const user = userEvent.setup(); + loadDataset(); + useClassStore.setState({ + classes: [ + { classId: 1, label: 'A', color: '#f00', isVisible: true }, + { classId: 2, label: 'B', color: '#0f0', isVisible: true }, + ], + }); + renderPage(); + + // Auto-activated to the first class (classId 1). + expect(screen.getByTestId('annotation-canvas')).toHaveAttribute('data-active-class-id', '1'); + + // Start a brush instance, then switch class — activeBrushShapeId must reset. + await user.click(screen.getByText('mock-new-brush')); + expect(screen.getByTestId('annotation-canvas')).toHaveAttribute('data-active-brush-shape-id', 'brush-1'); + + await user.click(screen.getByText('mock-activate-class-2')); + expect(screen.getByTestId('annotation-canvas')).toHaveAttribute('data-active-class-id', '2'); + expect(screen.getByTestId('annotation-canvas')).toHaveAttribute('data-active-brush-shape-id', 'null'); + }); + + it('clearing the active brush instance after a class delete does not change activeClassId', async () => { + const user = userEvent.setup(); + loadDataset(); + useClassStore.setState({ classes: [{ classId: 1, label: 'A', color: '#f00', isVisible: true }] }); + renderPage(); + + await user.click(screen.getByText('mock-new-brush')); + expect(screen.getByTestId('annotation-canvas')).toHaveAttribute('data-active-brush-shape-id', 'brush-1'); + + await user.click(screen.getByText('mock-delete-class-1')); + expect(screen.getByTestId('annotation-canvas')).toHaveAttribute('data-active-brush-shape-id', 'null'); + expect(screen.getByTestId('annotation-canvas')).toHaveAttribute('data-active-class-id', '1'); + }); + + it('save button reflects isDirty/isSaving and opens the save modal with the current summary via buildSavePayload', async () => { + const user = userEvent.setup(); + loadDataset(); + useSaveReturn.isDirty = true; + useSaveReturn.saveSummary = { shapeCount: 4, classCount: 2 }; + buildSavePayload.mockReturnValue({ classes: [], slices: {}, split_by_slice: {}, negative_slices: [] }); + renderPage(); + + const saveBtn = screen.getByRole('button', { name: /^Save$/ }); + expect(screen.getByText('Unsaved changes')).toBeInTheDocument(); + + await user.click(saveBtn); + const modal = screen.getByTestId('save-modal'); + expect(modal).toHaveAttribute('data-shape-count', '4'); + expect(modal).toHaveAttribute('data-class-count', '2'); + }); + + it('does not open the save modal when buildSavePayload returns null (nothing to save)', async () => { + const user = userEvent.setup(); + loadDataset(); + buildSavePayload.mockReturnValue(null); + renderPage(); + + await user.click(screen.getByRole('button', { name: /Saved/ })); + expect(screen.queryByTestId('save-modal')).not.toBeInTheDocument(); + }); + + it('confirming the save modal calls save() and closes the modal on success', async () => { + const user = userEvent.setup(); + loadDataset(); + buildSavePayload.mockReturnValue({ classes: [], slices: {}, split_by_slice: {}, negative_slices: [] }); + save.mockResolvedValueOnce(true); + renderPage(); + + await user.click(screen.getByRole('button', { name: /Saved/ })); + await user.click(screen.getByText('mock-confirm-save')); + + expect(save).toHaveBeenCalledWith({ annotatedBy: 'me', notes: 'n' }); + await waitFor(() => expect(screen.queryByTestId('save-modal')).not.toBeInTheDocument()); + }); + + it('keeps the save modal open when save() fails', async () => { + const user = userEvent.setup(); + loadDataset(); + buildSavePayload.mockReturnValue({ classes: [], slices: {}, split_by_slice: {}, negative_slices: [] }); + save.mockResolvedValueOnce(false); + renderPage(); + + await user.click(screen.getByRole('button', { name: /Saved/ })); + await user.click(screen.getByText('mock-confirm-save')); + + await waitFor(() => expect(save).toHaveBeenCalled()); + expect(screen.getByTestId('save-modal')).toBeInTheDocument(); + }); + + it('shows the version-history entry point only once versions exist, and preview/restore wire through', async () => { + const user = userEvent.setup(); + loadDataset(); + useSaveReturn.versions = [ + { version: 1, saved_at: '2024-01-01T00:00:00Z', shape_count: 1, class_count: 1 }, + { version: 2, saved_at: '2024-01-02T00:00:00Z', shape_count: 2, class_count: 1 }, + ]; + fetchVersionPayload.mockResolvedValue({ classes: [], slices: {}, split_by_slice: {}, negative_slices: [] }); + renderPage(); + + const historyBtn = screen.getByTitle('Version history'); + expect(historyBtn).toHaveTextContent('2'); + await user.click(historyBtn); + expect(screen.getByTestId('version-history-modal')).toBeInTheDocument(); + + // Preview closes the history modal and opens the version preview bar. + await user.click(screen.getByText('mock-preview-v2')); + expect(screen.queryByTestId('version-history-modal')).not.toBeInTheDocument(); + await waitFor(() => expect(fetchVersionPayload).toHaveBeenCalledWith(2)); + expect(await screen.findByText(/Previewing v/)).toBeInTheDocument(); + + // Exiting the preview bar clears previewVersion (bar disappears). + await user.click(screen.getByTitle('Exit preview')); + expect(screen.queryByText(/Previewing v/)).not.toBeInTheDocument(); + }); + + it('restoring a version from history calls restoreVersion', async () => { + const user = userEvent.setup(); + loadDataset(); + useSaveReturn.versions = [{ version: 1, saved_at: '2024-01-01T00:00:00Z', shape_count: 1, class_count: 1 }]; + renderPage(); + + await user.click(screen.getByTitle('Version history')); + await user.click(screen.getByText('mock-restore-v1')); + expect(restoreVersion).toHaveBeenCalledWith(1); + }); + + it('Insights: opens on click, and a QA focus jumps the slice and sets the focus region on the canvas', async () => { + const user = userEvent.setup(); + loadDataset(); + renderPage(); + + await user.click(screen.getByRole('button', { name: /Insights/ })); + expect(screen.getByTestId('insights-modal')).toBeInTheDocument(); + + await user.click(screen.getByText('mock-insight-focus')); + expect(useDatasetStore.getState().currentSlice).toBe(3); + const canvas = screen.getByTestId('annotation-canvas'); + const region = JSON.parse(canvas.getAttribute('data-focus-region') ?? 'null'); + expect(region).toMatchObject({ x: 1, y: 2, w: 3, h: 4 }); + + await user.click(screen.getByText('mock-close-insights')); + expect(screen.queryByTestId('insights-modal')).not.toBeInTheDocument(); + }); + + it('Export opens and closes the download modal', async () => { + const user = userEvent.setup(); + loadDataset(); + renderPage(); + + await user.click(screen.getByRole('button', { name: /Export/ })); + expect(screen.getByTestId('download-modal')).toBeInTheDocument(); + await user.click(screen.getByText('mock-close-download')); + expect(screen.queryByTestId('download-modal')).not.toBeInTheDocument(); + }); + + it('renders the denoise bake modal only when a source is open', () => { + loadDataset({ source: 'sample.zarr' }); + renderPage(); + expect(screen.getByTestId('denoise-bake-modal')).toBeInTheDocument(); + }); + + it('passes the correct callbacks to useKeybinds, and the delete-selected callback removes selected shapes on the current slice', () => { + loadDataset(); + const sourceKey = 'tiled:http://localhost:8000/api:sample.zarr'; + useAnnotationStore.setState({ + byImage: { [sourceKey]: { '0': [shape('s1', 1), shape('s2', 1)] } }, + }); + useToolStore.setState({ selectedShapeIds: ['s1'] }); + renderPage(); + + expect(useKeybinds).toHaveBeenCalled(); + const [activeClassId, onActivateClass, onNewBrushInstance, onDeleteSelected, onCancelDraft] = + useKeybinds.mock.calls[useKeybinds.mock.calls.length - 1]; + expect(typeof onActivateClass).toBe('function'); + expect(typeof onNewBrushInstance).toBe('function'); + expect(activeClassId).toBeNull(); + + act(() => onDeleteSelected()); + expect(useAnnotationStore.getState().byImage[sourceKey]['0'].map((s: Shape) => s.id)).toEqual(['s2']); + expect(useToolStore.getState().selectedShapeIds).toEqual([]); + + act(() => onCancelDraft()); + }); + + describe('pixel classifier orchestration', () => { + it('Train calls clf.train with the current slice shapes when "train across slices" is off', async () => { + const user = userEvent.setup(); + loadDataset(); + const sourceKey = 'tiled:http://localhost:8000/api:sample.zarr'; + const shapes = [shape('s1', 1)]; + useAnnotationStore.setState({ byImage: { [sourceKey]: { '0': shapes } } }); + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-clf-train')); + expect(clfTrain).toHaveBeenCalledWith(shapes); + expect(clfTrainAcrossSlices).not.toHaveBeenCalled(); + }); + + it('Train calls clf.trainAcrossSlices with every non-empty slice once "train across slices" is toggled on', async () => { + const user = userEvent.setup(); + loadDataset(); + const sourceKey = 'tiled:http://localhost:8000/api:sample.zarr'; + useAnnotationStore.setState({ + byImage: { [sourceKey]: { '0': [shape('s1', 1)], '1': [shape('s2', 1)], '2': [] } }, + }); + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-toggle-train-across')); + await user.click(screen.getByText('mock-clf-train')); + + expect(clfTrainAcrossSlices).toHaveBeenCalledWith({ + 0: [shape('s1', 1)], + 1: [shape('s2', 1)], + }); + expect(clfTrain).not.toHaveBeenCalled(); + }); + + it('Predict calls clf.predict with the current slice shapes', async () => { + const user = userEvent.setup(); + loadDataset(); + const sourceKey = 'tiled:http://localhost:8000/api:sample.zarr'; + const shapes = [shape('s1', 1)]; + useAnnotationStore.setState({ byImage: { [sourceKey]: { '0': shapes } } }); + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-clf-predict')); + expect(clfPredict).toHaveBeenCalledWith(shapes); + }); + + it('Commit vectorizes the predicted PNG into shapes, appends them, and dismisses the overlay', async () => { + const user = userEvent.setup(); + loadDataset(); + const sourceKey = 'tiled:http://localhost:8000/api:sample.zarr'; + usePixelClassifierReturn.commitUrl = 'blob:commit'; + usePixelClassifierReturn.model = { classIds: [1, 2] }; + loadLabelPng.mockResolvedValue({ data: new Uint8Array([1]), width: 1, height: 1 }); + const newShape = shape('predicted-1', 1); + labelMapToPolygonShapes.mockReturnValue([newShape]); + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-clf-commit')); + + await waitFor(() => + expect(useAnnotationStore.getState().byImage[sourceKey]?.['0']).toEqual([newShape]), + ); + expect(clfDismiss).toHaveBeenCalledTimes(1); + }); + + it('Apply-to-volume runs across every slice index up to the dataset\'s slice count', async () => { + const user = userEvent.setup(); + loadDataset(); // META.nSlices === 5 + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-apply-to-volume')); + expect(clfApplyAcrossVolume).toHaveBeenCalledWith([0, 1, 2, 3, 4]); + }); + + it('Commit-volume-apply records predicted-raster pointers only for slices with no existing shapes', async () => { + const user = userEvent.setup(); + loadDataset(); + const sourceKey = 'tiled:http://localhost:8000/api:sample.zarr'; + useAnnotationStore.setState({ byImage: { [sourceKey]: { '1': [shape('existing', 1)] } } }); + usePixelClassifierReturn.volumeApplyJob = { + done: 5, + total: 5, + result: { runs: { '0': 'run-a', '1': 'run-b', '2': 'run-c' } }, + }; + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + // Select a commit class first via the toggle so commitClassIds is non-empty. + await user.click(screen.getByText('mock-toggle-train-across')); // harmless toggle to exercise a click + await user.click(screen.getByText('mock-commit-volume-apply')); + + // commitClassIds defaults to [] until a model exists; with model null the + // handler no-ops (commitClassIds.length === 0) — so no pointers recorded. + expect(usePredictedRasterStore.getState().bySource[sourceKey]).toBeUndefined(); + }); + + it('Commit-volume-apply records pointers for un-annotated slices once a model supplies commit classes', async () => { + const user = userEvent.setup(); + loadDataset(); + const sourceKey = 'tiled:http://localhost:8000/api:sample.zarr'; + useAnnotationStore.setState({ byImage: { [sourceKey]: { '1': [shape('existing', 1)] } } }); + usePixelClassifierReturn.model = { classIds: [1, 2] }; + usePixelClassifierReturn.volumeApplyJob = { + done: 3, + total: 3, + result: { runs: { '0': 'run-a', '1': 'run-b', '2': 'run-c' } }, + }; + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-commit-volume-apply')); + + const pointers = usePredictedRasterStore.getState().bySource[sourceKey]; + // Slice 1 already has real shapes, so it's skipped; 0 and 2 get pointers. + expect(pointers).toEqual({ + '0': { runId: 'run-a', classIds: [1, 2] }, + '2': { runId: 'run-c', classIds: [1, 2] }, + }); + expect(clfResetVolumeApplyJob).toHaveBeenCalledTimes(1); + }); + + it('View in 3D navigates to /volume?mask=fast', async () => { + const user = userEvent.setup(); + loadDataset(); + usePixelClassifierReturn.model = { classIds: [1] }; + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-view-in-3d')); + expect(await screen.findByText('VOLUME PAGE')).toBeInTheDocument(); + }); + + it('Train a deep model navigates to /train', async () => { + const user = userEvent.setup(); + loadDataset(); + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-train-deep-model')); + expect(await screen.findByText('TRAIN PAGE')).toBeInTheDocument(); + }); + + it('Sync to Tiled alerts (and does not start the job) for a non-Tiled source', async () => { + const user = userEvent.setup(); + loadDataset({ kind: 'local', source: 'data/foo', serverUri: null }); + const alertSpy = vi.spyOn(window, 'alert').mockImplementation(() => {}); + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-sync-to-tiled')); + expect(alertSpy).toHaveBeenCalledWith('Pushing masks to Tiled only works for Tiled sources.'); + expect(startMaskSync).not.toHaveBeenCalled(); + }); + + it('Sync to Tiled starts the mask-sync job with the current slices/classes for a Tiled source', async () => { + const user = userEvent.setup(); + loadDataset(); + const sourceKey = 'tiled:http://localhost:8000/api:sample.zarr'; + const shapes = [shape('s1', 1)]; + useAnnotationStore.setState({ byImage: { [sourceKey]: { '0': shapes } } }); + useClassStore.setState({ classes: [{ classId: 1, label: 'A', color: '#f00', isVisible: true }] }); + renderPage(); + + await user.click(screen.getByRole('tab', { name: 'Predict' })); + await user.click(screen.getByText('mock-sync-to-tiled')); + + expect(startMaskSync).toHaveBeenCalledWith({ + sources: [ + expect.objectContaining({ + kind: 'tiled', + source: 'sample.zarr', + server_uri: 'http://localhost:8000/api', + slices: { '0': shapes }, + }), + ], + classes: [{ classId: 1, label: 'A', color: '#f00', isVisible: true }], + }); + }); + }); +}); diff --git a/frontend/src/app/pages/AnnotatePage.tsx b/frontend/src/app/pages/AnnotatePage.tsx index c9d1062..72e03d4 100644 --- a/frontend/src/app/pages/AnnotatePage.tsx +++ b/frontend/src/app/pages/AnnotatePage.tsx @@ -1,11 +1,13 @@ /** * AnnotatePage — react-konva canvas workspace with sidebar tools. */ -import { useState, useEffect, useCallback } from 'react'; +import { useState, useEffect, useCallback, useMemo, useRef } from 'react'; import { useNavigate } from 'react-router'; import { DownloadSimple, FloppyDisk, ClockCounterClockwise, CircleDashed, ChartBar } from '@phosphor-icons/react'; +import { cn } from '@/lib/utils'; import { useDatasetStore } from '@/stores/datasetStore'; import { useAnnotationStore } from '@/stores/annotationStore'; +import { usePredictedRasterStore } from '@/stores/predictedRasterStore'; import { useToolStore } from '@/stores/toolStore'; import { useClassStore } from '@/stores/classStore'; import { useDraftSync } from '@/hooks/useDraftSync'; @@ -13,28 +15,76 @@ import { clearHistory } from '@/hooks/editHistory'; import { useGuideLoad } from '@/hooks/useGuideSync'; import { useSave, type VersionPayload } from '@/hooks/useSave'; import { buildSourceKey } from '@/lib/sourceKey'; +import { API_BASE } from '@/config'; import { useKeybinds } from '@/hooks/useKeybinds'; import type { ColormapName } from '@/lib/colormaps'; import Toolbar from '@/components/annotate/Toolbar'; import ClassManager from '@/components/annotate/ClassManager'; +import LayersPanel from '@/components/annotate/LayersPanel'; import DisplayControls from '@/components/annotate/DisplayControls'; +import DenoisePanel from '@/components/annotate/DenoisePanel'; +import DenoiseBakeModal from '@/components/annotate/DenoiseBakeModal'; import SliceNavigator from '@/components/annotate/SliceNavigator'; import MaskToolsPanel from '@/components/annotate/MaskToolsPanel'; import MeasurementPanel from '@/components/annotate/MeasurementPanel'; +import FeatureChannelsPanel from '@/components/annotate/FeatureChannelsPanel'; +import SuggestLabelsPanel from '@/components/annotate/SuggestLabelsPanel'; +import PixelClassifierPanel from '@/components/annotate/PixelClassifierPanel'; import AnnotationCanvas from '@/components/annotate/AnnotationCanvas'; +import { useFeatureChannels } from '@/hooks/useFeatureChannels'; +import { usePixelClassifier } from '@/hooks/usePixelClassifier'; +import { useExportJob } from '@/hooks/useExportJob'; +import { useFeatureManifold } from '@/hooks/useFeatureManifold'; +import { loadLabelPng, labelMapToPolygonShapes } from '@/lib/pixelClf'; +import { ipredRunCommitUrl } from '@/lib/ipredApi'; +import type { Shape } from '@/stores/annotationStore'; import DebouncedSlider from '@/components/common/DebouncedSlider'; import DownloadModal from '@/components/annotate/DownloadModal'; import InsightsModal from '@/components/annotate/InsightsModal'; import VersionHistoryModal from '@/components/annotate/VersionHistoryModal'; import VersionPreviewBar from '@/components/annotate/VersionPreviewBar'; +import PerfOverlay from '@/components/annotate/PerfOverlay'; +import type { SamplerFit } from '@/components/annotate/AnnotationCanvas'; +import { initPerf } from '@/lib/perf'; +import { loadDisplayPrefs, saveDisplayPrefs } from '@/lib/displayPrefs'; import SaveModal from '@/components/annotate/SaveModal'; import type { SaveDraftPayload } from '@/hooks/useSave'; +type AnnotateStage = 'draw' | 'assist' | 'predict'; + +const STAGES: { id: AnnotateStage; label: string }[] = [ + { id: 'draw', label: 'Draw' }, + { id: 'assist', label: 'Assist' }, + { id: 'predict', label: 'Predict' }, +]; + +function StageSwitcher({ stage, onChange }: { stage: AnnotateStage; onChange: (s: AnnotateStage) => void }) { + return ( +
+ {STAGES.map((s) => ( + + ))} +
+ ); +} + /** Renders the annotation workspace: tool sidebar, canvas, and save/version/export flows. */ export default function AnnotatePage() { const navigate = useNavigate(); - const { source, kind, serverUri, meta } = useDatasetStore(); - const { removeShapes } = useAnnotationStore(); + const { source, kind, serverUri, meta, denoise, setDenoise } = useDatasetStore(); + const { removeShapes, addShapes, addShapesAcrossSlices, byImage, splitBySlice, negativeSlices } = useAnnotationStore(); const { selectedShapeIds, setSelectedShapeId, fillOpacity, setFillOpacity } = useToolStore(); const { classes } = useClassStore(); @@ -57,18 +107,45 @@ export default function AnnotatePage() { if (classes.length > 0) setActiveClassId(classes[0].classId); }, [classes, activeClassId]); - const [brightness, setBrightness] = useState(0); - const [contrast, setContrast] = useState(0); + // Cosmetic display sliders survive reload via localStorage (global viewer + // preference, not per-sample/draft) — read once, lazily, on first mount. + const [initialDisplayPrefs] = useState(() => loadDisplayPrefs()); + const [brightness, setBrightness] = useState(() => initialDisplayPrefs.brightness ?? 0); + const [contrast, setContrast] = useState(() => initialDisplayPrefs.contrast ?? 0); // Min/max levels window (0–255) + histogram of the current slice (client-side). - const [levelsLo, setLevelsLo] = useState(0); - const [levelsHi, setLevelsHi] = useState(255); + const [levelsLo, setLevelsLo] = useState(() => initialDisplayPrefs.levelsLo ?? 0); + const [levelsHi, setLevelsHi] = useState(() => initialDisplayPrefs.levelsHi ?? 255); const [histogramBins, setHistogramBins] = useState(null); // Display-only false-color map + gamma. - const [colormap, setColormap] = useState('gray'); - const [gamma, setGamma] = useState(1); - // Display-only nonlinear preprocessors (adaptive CLAHE / Sharpen). - const [clahe, setClahe] = useState(false); - const [sharpen, setSharpen] = useState(false); + const [colormap, setColormap] = useState(() => initialDisplayPrefs.colormap ?? 'gray'); + const [gamma, setGamma] = useState(() => initialDisplayPrefs.gamma ?? 1); + // Display-only nonlinear preprocessors (Gaussian blur / adaptive CLAHE / Sharpen). + const [clahe, setClahe] = useState(() => initialDisplayPrefs.clahe ?? false); + const [sharpen, setSharpen] = useState(() => initialDisplayPrefs.sharpen ?? false); + const [blur, setBlur] = useState(() => initialDisplayPrefs.blur ?? 0); + // Persist on every change — deliberately excludes `denoise`, which stays + // unpersisted (see datasetStore.ts): a strength tuned to one volume's noise + // level would silently mis-filter a different one carried over. + useEffect(() => { + saveDisplayPrefs({ brightness, contrast, levelsLo, levelsHi, gamma, colormap, clahe, sharpen, blur }); + }, [brightness, contrast, levelsLo, levelsHi, gamma, colormap, clahe, sharpen, blur]); + const [bakeOpen, setBakeOpen] = useState(false); + // Working resolution for the drawing tools (1x, 2x, 4x). Annotation coordinates + // stay native; upscaling only buys sub-pixel precision on small features. + const [upscale, setUpscale] = useState(1); + // Each level costs 4x the pixels in the base canvas AND in every tool field, so + // cap the offer at what this slice can afford (the canvas enforces the same limit). + // Bundled for the Toolbar's threshold band picker, which remaps the (base-space) + // histogram into displayed space so the plot matches the image and the band. + const thresholdDisplay = useMemo( + () => ({ brightness, contrast, levelsLo, levelsHi, gamma }), + [brightness, contrast, levelsLo, levelsHi, gamma], + ); + const maxUpscale = useMemo(() => { + if (!meta) return 4; + const px = meta.width * meta.height; + return px * 16 <= 64e6 ? 4 : px * 4 <= 64e6 ? 2 : 1; + }, [meta]); const [showDownload, setShowDownload] = useState(false); const [showInsights, setShowInsights] = useState(false); // Region to zoom to + highlight on the canvas (from an Insights QA flag). @@ -86,6 +163,21 @@ export default function AnnotatePage() { ? buildSourceKey(kind as 'tiled' | 'local', source, serverUri) : null; + // Dev-only timing HUD (?perf=1). Read once — the flag is fixed for the session. + const [perfOn] = useState(() => initPerf()); + const [stage, setStage] = useState('draw'); + + // Latest Sampler lasso fit — transient UI state, not persisted: only the band + // and blur it applies are durable. + const [samplerFit, setSamplerFit] = useState(null); + /** Put back the band and blur that were in force before the last fit. */ + const revertSamplerFit = useCallback(() => { + if (!samplerFit) return; + useToolStore.getState().setThresholdBand(samplerFit.previousBand[0], samplerFit.previousBand[1]); + setBlur(samplerFit.previousBlur); + setSamplerFit(null); + }, [samplerFit]); + // Crash-recovery autosave (local draft only, no Tiled sync) useDraftSync(sourceKey); // Undo/redo is per-sample: reset the region history + class-delete journal on switch. @@ -98,12 +190,244 @@ export default function AnnotatePage() { const { currentSlice, setSlice } = useDatasetStore(); + // Shapes on the current slice — shown in the perf HUD, since it is the variable + // that drives most of the costs it reports. Only subscribed when the HUD is on. + const currentShapeCount = useAnnotationStore((s) => + perfOn && sourceKey ? (s.byImage[sourceKey]?.[String(currentSlice)]?.length ?? 0) : 0, + ); + /** From an Insights QA flag: jump to its slice and zoom/highlight its region. */ const handleInsightFocus = useCallback((slice: number, bbox?: { x: number; y: number; w: number; h: number }) => { setSlice(slice); setFocusRegion(bbox ? { ...bbox, nonce: Date.now() } : null); }, [setSlice]); + const sliceShapes = sourceKey ? (byImage[sourceKey]?.[String(currentSlice)] ?? []) : []; + + const features = useFeatureChannels({ + source, + kind, + sliceIndex: currentSlice, + serverUri, + }); + + const clf = usePixelClassifier({ + featureJobId: features.job?.jobId ?? null, + // Sample/composition identity only — NOT the slice, so a trained model + // survives a slice change and can be applied to whichever slice is on + // screen (see usePixelClassifier's two reset effects for the split). + resetKey: sourceKey, + source, + kind, + serverUri, + sliceIndex: currentSlice, + onFeatureJobExpired: features.invalidateJob, + onFeatureReady: (info) => features.adoptFeatureBank(info), + }); + + const manifold = useFeatureManifold({ + featureJobId: features.job?.jobId ?? null, + onFeatureJobExpired: features.invalidateJob, + }); + + // ---- Batch iPred: multi-slice train + whole-volume apply ---- + const [trainAcrossSlices, setTrainAcrossSlices] = useState(false); + const [commitClassIds, setCommitClassIds] = useState([]); + const { setPointers: setPredictedPointers, clearSlice: clearPredictedSlice } = usePredictedRasterStore(); + const predictedPointers = usePredictedRasterStore((s) => s.bySource); + + // Every non-empty slice of this sample — the pool for "train across all + // annotated slices" and the source of truth for `annotatedSliceCount`. + const annotatedSlices = useMemo(() => { + const slices = sourceKey ? (byImage[sourceKey] ?? {}) : {}; + const out: Record = {}; + for (const [key, shapes] of Object.entries(slices)) { + if (shapes.length > 0) out[Number(key)] = shapes; + } + return out; + }, [sourceKey, byImage]); + const annotatedSliceCount = Object.keys(annotatedSlices).length; + const totalSliceCount = meta?.nSlices ?? 1; + + // Keep the commit class filter in sync with whichever model is current — + // default to "all classes" each time a (re)train finishes. + useEffect(() => { + setCommitClassIds(clf.model?.classIds ?? []); + }, [clf.model]); + + const toggleCommitClassId = useCallback((classId: number) => { + setCommitClassIds((prev) => + prev.includes(classId) ? prev.filter((c) => c !== classId) : [...prev, classId], + ); + }, []); + + const handleClfTrain = useCallback(() => { + if (trainAcrossSlices) { + void clf.trainAcrossSlices(annotatedSlices); + } else { + void clf.train(sliceShapes); + } + }, [clf, sliceShapes, trainAcrossSlices, annotatedSlices]); + + const handleClfPredict = useCallback(() => { + void clf.predict(sliceShapes); + }, [clf, sliceShapes]); + + const handleClfCommit = useCallback(async () => { + if (!sourceKey || !clf.commitUrl || !clf.model) return; + try { + const { data, width, height } = await loadLabelPng(clf.commitUrl); + const shapes = labelMapToPolygonShapes(data, width, height, commitClassIds, { + minRegion: 64, + preserveShapes: sliceShapes, + origin: 'predicted', + }); + if (shapes.length) addShapes(sourceKey, currentSlice, shapes); + clf.dismiss(); + } catch { + /* keep overlay; error surfaces via clf.error if needed */ + } + }, [sourceKey, clf, addShapes, currentSlice, sliceShapes, commitClassIds]); + + const handleApplyToVolume = useCallback(() => { + const sliceIndices = Array.from({ length: totalSliceCount }, (_, i) => i); + void clf.applyAcrossVolume(sliceIndices); + }, [clf, totalSliceCount]); + + const handleCancelVolumeApply = useCallback(async () => { + const jobId = clf.volumeApplyJob.jobId; + if (!jobId) return; + // Cooperative: the job stops at its next slice boundary (see DenoiseBakeModal + // for the same pattern against the same /api/export/cancel/{id} route). + await fetch(`${API_BASE}/api/export/cancel/${jobId}`, { method: 'POST' }); + }, [clf.volumeApplyJob.jobId]); + + // "Commit predicted shapes" no longer eagerly fetches + vectorizes every + // slice's commit.png into real Shape[] — that was the direct cause of the + // reported 297MB-draft/browser-OOM incident (139,004 shapes from one + // 690-slice volume-apply commit, 98.8% predicted-origin, all landing in + // the autosaved draft whether anyone ever looked at them or not). + // + // Instead this just records a lightweight pointer per slice + // (predictedRasterStore — {runId, classIds}, tens of bytes, NOT part of + // annotationStore/useDraftSync's autosave). The existing Predictions layer + // already renders a pointer's commit.png directly (usePixelClassifier's + // live-preview effect resolves it the same way it resolves a still-running + // job's per-slice result), so nothing is lost visually — a slice only ever + // becomes real, editable Shape[] when the user explicitly vectorizes it + // (see handleMakeSliceEditable below), one slice at a time. + // + // Slices that already have real shapes are left alone: creating a pointer + // for an already-annotated slice would show the run's raw (unedited) + // commit.png overlapping hand-drawn or already-vectorized content — the + // "make editable" path is what merges predicted-into-existing correctly + // (via preserveShapes), not this bulk commit step. + const handleCommitVolumeApply = useCallback(() => { + if (!sourceKey) return; + const result = clf.volumeApplyJob.result as { runs?: Record } | null; + const runs = result?.runs; + if (!runs || commitClassIds.length === 0) return; + const pointers: Record = {}; + for (const [sliceKey, runId] of Object.entries(runs)) { + if ((byImage[sourceKey]?.[sliceKey] ?? []).length > 0) continue; + pointers[sliceKey] = { runId, classIds: commitClassIds }; + } + if (Object.keys(pointers).length) setPredictedPointers(sourceKey, pointers); + clf.resetVolumeApplyJob(); + }, [sourceKey, clf, byImage, commitClassIds, setPredictedPointers]); + + // The only path that turns a predictedRasterStore pointer into real, + // editable Shape[] — one slice at a time, on explicit request. Reuses the + // exact same tracer the old eager-commit path used per slice, including + // `preserveShapes` so an already-partially-annotated slice merges rather + // than overwrites. + const [vectorizingSlice, setVectorizingSlice] = useState(false); + const handleMakeSliceEditable = useCallback(async () => { + if (!sourceKey) return; + const pointer = predictedPointers[sourceKey]?.[String(currentSlice)]; + if (!pointer) return; + setVectorizingSlice(true); + try { + const { data, width, height } = await loadLabelPng(ipredRunCommitUrl(pointer.runId)); + const existing = byImage[sourceKey]?.[String(currentSlice)] ?? []; + const shapes = labelMapToPolygonShapes(data, width, height, pointer.classIds, { + minRegion: 64, + preserveShapes: existing, + origin: 'predicted', + }); + if (shapes.length) addShapes(sourceKey, currentSlice, shapes); + clearPredictedSlice(sourceKey, String(currentSlice)); + } finally { + setVectorizingSlice(false); + } + }, [sourceKey, currentSlice, predictedPointers, byImage, addShapes, clearPredictedSlice]); + + // Push to Tiled and View in 3D are two INDEPENDENT actions (previously one + // combined "Push to Tiled + view in 3D" button that always navigated on + // success) — forcing a navigation to /volume right after a push meant that + // on a dataset with no volume pyramid built yet, you landed straight in + // "Build Volume" with no way back to Annotate short of the browser's own + // back button. Now pushing stays on this tab, and viewing in 3D is a + // separate click that works whether or not you just pushed (e.g. to look + // at a result pushed earlier in the session). + const maskSyncJob = useExportJob('mask-sync-view3d'); + const handleSyncToTiled = useCallback(() => { + if (!source || !kind || kind !== 'tiled') { + window.alert('Pushing masks to Tiled only works for Tiled sources.'); + return; + } + // predicted_slices carries any still-un-vectorized predicted regions + // (predictedRasterStore pointers) straight through — the backend fetches + // their commit.png directly from ipred and rasterizes it server-side, so + // "Push to Tiled" works correctly without first vectorizing every + // committed slice into Shape[] client-side (see the lazy-vectorization + // plan item). A slice already in `slices` (real shapes) doesn't need its + // pointer sent — build_mask_volumes already prefers real shapes anyway, + // but there's no reason to make the backend re-derive that here too. + const pointersForSource = predictedPointers[sourceKey!] ?? {}; + const predictedSlicesPayload = Object.fromEntries( + Object.entries(pointersForSource) + .filter(([sliceKey]) => !(byImage[sourceKey!]?.[sliceKey]?.length)) + .map(([sliceKey, p]) => [sliceKey, { run_id: p.runId, class_ids: p.classIds }]), + ); + maskSyncJob.startMaskSync({ + sources: [{ + kind, + source, + server_uri: serverUri ?? null, + slices: byImage[sourceKey!] ?? {}, + split_by_slice: splitBySlice[sourceKey!] ?? {}, + negative_slices: negativeSlices[sourceKey!] ?? [], + predicted_slices: predictedSlicesPayload, + }], + classes, + }); + }, [source, kind, serverUri, sourceKey, byImage, splitBySlice, negativeSlices, predictedPointers, classes, maskSyncJob]); + + const handleViewIn3D = useCallback(() => { + navigate('/volume?mask=fast'); + }, [navigate]); + + /** Isolate a Sampler-fitted intensity band in the 3D transfer function — + * see VolumePage's bandLo/bandHi query-param handling. */ + const handleSendBandTo3D = useCallback((lo: number, hi: number) => { + navigate(`/volume?bandLo=${lo}&bandHi=${hi}`); + }, [navigate]); + + const classLabelForId = useCallback( + (classId: number) => { + const c = classes.find((x) => x.classId === classId); + return c?.label?.trim() || `class ${classId}`; + }, + [classes], + ); + + const predictionClassColorById = useMemo(() => { + const m = new Map(); + for (const c of classes) m.set(c.classId, c.color); + return m; + }, [classes]); + // Load the previewed version's payload (cached) whenever the slider moves. useEffect(() => { if (previewVersion === null) { @@ -212,36 +536,180 @@ export default function AnnotatePage() { value={Math.round(fillOpacity * 100)} onChange={(v) => setFillOpacity(v / 100)} /> -
- -
- { setBrightness(0); setContrast(0); setLevelsLo(0); setLevelsHi(255); setColormap('gray'); setGamma(1); setClahe(false); setSharpen(false); }} - histogramBins={histogramBins} - levelsLo={levelsLo} - levelsHi={levelsHi} - onLevelsChange={(lo, hi) => { setLevelsLo(lo); setLevelsHi(hi); }} - onLevelsReset={() => { setLevelsLo(0); setLevelsHi(255); }} - colormap={colormap} - gamma={gamma} - onColormapChange={setColormap} - onGammaChange={setGamma} - clahe={clahe} - sharpen={sharpen} - onClaheChange={setClahe} - onSharpenChange={setSharpen} + -
+ -
- -
- -
+ + {stage === 'draw' && ( + <> + +
+ { setBrightness(0); setContrast(0); setLevelsLo(0); setLevelsHi(255); setColormap('gray'); setGamma(1); setClahe(false); setSharpen(false); setBlur(0); setUpscale(1); }} + histogramBins={histogramBins} + levelsLo={levelsLo} + levelsHi={levelsHi} + onLevelsChange={(lo, hi) => { setLevelsLo(lo); setLevelsHi(hi); }} + onLevelsReset={() => { setLevelsLo(0); setLevelsHi(255); }} + colormap={colormap} + gamma={gamma} + onColormapChange={setColormap} + onGammaChange={setGamma} + clahe={clahe} + sharpen={sharpen} + onClaheChange={setClahe} + onSharpenChange={setSharpen} + blur={blur} + onBlurChange={setBlur} + upscale={upscale} + onUpscaleChange={setUpscale} + maxUpscale={maxUpscale} + denoiseSlot={ + setBakeOpen(true) : undefined} + /> + } + /> + + + + )} + + {stage === 'assist' && ( + <> + + { void manifold.sample(); }} + onManifoldDismiss={manifold.dismiss} + manifoldRoiShapeCount={manifold.placementMask?.length ?? 0} + canCaptureManifoldRoi={selectedShapeIds.length > 0} + onCaptureManifoldRoi={() => { + const selected = sliceShapes.filter((s) => selectedShapeIds.includes(s.id)); + manifold.setPlacementMaskFromShapes(selected); + }} + onClearManifoldRoi={manifold.clearPlacementMask} + /> + + )} + + {stage === 'predict' && ( + 0} + training={clf.training} + predicting={clf.predicting} + params={clf.params} + onParamsChange={clf.setParams} + model={clf.model} + hasPrediction={!!clf.commitUrl} + predictCounts={clf.predictCounts} + probaClassIndex={clf.probaClassIndex} + activeProbaClassId={clf.activeProbaClassId} + activeProbaThreshold={clf.activeProbaThreshold} + classLabelForId={classLabelForId} + onCycleProbaClass={clf.cycleProbaClass} + onProbaThresholdChange={clf.setProbaThreshold} + error={clf.error} + onTrain={handleClfTrain} + onPredict={handleClfPredict} + onCommit={() => { void handleClfCommit(); }} + onDismiss={clf.dismiss} + commitLabel="Commit singletons" + featureSetupId={features.job?.setupId ?? null} + trainerId={clf.trainerId} + annotatedSliceCount={annotatedSliceCount} + trainAcrossSlices={trainAcrossSlices} + onTrainAcrossSlicesChange={setTrainAcrossSlices} + multiTraining={clf.multiTraining} + multiTrainProgress={{ done: clf.multiTrainJob.done, total: clf.multiTrainJob.total }} + totalSliceCount={totalSliceCount} + commitClassIds={commitClassIds} + onToggleCommitClassId={toggleCommitClassId} + volumeApplying={clf.volumeApplying} + volumeApplyProgress={{ done: clf.volumeApplyJob.done, total: clf.volumeApplyJob.total }} + volumeApplyResult={ + clf.volumeApplyJob.result + ? (() => { + const r = clf.volumeApplyJob.result as { + runs?: Record; + errors?: unknown[]; + cancelled?: boolean; + }; + return { + runCount: Object.keys(r.runs ?? {}).length, + errorCount: (r.errors ?? []).length, + cancelled: !!r.cancelled, + }; + })() + : null + } + onApplyToVolume={handleApplyToVolume} + onCommitVolumeApply={handleCommitVolumeApply} + onCancelVolumeApply={() => { void handleCancelVolumeApply(); }} + onDismissVolumeApply={clf.resetVolumeApplyJob} + hasPredictedPointerOnCurrentSlice={ + !!sourceKey && !!predictedPointers[sourceKey]?.[String(currentSlice)] + } + vectorizingSlice={vectorizingSlice} + onMakeSliceEditable={() => { void handleMakeSliceEditable(); }} + hasAnyPredictedPointers={ + !!sourceKey && Object.keys(predictedPointers[sourceKey] ?? {}).length > 0 + } + onTrainDeepModel={() => navigate('/train')} + onSyncToTiled={handleSyncToTiled} + syncingToTiled={maskSyncJob.state.status === 'running'} + syncedToTiled={maskSyncJob.state.status === 'done'} + syncToTiledError={maskSyncJob.state.status === 'error' ? maskSyncJob.state.error : null} + onViewIn3D={handleViewIn3D} + /> + )} {/* Save button + status */}
@@ -314,13 +782,28 @@ export default function AnnotatePage() { gamma={gamma} clahe={clahe} sharpen={sharpen} + blur={blur} + upscale={upscale} onHistogram={setHistogramBins} + onSamplerFit={setSamplerFit} + onBlurChange={setBlur} activeClassId={activeClassId} activeBrushShapeId={activeBrushShapeId} onNewBrushInstance={setActiveBrushShapeId} previewShapes={previewShapes} previewClasses={previewPayload?.classes ?? null} focusRegion={focusRegion} + featureChannelUrl={features.channelUrl} + probaOverlayUrl={clf.probaUrl} + clfCommitUrl={clf.commitUrl} + clfStatusUrl={clf.statusUrl} + predictionClassColorById={predictionClassColorById} + manifoldHeatmapUrl={manifold.heatmapUrl} + manifoldHeatmapOpacity={manifold.heatmapOpacity} + manifoldShowHeatmap={manifold.showHeatmap} + manifoldMarkers={manifold.points} + manifoldShowMarkers={manifold.showMarkers} + manifoldBoxSize={manifold.meta?.boxSize ?? 64} /> {previewVersion !== null && ( )} + {perfOn && }
@@ -368,6 +852,17 @@ export default function AnnotatePage() { onClose={() => setShowVersionHistory(false)} /> )} + {source && ( + setBakeOpen(false)} + source={source} + serverUri={serverUri} + denoise={denoise} + nSlices={meta?.nSlices ?? 1} + methodLabel={denoise.method} + /> + )} ); } diff --git a/frontend/src/app/pages/BrowsePage.test.tsx b/frontend/src/app/pages/BrowsePage.test.tsx new file mode 100644 index 0000000..1619126 --- /dev/null +++ b/frontend/src/app/pages/BrowsePage.test.tsx @@ -0,0 +1,232 @@ +/** + * BrowsePage — orchestration tests: connection banner, kind switch (tiled vs + * local), navigation to /connect, and server-list wiring into ColumnBrowser. + * + * ColumnBrowser and LocalSampleBrowser are mocked out: they are separate, + * already-tested units with their own fetch-driven state machines (see + * ColumnBrowser.test.tsx / LocalSampleBrowser.test.tsx). BrowsePage only needs + * to exercise the props it wires into them and its own connection/navigation + * logic, so useConnectionStore and useNavigate (via a real MemoryRouter) stay + * real. + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { MemoryRouter, Routes, Route } from 'react-router'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import BrowsePage from './BrowsePage'; +import { useConnectionStore } from '@/stores/connectionStore'; + +const openLocalFile = vi.fn(); +vi.mock('@/hooks/useOpenInAnnotate', () => ({ + useOpenInAnnotate: () => ({ openTiledArray: vi.fn(), openLocalFile }), +})); + +vi.mock('@/components/Browse/ColumnBrowser', () => ({ + default: (props: any) => ( +
+ serverUri:{props.serverUri} + containerPath:{String(props.containerPath)} + focusPath:{String(props.focusPath)} + selectedServerUri:{props.selectedServerUri} + servers:{props.servers.map((s: any) => s.name).join(',')} + + +
+ ), +})); + +vi.mock('@/components/Browse/LocalSampleBrowser', () => ({ + default: (props: any) => ( +
+ root:{props.root} + rel:{props.rel} + +
+ ), +})); + +const initialConnectionState = useConnectionStore.getState(); + +function jsonResponse(body: unknown, ok = true) { + return { + ok, + status: ok ? 200 : 500, + json: async () => body, + text: async () => (typeof body === 'string' ? body : JSON.stringify(body)), + } as Response; +} + +const SERVERS = [ + { name: 'Server A', uri: 'http://a', has_api_key: false }, + { name: 'Server B', uri: 'http://b', has_api_key: false }, +]; + +function makeFetchMock() { + return vi.fn(async (input: RequestInfo | URL) => { + const url = String(input); + if (url.includes('/api/config/servers')) return jsonResponse(SERVERS); + return jsonResponse({}, false); + }); +} + +function renderPage() { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render( + + + + } /> + CONNECT PAGE
} /> + + + , + ); +} + +beforeEach(() => { + useConnectionStore.setState(initialConnectionState, true); +}); + +afterEach(() => { + cleanup(); + vi.clearAllMocks(); +}); + +describe('BrowsePage', () => { + it('shows a not-connected prompt and navigates to Connect', async () => { + global.fetch = makeFetchMock(); + const user = userEvent.setup(); + renderPage(); + + expect(screen.getByText('No dataset connected.')).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: 'Go to Connect' })); + expect(await screen.findByText('CONNECT PAGE')).toBeInTheDocument(); + }); + + it('renders the tiled ColumnBrowser with connection details when kind is tiled', async () => { + global.fetch = makeFetchMock(); + useConnectionStore.getState().setConnection({ + kind: 'tiled', + serverUri: 'http://a', + browseContainerPath: 'browse/foo', + browseFocusPath: 'browse/foo/bar', + label: 'Server A', + sampleCount: 5, + }); + renderPage(); + + expect(screen.getByText('Server A')).toBeInTheDocument(); + expect(screen.getByText(/5 samples/)).toBeInTheDocument(); + + const columnBrowser = await screen.findByTestId('column-browser'); + expect(columnBrowser).toHaveTextContent('serverUri:http://a'); + expect(columnBrowser).toHaveTextContent('containerPath:browse/foo'); + expect(columnBrowser).toHaveTextContent('focusPath:browse/foo/bar'); + expect(columnBrowser).toHaveTextContent('selectedServerUri:http://a'); + await screen.findByText('servers:Server A,Server B'); + expect(screen.queryByTestId('local-browser')).not.toBeInTheDocument(); + }); + + it('renders the LocalSampleBrowser with root/rel when kind is local', async () => { + global.fetch = makeFetchMock(); + useConnectionStore.getState().setConnection({ + kind: 'local', + localRoot: '/data', + localRel: 'sub', + label: '/data/sub', + sampleCount: 3, + }); + renderPage(); + + const localBrowser = await screen.findByTestId('local-browser'); + expect(localBrowser).toHaveTextContent('root:/data'); + expect(localBrowser).toHaveTextContent('rel:sub'); + expect(screen.queryByTestId('column-browser')).not.toBeInTheDocument(); + // Servers query is disabled (enabled: kind === 'tiled') for local connections. + expect(global.fetch).not.toHaveBeenCalled(); + }); + + it('shows singular sample count text when there is exactly one sample', async () => { + global.fetch = makeFetchMock(); + useConnectionStore.getState().setConnection({ + kind: 'local', + localRoot: '/data', + localRel: '', + label: '/data', + sampleCount: 1, + }); + renderPage(); + expect(await screen.findByText(/1 sample$/)).toBeInTheDocument(); + }); + + it('clicking "Change" navigates to /connect', async () => { + global.fetch = makeFetchMock(); + useConnectionStore.getState().setConnection({ + kind: 'local', + localRoot: '/data', + localRel: '', + label: '/data', + sampleCount: 1, + }); + const user = userEvent.setup(); + renderPage(); + + await user.click(screen.getByRole('button', { name: 'Change' })); + expect(await screen.findByText('CONNECT PAGE')).toBeInTheDocument(); + }); + + it('opening a local sample calls openLocalFile via useOpenInAnnotate', async () => { + global.fetch = makeFetchMock(); + useConnectionStore.getState().setConnection({ + kind: 'local', + localRoot: '/data', + localRel: '', + label: '/data', + sampleCount: 1, + }); + const user = userEvent.setup(); + renderPage(); + + await user.click(await screen.findByText('open-local')); + expect(openLocalFile).toHaveBeenCalledWith('rel/path.tif'); + }); + + it('changing the server in ColumnBrowser updates the connection store', async () => { + global.fetch = makeFetchMock(); + useConnectionStore.getState().setConnection({ + kind: 'tiled', + serverUri: 'http://a', + label: 'Server A', + sampleCount: 5, + }); + const user = userEvent.setup(); + renderPage(); + + await screen.findByTestId('column-browser'); + await user.click(screen.getByText('change-server')); + + expect(useConnectionStore.getState()).toMatchObject({ + kind: 'tiled', + serverUri: 'http://b', + label: 'Server B', + }); + }); + + it('changing the annotation filter is reflected back into ColumnBrowser', async () => { + global.fetch = makeFetchMock(); + useConnectionStore.getState().setConnection({ + kind: 'tiled', + serverUri: 'http://a', + label: 'Server A', + sampleCount: 5, + }); + const user = userEvent.setup(); + renderPage(); + + await screen.findByTestId('column-browser'); + await user.click(screen.getByText('change-annotation')); + // Re-rendering with the new filter doesn't crash and the mock still shows. + expect(await screen.findByTestId('column-browser')).toBeInTheDocument(); + }); +}); diff --git a/frontend/src/app/pages/ConnectPage.test.tsx b/frontend/src/app/pages/ConnectPage.test.tsx new file mode 100644 index 0000000..2cef527 --- /dev/null +++ b/frontend/src/app/pages/ConnectPage.test.tsx @@ -0,0 +1,328 @@ +/** + * ConnectPage — component tests for server selection, connect/verify flows + * (Tiled + Local), navigation, and the optional dataset-container picker. + * + * IngestDropzone/ZarrLoader and useOpenInAnnotate are mocked: they are + * separate, already-complex units with their own fetch/upload logic, and + * ConnectPage only needs to exercise the callbacks it wires into them + * (browseIngested / annotateIngested / annotateZarr). + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen, waitFor, within } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { MemoryRouter, Routes, Route } from 'react-router'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import ConnectPage from './ConnectPage'; +import { useConnectionStore } from '@/stores/connectionStore'; + +const openTiledArray = vi.fn(); + +vi.mock('@/hooks/useOpenInAnnotate', () => ({ + useOpenInAnnotate: () => ({ openTiledArray }), +})); + +vi.mock('@/components/Ingest/IngestDropzone', () => ({ + default: (props: { onBrowse?: (c: string, n: number) => void; onAnnotate?: (c: string, k: string) => void }) => ( +
+ + +
+ ), +})); + +vi.mock('@/components/Ingest/ZarrLoader', () => ({ + default: (props: { onAnnotate?: (p: string) => void }) => ( +
+ +
+ ), +})); + +const initialConnectionState = useConnectionStore.getState(); + +function jsonResponse(body: unknown, ok = true) { + return { + ok, + json: async () => body, + text: async () => (typeof body === 'string' ? body : JSON.stringify(body)), + } as Response; +} + +const SERVERS = [ + { name: 'Local Tiled', uri: 'http://localhost:8000/api', has_api_key: false }, + { name: 'Remote Tiled', uri: 'http://remote:8000/api', has_api_key: true }, +]; + +/** Default fetch router; tests override specific endpoints as needed. */ +function makeFetchMock(overrides: Partial Response | Promise>> = {}) { + return vi.fn(async (input: RequestInfo | URL) => { + const url = String(input); + if (overrides.servers && url.includes('/api/config/servers')) return overrides.servers(url); + if (url.includes('/api/config/servers')) return jsonResponse(SERVERS); + + if (overrides.tiledList && url.includes('/api/tiled/list')) return overrides.tiledList(url); + if (url.includes('/api/tiled/list')) return jsonResponse([]); + + if (overrides.localList && url.includes('/api/local/list')) return overrides.localList(url); + if (url.includes('/api/local/list')) return jsonResponse([]); + + if (overrides.summary && url.includes('/api/connect/summary')) return overrides.summary(url); + if (url.includes('/api/connect/summary')) return jsonResponse({ kind: 'tiled', server_uri: SERVERS[0].uri, label: 'Local Tiled', sample_count: 5 }); + + throw new Error(`Unhandled fetch: ${url}`); + }); +} + +function renderPage() { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render( + + + + } /> + BROWSE PAGE
} /> + + + , + ); +} + +describe('ConnectPage', () => { + beforeEach(() => { + useConnectionStore.setState(initialConnectionState, true); + }); + + afterEach(() => { + cleanup(); + vi.clearAllMocks(); + }); + + it('shows a loading state before the server list arrives, then lists servers', async () => { + global.fetch = makeFetchMock(); + renderPage(); + + expect(screen.getByText(/Loading servers…/)).toBeInTheDocument(); + + await screen.findByRole('combobox', { name: 'Tiled server' }); + const select = screen.getByRole('combobox', { name: 'Tiled server' }) as HTMLSelectElement; + expect(within(select).getByText(/Local Tiled/)).toBeInTheDocument(); + expect(within(select).getByText(/Remote Tiled/)).toBeInTheDocument(); + // Auto-selects the first server once the list loads. + await waitFor(() => expect(select.value).toBe(SERVERS[0].uri)); + }); + + it('connects to the selected Tiled server (verify) and reveals "Go to Browse" without navigating', async () => { + global.fetch = makeFetchMock(); + const user = userEvent.setup(); + renderPage(); + + await screen.findByRole('combobox', { name: 'Tiled server' }); + await user.click(screen.getByRole('button', { name: /Connect/ })); + + await screen.findByText(/Connected — 5 samples found/); + expect(useConnectionStore.getState()).toMatchObject({ + kind: 'tiled', + serverUri: SERVERS[0].uri, + label: 'Local Tiled', + sampleCount: 5, + }); + + // Verify does not navigate away; it reveals a separate action instead. + expect(screen.queryByText('BROWSE PAGE')).not.toBeInTheDocument(); + const goToBrowse = await screen.findByRole('button', { name: /Go to Browse/ }); + + await user.click(goToBrowse); + expect(await screen.findByText('BROWSE PAGE')).toBeInTheDocument(); + }); + + it('shows a failure message when connecting to Tiled fails', async () => { + global.fetch = makeFetchMock({ + summary: async () => jsonResponse('server unreachable', false), + }); + const user = userEvent.setup(); + renderPage(); + + await screen.findByRole('combobox', { name: 'Tiled server' }); + await user.click(screen.getByRole('button', { name: /Connect/ })); + + await screen.findByText(/Failed: Error: server unreachable/); + expect(screen.queryByRole('button', { name: /Go to Browse/ })).not.toBeInTheDocument(); + expect(useConnectionStore.getState().kind).toBeNull(); + }); + + it('switching Tiled server resets the verified/connected state', async () => { + global.fetch = makeFetchMock(); + const user = userEvent.setup(); + renderPage(); + + await screen.findByRole('combobox', { name: 'Tiled server' }); + await user.click(screen.getByRole('button', { name: /Connect/ })); + await screen.findByText(/Connected —/); + + const select = screen.getByRole('combobox', { name: 'Tiled server' }); + await user.selectOptions(select, SERVERS[1].uri); + + expect(screen.queryByRole('button', { name: /Go to Browse/ })).not.toBeInTheDocument(); + expect(screen.queryByText(/Connected —/)).not.toBeInTheDocument(); + }); + + it('lets you browse into a sub-container and connects with that container selected', async () => { + global.fetch = makeFetchMock({ + tiledList: async (url) => { + if (url.includes('path=sub')) return jsonResponse([]); + return jsonResponse([{ name: 'sub', path: 'sub', is_dir: true, is_array: false }]); + }, + summary: async (url) => { + expect(url).toContain('container_path=sub'); + return jsonResponse({ kind: 'tiled', server_uri: SERVERS[0].uri, label: 'Local Tiled', sample_count: 2 }); + }, + }); + const user = userEvent.setup(); + renderPage(); + + await screen.findByRole('combobox', { name: 'Tiled server' }); + await user.click(screen.getByRole('button', { name: /Dataset to view/ })); + await user.click(await screen.findByText('sub')); + await user.click(await screen.findByText('Browse "sub"')); + expect(screen.getByText(/Browse will show/)).toHaveTextContent('sub'); + + await user.click(screen.getByRole('button', { name: /Connect/ })); + await screen.findByText(/Connected — 2 samples found/); + expect(useConnectionStore.getState().browseContainerPath).toBe('sub'); + }); + + it('local mode: grants a root, browses folders, selects one, and connects (navigates to Browse)', async () => { + global.fetch = makeFetchMock({ + localList: async (url) => { + if (url.includes('rel=data')) return jsonResponse([]); + return jsonResponse([{ name: 'data', path: 'data', is_dir: true, size: null }]); + }, + summary: async () => jsonResponse({ kind: 'local', server_uri: null, label: '/root/data', sample_count: 7 }), + }); + const user = userEvent.setup(); + renderPage(); + + await user.click(screen.getByRole('button', { name: /Local Folder/ })); + await user.type(screen.getByPlaceholderText('/absolute/path/to/data'), '/root'); + await user.click(screen.getByRole('button', { name: 'Grant' })); + + await screen.findByText('data'); + await user.click(screen.getByText('data')); + await user.click(screen.getByText(/Use "data" as dataset folder/)); + + const connectBtn = screen.getByRole('button', { name: 'Connect' }); + expect(connectBtn).not.toBeDisabled(); + await user.click(connectBtn); + + await screen.findByText(/Connected — 7 samples found/); + expect(await screen.findByText('BROWSE PAGE')).toBeInTheDocument(); + expect(useConnectionStore.getState()).toMatchObject({ + kind: 'local', + localRoot: '/root', + localRel: 'data', + sampleCount: 7, + }); + }); + + it('local mode: Connect is disabled until a root is granted and a folder is selected', async () => { + // Note: selecting "this root" itself sets selectedFolder to '' (the + // relative path of the root), which is falsy in the component's + // `canConnect` check — so Connect stays disabled unless an actual + // sub-folder is chosen. This test documents that real behavior. + global.fetch = makeFetchMock({ + localList: async (url) => { + if (url.includes('rel=data')) return jsonResponse([]); + return jsonResponse([{ name: 'data', path: 'data', is_dir: true, size: null }]); + }, + }); + const user = userEvent.setup(); + renderPage(); + + await user.click(screen.getByRole('button', { name: /Local Folder/ })); + expect(screen.getByRole('button', { name: 'Connect' })).toBeDisabled(); + + await user.type(screen.getByPlaceholderText('/absolute/path/to/data'), '/root'); + await user.click(screen.getByRole('button', { name: 'Grant' })); + await screen.findByText('data'); + // Root granted but no folder explicitly selected yet. + expect(screen.getByRole('button', { name: 'Connect' })).toBeDisabled(); + // Selecting "this root" does not count as a folder selection either. + await user.click(screen.getByText(/Use "this root" as dataset folder/)); + expect(screen.getByRole('button', { name: 'Connect' })).toBeDisabled(); + + await user.click(screen.getByText('data')); + await user.click(screen.getByText(/Use "data" as dataset folder/)); + expect(screen.getByRole('button', { name: 'Connect' })).not.toBeDisabled(); + }); + + it('local mode: shows a failure message when connect fails', async () => { + global.fetch = makeFetchMock({ + localList: async (url) => { + if (url.includes('rel=data')) return jsonResponse([]); + return jsonResponse([{ name: 'data', path: 'data', is_dir: true, size: null }]); + }, + summary: async () => jsonResponse('bad path', false), + }); + const user = userEvent.setup(); + renderPage(); + + await user.click(screen.getByRole('button', { name: /Local Folder/ })); + await user.type(screen.getByPlaceholderText('/absolute/path/to/data'), '/root'); + await user.click(screen.getByRole('button', { name: 'Grant' })); + await screen.findByText('data'); + await user.click(screen.getByText('data')); + await user.click(screen.getByText(/Use "data" as dataset folder/)); + await user.click(screen.getByRole('button', { name: 'Connect' })); + + await screen.findByText(/Failed: Error: bad path/); + }); + + it('browsing an ingested dataset sets the connection (no container, focus on the dataset) and navigates', async () => { + global.fetch = makeFetchMock(); + const user = userEvent.setup(); + renderPage(); + + await screen.findByRole('combobox', { name: 'Tiled server' }); + await user.click(screen.getByText('mock-ingest-browse')); + + expect(await screen.findByText('BROWSE PAGE')).toBeInTheDocument(); + expect(useConnectionStore.getState()).toMatchObject({ + kind: 'tiled', + browseContainerPath: null, + browseFocusPath: 'browse/foo', + }); + }); + + it('annotating an ingested sample sets the connection (stays put) and opens it in Annotate', async () => { + global.fetch = makeFetchMock(); + const user = userEvent.setup(); + renderPage(); + + await screen.findByRole('combobox', { name: 'Tiled server' }); + await user.click(screen.getByText('mock-ingest-annotate')); + + // annotateIngested does not navigate to Browse. + expect(screen.queryByText('BROWSE PAGE')).not.toBeInTheDocument(); + expect(useConnectionStore.getState()).toMatchObject({ + kind: 'tiled', + browseContainerPath: 'browse/foo', + }); + expect(openTiledArray).toHaveBeenCalledWith('browse/foo/img1', SERVERS[0].uri); + }); + + it('annotating a registered Zarr level derives the container path and opens it in Annotate', async () => { + global.fetch = makeFetchMock(); + const user = userEvent.setup(); + renderPage(); + + await screen.findByRole('combobox', { name: 'Tiled server' }); + await user.click(screen.getByText('mock-zarr-annotate')); + + expect(screen.queryByText('BROWSE PAGE')).not.toBeInTheDocument(); + expect(useConnectionStore.getState()).toMatchObject({ + kind: 'tiled', + browseContainerPath: 'browse/vol', + }); + expect(openTiledArray).toHaveBeenCalledWith('browse/vol/multiscale/level_0/array', SERVERS[0].uri); + }); +}); diff --git a/frontend/src/app/pages/ConnectPage.tsx b/frontend/src/app/pages/ConnectPage.tsx index 07e0abb..08c0de6 100644 --- a/frontend/src/app/pages/ConnectPage.tsx +++ b/frontend/src/app/pages/ConnectPage.tsx @@ -27,6 +27,7 @@ import { API_BASE } from '@/config'; import { useConnectionStore } from '@/stores/connectionStore'; import { useOpenInAnnotate } from '@/hooks/useOpenInAnnotate'; import IngestDropzone from '@/components/Ingest/IngestDropzone'; +import ZarrLoader from '@/components/Ingest/ZarrLoader'; interface ServerInfo { name: string; @@ -147,6 +148,13 @@ export default function ConnectPage() { */ const browseIngested = (containerPath: string) => connectTiled(null, true, containerPath); + /** Open a registered Zarr level (a full Tiled path) straight in Annotate. */ + const annotateZarr = (tiledPath: string) => { + const containerPath = tiledPath.split('/').slice(0, -3).join('/') || 'browse'; + connectTiled(containerPath, false); + void openTiledArray(tiledPath, selectedServerUri); + }; + // Jump straight from ingest to the Annotate tab for the first uploaded sample. const annotateIngested = (containerPath: string, firstKey: string) => { connectTiled(containerPath, false); // set connection context, don't navigate to Browse @@ -367,11 +375,28 @@ export default function ConnectPage() {

Load / Ingest Datasets

{selectedServerUri ? ( - + <> + + + {/* Zarr volumes are too large to upload; they are registered in + place instead, so this sits beside the dropzone rather than + inside it. */} +
+
+ +

Load a Zarr volume

+
+ +
+ ) : (

Select a server above to ingest data.

)} diff --git a/frontend/src/app/pages/TrainPage.tsx b/frontend/src/app/pages/TrainPage.tsx new file mode 100644 index 0000000..042832e --- /dev/null +++ b/frontend/src/app/pages/TrainPage.tsx @@ -0,0 +1,335 @@ +/** + * TrainPage — fine-tune a dlsia TUNet model on this session's annotated + * samples, then run inference with any saved run: preview the predicted + * masks, import them as editable annotations, or push them to Tiled. + * + * dlsia TUNet only — DINOv3 LoRA is deferred (see Phase 5.5 in the + * integration plan), so there is no model-family picker here. + */ +import { useEffect, useMemo, useRef, useState } from 'react'; +import { useSearchParams } from 'react-router'; +import { v4 as uuidv4 } from 'uuid'; +import { Brain } from '@phosphor-icons/react'; +import { API_BASE } from '@/config'; +import { useDatasetStore } from '@/stores/datasetStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { useClassStore } from '@/stores/classStore'; +import { buildSourceKey } from '@/lib/sourceKey'; +import { useTrainCapability } from '@/hooks/useTrainCapability'; +import { useTrainRuns } from '@/hooks/useTrainRuns'; +import { useExportJob } from '@/hooks/useExportJob'; +import { useDraftSync } from '@/hooks/useDraftSync'; +import { useImageSlice } from '@/hooks/useImageSlice'; +import { gatherTrainingSources, listAnnotatedSourceKeys } from '@/lib/gatherTrainingSources'; +import { buildModelConfig, trainConfigSignature, validateTrainConfig } from '@/lib/trainModelConfig'; +import { trainDenoisePayload } from '@/lib/trainDenoiseOption'; +import { isSegmentationRun } from '@/lib/runCompatibility'; +import { remapPredictedShapes, type RunClass } from '@/lib/importPredictions'; +import type { Shape } from '@/stores/annotationStore'; +import CapabilityBanner from '@/components/train/CapabilityBanner'; +import TrainingDataPanel from '@/components/train/TrainingDataPanel'; +import HyperparamsPanel, { type HyperparamsState } from '@/components/train/HyperparamsPanel'; +import JobProgressBar from '@/components/train/JobProgressBar'; +import RunsPanel from '@/components/train/RunsPanel'; +import TrainDenoiseToggle from '@/components/train/TrainDenoiseToggle'; +import InferencePanel from '@/components/train/InferencePanel'; + +const DEFAULT_HYPERPARAMS: HyperparamsState = { + // 512: tile count scales with 1/size², so a bigger window means far fewer + // forward passes per slice and more context in each. + epochs: 60, lr: 1e-3, batch_size: 4, image_size: 512, flip_augment: true, tiling: true, + depth: 4, base_channels: 8, growth_rate: 1.5, +}; + +export default function TrainPage() { + const { source, kind, serverUri, meta, currentSlice, renderOpts } = useDatasetStore(); + const { byImage, splitBySlice, negativeSlices, addShapes } = useAnnotationStore(); + const { classes, setClasses } = useClassStore(); + + // The Annotate tab's denoise setting (Phase 2's DenoisePanel) — only ever + // read here, and only when "Train on denoised input" is ticked below. + const denoise = useDatasetStore((s) => s.denoise); + + const { capability } = useTrainCapability(); + const { runs: allRuns, invalidate: refreshRuns, deleteRun } = useTrainRuns(); + // This tab's Runs/Inference card is segmentation-only (it imports predicted + // SHAPES) — a saved denoiser run has no class list to remap predictions + // against, so it's excluded here. + const runs = useMemo(() => allRuns.filter(isSegmentationRun), [allRuns]); + // persistKey: survives switching to another tab and back mid-job — see + // useExportJob's docstring. Fixed keys (not scoped to a sample) are correct + // here: once submitted, a job is already bound to its own run_id server-side, + // independent of anything that changes in this page afterward, and only one + // training/probe job can run at a time regardless. + const { state: trainJob, startJob: startTrainJob } = useExportJob('train:start'); + const { state: probeJob, startJob: startProbeJob, reset: resetProbeJob } = useExportJob('train:probe'); + // Set at handleEstimateBatch time, compared against the current config when + // the probe finishes — see the adoption effect below. + const probeConfigSignature = useRef(null); + + const [selectedKeys, setSelectedKeys] = useState>(new Set()); + const [hyperparams, setHyperparams] = useState(DEFAULT_HYPERPARAMS); + const [runName, setRunName] = useState(''); + // Off by default: this one changes what the model LEARNS, not just what's on + // screen, so it has to be asked for explicitly. + const [trainOnDenoised, setTrainOnDenoised] = useState(false); + const [selectedRunId, setSelectedRunId] = useState(null); + const [dataError, setDataError] = useState(null); + + // "Run inference" from Annotate/Browse jumps straight to the Runs/Inference + // card via ?focus=infer, since that's the section they actually asked for. + const [searchParams] = useSearchParams(); + const inferenceSectionRef = useRef(null); + useEffect(() => { + if (searchParams.get('focus') === 'infer') { + inferenceSectionRef.current?.scrollIntoView({ behavior: 'smooth', block: 'start' }); + } + // eslint-disable-next-line react-hooks/exhaustive-deps + }, []); + + const isTiledSource = kind === 'tiled'; + const sourceKey = source && kind ? buildSourceKey(kind as 'tiled' | 'local', source, serverUri) : null; + // "Import as annotations" below writes shapes (and possibly new classes) + // straight into the store for this sourceKey — without this, nothing here + // autosaves them: AnnotatePage is the only OTHER place that mounts + // useDraftSync, and routes render exactly one page at a time, so navigating + // to Annotate afterward would establish a "clean" baseline that already + // includes the import, and the debounced PUT would never fire. This closes + // that gap the same way every edit in Annotate is already persisted. + useDraftSync(sourceKey); + const candidates = useMemo(() => listAnnotatedSourceKeys(byImage), [byImage]); + + // Auto-select the currently open sample the first time it becomes available. + useEffect(() => { + if (sourceKey && candidates.some((c) => c.sourceKey === sourceKey)) { + setSelectedKeys((prev) => (prev.size === 0 ? new Set([sourceKey]) : prev)); + } + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [sourceKey, candidates.length]); + + const toggleKey = (key: string) => { + setSelectedKeys((prev) => { + const next = new Set(prev); + if (next.has(key)) next.delete(key); else next.add(key); + return next; + }); + }; + + const handleHyperparamsChange = (updates: Partial) => + setHyperparams((prev) => ({ ...prev, ...updates })); + + const handleStartTraining = () => { + setDataError(null); + let sources; + try { + sources = gatherTrainingSources(Array.from(selectedKeys), byImage, splitBySlice, negativeSlices); + } catch (err) { + setDataError(err instanceof Error ? err.message : String(err)); + return; + } + // Catch these here; the server would otherwise reject them with a validation + // error that names the schema field rather than the visible control. + const configError = validateTrainConfig(hyperparams); + if (configError) { + setDataError(configError); + return; + } + const model = buildModelConfig(hyperparams); + + void startTrainJob('/api/train/start', { + sources, + classes, + model, + run_name: runName.trim() || null, + // Adds nothing at all when the checkbox is off, so an un-denoised run is + // byte-identical to what this page sent before the option existed. + ...trainDenoisePayload(trainOnDenoised, denoise), + }); + }; + + // Refresh the runs list once a training job finishes. + useEffect(() => { + if (trainJob.status === 'done') refreshRuns(); + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [trainJob.status]); + + // Shared registry — same cancel route every background job in this app uses. + const handleCancelTrain = () => { + if (trainJob.jobId) void fetch(`${API_BASE}/api/export/cancel/${trainJob.jobId}`, { method: 'POST' }); + }; + + const handleCancelProbe = () => { + if (probeJob.jobId) void fetch(`${API_BASE}/api/export/cancel/${probeJob.jobId}`, { method: 'POST' }); + }; + + // Both Start and Estimate hit the same server-side gates (ML_LOCK, torch, dlsia) + // and must agree on when they're blocked — checking a different subset per + // button just means the one that under-checks fires a request the server was + // always going to 409/503 back. `capability.busy` is the server's own ML_LOCK + // state (catches jobs from other tabs/clients); the two job statuses cover this + // tab's own in-flight request before that poll has caught up. + const trainOrProbeRunning = trainJob.status === 'running' || probeJob.status === 'running'; + const deviceBusy = trainOrProbeRunning || capability.busy; + const needsDlsia = !capability.dlsia.available; + const startDisabled = deviceBusy || !capability.torch_available || needsDlsia; + const estimateDisabledReason = trainOrProbeRunning + ? 'A training run or another batch-size probe is already using the device' + : capability.busy + ? 'Another training or inference job is using the device' + : !capability.torch_available + ? 'Estimating is unavailable: torch is not installed on this server' + : needsDlsia + ? 'Estimating is unavailable: dlsia is not installed on this server' + : null; + + /** Measure the largest batch size this model config fits, then adopt it. + * Reuses the shared export-job plumbing, so progress lines stream in as the + * probe steps 1, 2, 4, 8… (see backend/batch_probe.py). */ + const handleEstimateBatch = () => { + setDataError(null); + const configError = validateTrainConfig(hyperparams); + if (configError) { + setDataError(configError); + return; + } + const model = buildModelConfig(hyperparams); + probeConfigSignature.current = trainConfigSignature(hyperparams); + void startProbeJob('/api/train/estimate-batch', { + model, + // Head width barely moves memory, so the probe works before any classes exist. + n_classes: Math.max(1, classes.length || 2), + }); + }; + + // Adopt the measured value once the probe finishes — but only if the config + // it measured is still the one that would be submitted. The user is free to + // change patch-size/tiling while a probe is running; if they did, this + // result describes a different memory footprint and must not silently + // overwrite batch_size (a cancelled probe's partial measurement is excluded + // the same way — see batch_probe.py's `cancelled` flag). + useEffect(() => { + const suggested = probeJob.result?.suggested_batch_size; + const cancelled = probeJob.result?.cancelled === true; + const currentSignature = trainConfigSignature(hyperparams); + if ( + probeJob.status === 'done' && typeof suggested === 'number' && !cancelled + && probeConfigSignature.current === currentSignature + ) { + setHyperparams((prev) => ({ ...prev, batch_size: suggested })); + } + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [probeJob.status, probeJob.result]); + + /** Latest probe line: its own error, the streaming log tail, or the summary. */ + const batchEstimateNote = useMemo(() => { + if (probeJob.status === 'error') return probeJob.error ?? 'Could not estimate a batch size.'; + if (probeJob.status === 'running') return probeJob.log.at(-1) ?? 'Loading the model…'; + if (probeJob.status === 'done') { + const note = probeJob.result?.note; + return typeof note === 'string' ? note : null; + } + return null; + }, [probeJob.status, probeJob.error, probeJob.log, probeJob.result]); + + const handleDeleteRun = async (runId: string) => { + setDataError(null); + const ok = await deleteRun(runId); + if (!ok) { + setDataError('Failed to delete run.'); + return; + } + if (selectedRunId === runId) setSelectedRunId(null); + }; + + const sliceQuery = useImageSlice(source, kind, currentSlice, renderOpts, serverUri); + const baseImageUrl = sliceQuery.data ?? null; + + const handleImportPredictions = (runClasses: RunClass[], slices: Record) => { + if (!sourceKey) return; + const remapped = remapPredictedShapes(runClasses, classes, slices as Record); + setClasses(remapped.classes); + for (const [sliceKey, shapes] of Object.entries(remapped.slices)) { + // Predicted shape ids are deterministic per run+slice (see infer_jobs.py), + // so importing the same run onto the same slice twice would otherwise + // collide with the previous import's ids — corrupting the shape list + // (duplicate React/Konva keys) and making the canvas unresponsive. + const freshIds = shapes.map((shape) => ({ ...shape, id: uuidv4() })); + addShapes(sourceKey, Number(sliceKey), freshIds); + } + }; + + return ( +
+
+ +

Train

+
+ + + +
+ + + + + + + {dataError &&

{dataError}

} + +
+ + {trainJob.status === 'running' && ( + + )} +
+ + + {trainJob.status === 'done' && ( +

+ Saved run {String(trainJob.result?.run_id ?? '')} + {typeof trainJob.result?.val_miou === 'number' && ` — val mIoU ${trainJob.result.val_miou.toFixed(3)}`}. +

+ )} +
+ +
+ + +
+
+ ); +} diff --git a/frontend/src/app/pages/VolumePage.tsx b/frontend/src/app/pages/VolumePage.tsx new file mode 100644 index 0000000..c5d468f --- /dev/null +++ b/frontend/src/app/pages/VolumePage.tsx @@ -0,0 +1,258 @@ +/** + * VolumePage — 3D view of the open dataset, rendered straight from Tiled. + * + * The volume is streamed by the vendored WebGPU renderer directly off Tiled's + * `/zarr/v2` router (see `lib/zarrUrl.ts`) — there is no export step and no + * downsampled-volume endpoint in between, so whatever is open in Annotate is + * what renders here. Rendering controls (transfer function, crop box, slice + * planes, lighting) come from the renderer's own HUD, docked on the right. + * + * Which node actually holds the volume is asked of the backend rather than + * guessed from the path: a registered Zarr volume is one already, a TIFF stack's + * lives in a `__volume` sidecar, and a stack nobody has built one for has none. + * See `backend/volume_nodes.py`. + * + * Annotation overlay: two independent mask/annotation layers ("Fast (iPred)" + * and "Deep (dlsia)" — see `MaskLayersPanel`), backed by the vendored + * renderer's own two-slot mask API and the `__masks[_deep]` Tiled + * containers `tiled_mask_sync.write_masks_to_tiled` writes. + */ +import { useEffect, useState } from 'react'; +import { useNavigate, useSearchParams } from 'react-router'; +import { useQuery } from '@tanstack/react-query'; +import { Cube } from '@phosphor-icons/react'; +import { API_BASE } from '@/config'; +import { useDatasetStore } from '@/stores/datasetStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { buildSourceKey } from '@/lib/sourceKey'; +import { buildZarrUrl, describeUnavailable, type ZarrUnavailable } from '@/lib/zarrUrl'; +import type { ServerInfo } from '@/types/server'; +import VolumeViewer, { webGpuAvailability, type WebGpuViewerInstance } from '@/components/volume/VolumeViewer'; +import BuildVolumePanel from '@/components/volume/BuildVolumePanel'; +import MaskLayersPanel from '@/components/volume/MaskLayersPanel'; +import RebuildVolumeControl from '@/components/volume/RebuildVolumeControl'; +import { buildBandOpacityCurve, mapByteBandToViewerDomain } from '@/lib/bandTransferFunction'; + +/** Shape of `GET /api/volume/resolve`. */ +interface VolumeNode { + path: string | null; + mode: 'self' | 'ancestor' | 'sidecar' | 'none'; + source_dir: string | null; + message: string; +} + +/** Centered message panel — every non-rendering state uses this shape. */ +function Notice({ + title, + detail, + hint, + showReconnect, +}: { + title: string; + detail: string; + hint?: string; + /** Shows a "Go to Connect" button instead of leaving a dead-end message. */ + showReconnect?: boolean; +}) { + const navigate = useNavigate(); + return ( +
+
+ +

{title}

+

{detail}

+ {hint &&

{hint}

} + {showReconnect && ( + + )} +
+
+ ); +} + +export default function VolumePage() { + const { kind, source, serverUri, meta } = useDatasetStore(); + const byImage = useAnnotationStore((s) => s.byImage); + const sourceKey = source && kind ? buildSourceKey(kind as 'tiled' | 'local', source, serverUri) : null; + const [bootError, setBootError] = useState(null); + const [viewerInstance, setViewerInstance] = useState(null); + // Bumped after a rebuild — folded into VolumeViewer's `key` below to force + // a remount even though the resolved path (and therefore `url`) doesn't + // change on a rebuild the way it does for a brand-new volume. + const [rebuildNonce, setRebuildNonce] = useState(0); + // "View in 3D" hand-offs from Train/Annotate arrive as `?mask=fast|deep` so + // the relevant layer loads automatically instead of requiring a second + // manual click (see MaskLayersPanel's `autoLoadSlot`). + const [searchParams] = useSearchParams(); + const maskParam = searchParams.get('mask'); + const autoLoadSlot = maskParam === 'fast' ? 0 : maskParam === 'deep' ? 1 : undefined; + + // A new dataset deserves a fresh attempt — otherwise one bad volume leaves the + // page stuck on its error for every dataset opened afterwards. + useEffect(() => setBootError(null), [source, serverUri]); + + // "View band in 3D" hand-off from Annotate's Sampler tool (a fitted + // intensity band, in the 2D canvas's own 0-255 byte space) arrives as + // `?bandLo=&bandHi=` — isolate it in the transfer function once the viewer + // is ready. See bandTransferFunction.ts and the plan's #21 threshold-fit + // bridge. + // + // The byte band does NOT translate to the viewer's [0,1] domain by a plain + // /255 — the 2D canvas's bytes are normalized against `meta.globalValueRange` + // (a percentile stretch), while the 3D viewer normalizes raw voxels against + // its OWN, separately-estimated `getValueRange()` — different statistics + // computed from different data, typically close but never identical. + // `mapByteBandToViewerDomain` converts byte -> raw physical value -> the + // viewer's actual domain, rather than assuming the two normalizations agree + // (confirmed live: assuming they did produced a wildly wrong band). + // + // The vendored viewer restores a per-sample cached transfer function from + // localStorage the moment the first real data level streams in (a one-shot + // "boot restore" internal to the renderer, with no public event exposed to + // detect it — WebGpuViewerInstance has no "level loaded"/histogram-ready + // signal) — which races with and can silently clobber the band we just + // set. There's no way to sequence after that restore from the app today, + // so re-apply on a bounded retry schedule that comfortably outlasts a + // typical first-level load; each retry is cheap (a state merge + + // re-render), so the redundant calls cost nothing. + const bandLoParam = searchParams.get('bandLo'); + const bandHiParam = searchParams.get('bandHi'); + useEffect(() => { + if (!viewerInstance || bandLoParam == null || bandHiParam == null) return; + const loByte = Number(bandLoParam); + const hiByte = Number(bandHiParam); + if (!Number.isFinite(loByte) || !Number.isFinite(hiByte) || hiByte <= loByte) return; + const apply = () => { + const viewerRange = viewerInstance.getValueRange(); + const [lo01, hi01] = mapByteBandToViewerDomain(loByte, hiByte, meta?.globalValueRange, viewerRange); + viewerInstance.setRendering({ ...viewerInstance.getRendering(), opacityPoints: buildBandOpacityCurve(lo01, hi01) }); + }; + const retryDelaysMs = [0, 200, 500, 1000, 2000, 4000, 8000]; + const timers = retryDelaysMs.map((ms) => setTimeout(apply, ms)); + return () => timers.forEach(clearTimeout); + }, [viewerInstance, bandLoParam, bandHiParam, meta?.globalValueRange]); + + // Fallback only: a dataset opened against the default local Tiled carries a + // null serverUri, and the resolved address (whatever port start_all.sh + // actually bound) lives here. Never assume 8010. + const { data: servers = [] } = useQuery({ + queryKey: ['servers'], + queryFn: async () => { + const res = await fetch(`${API_BASE}/api/config/servers`); + if (!res.ok) throw new Error('Failed to load servers'); + return res.json(); + }, + enabled: kind === 'tiled' && !serverUri, + }); + + const resolvedUri = serverUri ?? servers[0]?.uri ?? null; + + const { data: node, isLoading: resolving, isError: resolveError } = useQuery({ + queryKey: ['volume-node', resolvedUri, source], + queryFn: async () => { + const params = new URLSearchParams({ source: source! }); + if (resolvedUri) params.set('server_uri', resolvedUri); + const res = await fetch(`${API_BASE}/api/volume/resolve?${params}`); + if (!res.ok) throw new Error('Failed to resolve the volume node'); + return res.json(); + }, + enabled: kind === 'tiled' && !!source, + retry: 1, + }); + + const availability = webGpuAvailability(); + if (!availability.ok) { + return ; + } + + // A genuine failure to reach Tiled (server down, network drop) previously + // fell through to the generic "resolving" reason forever — a dead end with + // no way back to Connect short of the browser's own back button. + if (kind === 'tiled' && resolveError) { + return ( + + ); + } + + // Resolution only applies to Tiled sources; everything else already has a + // reason of its own and should not sit on a spinner. + let reason: ZarrUnavailable | null = null; + if (kind === 'tiled' && source && (resolving || !node)) reason = 'resolving'; + else if (node && node.mode === 'none') reason = 'no-volume'; + + const { url, reason: urlReason } = buildZarrUrl(kind, node?.path ?? null, resolvedUri); + reason = reason ?? (url ? null : urlReason); + + // A stack with no volume is not an error state — it is one the user can + // resolve in place, so offer the build rather than describing an API call. + if (reason === 'no-volume' && source) { + return ; + } + + if (!url || reason) { + return ( + + ); + } + + if (bootError) { + return ( + + ); + } + + return ( +
+ setBootError(e instanceof Error ? e.message : String(e))} + /> + + {node?.mode === 'sidecar' && source && ( + // Bottom-left: the vendor's own docked HUD (Data/TF/Render tabs) spans + // the whole right column, and MaskLayersPanel already owns top-left — + // this is the one corner clear of both. +
+ setRebuildNonce((n) => n + 1)} + /> +
+ )} +
+ ); +} diff --git a/frontend/src/components/Browse/BrowseColumn.test.tsx b/frontend/src/components/Browse/BrowseColumn.test.tsx new file mode 100644 index 0000000..9dbfa4f --- /dev/null +++ b/frontend/src/components/Browse/BrowseColumn.test.tsx @@ -0,0 +1,138 @@ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import BrowseColumn from './BrowseColumn'; +import type { ColumnState } from './hooks/useBrowseData'; + +afterEach(() => { + cleanup(); +}); + +function makeColumn(overrides: Partial = {}): ColumnState { + return { + field: 'sample_type', + values: [], + loading: false, + error: null, + selected: null, + ...overrides, + }; +} + +const baseProps = { + colIndex: 0, + facets: ['sample_type', 'technique', 'beamline'], + width: 200, + onFieldChange: vi.fn(), + onSelect: vi.fn(), + onRemove: vi.fn(), + isLast: true, +}; + +describe('BrowseColumn', () => { + it('renders all facets as select options', () => { + render(); + const select = screen.getByTitle('sample_type') as HTMLSelectElement; + expect(select.value).toBe('sample_type'); + for (const facet of baseProps.facets) { + expect(screen.getByRole('option', { name: facet })).toBeInTheDocument(); + } + }); + + it('shows a loading state', () => { + render(); + expect(screen.getByText('Loading…')).toBeInTheDocument(); + }); + + it('shows an error state', () => { + render(); + expect(screen.getByText('boom')).toBeInTheDocument(); + }); + + it('shows a "No values" empty state when not loading/erroring and empty', () => { + render(); + expect(screen.getByText('No values')).toBeInTheDocument(); + }); + + it('renders values with their counts', () => { + render( + , + ); + expect(screen.getByText('gold')).toBeInTheDocument(); + expect(screen.getByText('12')).toBeInTheDocument(); + expect(screen.getByText('silver')).toBeInTheDocument(); + expect(screen.getByText('3')).toBeInTheDocument(); + }); + + it('applies formatValue to displayed values but not to onSelect', async () => { + const onSelect = vi.fn(); + const user = userEvent.setup(); + const formatValue = (field: string, value: string) => `${field}:${value}`; + render( + , + ); + expect(screen.getByText('sample_type:gold')).toBeInTheDocument(); + await user.click(screen.getByText('sample_type:gold')); + expect(onSelect).toHaveBeenCalledWith(0, 'gold'); + }); + + it('clicking an unselected value selects it', async () => { + const onSelect = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByText('gold')); + expect(onSelect).toHaveBeenCalledWith(0, 'gold'); + }); + + it('clicking the already-selected value clears it (passes null)', async () => { + const onSelect = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByText('gold')); + expect(onSelect).toHaveBeenCalledWith(0, null); + }); + + it('changing the select calls onFieldChange with the column index and new field', async () => { + const onFieldChange = vi.fn(); + const user = userEvent.setup(); + render(); + await user.selectOptions(screen.getByTitle('sample_type'), 'technique'); + expect(onFieldChange).toHaveBeenCalledWith(0, 'technique'); + }); + + it('clicking remove calls onRemove with the column index', async () => { + const onRemove = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByTitle('Remove column')); + expect(onRemove).toHaveBeenCalledWith(0); + }); +}); diff --git a/frontend/src/components/Browse/BrowseDetailPanel.test.tsx b/frontend/src/components/Browse/BrowseDetailPanel.test.tsx new file mode 100644 index 0000000..4049ac5 --- /dev/null +++ b/frontend/src/components/Browse/BrowseDetailPanel.test.tsx @@ -0,0 +1,176 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import BrowseDetailPanel from './BrowseDetailPanel'; +import type { BrowseItem } from './hooks/useBrowseData'; + +function makeItem(overrides: Partial = {}): BrowseItem { + return { + path: 'ds/a', + sample: 'sample-a', + metadata: {}, + ...overrides, + }; +} + +/** Stub window.Image so the thumbnail-loading effect resolves synchronously + * and predictably instead of depending on jsdom's (nonexistent) image decoding. */ +let imageInstances: FakeImage[] = []; +class FakeImage { + onload: (() => void) | null = null; + onerror: (() => void) | null = null; + _src = ''; + constructor() { + imageInstances.push(this); + } + set src(value: string) { + this._src = value; + } + get src() { + return this._src; + } +} + +beforeEach(() => { + imageInstances = []; + vi.stubGlobal('Image', FakeImage); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('BrowseDetailPanel', () => { + it('renders the sample name and path in the header', () => { + render(); + expect(screen.getByText('my-sample')).toBeInTheDocument(); + expect(screen.getByText('a/b/c')).toBeInTheDocument(); + }); + + it('clicking the close button calls onClose', async () => { + const onClose = vi.fn(); + const user = userEvent.setup(); + render(); + // The close button is the only icon-only button in the header without a title. + const buttons = screen.getAllByRole('button'); + await user.click(buttons[0]); + expect(onClose).toHaveBeenCalled(); + }); + + it('shows the loading preview state before the thumbnail resolves', () => { + render(); + expect(screen.getByText('Loading preview…')).toBeInTheDocument(); + }); + + it('shows the image once the thumbnail loads successfully', async () => { + render(); + expect(imageInstances).toHaveLength(1); + act(() => { + imageInstances[0].onload?.(); + }); + const img = await screen.findByAltText('Array preview'); + expect(img).toHaveAttribute('src', expect.stringContaining('tiled_path=ds%2Fa')); + expect(img.getAttribute('src')).toContain('server_uri=http%3A%2F%2Fserver'); + }); + + it('renders nothing for the preview area when the thumbnail fails to load', async () => { + render(); + act(() => { + imageInstances[0].onerror?.(); + }); + expect(screen.queryByText('Loading preview…')).not.toBeInTheDocument(); + expect(screen.queryByAltText('Array preview')).not.toBeInTheDocument(); + }); + + it('refetches the thumbnail when the item path changes', () => { + const { rerender } = render(); + expect(imageInstances).toHaveLength(1); + rerender(); + expect(imageInstances).toHaveLength(2); + expect(imageInstances[1].src).toContain('tiled_path=b'); + }); + + it('does not render the Open in Annotate button when onOpenInAnnotate is omitted', () => { + render(); + expect(screen.queryByText('Open in Annotate')).not.toBeInTheDocument(); + }); + + it('clicking Open in Annotate calls the callback', async () => { + const onOpenInAnnotate = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByText('Open in Annotate')); + expect(onOpenInAnnotate).toHaveBeenCalled(); + }); + + it('groups known metadata keys into their labelled sections', () => { + render( + , + ); + expect(screen.getByText('Experiment')).toBeInTheDocument(); + expect(screen.getByText('Identity')).toBeInTheDocument(); + expect(screen.getByText('PI')).toBeInTheDocument(); + expect(screen.getByText('Dr. Smith')).toBeInTheDocument(); + expect(screen.getByText('ThinFilmID')).toBeInTheDocument(); + expect(screen.getByText('TF-1')).toBeInTheDocument(); + }); + + it('maps aliased keys to friendlier display labels', () => { + render( + , + ); + expect(screen.getByText('Annotated')).toBeInTheDocument(); + expect(screen.getByText('Shape count')).toBeInTheDocument(); + expect(screen.getByText('4')).toBeInTheDocument(); + }); + + it('puts unrecognised metadata keys into an "Other" section, skipping noisy/internal keys', () => { + render( + , + ); + expect(screen.getByText('Other')).toBeInTheDocument(); + expect(screen.getByText('custom_field')).toBeInTheDocument(); + expect(screen.getByText('custom-value')).toBeInTheDocument(); + expect(screen.queryByText('vae_embedding')).not.toBeInTheDocument(); + expect(screen.queryByText('thinfilm_internal')).not.toBeInTheDocument(); + expect(screen.queryByText('Sample foo')).not.toBeInTheDocument(); + }); + + it('formats array metadata values as an item-count summary', () => { + render( + , + ); + expect(screen.getByText('[3 items]')).toBeInTheDocument(); + }); + + it('formats null/undefined-ish metadata values as an em dash', () => { + render( + , + ); + // null values are filtered out before display, so no "Other" section appears. + expect(screen.queryByText('Other')).not.toBeInTheDocument(); + }); + + it('renders no metadata sections when metadata is empty', () => { + render(); + expect(screen.queryByText('Identity')).not.toBeInTheDocument(); + expect(screen.queryByText('Other')).not.toBeInTheDocument(); + }); +}); diff --git a/frontend/src/components/Browse/ColumnBrowser.test.tsx b/frontend/src/components/Browse/ColumnBrowser.test.tsx new file mode 100644 index 0000000..a81a62e --- /dev/null +++ b/frontend/src/components/Browse/ColumnBrowser.test.tsx @@ -0,0 +1,310 @@ +/** + * ColumnBrowser — tests the Miller-column orchestration: toolbar wiring, the + * auto "show all" on connect, adding columns, item/slice selection routing, + * open-in-Annotate, and connection-error display. + * + * BrowseColumn/ItemsColumn/SlicesColumn/BrowseDetailPanel are mocked out — they + * are separate, already-tested units (owned by a parallel test-writing pass); + * ColumnBrowser only needs to exercise the props/callbacks it wires into them. + * useBrowseData itself is real, so the actual fetch-driven state machine + * (facets → show-all → items) still runs end to end against a stubbed fetch. + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { MemoryRouter } from 'react-router'; +import ColumnBrowser from './ColumnBrowser'; +import type { BrowseItem } from './hooks/useBrowseData'; +import type { ServerInfo } from '@/types/server'; + +const openTiledArray = vi.fn(async () => {}); +vi.mock('@/hooks/useOpenInAnnotate', () => ({ + useOpenInAnnotate: () => ({ openTiledArray, openLocalFile: vi.fn() }), +})); + +vi.mock('./BrowseColumn', () => ({ + default: (props: any) => ( +
+ field:{props.column.field} + + +
+ ), +})); + +vi.mock('./ItemsColumn', () => ({ + default: (props: any) => ( +
+ {props.items.map((it: BrowseItem) => ( + + ))} +
+ ), +})); + +vi.mock('./SlicesColumn', () => ({ + default: (props: any) => ( +
+ slices-of:{props.dataset.sample} + +
+ ), +})); + +vi.mock('./BrowseDetailPanel', () => ({ + default: (props: any) => ( +
+ detail:{props.item.sample} + +
+ ), +})); + +function jsonResponse(body: unknown, ok = true) { + return { + ok, + status: ok ? 200 : 500, + json: async () => body, + text: async () => (typeof body === 'string' ? body : JSON.stringify(body)), + } as Response; +} + +function makeFetchMock( + overrides: Record Response | Promise> = {}, +) { + return vi.fn(async (input: RequestInfo | URL) => { + const url = new URL(String(input), 'http://localhost'); + for (const [path, handler] of Object.entries(overrides)) { + if (url.pathname === path) return handler(url); + } + return jsonResponse({}, false); + }); +} + +const SERVERS: ServerInfo[] = [ + { name: 'Server A', uri: 'http://a', has_api_key: false }, + { name: 'Server B', uri: 'http://b', has_api_key: false }, +]; + +function renderBrowser(overrides: Partial> = {}) { + const props = { + serverUri: 'http://a', + containerPath: null, + focusPath: null, + servers: SERVERS, + selectedServerUri: 'http://a', + onServerChange: vi.fn(), + annotationFilter: 'all' as const, + onAnnotationFilterChange: vi.fn(), + ...overrides, + }; + return { ...render(, { wrapper: MemoryRouter }), props }; +} + +beforeEach(() => { + vi.spyOn(console, 'warn').mockImplementation(() => {}); +}); + +afterEach(() => { + cleanup(); + vi.clearAllMocks(); +}); + +describe('ColumnBrowser', () => { + it('renders the toolbar with server + annotation selects', async () => { + global.fetch = makeFetchMock({ '/api/browse/facets': () => jsonResponse({ facets: [] }) }); + renderBrowser(); + expect(screen.getByText('Metadata Browser')).toBeInTheDocument(); + expect(screen.getByRole('combobox', { name: 'Server' })).toBeInTheDocument(); + expect(screen.getByRole('combobox', { name: 'Annotation' })).toBeInTheDocument(); + // Let the facets fetch (and its show-all effect) settle before the test ends. + await waitFor(() => expect(screen.getByTestId('items-column')).toBeInTheDocument()); + }); + + it('auto shows every sample once connected, rendering the items column', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['field1'] }), + '/api/browse/items': () => + jsonResponse({ items: [{ path: 'p1', sample: 'sample-1', metadata: {} }], total: 1 }), + }); + renderBrowser(); + + expect(await screen.findByTestId('items-column')).toBeInTheDocument(); + // showAll() sets showingAll=true and kicks off loadItems({}) without + // waiting for it — ColumnBrowser mounts ItemsColumn as soon as + // showingAll flips, independent of whether the items fetch has actually + // resolved. Finding the (possibly still-empty) column is not proof the + // sample has rendered; only findByText's own retry-until-appears + // semantics correctly wait for the item itself. + expect(await screen.findByText('item:sample-1')).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /All samples/ })).toHaveClass('bg-slate-600'); + }); + + it('renders samples when a pending items response completes', async () => { + let resolveItems!: (response: Response) => void; + const pendingItems = new Promise((resolve) => { + resolveItems = resolve; + }); + const itemsHandler = vi.fn(() => pendingItems); + + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['field1'] }), + '/api/browse/items': itemsHandler, + }); + + renderBrowser(); + + try { + expect(await screen.findByTestId('items-column')).toBeInTheDocument(); + expect(itemsHandler).toHaveBeenCalledTimes(1); + + // The container exists, but the response is deliberately still pending + // — this is the intermediate state the race above was missing. + expect(screen.queryByText('item:sample-1')).not.toBeInTheDocument(); + } finally { + // Release the request even if an assertion fails, so its timeout and + // pending updates do not leak into subsequent tests. + await act(async () => { + resolveItems( + jsonResponse({ items: [{ path: 'p1', sample: 'sample-1', metadata: {} }], total: 1 }), + ); + await pendingItems; + }); + } + + expect(await screen.findByText('item:sample-1')).toBeInTheDocument(); + }); + + it('shows the disconnected banner when facets cannot be reached, and no items column', async () => { + global.fetch = makeFetchMock({ '/api/browse/facets': () => jsonResponse('down', false) }); + renderBrowser(); + + expect(await screen.findByText(/Cannot reach the API server/)).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /go to connect/i })).toBeInTheDocument(); + expect(screen.queryByTestId('items-column')).not.toBeInTheDocument(); + }); + + it('Add column adds the first unused facet, and adding again picks the next one', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['field1', 'field2'] }), + '/api/browse/items': () => jsonResponse({ items: [], total: 0 }), + '/api/browse/column': () => jsonResponse({ values: [] }), + }); + const user = userEvent.setup(); + renderBrowser(); + await screen.findByTestId('items-column'); + + await user.click(screen.getByRole('button', { name: 'Add column' })); + expect(await screen.findByTestId('browse-column-0')).toHaveTextContent('field:field1'); + + await user.click(screen.getByRole('button', { name: 'Add column' })); + expect(await screen.findByTestId('browse-column-1')).toHaveTextContent('field:field2'); + }); + + it('selecting a single-image sample opens the detail panel', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/items': () => + jsonResponse({ items: [{ path: 'p1', sample: 'single-1', metadata: {} }], total: 1 }), + }); + const user = userEvent.setup(); + renderBrowser(); + + await user.click(await screen.findByText('item:single-1')); + expect(await screen.findByTestId('detail-panel')).toHaveTextContent('detail:single-1'); + + await user.click(screen.getByText('close-detail')); + await waitFor(() => expect(screen.queryByTestId('detail-panel')).not.toBeInTheDocument()); + }); + + it('selecting a multi-slice sample expands the slices column instead of the detail panel', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/items': () => + jsonResponse({ + items: [{ path: 'vol1', sample: 'volume-1', metadata: {}, n_slices: 5 }], + total: 1, + }), + '/api/browse/slices': () => + jsonResponse({ items: [{ path: 'vol1/0', sample: '0', metadata: {} }] }), + }); + const user = userEvent.setup(); + renderBrowser(); + + await user.click(await screen.findByText('item:volume-1')); + expect(await screen.findByTestId('slices-column')).toHaveTextContent('slices-of:volume-1'); + expect(screen.queryByTestId('detail-panel')).not.toBeInTheDocument(); + + await user.click(screen.getByText('close-slices')); + await waitFor(() => expect(screen.queryByTestId('slices-column')).not.toBeInTheDocument()); + }); + + it('Open in Annotate opens the selected item via useOpenInAnnotate', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/items': () => + jsonResponse({ items: [{ path: 'p1', sample: 'single-1', metadata: {} }], total: 1 }), + }); + const user = userEvent.setup(); + renderBrowser(); + + await user.click(await screen.findByText('item:single-1')); + await user.click(screen.getByRole('button', { name: /Open in Annotate/ })); + + await waitFor(() => expect(openTiledArray).toHaveBeenCalledWith('p1', 'http://a')); + }); + + it('auto-selects the item matching focusPath once the full list has loaded', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/items': () => + jsonResponse({ + items: [ + { path: 'p1', sample: 'first', metadata: {} }, + { path: 'p2', sample: 'second', metadata: {} }, + ], + total: 2, + }), + }); + renderBrowser({ focusPath: 'p2' }); + + expect(await screen.findByTestId('detail-panel')).toHaveTextContent('detail:second'); + }); + + it('calls onServerChange when the server select changes', async () => { + global.fetch = makeFetchMock({ '/api/browse/facets': () => jsonResponse({ facets: [] }) }); + const onServerChange = vi.fn(); + const user = userEvent.setup(); + renderBrowser({ onServerChange }); + + await user.selectOptions(screen.getByRole('combobox', { name: 'Server' }), 'http://b'); + expect(onServerChange).toHaveBeenCalledWith('http://b'); + }); + + it('calls onAnnotationFilterChange when the annotation select changes', async () => { + global.fetch = makeFetchMock({ '/api/browse/facets': () => jsonResponse({ facets: [] }) }); + const onAnnotationFilterChange = vi.fn(); + const user = userEvent.setup(); + renderBrowser({ onAnnotationFilterChange }); + + await user.selectOptions(screen.getByRole('combobox', { name: 'Annotation' }), 'annotated'); + expect(onAnnotationFilterChange).toHaveBeenCalledWith('annotated'); + }); + + it('Refresh triggers another items fetch while showing all', async () => { + const fetchMock = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/items': () => jsonResponse({ items: [], total: 0 }), + }); + global.fetch = fetchMock; + const user = userEvent.setup(); + renderBrowser(); + + await screen.findByTestId('items-column'); + const callsBefore = fetchMock.mock.calls.length; + await user.click(screen.getByRole('button', { name: 'Refresh' })); + await waitFor(() => expect(fetchMock.mock.calls.length).toBeGreaterThan(callsBefore)); + }); +}); diff --git a/frontend/src/components/Browse/ColumnBrowser.tsx b/frontend/src/components/Browse/ColumnBrowser.tsx index 9e6f315..03b049b 100644 --- a/frontend/src/components/Browse/ColumnBrowser.tsx +++ b/frontend/src/components/Browse/ColumnBrowser.tsx @@ -1,4 +1,5 @@ import React, { useCallback, useEffect, useMemo, useRef, useState } from 'react'; +import { useNavigate } from 'react-router'; import { ArrowsClockwise, PencilSimple, Plus, Stack } from '@phosphor-icons/react'; import BrowseColumn from './BrowseColumn'; import BrowseDetailPanel from './BrowseDetailPanel'; @@ -72,6 +73,7 @@ export default function ColumnBrowser({ }: ColumnBrowserProps) { const { state, actions } = useBrowseData(serverUri, 'All', undefined, containerPath); const { openTiledArray } = useOpenInAnnotate(); + const navigate = useNavigate(); const scrollRef = useRef(null); const [columnWidths, setColumnWidths] = useState([]); @@ -237,8 +239,14 @@ export default function ColumnBrowser({ )} {state.connectionStatus === 'disconnected' && ( -
- Cannot reach the API server. Make sure the backend (port 8002) and Tiled server are running. +
+ Cannot reach the API server. Make sure the backend (port 8002) and Tiled server are running. +
)} diff --git a/frontend/src/components/Browse/ItemsColumn.test.tsx b/frontend/src/components/Browse/ItemsColumn.test.tsx new file mode 100644 index 0000000..82740a1 --- /dev/null +++ b/frontend/src/components/Browse/ItemsColumn.test.tsx @@ -0,0 +1,169 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import ItemsColumn from './ItemsColumn'; +import { useRatingStore } from '@/stores/ratingStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { buildSourceKey } from '@/lib/sourceKey'; +import type { BrowseItem } from './hooks/useBrowseData'; + +const initialRatingState = useRatingStore.getState(); + +function makeItem(overrides: Partial = {}): BrowseItem { + return { + path: 'ds/a', + sample: 'sample-a', + metadata: {}, + ...overrides, + }; +} + +function renderWithClient(ui: React.ReactElement) { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render({ui}); +} + +beforeEach(() => { + useRatingStore.setState(initialRatingState, true); + useAnnotationStore.getState().reset(); + window.localStorage.clear(); + vi.stubGlobal( + 'fetch', + vi.fn().mockResolvedValue({ ok: true, json: async () => [] }), + ); + // jsdom doesn't implement scrollIntoView. + window.HTMLElement.prototype.scrollIntoView = vi.fn(); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +const baseProps = { + items: [] as BrowseItem[], + total: 0, + loading: false, + selectedItem: null, + onSelect: vi.fn(), + onOpenInAnnotate: vi.fn(), + width: 220, + serverUri: 'http://server', + annotationFilter: 'all' as const, +}; + +describe('ItemsColumn', () => { + it('shows a loading state', () => { + renderWithClient(); + expect(screen.getByText('Loading…')).toBeInTheDocument(); + }); + + it('shows an empty state when there are no matching samples', () => { + renderWithClient(); + expect(screen.getByText('No matching samples')).toBeInTheDocument(); + }); + + it('renders items and the total count badge', () => { + const items = [makeItem({ path: 'a', sample: 'a' }), makeItem({ path: 'b', sample: 'b' })]; + renderWithClient(); + expect(screen.getByText('a')).toBeInTheDocument(); + expect(screen.getByText('b')).toBeInTheDocument(); + expect(screen.getByText('2')).toBeInTheDocument(); + }); + + it('clicking a row selects it, clicking again clears the selection', async () => { + const onSelect = vi.fn(); + const user = userEvent.setup(); + const item = makeItem({ path: 'a', sample: 'a' }); + const { rerender } = renderWithClient( + + + , + ); + await user.click(screen.getByText('a')); + expect(onSelect).toHaveBeenCalledWith(item); + + onSelect.mockClear(); + rerender( + + + , + ); + await user.click(screen.getByText('a')); + expect(onSelect).toHaveBeenCalledWith(null); + }); + + it('clicking the open-in-annotate button calls onOpenInAnnotate with the item and does not select it', async () => { + const onOpenInAnnotate = vi.fn(); + const onSelect = vi.fn(); + const user = userEvent.setup(); + const item = makeItem({ path: 'a', sample: 'a' }); + renderWithClient( + , + ); + await user.click(screen.getByTitle('Open in Annotate')); + expect(onOpenInAnnotate).toHaveBeenCalledWith(item); + expect(onSelect).not.toHaveBeenCalled(); + }); + + it('shows a volume badge and caret for multi-slice items, not for single-slice items', () => { + const volume = makeItem({ path: 'v', sample: 'vol', n_slices: 5 }); + const single = makeItem({ path: 's', sample: 'single', n_slices: 1 }); + renderWithClient(); + expect(screen.getByText('5')).toBeInTheDocument(); + expect(screen.getByTitle('5 slices — click to browse them')).toBeInTheDocument(); + }); + + it('shows the annotated badge when metadata marks the item annotated', () => { + const item = makeItem({ path: 'a', sample: 'a', metadata: { studio_annotated: 'yes' } }); + renderWithClient(); + expect(screen.getByText('annotated')).toBeInTheDocument(); + }); + + it('does not show the annotated badge for unannotated items', () => { + const item = makeItem({ path: 'a', sample: 'a' }); + renderWithClient(); + expect(screen.queryByText('annotated')).not.toBeInTheDocument(); + }); + + it('filters out unannotated items when annotationFilter is "annotated"', () => { + const annotated = makeItem({ path: 'a', sample: 'a', metadata: { studio_annotated: 'yes' } }); + const plain = makeItem({ path: 'b', sample: 'b' }); + renderWithClient( + , + ); + expect(screen.getByText('a')).toBeInTheDocument(); + expect(screen.queryByText('b')).not.toBeInTheDocument(); + // filtered count differs from total, so shows "1 / 2" + expect(screen.getByText('1 / 2')).toBeInTheDocument(); + }); + + it('filters out annotated items when annotationFilter is "unannotated"', () => { + const annotated = makeItem({ path: 'a', sample: 'a', metadata: { studio_annotated: 'yes' } }); + const plain = makeItem({ path: 'b', sample: 'b' }); + renderWithClient( + , + ); + expect(screen.queryByText('a')).not.toBeInTheDocument(); + expect(screen.getByText('b')).toBeInTheDocument(); + }); + + it('filters by minimum star rating', async () => { + const item = makeItem({ path: 'a', sample: 'a' }); + const sourceKey = buildSourceKey('tiled', item.path, baseProps.serverUri); + useRatingStore.setState({ ratings: { [sourceKey]: 1 } }); + const user = userEvent.setup(); + renderWithClient(); + + expect(screen.getByText('a')).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: '★★' })); + expect(screen.queryByText('a')).not.toBeInTheDocument(); + expect(screen.getByText('No samples match the current filters.')).toBeInTheDocument(); + }); + + it('shows the "no matches" filtered message distinctly from the plain empty state', () => { + renderWithClient(); + expect(screen.getByText('No samples match the current filters.')).toBeInTheDocument(); + }); +}); diff --git a/frontend/src/components/Browse/LocalSampleBrowser.test.tsx b/frontend/src/components/Browse/LocalSampleBrowser.test.tsx new file mode 100644 index 0000000..fef6241 --- /dev/null +++ b/frontend/src/components/Browse/LocalSampleBrowser.test.tsx @@ -0,0 +1,221 @@ +/** + * LocalSampleBrowser — flat local-folder sample list: loading/error states, + * annotation + star filters, rating persistence, joinPath, and Annotate wiring. + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen, waitFor, within } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import LocalSampleBrowser, { joinPath } from './LocalSampleBrowser'; +import { useRatingStore } from '@/stores/ratingStore'; +import { buildSourceKey } from '@/lib/sourceKey'; + +const initialRatingState = useRatingStore.getState(); + +function jsonResponse(body: unknown, ok = true) { + return { + ok, + status: ok ? 200 : 500, + json: async () => body, + text: async () => (typeof body === 'string' ? body : JSON.stringify(body)), + } as Response; +} + +function makeFetchMock( + overrides: Record Response | Promise> = {}, +) { + return vi.fn(async (input: RequestInfo | URL) => { + const url = new URL(String(input), 'http://localhost'); + for (const [path, handler] of Object.entries(overrides)) { + if (url.pathname === path) return handler(url); + } + if (url.pathname === '/api/annotations/drafts') return jsonResponse([]); + return jsonResponse({}, false); + }); +} + +const SAMPLES = [ + { name: 'a.tif', path: 'a.tif' }, + { name: 'b.tif', path: 'sub/b.tif' }, +]; + +function renderBrowser( + overrides: Partial> = {}, +) { + const props = { + root: '/data', + rel: '', + onOpenInAnnotate: vi.fn(), + annotationFilter: 'all' as const, + onAnnotationFilterChange: vi.fn(), + ...overrides, + }; + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return { + ...render( + + + , + ), + props, + }; +} + +beforeEach(() => { + useRatingStore.setState(initialRatingState, true); + window.localStorage.clear(); +}); + +afterEach(() => { + cleanup(); + vi.restoreAllMocks(); +}); + +describe('joinPath', () => { + it('returns root unchanged when rel is empty', () => { + expect(joinPath('/data', '')).toBe('/data'); + }); + + it('joins root and rel, normalising duplicate slashes', () => { + expect(joinPath('/data/', '/sub/file.tif')).toBe('/data/sub/file.tif'); + expect(joinPath('/data', 'sub/file.tif')).toBe('/data/sub/file.tif'); + }); +}); + +describe('LocalSampleBrowser', () => { + it('shows a loading state before samples arrive', async () => { + let resolveFetch!: (v: Response) => void; + global.fetch = vi.fn((input: RequestInfo | URL) => { + const url = String(input); + if (url.includes('/api/local/samples')) { + return new Promise((resolve) => { + resolveFetch = resolve; + }); + } + return Promise.resolve(jsonResponse([])); + }) as unknown as typeof fetch; + + renderBrowser(); + expect(screen.getByText('Loading samples…')).toBeInTheDocument(); + + resolveFetch(jsonResponse({ items: [], total: 0 })); + await waitFor(() => expect(screen.queryByText('Loading samples…')).not.toBeInTheDocument()); + }); + + it('shows an error message when the samples request fails', async () => { + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse('disk unavailable', false), + }); + renderBrowser(); + expect(await screen.findByText('disk unavailable')).toBeInTheDocument(); + }); + + it('lists samples with a total count and the current folder path', async () => { + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse({ items: SAMPLES, total: 2 }), + }); + renderBrowser({ rel: 'sub' }); + + expect(await screen.findByText('a.tif')).toBeInTheDocument(); + expect(screen.getByText('b.tif')).toBeInTheDocument(); + expect(screen.getByText('2')).toBeInTheDocument(); + expect(screen.getByText('/data/sub')).toBeInTheDocument(); + }); + + it('shows an empty-folder message when there are no samples', async () => { + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse({ items: [], total: 0 }), + }); + renderBrowser(); + expect(await screen.findByText('No image files found in this folder.')).toBeInTheDocument(); + }); + + it('shows an annotated badge for samples with drafts, based on their source key', async () => { + const sk = buildSourceKey('local', joinPath('/data', 'a.tif')); + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse({ items: SAMPLES, total: 2 }), + '/api/annotations/drafts': () => jsonResponse([{ source_key: sk, has_annotations: true }]), + }); + renderBrowser(); + + await screen.findByText('a.tif'); + const row = screen.getByText('a.tif').closest('.group') as HTMLElement; + expect(row).toHaveTextContent('annotated'); + const otherRow = screen.getByText('b.tif').closest('.group') as HTMLElement; + expect(otherRow).not.toHaveTextContent('annotated'); + }); + + it('filters by the annotation-status select via the onAnnotationFilterChange callback', async () => { + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse({ items: SAMPLES, total: 2 }), + }); + const onAnnotationFilterChange = vi.fn(); + const user = userEvent.setup(); + renderBrowser({ onAnnotationFilterChange }); + + await screen.findByText('a.tif'); + await user.selectOptions(screen.getByRole('combobox'), 'annotated'); + expect(onAnnotationFilterChange).toHaveBeenCalledWith('annotated'); + }); + + it('unannotated filter hides samples with drafts', async () => { + const sk = buildSourceKey('local', joinPath('/data', 'a.tif')); + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse({ items: SAMPLES, total: 2 }), + '/api/annotations/drafts': () => jsonResponse([{ source_key: sk, has_annotations: true }]), + }); + renderBrowser({ annotationFilter: 'unannotated' }); + + await screen.findByText('b.tif'); + expect(screen.queryByText('a.tif')).not.toBeInTheDocument(); + expect(screen.getByText('1 / 2')).toBeInTheDocument(); + }); + + it('star filter hides samples below the chosen minimum rating', async () => { + const skA = buildSourceKey('local', joinPath('/data', 'a.tif')); + useRatingStore.setState({ ratings: { [skA]: 2 } }); + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse({ items: SAMPLES, total: 2 }), + }); + const user = userEvent.setup(); + renderBrowser(); + + await screen.findByText('a.tif'); + await user.click(screen.getByText('★★')); + expect(screen.getByText('a.tif')).toBeInTheDocument(); + expect(screen.queryByText('b.tif')).not.toBeInTheDocument(); + + await user.click(screen.getByText('All')); + expect(screen.getByText('b.tif')).toBeInTheDocument(); + }); + + it('clicking a star rates the sample and persists it in the rating store', async () => { + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse({ items: SAMPLES, total: 2 }), + }); + const user = userEvent.setup(); + renderBrowser(); + + await screen.findByText('a.tif'); + const row = screen.getByText('a.tif').closest('.group') as HTMLElement; + await user.click(within(row).getByLabelText('2 stars')); + + const sk = buildSourceKey('local', joinPath('/data', 'a.tif')); + expect(useRatingStore.getState().ratings[sk]).toBe(2); + }); + + it('clicking Annotate calls onOpenInAnnotate with the absolute path', async () => { + global.fetch = makeFetchMock({ + '/api/local/samples': () => jsonResponse({ items: SAMPLES, total: 2 }), + }); + const onOpenInAnnotate = vi.fn(); + const user = userEvent.setup(); + renderBrowser({ onOpenInAnnotate }); + + await screen.findByText('b.tif'); + const row = screen.getByText('b.tif').closest('.group') as HTMLElement; + await user.click(within(row).getByRole('button', { name: /Annotate/ })); + + expect(onOpenInAnnotate).toHaveBeenCalledWith('/data/sub/b.tif'); + }); +}); diff --git a/frontend/src/components/Browse/ResizeDivider.test.tsx b/frontend/src/components/Browse/ResizeDivider.test.tsx new file mode 100644 index 0000000..bfdac82 --- /dev/null +++ b/frontend/src/components/Browse/ResizeDivider.test.tsx @@ -0,0 +1,71 @@ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, fireEvent, render, screen } from '@testing-library/react'; +import ResizeDivider, { MAX_WIDTH, MIN_WIDTH } from './ResizeDivider'; + +afterEach(() => { + cleanup(); +}); + +describe('ResizeDivider', () => { + it('renders a vertical separator', () => { + render(); + expect(screen.getByRole('separator')).toHaveAttribute('aria-orientation', 'vertical'); + }); + + it('dragging right grows the column when resizeRight is false', () => { + const onResize = vi.fn(); + render(); + const divider = screen.getByRole('separator'); + + fireEvent.mouseDown(divider, { clientX: 100 }); + act(() => { + fireEvent.mouseMove(document, { clientX: 150 }); + }); + expect(onResize).toHaveBeenCalledWith(250); + }); + + it('dragging right shrinks the column when resizeRight is true', () => { + const onResize = vi.fn(); + render(); + const divider = screen.getByRole('separator'); + + fireEvent.mouseDown(divider, { clientX: 100 }); + act(() => { + fireEvent.mouseMove(document, { clientX: 150 }); + }); + expect(onResize).toHaveBeenCalledWith(150); + }); + + it('clamps to minWidth/maxWidth', () => { + const onResize = vi.fn(); + render(); + const divider = screen.getByRole('separator'); + + fireEvent.mouseDown(divider, { clientX: 100 }); + act(() => { + fireEvent.mouseMove(document, { clientX: -10000 }); + }); + expect(onResize).toHaveBeenLastCalledWith(MIN_WIDTH); + + act(() => { + fireEvent.mouseMove(document, { clientX: 10000 }); + }); + expect(onResize).toHaveBeenLastCalledWith(MAX_WIDTH); + }); + + it('stops tracking mouse movement after mouseup', () => { + const onResize = vi.fn(); + render(); + const divider = screen.getByRole('separator'); + + fireEvent.mouseDown(divider, { clientX: 100 }); + act(() => { + fireEvent.mouseUp(document); + }); + onResize.mockClear(); + act(() => { + fireEvent.mouseMove(document, { clientX: 300 }); + }); + expect(onResize).not.toHaveBeenCalled(); + }); +}); diff --git a/frontend/src/components/Browse/SlicesColumn.test.tsx b/frontend/src/components/Browse/SlicesColumn.test.tsx new file mode 100644 index 0000000..fdb114b --- /dev/null +++ b/frontend/src/components/Browse/SlicesColumn.test.tsx @@ -0,0 +1,161 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import SlicesColumn from './SlicesColumn'; +import { useRatingStore } from '@/stores/ratingStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { buildSourceKey } from '@/lib/sourceKey'; +import type { BrowseItem } from './hooks/useBrowseData'; + +const initialRatingState = useRatingStore.getState(); + +function makeItem(overrides: Partial = {}): BrowseItem { + return { + path: 'ds/a', + sample: 'ds_0001', + metadata: {}, + ...overrides, + }; +} + +function renderWithClient(ui: React.ReactElement) { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render({ui}); +} + +beforeEach(() => { + useRatingStore.setState(initialRatingState, true); + useAnnotationStore.getState().reset(); + window.localStorage.clear(); + vi.stubGlobal( + 'fetch', + vi.fn().mockResolvedValue({ ok: true, json: async () => [] }), + ); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +const dataset = makeItem({ path: 'ds', sample: 'ds' }); + +const baseProps = { + dataset, + slices: [] as BrowseItem[], + loading: false, + selectedItem: null, + onSelect: vi.fn(), + onOpenInAnnotate: vi.fn(), + onClose: vi.fn(), + width: 200, + serverUri: 'http://server', +}; + +describe('SlicesColumn', () => { + it('shows a loading state', () => { + renderWithClient(); + expect(screen.getByText('Loading slices…')).toBeInTheDocument(); + }); + + it('shows an empty state when there are no slices', () => { + renderWithClient(); + expect(screen.getByText('No slices found.')).toBeInTheDocument(); + }); + + it('renders the slice count badge and the "open all" button count', () => { + const slices = [makeItem({ path: 'ds/1', sample: 'ds_0001' }), makeItem({ path: 'ds/2', sample: 'ds_0002' })]; + renderWithClient(); + expect(screen.getByText('2')).toBeInTheDocument(); + expect(screen.getByText(/Open all 2 slices as a volume/)).toBeInTheDocument(); + }); + + it('labels a slice by its image_number metadata when present', () => { + const slice = makeItem({ path: 'ds/1', sample: 'ds_weird', metadata: { image_number: 7 } }); + renderWithClient(); + expect(screen.getByText('Slice 7')).toBeInTheDocument(); + }); + + it('labels a slice by its trailing suffix when it shares the dataset prefix', () => { + const slice = makeItem({ path: 'ds/1', sample: 'ds_0003' }); + renderWithClient(); + expect(screen.getByText('Slice 0003')).toBeInTheDocument(); + }); + + it('falls back to the raw sample name when no image_number and no shared prefix', () => { + const slice = makeItem({ path: 'ds/1', sample: 'totally_different' }); + renderWithClient(); + expect(screen.getByText('totally_different')).toBeInTheDocument(); + }); + + it('clicking the back button calls onClose', async () => { + const onClose = vi.fn(); + const user = userEvent.setup(); + renderWithClient(); + await user.click(screen.getByTitle('Back to samples')); + expect(onClose).toHaveBeenCalled(); + }); + + it('clicking "open all as volume" calls onOpenInAnnotate with the dataset', async () => { + const onOpenInAnnotate = vi.fn(); + const user = userEvent.setup(); + renderWithClient(); + await user.click(screen.getByText(/Open all/)); + expect(onOpenInAnnotate).toHaveBeenCalledWith(dataset); + }); + + it('clicking a slice row selects it, clicking again clears the selection', async () => { + const onSelect = vi.fn(); + const user = userEvent.setup(); + const slice = makeItem({ path: 'ds/1', sample: 'ds_0001' }); + const { rerender } = renderWithClient( + , + ); + await user.click(screen.getByText('Slice 0001')); + expect(onSelect).toHaveBeenCalledWith(slice); + + onSelect.mockClear(); + rerender( + + + , + ); + await user.click(screen.getByText('Slice 0001')); + expect(onSelect).toHaveBeenCalledWith(null); + }); + + it('clicking the per-slice open-in-annotate button calls onOpenInAnnotate with that slice, not the dataset', async () => { + const onOpenInAnnotate = vi.fn(); + const onSelect = vi.fn(); + const user = userEvent.setup(); + const slice = makeItem({ path: 'ds/1', sample: 'ds_0001' }); + renderWithClient( + , + ); + await user.click(screen.getByTitle('Open this slice in Annotate')); + expect(onOpenInAnnotate).toHaveBeenCalledWith(slice); + expect(onSelect).not.toHaveBeenCalled(); + }); + + it('shows the annotated badge when the slice is flagged annotated', () => { + const slice = makeItem({ path: 'ds/1', sample: 'ds_0001', metadata: { studio_annotated: 'yes' } }); + renderWithClient(); + expect(screen.getByText('annotated')).toBeInTheDocument(); + }); + + it('does not show the annotated badge for an unannotated slice', () => { + const slice = makeItem({ path: 'ds/1', sample: 'ds_0001' }); + renderWithClient(); + expect(screen.queryByText('annotated')).not.toBeInTheDocument(); + }); + + it('renders a star rating control reflecting the stored rating for the slice', () => { + const slice = makeItem({ path: 'ds/1', sample: 'ds_0001' }); + const sourceKey = buildSourceKey('tiled', slice.path, baseProps.serverUri); + useRatingStore.setState({ ratings: { [sourceKey]: 2 } }); + renderWithClient(); + // StarRating renders 3 star buttons per row. + expect(screen.getAllByRole('button', { name: /star/ })).toHaveLength(3); + }); +}); diff --git a/frontend/src/components/Browse/StarRating.test.tsx b/frontend/src/components/Browse/StarRating.test.tsx new file mode 100644 index 0000000..e972aa7 --- /dev/null +++ b/frontend/src/components/Browse/StarRating.test.tsx @@ -0,0 +1,59 @@ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import StarRating from './StarRating'; + +afterEach(() => { + cleanup(); +}); + +describe('StarRating', () => { + it('renders three stars', () => { + render(); + expect(screen.getAllByRole('button')).toHaveLength(3); + }); + + it('fills stars up to and including the current value', () => { + render(); + const buttons = screen.getAllByRole('button'); + expect(buttons[0].querySelector('svg')).toHaveClass('text-amber-400'); + expect(buttons[1].querySelector('svg')).toHaveClass('text-amber-400'); + expect(buttons[2].querySelector('svg')).toHaveClass('text-slate-600'); + }); + + it('clicking a star calls onChange with that star value', async () => { + const onChange = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('2 stars')); + expect(onChange).toHaveBeenCalledWith(2); + }); + + it('clicking the already-set star clears the rating to 0', async () => { + const onChange = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('2 stars')); + expect(onChange).toHaveBeenCalledWith(0); + }); + + it('readonly disables all buttons and omits the title', () => { + render(); + for (const button of screen.getAllByRole('button')) { + expect(button).toBeDisabled(); + expect(button).not.toHaveAttribute('title'); + } + }); + + it('click does not propagate to a parent handler', async () => { + const onParentClick = vi.fn(); + const user = userEvent.setup(); + render( +
+ +
, + ); + await user.click(screen.getByLabelText('1 star')); + expect(onParentClick).not.toHaveBeenCalled(); + }); +}); diff --git a/frontend/src/components/Browse/hooks/useBrowseData.test.ts b/frontend/src/components/Browse/hooks/useBrowseData.test.ts new file mode 100644 index 0000000..a5969dc --- /dev/null +++ b/frontend/src/components/Browse/hooks/useBrowseData.test.ts @@ -0,0 +1,439 @@ +/** + * useBrowseData — hook tests for the Browse data-fetching/state layer. + * Fetch is stubbed per-endpoint; renderHook drives the hook directly (no DOM). + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, renderHook, waitFor } from '@testing-library/react'; +import { useBrowseData, type BrowseItem } from './useBrowseData'; + +function jsonResponse(body: unknown, ok = true) { + return { + ok, + status: ok ? 200 : 500, + json: async () => body, + text: async () => (typeof body === 'string' ? body : JSON.stringify(body)), + } as Response; +} + +/** Routes fetch calls to per-endpoint handlers by pathname; unmatched calls 404. */ +function makeFetchMock( + overrides: Record Response | Promise> = {}, +) { + return vi.fn(async (input: RequestInfo | URL) => { + const url = new URL(String(input), 'http://localhost'); + for (const [path, handler] of Object.entries(overrides)) { + if (url.pathname === path) return handler(url); + } + return jsonResponse({}, false); + }); +} + +const item = (path: string, extra: Partial = {}): BrowseItem => ({ + path, + sample: path, + metadata: {}, + ...extra, +}); + +beforeEach(() => { + vi.spyOn(console, 'warn').mockImplementation(() => {}); +}); + +afterEach(() => { + cleanupFakeTimers(); + vi.restoreAllMocks(); +}); + +function cleanupFakeTimers() { + if (vi.isFakeTimers()) vi.useRealTimers(); +} + +describe('useBrowseData', () => { + it('starts with facetsLoading true and loads facets on mount, marking connected', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['sample_name', 'technique'] }), + }); + + const { result } = renderHook(() => useBrowseData('http://server', 'All')); + + expect(result.current.state.facetsLoading).toBe(true); + expect(result.current.state.connectionStatus).toBe('loading'); + + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + expect(result.current.state.facets).toEqual(['sample_name', 'technique']); + expect(result.current.state.connectionStatus).toBe('connected'); + }); + + it('marks the connection disconnected when the facets request fails', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse('boom', false), + }); + + const { result } = renderHook(() => useBrowseData('http://server', 'All')); + + await waitFor(() => expect(result.current.state.connectionStatus).toBe('disconnected')); + expect(result.current.state.facetsLoading).toBe(false); + }); + + it('requests facets with technique/container_path/server params', async () => { + const fetchMock = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + }); + global.fetch = fetchMock; + + renderHook(() => useBrowseData('http://server', 'SAXS', 'secret-key', 'browse/foo')); + + await waitFor(() => expect(fetchMock).toHaveBeenCalled()); + const calledUrl = new URL(String(fetchMock.mock.calls[0][0]), 'http://localhost'); + expect(calledUrl.searchParams.get('technique')).toBe('SAXS'); + expect(calledUrl.searchParams.get('container_path')).toBe('browse/foo'); + expect(calledUrl.searchParams.get('server_uri')).toBe('http://server'); + expect(calledUrl.searchParams.get('server_api_key')).toBe('secret-key'); + // technique !== 'All' so no forced refresh. + expect(calledUrl.searchParams.get('refresh')).toBeNull(); + }); + + it('forces a refresh for technique "All"', async () => { + const fetchMock = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + }); + global.fetch = fetchMock; + + renderHook(() => useBrowseData('http://server', 'All')); + + await waitFor(() => expect(fetchMock).toHaveBeenCalled()); + const calledUrl = new URL(String(fetchMock.mock.calls[0][0]), 'http://localhost'); + expect(calledUrl.searchParams.get('refresh')).toBe('true'); + }); + + it('addColumn appends a loading column then fills in its values', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['field1'] }), + '/api/browse/column': () => + jsonResponse({ values: [{ value: 'v1', count: 3, sample_paths: [] }] }), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + act(() => result.current.actions.addColumn('field1')); + expect(result.current.state.columns).toHaveLength(1); + expect(result.current.state.columns[0]).toMatchObject({ field: 'field1', loading: true }); + + await waitFor(() => expect(result.current.state.columns[0].loading).toBe(false)); + expect(result.current.state.columns[0].values).toEqual([ + { value: 'v1', count: 3, sample_paths: [] }, + ]); + }); + + it('records a column error when the column request fails', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['field1'] }), + '/api/browse/column': () => jsonResponse('nope', false), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + act(() => result.current.actions.addColumn('field1')); + await waitFor(() => expect(result.current.state.columns[0].loading).toBe(false)); + expect(result.current.state.columns[0].error).toBe('HTTP 500'); + }); + + it('selectValue on the last column loads leaf items', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['field1'] }), + '/api/browse/column': () => + jsonResponse({ values: [{ value: 'v1', count: 1, sample_paths: [] }] }), + '/api/browse/items': () => + jsonResponse({ items: [item('p1', { sample: 's1' })], total: 1 }), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + act(() => result.current.actions.addColumn('field1')); + await waitFor(() => expect(result.current.state.columns[0].loading).toBe(false)); + + act(() => result.current.actions.selectValue(0, 'v1')); + expect(result.current.state.columns[0].selected).toBe('v1'); + + await waitFor(() => expect(result.current.state.itemsLoading).toBe(false)); + expect(result.current.state.items).toHaveLength(1); + expect(result.current.state.items[0].path).toBe('p1'); + expect(result.current.state.itemsTotal).toBe(1); + }); + + it('selectValue(null) clears the selection and items without fetching items', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['field1'] }), + '/api/browse/column': () => + jsonResponse({ values: [{ value: 'v1', count: 1, sample_paths: [] }] }), + '/api/browse/items': () => jsonResponse({ items: [item('p1')], total: 1 }), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + act(() => result.current.actions.addColumn('field1')); + await waitFor(() => expect(result.current.state.columns[0].loading).toBe(false)); + act(() => result.current.actions.selectValue(0, 'v1')); + await waitFor(() => expect(result.current.state.items).toHaveLength(1)); + + act(() => result.current.actions.selectValue(0, null)); + expect(result.current.state.columns[0].selected).toBeNull(); + expect(result.current.state.items).toHaveLength(0); + expect(result.current.state.itemsTotal).toBe(0); + }); + + it('selecting a value with a next column loads that column instead of items', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['f1', 'f2'] }), + '/api/browse/column': (url) => { + const field = url.searchParams.get('field'); + if (field === 'f1') return jsonResponse({ values: [{ value: 'a', count: 2, sample_paths: [] }] }); + return jsonResponse({ values: [{ value: 'b', count: 1, sample_paths: [] }] }); + }, + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + act(() => result.current.actions.addColumn('f1')); + await waitFor(() => expect(result.current.state.columns[0].loading).toBe(false)); + act(() => result.current.actions.addColumn('f2')); + await waitFor(() => expect(result.current.state.columns[1]?.loading).toBe(false)); + + act(() => result.current.actions.selectValue(0, 'a')); + // Second column reloads (filtered by field1=a) rather than leaf items loading. + await waitFor(() => expect(result.current.state.columns[1].loading).toBe(false)); + expect(result.current.state.columns[1].selected).toBeNull(); + expect(result.current.state.columns[1].values).toEqual([{ value: 'b', count: 1, sample_paths: [] }]); + expect(result.current.state.itemsLoading).toBe(false); + expect(result.current.state.items).toHaveLength(0); + }); + + it('removeColumn drops that column and everything after it', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['f1', 'f2'] }), + '/api/browse/column': () => jsonResponse({ values: [] }), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + act(() => result.current.actions.addColumn('f1')); + await waitFor(() => expect(result.current.state.columns[0].loading).toBe(false)); + act(() => result.current.actions.addColumn('f2')); + await waitFor(() => expect(result.current.state.columns[1]?.loading).toBe(false)); + + act(() => result.current.actions.removeColumn(1)); + expect(result.current.state.columns).toHaveLength(1); + + act(() => result.current.actions.removeColumn(0)); + expect(result.current.state.columns).toHaveLength(0); + }); + + it('changeColumnField replaces a column and drops later ones', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['f1', 'f2', 'f3'] }), + '/api/browse/column': (url) => + jsonResponse({ values: [{ value: `v-${url.searchParams.get('field')}`, count: 1, sample_paths: [] }] }), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + act(() => result.current.actions.addColumn('f1')); + await waitFor(() => expect(result.current.state.columns[0].loading).toBe(false)); + act(() => result.current.actions.addColumn('f2')); + await waitFor(() => expect(result.current.state.columns[1]?.loading).toBe(false)); + + act(() => result.current.actions.changeColumnField(1, 'f3')); + await waitFor(() => expect(result.current.state.columns[1]?.loading).toBe(false)); + expect(result.current.state.columns).toHaveLength(2); + expect(result.current.state.columns[1].field).toBe('f3'); + expect(result.current.state.columns[1].values[0].value).toBe('v-f3'); + }); + + it('showAll clears columns, sets showingAll, and loads every item', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['f1'] }), + '/api/browse/items': () => jsonResponse({ items: [item('a'), item('b')], total: 2 }), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + act(() => result.current.actions.showAll()); + expect(result.current.state.showingAll).toBe(true); + expect(result.current.state.columns).toHaveLength(0); + + await waitFor(() => expect(result.current.state.itemsLoading).toBe(false)); + expect(result.current.state.items).toHaveLength(2); + expect(result.current.state.itemsTotal).toBe(2); + }); + + it('items fetch failure clears items and total', async () => { + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/items': () => jsonResponse('nope', false), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + act(() => result.current.actions.showAll()); + + await waitFor(() => expect(result.current.state.itemsLoading).toBe(false)); + expect(result.current.state.items).toEqual([]); + expect(result.current.state.itemsTotal).toBe(0); + }); + + it('expandSample loads slices for a multi-image item, and null collapses it', async () => { + const dataset = item('vol1', { sample: 'vol1', n_slices: 3 }); + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/slices': () => + jsonResponse({ items: [item('vol1/0'), item('vol1/1')] }), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + act(() => result.current.actions.expandSample(dataset)); + expect(result.current.state.expandedSample).toEqual(dataset); + expect(result.current.state.slicesLoading).toBe(true); + + await waitFor(() => expect(result.current.state.slicesLoading).toBe(false)); + expect(result.current.state.slices).toHaveLength(2); + + act(() => result.current.actions.expandSample(null)); + expect(result.current.state.expandedSample).toBeNull(); + expect(result.current.state.slices).toHaveLength(0); + }); + + it('slices fetch failure clears the slices list', async () => { + const dataset = item('vol1', { n_slices: 2 }); + global.fetch = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/slices': () => jsonResponse('boom', false), + }); + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + act(() => result.current.actions.expandSample(dataset)); + + await waitFor(() => expect(result.current.state.slicesLoading).toBe(false)); + expect(result.current.state.slices).toEqual([]); + }); + + it('selectItem sets and clears the selected leaf item', async () => { + global.fetch = makeFetchMock({ '/api/browse/facets': () => jsonResponse({ facets: [] }) }); + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + const it1 = item('p1'); + act(() => result.current.actions.selectItem(it1)); + expect(result.current.state.selectedItem).toEqual(it1); + act(() => result.current.actions.selectItem(null)); + expect(result.current.state.selectedItem).toBeNull(); + }); + + it('refresh reloads columns and items using the latest state', async () => { + const fetchMock = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['f1'] }), + '/api/browse/column': () => jsonResponse({ values: [{ value: 'a', count: 1, sample_paths: [] }] }), + '/api/browse/items': () => jsonResponse({ items: [item('a')], total: 1 }), + }); + global.fetch = fetchMock; + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + act(() => result.current.actions.addColumn('f1')); + await waitFor(() => expect(result.current.state.columns[0].loading).toBe(false)); + act(() => result.current.actions.selectValue(0, 'a')); + await waitFor(() => expect(result.current.state.itemsLoading).toBe(false)); + + const callsBefore = fetchMock.mock.calls.length; + act(() => result.current.actions.refresh()); + await waitFor(() => expect(fetchMock.mock.calls.length).toBeGreaterThan(callsBefore)); + // One call for the (only) column, one for items. + expect(fetchMock.mock.calls.length).toBe(callsBefore + 2); + }); + + it('refresh with showingAll re-loads items with no filters', async () => { + const fetchMock = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: [] }), + '/api/browse/items': () => jsonResponse({ items: [item('a')], total: 1 }), + }); + global.fetch = fetchMock; + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + act(() => result.current.actions.showAll()); + await waitFor(() => expect(result.current.state.itemsLoading).toBe(false)); + + const callsBefore = fetchMock.mock.calls.length; + act(() => result.current.actions.refresh()); + await waitFor(() => expect(fetchMock.mock.calls.length).toBe(callsBefore + 1)); + }); + + it('refresh is a no-op when there are no columns and showingAll is false', async () => { + const fetchMock = makeFetchMock({ '/api/browse/facets': () => jsonResponse({ facets: [] }) }); + global.fetch = fetchMock; + + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + const callsBefore = fetchMock.mock.calls.length; + act(() => result.current.actions.refresh()); + expect(fetchMock.mock.calls.length).toBe(callsBefore); + }); + + it('sets up a facets poll interval and clears it on unmount', () => { + global.fetch = makeFetchMock({ '/api/browse/facets': () => jsonResponse({ facets: [] }) }); + const setSpy = vi.spyOn(global, 'setInterval'); + const clearSpy = vi.spyOn(global, 'clearInterval'); + + const { unmount } = renderHook(() => useBrowseData('uri', 'tech')); + expect(setSpy).toHaveBeenCalledWith(expect.any(Function), 30_000); + + unmount(); + expect(clearSpy).toHaveBeenCalled(); + }); + + it('re-fetches and resets state when serverUri/technique change', async () => { + const fetchMock = makeFetchMock({ + '/api/browse/facets': () => jsonResponse({ facets: ['f1'] }), + '/api/browse/items': () => jsonResponse({ items: [item('a')], total: 1 }), + }); + global.fetch = fetchMock; + + const { result, rerender } = renderHook( + ({ serverUri, technique }) => useBrowseData(serverUri, technique), + { initialProps: { serverUri: 'uri1', technique: 'All' } }, + ); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + act(() => result.current.actions.showAll()); + await waitFor(() => expect(result.current.state.items).toHaveLength(1)); + + rerender({ serverUri: 'uri2', technique: 'All' }); + // Reset wipes items/columns/showingAll immediately. + expect(result.current.state.items).toHaveLength(0); + expect(result.current.state.showingAll).toBe(false); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + }); + + it('loadFacets({silent: true}) does not toggle facetsLoading', async () => { + global.fetch = makeFetchMock({ '/api/browse/facets': () => jsonResponse({ facets: ['x'] }) }); + const { result } = renderHook(() => useBrowseData('uri', 'tech')); + await waitFor(() => expect(result.current.state.facetsLoading).toBe(false)); + + let sawLoadingTrue = false; + await act(async () => { + const p = result.current.actions.loadFacets({ silent: true }); + sawLoadingTrue = result.current.state.facetsLoading; + await p; + }); + expect(sawLoadingTrue).toBe(false); + expect(result.current.state.facetsLoading).toBe(false); + expect(result.current.state.facets).toEqual(['x']); + }); +}); diff --git a/frontend/src/components/CompositionPanel/CompositionPanel.test.tsx b/frontend/src/components/CompositionPanel/CompositionPanel.test.tsx new file mode 100644 index 0000000..4ccf439 --- /dev/null +++ b/frontend/src/components/CompositionPanel/CompositionPanel.test.tsx @@ -0,0 +1,243 @@ +/** + * CompositionPanel — smoke tests for concat preview wiring. + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, fireEvent, render, screen, waitFor, within } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import CompositionPanel from './index'; +import { useIpredStore } from '@/stores/ipredStore'; + +vi.mock('@/lib/ipredApi', () => ({ + listIpredModules: vi.fn(async () => [ + { + id: 'skimage_multiscale', + name: 'Skimage multiscale', + description: 'skimage', + runtime: 'numpy', + ready: true, + accepts_input_from: false, + produces_channels: true, + produces_embedding: false, + params_schema: { sigma_min: { default: 1 } }, + }, + { + id: 'tomojepa', + name: 'TomoJEPA', + description: 'mark', + runtime: 'torch', + ready: true, + accepts_input_from: true, + produces_channels: false, + produces_embedding: true, + params_schema: { weights_id: { default: 'mark25' }, input_size: { default: 512 } }, + }, + { + id: 'pca', + name: 'PCA reduce', + description: 'pca', + runtime: 'numpy', + ready: true, + accepts_input_from: true, + produces_channels: true, + produces_embedding: false, + params_schema: { dims: { default: 64 } }, + }, + ]), + listIpredCompositions: vi.fn(async () => [ + { + id: 'comp-skimage', + name: 'Skimage multiscale', + builtin: true, + nodes: [{ id: 'n1', module: 'skimage_multiscale', params: {} }], + outputs: ['n1'], + }, + ]), + getIpredComposition: vi.fn(async (id: string) => ({ + id, + name: 'Skimage multiscale', + builtin: true, + nodes: [{ id: 'n1', module: 'skimage_multiscale', params: { sigma_min: 1 } }], + outputs: ['n1'], + preview_labels: ['intensity σ=1'], + })), + previewIpredComposition: vi.fn(async () => ({ + preview_labels: ['intensity σ=1', 'pca0'], + })), + upsertIpredComposition: vi.fn(async (p) => ({ + id: 'comp-custom', + name: p.name, + nodes: p.nodes, + outputs: p.outputs, + })), +})); + +describe('CompositionPanel', () => { + beforeEach(() => { + useIpredStore.setState({ preferredCompositionId: 'comp-skimage' }); + }); + + afterEach(() => { + cleanup(); + vi.clearAllMocks(); + }); + + it('shows concat preview from loaded composition', async () => { + render(); + await waitFor(() => { + expect(screen.getByTestId('ipred-concat-preview')).toHaveTextContent('intensity'); + }); + }); + + it('adds a module from the catalog', async () => { + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-module-tomojepa'); + await user.click(screen.getByTestId('ipred-module-tomojepa')); + await waitFor(() => { + expect(screen.getByTestId('ipred-node-n2')).toBeTruthy(); + }); + }); + + it('a ready module button is enabled; disables when not ready', async () => { + render(); + await waitFor(() => expect(screen.getByTestId('ipred-module-tomojepa')).toBeEnabled()); + // pca is also ready per the mock catalog. + expect(screen.getByTestId('ipred-module-pca')).toBeEnabled(); + }); + + it('removing a node updates the graph summary and clears it from outputs', async () => { + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-node-n1'); + await user.click(within(screen.getByTestId('ipred-node-n1')).getByTitle('Remove')); + await waitFor(() => expect(screen.queryByTestId('ipred-node-n1')).not.toBeInTheDocument()); + expect(screen.getByText(/Graph: empty/)).toBeInTheDocument(); + }); + + it('adding a tomojepa node shows its weights/input-size/resize controls', async () => { + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-module-tomojepa'); + await user.click(screen.getByTestId('ipred-module-tomojepa')); + const node = await screen.findByTestId('ipred-node-n2'); + expect(within(node).getByText('Weights')).toBeInTheDocument(); + expect(within(node).getByText('Input size')).toBeInTheDocument(); + expect(within(node).getByText('Resize')).toBeInTheDocument(); + // accepts_input_from -> shows the "Input from" selector too. + expect(within(node).getByText('Input from')).toBeInTheDocument(); + }); + + it('changing a tomojepa node input_size updates its param', async () => { + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-module-tomojepa'); + await user.click(screen.getByTestId('ipred-module-tomojepa')); + const node = await screen.findByTestId('ipred-node-n2'); + const input = within(node).getByDisplayValue('512'); + fireEvent.change(input, { target: { value: '256' } }); + expect(within(node).getByDisplayValue('256')).toBeInTheDocument(); + }); + + it('toggling a channel-producing node in/out of the output bank', async () => { + const user = userEvent.setup(); + render(); + const node = await screen.findByTestId('ipred-node-n1'); + const checkbox = within(node).getByLabelText('In bank concat') as HTMLInputElement; + expect(checkbox.checked).toBe(true); // n1 (skimage_multiscale) produces channels, in the builtin comp + await user.click(checkbox); + expect(checkbox.checked).toBe(false); + expect(screen.getByText(/Outputs order: none/)).toBeInTheDocument(); + }); + + it('reordering outputs with the up/down arrows', async () => { + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-module-pca'); + await user.click(screen.getByTestId('ipred-module-pca')); // n2, also produces_channels -> added to outputs + await screen.findByTestId('ipred-node-n2'); + expect(screen.getByText(/Outputs order: n1 → n2/)).toBeInTheDocument(); + + const node2 = screen.getByTestId('ipred-node-n2'); + await user.click(within(node2).getByText('↑')); + expect(screen.getByText(/Outputs order: n2 → n1/)).toBeInTheDocument(); + }); + + it('switching the active composition via the dropdown', async () => { + const { getIpredComposition } = await import('@/lib/ipredApi'); + (getIpredComposition as any).mockImplementation(async (id: string) => { + if (id === 'comp-skimage') { + return { + id, name: 'Skimage multiscale', builtin: true, + nodes: [{ id: 'n1', module: 'skimage_multiscale', params: {} }], + outputs: ['n1'], preview_labels: ['intensity σ=1'], + }; + } + return { id, name: 'Other comp', builtin: false, nodes: [], outputs: [], preview_labels: [] }; + }); + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-composition-select'); + // Only one composition in the mocked list, so just confirm the select wiring + // calls through to getIpredComposition with the chosen id. + const select = screen.getByTestId('ipred-composition-select') as HTMLSelectElement; + fireEvent.change(select, { target: { value: 'comp-skimage' } }); + await waitFor(() => expect(getIpredComposition).toHaveBeenCalledWith('comp-skimage')); + }); + + it('save() persists the composition and re-selects it', async () => { + const { upsertIpredComposition } = await import('@/lib/ipredApi'); + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-node-n1'); + await user.click(screen.getByRole('button', { name: /Save/ })); + await waitFor(() => expect(upsertIpredComposition).toHaveBeenCalled()); + const [payload] = (upsertIpredComposition as any).mock.calls[0]; + expect(payload.composition_id).toBe('comp-skimage'); + }); + + it('clone & save omits composition_id and renames with a " copy" suffix', async () => { + const { upsertIpredComposition } = await import('@/lib/ipredApi'); + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-node-n1'); + await user.click(screen.getByRole('button', { name: /Clone & save/ })); + await waitFor(() => expect(upsertIpredComposition).toHaveBeenCalled()); + const [payload] = (upsertIpredComposition as any).mock.calls[0]; + expect(payload.composition_id).toBeUndefined(); + expect(payload.name).toMatch(/ copy$/); + }); + + it('save() surfaces an error message on failure', async () => { + const { upsertIpredComposition } = await import('@/lib/ipredApi'); + (upsertIpredComposition as any).mockRejectedValueOnce(new Error('save failed')); + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-node-n1'); + await user.click(screen.getByRole('button', { name: /^Save/ })); + expect(await screen.findByTestId('ipred-composition-error')).toHaveTextContent('save failed'); + }); + + it('Reload re-fetches modules/compositions', async () => { + const { listIpredModules } = await import('@/lib/ipredApi'); + const user = userEvent.setup(); + render(); + await screen.findByTestId('ipred-node-n1'); + const callsBefore = (listIpredModules as any).mock.calls.length; + await user.click(screen.getByRole('button', { name: /Reload/ })); + await waitFor(() => expect((listIpredModules as any).mock.calls.length).toBeGreaterThan(callsBefore)); + }); + + it('an empty preview shows the placeholder text', async () => { + const { previewIpredComposition } = await import('@/lib/ipredApi'); + (previewIpredComposition as any).mockResolvedValue({ preview_labels: [] }); + const user = userEvent.setup(); + render(); + const node = await screen.findByTestId('ipred-node-n1'); + await user.click(within(node).getByLabelText('In bank concat')); + await waitFor(() => { + expect(screen.getByTestId('ipred-concat-preview')).toHaveTextContent( + 'add channel-producing nodes and mark them as outputs', + ); + }); + }); +}); diff --git a/frontend/src/components/CompositionPanel/index.tsx b/frontend/src/components/CompositionPanel/index.tsx new file mode 100644 index 0000000..28c76e1 --- /dev/null +++ b/frontend/src/components/CompositionPanel/index.tsx @@ -0,0 +1,505 @@ +/** + * CompositionPanel — modular feature composition window for ipred. + */ +import { useCallback, useEffect, useMemo, useState } from 'react'; +import { + ArrowsClockwise, + CircleNotch, + FloppyDisk, + Plus, + Trash, +} from '@phosphor-icons/react'; +import { useIpredStore } from '@/stores/ipredStore'; +import { + getIpredComposition, + listIpredCompositions, + listIpredModules, + previewIpredComposition, + upsertIpredComposition, + type CompositionDoc, + type CompositionNode, + type FeatureModuleInfo, +} from '@/lib/ipredApi'; +import { cn } from '@/lib/utils'; + +const btn = + 'px-2 py-1 rounded border text-[11px] transition-colors disabled:opacity-40 ' + + 'bg-white text-gray-700 border-gray-200 hover:bg-sky-50 hover:border-sky-300'; + +const section = + 'rounded-md border border-gray-200 bg-white p-4 flex flex-col gap-3'; + +function defaultParams(mod: FeatureModuleInfo): Record { + const out: Record = {}; + for (const [key, schema] of Object.entries(mod.params_schema ?? {})) { + const s = schema as { default?: unknown }; + if (s && 'default' in s) out[key] = s.default; + } + return out; +} + +function newNodeId(nodes: CompositionNode[]): string { + let i = nodes.length + 1; + const ids = new Set(nodes.map((n) => n.id)); + while (ids.has(`n${i}`)) i += 1; + return `n${i}`; +} + +/** Build / edit ordered module composition and show concat preview. */ +export default function CompositionPanel() { + const preferredCompositionId = useIpredStore((s) => s.preferredCompositionId); + const setPreferredCompositionId = useIpredStore((s) => s.setPreferredCompositionId); + + const [modules, setModules] = useState([]); + const [compositions, setCompositions] = useState([]); + const [name, setName] = useState('Custom composition'); + const [nodes, setNodes] = useState([]); + const [outputs, setOutputs] = useState([]); + const [previewLabels, setPreviewLabels] = useState([]); + const [editingId, setEditingId] = useState(null); + const [busy, setBusy] = useState(false); + const [error, setError] = useState(null); + + const moduleById = useMemo(() => { + const m = new Map(); + for (const x of modules) m.set(x.id, x); + return m; + }, [modules]); + + const refresh = useCallback(async () => { + setError(null); + try { + const [mods, comps] = await Promise.all([ + listIpredModules(), + listIpredCompositions(), + ]); + setModules(mods); + setCompositions(comps.sort((a, b) => a.name.localeCompare(b.name))); + const preferred = + preferredCompositionId && + comps.find((c) => c.id === preferredCompositionId); + const pick = preferred ?? comps.find((c) => c.id === 'comp-skimage-slimsam') ?? comps[0]; + if (pick) { + const detail = await getIpredComposition(pick.id); + setEditingId(detail.id); + setName(detail.name); + setNodes(detail.nodes ?? []); + setOutputs(detail.outputs ?? []); + setPreviewLabels(detail.preview_labels ?? []); + if (pick.id !== preferredCompositionId) { + setPreferredCompositionId(pick.id); + } + } + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + } + }, [preferredCompositionId, setPreferredCompositionId]); + + useEffect(() => { + void refresh(); + // Intentionally once on mount + when preference externally cleared + // eslint-disable-next-line react-hooks/exhaustive-deps + }, []); + + const selectComposition = async (id: string) => { + setError(null); + setPreferredCompositionId(id); + try { + const detail = await getIpredComposition(id); + setEditingId(detail.id); + setName(detail.name); + setNodes(detail.nodes ?? []); + setOutputs(detail.outputs ?? []); + setPreviewLabels(detail.preview_labels ?? []); + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + } + }; + + const refreshPreview = async ( + nextNodes: CompositionNode[], + nextOutputs: string[], + ) => { + try { + const r = await previewIpredComposition({ + name, + nodes: nextNodes, + outputs: nextOutputs, + }); + setPreviewLabels(r.preview_labels ?? []); + } catch { + setPreviewLabels([]); + } + }; + + const addModule = (moduleId: string) => { + const mod = moduleById.get(moduleId); + if (!mod) return; + const id = newNodeId(nodes); + const node: CompositionNode = { + id, + module: moduleId, + params: defaultParams(mod), + }; + if (mod.accepts_input_from && nodes.length > 0) { + // Prefer last node that can feed an image / emb + node.input_from = nodes[nodes.length - 1].id; + } + const nextNodes = [...nodes, node]; + const nextOutputs = + mod.produces_channels && !outputs.includes(id) + ? [...outputs, id] + : outputs; + setNodes(nextNodes); + setOutputs(nextOutputs); + void refreshPreview(nextNodes, nextOutputs); + }; + + const removeNode = (nid: string) => { + const nextNodes = nodes + .filter((n) => n.id !== nid) + .map((n) => + n.input_from === nid ? { ...n, input_from: undefined } : n, + ); + const nextOutputs = outputs.filter((o) => o !== nid); + setNodes(nextNodes); + setOutputs(nextOutputs); + void refreshPreview(nextNodes, nextOutputs); + }; + + const updateNodeParams = (nid: string, patch: Record) => { + const nextNodes = nodes.map((n) => + n.id === nid ? { ...n, params: { ...(n.params ?? {}), ...patch } } : n, + ); + setNodes(nextNodes); + void refreshPreview(nextNodes, outputs); + }; + + const setInputFrom = (nid: string, inputFrom: string) => { + const nextNodes = nodes.map((n) => + n.id === nid + ? { ...n, input_from: inputFrom || undefined } + : n, + ); + setNodes(nextNodes); + void refreshPreview(nextNodes, outputs); + }; + + const toggleOutput = (nid: string) => { + const next = outputs.includes(nid) + ? outputs.filter((o) => o !== nid) + : [...outputs, nid]; + setOutputs(next); + void refreshPreview(nodes, next); + }; + + const moveOutput = (nid: string, dir: -1 | 1) => { + const i = outputs.indexOf(nid); + if (i < 0) return; + const j = i + dir; + if (j < 0 || j >= outputs.length) return; + const next = [...outputs]; + [next[i], next[j]] = [next[j], next[i]]; + setOutputs(next); + void refreshPreview(nodes, next); + }; + + const save = async (asClone: boolean) => { + setBusy(true); + setError(null); + try { + const saved = await upsertIpredComposition({ + name: asClone ? `${name} copy` : name, + nodes, + outputs, + composition_id: asClone || !editingId ? undefined : editingId, + }); + setPreferredCompositionId(saved.id); + setEditingId(saved.id); + setName(saved.name); + await refresh(); + await selectComposition(saved.id); + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + } finally { + setBusy(false); + } + }; + + const chainSummary = nodes + .map((n) => { + const label = moduleById.get(n.module)?.name ?? n.module; + const arrow = n.input_from ? `←${n.input_from}` : ''; + return `${n.id}:${label}${arrow}`; + }) + .join(' · '); + + return ( +
+
+

Composition

+ +
+ +

+ String modules together. Outputs order is how channels + concatenate into the feature bank. +

+ + + +
+
+ + Modules + + {modules.map((m) => ( + + ))} +
+ +
+ + +

+ Graph: {chainSummary || 'empty'} +

+ +
+ {nodes.map((n) => { + const mod = moduleById.get(n.module); + return ( +
+
+ + {n.id} · {mod?.name ?? n.module} + + +
+ {mod?.accepts_input_from && ( + + )} + {n.module === 'tomojepa' && ( + + )} + {(n.module === 'tomojepa' || n.module === 'pca' || n.module === 'clahe') && ( +
+ {n.module === 'tomojepa' && ( + <> + + + + )} + {n.module === 'pca' && ( + + )} + {n.module === 'clahe' && ( + + )} +
+ )} +
+ + {outputs.includes(n.id) && ( + + + + + )} +
+
+ ); + })} +
+
+
+ +
+ + Bank will concatenate + +

+ {previewLabels.length + ? previewLabels.join(' · ') + : '— add channel-producing nodes and mark them as outputs —'} +

+

+ Outputs order: {outputs.join(' → ') || 'none'} +

+
+ +
+ + +
+ + {error && ( +

+ {error} +

+ )} +
+ ); +} diff --git a/frontend/src/components/HubAppLayout.tsx b/frontend/src/components/HubAppLayout.tsx index e7c4427..2883861 100644 --- a/frontend/src/components/HubAppLayout.tsx +++ b/frontend/src/components/HubAppLayout.tsx @@ -2,6 +2,7 @@ import HubHeader from "@/components/HubHeader"; import HubMainContent from "@/components/HubMainContent"; import HubSidebar from "@/components/HubSidebar"; import { cn } from "@/lib/utils"; +import { useConnectionHealth } from "@/hooks/useConnectionHealth"; import { RouteItem } from "@/types/navigationRouterTypes"; @@ -37,6 +38,11 @@ export default function HubAppLayout ( { onOpenTabSelector }: HubAppLayoutProps) { + // Drives connectionStore's `status` field from anywhere in the app, so + // HubHeader's indicator reflects live Tiled reachability regardless of + // which route is active. + useConnectionHealth(); + return (
{/* Sidebar: fixed overlay so it always receives clicks above any route content */} diff --git a/frontend/src/components/HubHeader.tsx b/frontend/src/components/HubHeader.tsx index 303f2ca..7a1e3cb 100644 --- a/frontend/src/components/HubHeader.tsx +++ b/frontend/src/components/HubHeader.tsx @@ -1,6 +1,8 @@ +import { useNavigate } from 'react-router'; import alsLogo from '@/assets/alsLogo.png'; import { cn } from '@/lib/utils'; -import { Gear } from '@phosphor-icons/react'; +import { Gear, Warning } from '@phosphor-icons/react'; +import { useConnectionStore } from '@/stores/connectionStore'; export type HubHeaderProps = { title?: string; @@ -9,6 +11,37 @@ export type HubHeaderProps = { titleClassName?: string; onOpenTabSelector?: () => void; } + +/** Small persistent Tiled-connection indicator; hidden for local/no connection. */ +function ConnectionStatus() { + const kind = useConnectionStore((s) => s.kind); + const status = useConnectionStore((s) => s.status); + const navigate = useNavigate(); + if (kind !== 'tiled' || status === 'unknown') return null; + + if (status === 'error') { + return ( + + ); + } + return ( + + + Tiled connected + + ); +} + /** HubHeader — top app bar with logo, title, and an optional "Change Tabs" button. */ export default function HubHeader({title="ALS COMPUTING HUB", logoUrl=alsLogo, className, titleClassName, onOpenTabSelector}: HubHeaderProps) { return ( @@ -17,16 +50,19 @@ export default function HubHeader({title="ALS COMPUTING HUB", logoUrl=alsLogo, c ALS logo

{title}

+
+ {onOpenTabSelector && ( )} +
) } \ No newline at end of file diff --git a/frontend/src/components/Ingest/ConflictDialog.test.tsx b/frontend/src/components/Ingest/ConflictDialog.test.tsx new file mode 100644 index 0000000..3fd09f7 --- /dev/null +++ b/frontend/src/components/Ingest/ConflictDialog.test.tsx @@ -0,0 +1,121 @@ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import ConflictDialog, { type IngestConflict } from './ConflictDialog'; + +function makeConflicts(n: number): IngestConflict[] { + return Array.from({ length: n }, (_, i) => ({ filename: `img${i}.tif`, key: `img${i}` })); +} + +function renderDialog(overrides: Partial> = {}) { + const props = { + containerPath: 'browse/foo', + totalFiles: 5, + conflicts: makeConflicts(2), + existingCount: 3, + suggestedContainerPath: 'browse/foo_2', + onReplace: vi.fn(), + onSkip: vi.fn(), + onNewDataset: vi.fn(), + onBrowseExisting: vi.fn(), + onCancel: vi.fn(), + ...overrides, + }; + const utils = render(); + return { ...utils, props }; +} + +afterEach(() => { + cleanup(); +}); + +describe('ConflictDialog', () => { + it('renders the summary sentence with container path and existing count', () => { + renderDialog(); + expect( + screen.getByText((_, el) => el?.textContent === '2 of 5 images are already in browse/foo (3 samples there now).') + ).toBeInTheDocument(); + }); + + it('uses singular wording when totalFiles is 1 and conflicts is 1', () => { + renderDialog({ totalFiles: 1, conflicts: makeConflicts(1), existingCount: 1 }); + expect( + screen.getByText((_, el) => el?.textContent === '1 of 1 image is already in browse/foo (1 sample there now).') + ).toBeInTheDocument(); + }); + + it('omits the existing-count parenthetical when existingCount is 0', () => { + renderDialog({ existingCount: 0 }); + expect( + screen.getByText((_, el) => el?.textContent === '2 of 5 images are already in browse/foo.') + ).toBeInTheDocument(); + }); + + it('lists conflicting filenames, capped at MAX_LISTED=5 with an overflow line', () => { + renderDialog({ conflicts: makeConflicts(7), totalFiles: 10 }); + for (let i = 0; i < 5; i++) { + expect(screen.getByText(`img${i}.tif`)).toBeInTheDocument(); + } + expect(screen.queryByText('img5.tif')).not.toBeInTheDocument(); + expect(screen.getByText('…and 2 more')).toBeInTheDocument(); + }); + + it('shows the new-image count on the Skip button when not all files conflict', () => { + renderDialog({ totalFiles: 5, conflicts: makeConflicts(2) }); + expect(screen.getByText('Ingests only the 3 new images')).toBeInTheDocument(); + }); + + it('disables Skip and shows the all-conflict message when every file conflicts', () => { + renderDialog({ totalFiles: 3, conflicts: makeConflicts(3) }); + const skipBtn = screen.getByRole('button', { name: /Skip the duplicates/ }); + expect(skipBtn).toBeDisabled(); + expect(screen.getByText('Nothing new to add — every image is already here')).toBeInTheDocument(); + }); + + it('shows the suggested container path on the "new dataset" option', () => { + renderDialog({ suggestedContainerPath: 'browse/foo_3' }); + expect(screen.getByText('browse/foo_3')).toBeInTheDocument(); + }); + + it('calls onReplace, onSkip, onNewDataset, onBrowseExisting when their buttons are clicked', async () => { + const user = userEvent.setup(); + const { props } = renderDialog(); + await user.click(screen.getByRole('button', { name: /Replace the existing images/ })); + expect(props.onReplace).toHaveBeenCalledTimes(1); + + await user.click(screen.getByRole('button', { name: /Skip the duplicates/ })); + expect(props.onSkip).toHaveBeenCalledTimes(1); + + await user.click(screen.getByRole('button', { name: /Ingest to a new dataset/ })); + expect(props.onNewDataset).toHaveBeenCalledTimes(1); + + await user.click(screen.getByRole('button', { name: /Browse the existing dataset/ })); + expect(props.onBrowseExisting).toHaveBeenCalledTimes(1); + }); + + it('calls onCancel from the X button, the Cancel button, and clicking the backdrop', async () => { + const user = userEvent.setup(); + const { props } = renderDialog(); + + await user.click(screen.getByLabelText('Close')); + expect(props.onCancel).toHaveBeenCalledTimes(1); + + await user.click(screen.getByText('Cancel')); + expect(props.onCancel).toHaveBeenCalledTimes(2); + }); + + it('does not call onCancel when clicking inside the dialog panel', async () => { + const user = userEvent.setup(); + const { props } = renderDialog(); + await user.click(screen.getByText('This dataset already has these images')); + expect(props.onCancel).not.toHaveBeenCalled(); + }); + + it('calls onCancel when clicking the backdrop directly', async () => { + const user = userEvent.setup(); + const { props, container } = renderDialog(); + const backdrop = container.firstChild as HTMLElement; + await user.click(backdrop); + expect(props.onCancel).toHaveBeenCalledTimes(1); + }); +}); diff --git a/frontend/src/components/Ingest/IngestDropzone.test.tsx b/frontend/src/components/Ingest/IngestDropzone.test.tsx new file mode 100644 index 0000000..2ca584f --- /dev/null +++ b/frontend/src/components/Ingest/IngestDropzone.test.tsx @@ -0,0 +1,457 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, fireEvent, render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import IngestDropzone from './IngestDropzone'; + +function jsonResponse(body: unknown, ok = true, status = 200) { + return { ok, status, json: async () => body, text: async () => (typeof body === 'string' ? body : JSON.stringify(body)) } as Response; +} + +function makeFile(name: string, content = 'x') { + return new File([content], name, { type: 'image/tiff' }); +} + +/** The dropzone renders two hidden — plain picker, then folder picker. */ +function getInputs(container: HTMLElement) { + const inputs = container.querySelectorAll('input[type="file"]'); + return { fileInput: inputs[0] as HTMLInputElement, dirInput: inputs[1] as HTMLInputElement }; +} + +function renderDropzone(overrides: Partial> = {}) { + const props = { + serverUri: 'http://tiled.example', + onBrowse: vi.fn(), + onAnnotate: vi.fn(), + ...overrides, + }; + const utils = render(); + return { ...utils, props }; +} + +/** Route fetch calls by endpoint fragment; each handler may be sync or async. */ +function makeFetchMock(overrides: { + preflight?: (body: any) => Response | Promise; + upload?: (fd: FormData) => Response | Promise; + status?: (jobId: string, call: number) => Response | Promise; +} = {}) { + let statusCalls = 0; + return vi.fn(async (input: RequestInfo | URL, opts?: RequestInit) => { + const url = String(input); + if (url.includes('/api/ingest/preflight')) { + const body = JSON.parse(opts!.body as string); + if (overrides.preflight) return overrides.preflight(body); + return jsonResponse({ container_exists: false, existing_count: 0, conflicts: [], suggested_container_path: `${body.container_path}_2` }); + } + if (url.includes('/api/ingest/upload')) { + const fd = opts!.body as FormData; + if (overrides.upload) return overrides.upload(fd); + return jsonResponse({ job_id: 'job-1' }); + } + if (url.includes('/api/ingest/status/')) { + const jobId = url.split('/api/ingest/status/')[1]; + statusCalls++; + if (overrides.status) return overrides.status(jobId, statusCalls); + return jsonResponse({ state: 'done', total: 1, done: 1, failed: 0, skipped: 0, errors: [], container_path: 'browse/foo' }); + } + throw new Error(`Unhandled fetch: ${url}`); + }); +} + +beforeEach(() => { + vi.stubGlobal('fetch', makeFetchMock()); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('IngestDropzone', () => { + it('shows an error when picking only unsupported files', async () => { + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + // userEvent.upload() enforces the input's `accept` attribute (as a real browser + // would), so an unsupported extension never reaches the file list. Use + // fireEvent.change with a directly-assigned FileList to exercise the + // component's own isSupported() rejection instead. + const file = makeFile('notes.txt'); + Object.defineProperty(fileInput, 'files', { value: [file], writable: true }); + fireEvent.change(fileInput); + expect(await screen.findByText('No supported image files (TIFF, PNG, JPG, NPY).')).toBeInTheDocument(); + }); + + it('picks a supported file, uploads with no conflicts, polls to completion, and shows Browse/Annotate actions', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + status: (_id, call) => + call === 1 + ? jsonResponse({ state: 'running', total: 1, done: 0, failed: 0, skipped: 0, errors: [], container_path: 'browse/scan1' }) + : jsonResponse({ state: 'done', total: 1, done: 1, failed: 0, skipped: 0, errors: [], container_path: 'browse/scan1' }), + }), + ); + const { container, props } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('scan1.tif')); + + await waitFor(() => expect(screen.getByText(/processed…/)).toBeInTheDocument()); + await waitFor( + () => expect(screen.getByText(/Ingested 1 of 1/)).toBeInTheDocument(), + { timeout: 3000 }, + ); + + const browseBtn = screen.getByRole('button', { name: 'Browse this dataset' }); + await user.click(browseBtn); + expect(props.onBrowse).toHaveBeenCalledWith('browse/scan1', 1); + + const annotateBtn = screen.getByRole('button', { name: 'Open in Annotate' }); + await user.click(annotateBtn); + expect(props.onAnnotate).toHaveBeenCalledWith('browse/scan1', 'scan1'); + }, 10000); + + it('derives the destination container from the single file name when left at the default', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + preflight: (body) => { + expect(body.container_path).toBe('browse/myscan'); + return jsonResponse({ container_exists: false, existing_count: 0, conflicts: [], suggested_container_path: 'browse/myscan_2' }); + }, + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('myscan.tif')); + await waitFor(() => expect(screen.getByText(/processed…|Ingested/)).toBeInTheDocument()); + }); + + it('shows the destination field only after expanding "Save uploaded images to"', async () => { + renderDropzone(); + expect(screen.queryByPlaceholderText('browse/my_dataset')).not.toBeInTheDocument(); + const user = userEvent.setup(); + await user.click(screen.getByText(/Save uploaded images to/)); + expect(screen.getByPlaceholderText('browse/my_dataset')).toBeInTheDocument(); + }); + + it('shows the ConflictDialog when preflight reports conflicts, and Replace resumes the upload with on_conflict=replace', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + preflight: () => + jsonResponse({ + container_exists: true, + existing_count: 3, + conflicts: [{ filename: 'dup.tif', key: 'dup' }], + suggested_container_path: 'browse/dup_2', + }), + upload: (fd) => { + expect(fd.get('on_conflict')).toBe('replace'); + return jsonResponse({ job_id: 'job-2' }); + }, + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('dup.tif')); + + expect(await screen.findByText('This dataset already has these images')).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: /Replace the existing image/ })); + + await waitFor(() => expect(screen.queryByText('This dataset already has these images')).not.toBeInTheDocument()); + await waitFor(() => expect(screen.getByText(/processed…|Ingested/)).toBeInTheDocument()); + }); + + it('Cancel on the ConflictDialog dismisses it without uploading', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + preflight: () => + jsonResponse({ + container_exists: true, + existing_count: 1, + conflicts: [{ filename: 'dup.tif', key: 'dup' }], + suggested_container_path: 'browse/dup_2', + }), + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('dup.tif')); + + expect(await screen.findByText('This dataset already has these images')).toBeInTheDocument(); + await user.click(screen.getByText('Cancel')); + expect(screen.queryByText('This dataset already has these images')).not.toBeInTheDocument(); + expect(screen.queryByText(/processed…/)).not.toBeInTheDocument(); + }); + + it('"Browse the existing dataset" on the conflict dialog calls onBrowse with the existing container and count', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + preflight: () => + jsonResponse({ + container_exists: true, + existing_count: 5, + conflicts: [{ filename: 'dup.tif', key: 'dup' }], + suggested_container_path: 'browse/dup_2', + }), + }), + ); + const { container, props } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('dup.tif')); + + expect(await screen.findByText('This dataset already has these images')).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: /Browse the existing dataset/ })); + expect(props.onBrowse).toHaveBeenCalledWith('browse/dup', 5); + expect(screen.queryByText('This dataset already has these images')).not.toBeInTheDocument(); + }); + + it('"Ingest to a new dataset" updates the destination and uploads to the suggested path', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + preflight: () => + jsonResponse({ + container_exists: true, + existing_count: 1, + conflicts: [{ filename: 'dup.tif', key: 'dup' }], + suggested_container_path: 'browse/dup_2', + }), + upload: (fd) => { + expect(fd.get('container_path')).toBe('browse/dup_2'); + expect(fd.get('on_conflict')).toBe('fail'); + return jsonResponse({ job_id: 'job-3' }); + }, + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('dup.tif')); + + expect(await screen.findByText('This dataset already has these images')).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: /Ingest to a new dataset/ })); + await waitFor(() => expect(screen.getByText(/processed…|Ingested/)).toBeInTheDocument()); + await user.click(screen.getByText(/Save uploaded images to/)); + expect(screen.getByDisplayValue('browse/dup_2')).toBeInTheDocument(); + }); + + it('falls back to uploading anyway when the preflight request itself fails', async () => { + const warnSpy = vi.spyOn(console, 'warn').mockImplementation(() => {}); + vi.stubGlobal( + 'fetch', + makeFetchMock({ + preflight: () => jsonResponse('boom', false, 500), + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('scan.tif')); + + await waitFor(() => expect(screen.getByText(/processed…|Ingested/)).toBeInTheDocument()); + expect(warnSpy).toHaveBeenCalled(); + warnSpy.mockRestore(); + }); + + it('shows an error message when starting the upload itself fails', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + upload: () => jsonResponse('server exploded', false, 500), + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('scan.tif')); + + expect(await screen.findByText(/Could not start the upload: server exploded/)).toBeInTheDocument(); + }); + + it('shows the fatal-failure state ("Nothing was ingested") when the job errors out', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + status: () => + jsonResponse({ state: 'error', total: 2, done: 0, failed: 2, skipped: 0, errors: [ + { filename: 'a.tif', kind: 'unreadable', message: '' }, + { filename: 'b.tif', kind: 'unreadable', message: '' }, + ], container_path: 'browse/bad' }), + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, [makeFile('a.tif'), makeFile('b.tif')]); + + await waitFor( + () => expect(screen.getByText(/Nothing was ingested/)).toBeInTheDocument(), + { timeout: 3000 }, + ); + expect(screen.getByText(/2 of 2 failed/)).toBeInTheDocument(); + expect(screen.getByText(/2 images could not be read as images/)).toBeInTheDocument(); + }, 10000); + + it('expands and collapses the filename list for a multi-file error group', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + status: () => + jsonResponse({ state: 'done', total: 2, done: 0, failed: 2, skipped: 0, errors: [ + { filename: 'a.tif', kind: 'unreadable', message: '' }, + { filename: 'b.tif', kind: 'unreadable', message: '' }, + ], container_path: 'browse/bad' }), + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, [makeFile('a.tif'), makeFile('b.tif')]); + + await waitFor( + () => expect(screen.getByText(/2 images could not be read as images/)).toBeInTheDocument(), + { timeout: 3000 }, + ); + expect(screen.queryByText('a.tif')).not.toBeInTheDocument(); + await user.click(screen.getByText('show filenames')); + expect(screen.getByText('a.tif')).toBeInTheDocument(); + expect(screen.getByText('b.tif')).toBeInTheDocument(); + await user.click(screen.getByText('hide filenames')); + expect(screen.queryByText('a.tif')).not.toBeInTheDocument(); + }, 10000); + + it('post-upload backstop: shows ConflictDialog when a completed job reports conflict errors', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + status: () => + jsonResponse({ + state: 'done', + total: 1, + done: 0, + failed: 1, + skipped: 0, + errors: [{ filename: 'dup.tif', kind: 'conflict', message: '' }], + container_path: 'browse/dup', + }), + preflight: () => + jsonResponse({ container_exists: false, existing_count: 0, conflicts: [], suggested_container_path: 'browse/dup_2' }), + }), + ); + const { container } = renderDropzone(); + const { fileInput } = getInputs(container); + const user = userEvent.setup(); + await user.upload(fileInput, makeFile('dup.tif')); + + await waitFor( + () => expect(screen.getByText('This dataset already has these images')).toBeInTheDocument(), + { timeout: 3000 }, + ); + }, 10000); + + it('picking a folder derives the destination from the folder name via webkitRelativePath', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + preflight: (body) => { + expect(body.container_path).toBe('browse/myfolder'); + return jsonResponse({ container_exists: false, existing_count: 0, conflicts: [], suggested_container_path: 'browse/myfolder_2' }); + }, + }), + ); + const { container } = renderDropzone(); + const { dirInput } = getInputs(container); + const file = makeFile('a.tif'); + Object.defineProperty(file, 'webkitRelativePath', { value: 'myfolder/a.tif' }); + const user = userEvent.setup(); + await user.upload(dirInput, file); + await waitFor(() => expect(screen.getByText(/processed…|Ingested/)).toBeInTheDocument()); + }); + + it('dropping a plain FileList (no directory entries) uploads the supported files', async () => { + const { container } = renderDropzone(); + const dropzone = screen.getByText(/Drag an image file or folder/).closest('[role="button"]') as HTMLElement; + const file = makeFile('dropped.tif'); + const dataTransfer = { items: [], files: [file] }; + fireEvent.drop(dropzone, { dataTransfer }); + await waitFor(() => expect(screen.getByText(/processed…|Ingested/)).toBeInTheDocument()); + }); + + it('dropping a directory via webkitGetAsEntry recurses and uploads all nested files', async () => { + vi.stubGlobal( + 'fetch', + makeFetchMock({ + preflight: (body) => { + expect(body.names.sort()).toEqual(['nested.tif', 'top.tif']); + expect(body.container_path).toBe('browse/mydir'); + return jsonResponse({ container_exists: false, existing_count: 0, conflicts: [], suggested_container_path: 'browse/mydir_2' }); + }, + }), + ); + const topFile = makeFile('top.tif'); + const nestedFile = makeFile('nested.tif'); + + const nestedFileEntry = { + isFile: true, + isDirectory: false, + file: (cb: (f: File) => void) => cb(nestedFile), + }; + let subdirRead = false; + const subdirEntry = { + isFile: false, + isDirectory: true, + name: 'subdir', + createReader: () => ({ + readEntries: (cb: (entries: any[]) => void) => { + if (subdirRead) return cb([]); + subdirRead = true; + cb([nestedFileEntry]); + }, + }), + }; + const topFileEntry = { + isFile: true, + isDirectory: false, + file: (cb: (f: File) => void) => cb(topFile), + }; + let rootRead = false; + const rootDirEntry = { + isFile: false, + isDirectory: true, + name: 'mydir', + createReader: () => ({ + readEntries: (cb: (entries: any[]) => void) => { + if (rootRead) return cb([]); + rootRead = true; + cb([topFileEntry, subdirEntry]); + }, + }), + }; + + const { container } = renderDropzone(); + const dropzone = screen.getByText(/Drag an image file or folder/).closest('[role="button"]') as HTMLElement; + const dataTransfer = { + items: [{ webkitGetAsEntry: () => rootDirEntry }], + files: [], + }; + fireEvent.drop(dropzone, { dataTransfer }); + await waitFor(() => expect(screen.getByText(/processed…|Ingested/)).toBeInTheDocument()); + }); + + it('sets and clears the dragging visual state on dragenter/dragleave', () => { + const { container } = renderDropzone(); + const dropzone = screen.getByText(/Drag an image file or folder/).closest('[role="button"]') as HTMLElement; + expect(dropzone.className).toContain('border-white/20'); + fireEvent.dragEnter(dropzone); + expect(dropzone.className).toContain('border-sky-400'); + fireEvent.dragLeave(dropzone); + expect(dropzone.className).toContain('border-white/20'); + }); +}); diff --git a/frontend/src/components/Ingest/ZarrLoader.test.tsx b/frontend/src/components/Ingest/ZarrLoader.test.tsx new file mode 100644 index 0000000..7b4c7de --- /dev/null +++ b/frontend/src/components/Ingest/ZarrLoader.test.tsx @@ -0,0 +1,559 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import ZarrLoader from './ZarrLoader'; + +function jsonResponse(body: unknown, ok = true, status = 200) { + return { ok, status, json: async () => body, text: async () => JSON.stringify(body) } as Response; +} + +const ZARR_INFO = { + name: 'volume.zarr', + path: '/data/volume.zarr', + levels: [ + { path: '0', shape: [100, 512, 512], dtype: 'uint16', n_slices: 100, height: 512, width: 512, downsample: [1, 1, 1] }, + { path: '1', shape: [100, 256, 256], dtype: 'uint16', n_slices: 100, height: 256, width: 256, downsample: [1, 2, 2] }, + ], + full_shape: [100, 512, 512], + dtype: 'uint16', + voxel_size: [1.5, 1.5, 1.5], + voxel_unit: 'um', +}; + +function renderLoader(overrides: Partial> = {}) { + const props = { + serverUri: 'http://tiled.example', + onBrowse: vi.fn(), + onAnnotate: vi.fn(), + ...overrides, + }; + const utils = render(); + return { ...utils, props }; +} + +async function typePathAndInspect(user: ReturnType, path = '/data/volume.zarr') { + const input = screen.getByPlaceholderText('/absolute/path/to/volume.zarr'); + await user.type(input, path); + await user.click(screen.getByRole('button', { name: /Inspect/ })); +} + +beforeEach(() => { + vi.stubGlobal('fetch', vi.fn()); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('ZarrLoader', () => { + it('disables Inspect until a path is entered', () => { + renderLoader(); + expect(screen.getByRole('button', { name: /Inspect/ })).toBeDisabled(); + }); + + it('inspects a path, shows volume info and resolution levels, no conflict', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse(ZARR_INFO); + if (url.includes('/api/zarr/preflight')) return jsonResponse({ exists: false }); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + await typePathAndInspect(user); + + expect(await screen.findByText('volume.zarr')).toBeInTheDocument(); + expect(screen.getByText(/100 × 512 × 512 · uint16 · 1.5 um\/voxel/)).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /Load volume/ })).not.toBeDisabled(); + }); + + it('inspecting via Enter key also triggers the request', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse(ZARR_INFO); + if (url.includes('/api/zarr/preflight')) return jsonResponse({ exists: false }); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + const input = screen.getByPlaceholderText('/absolute/path/to/volume.zarr'); + await user.type(input, '/data/volume.zarr{Enter}'); + expect(await screen.findByText('volume.zarr')).toBeInTheDocument(); + }); + + it('shows an error message when inspect fails with a JSON detail', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse({ detail: 'no such path' }, false, 404); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + await typePathAndInspect(user); + expect(await screen.findByText('no such path')).toBeInTheDocument(); + }); + + it('switching resolution level updates the coarse-level warning', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse(ZARR_INFO); + if (url.includes('/api/zarr/preflight')) return jsonResponse({ exists: false }); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + await typePathAndInspect(user); + await screen.findByText('volume.zarr'); + + // Full res (level 0) selected by default: no coarse warning. + expect(screen.queryByText(/Annotations are stored in full-resolution/)).not.toBeInTheDocument(); + + await user.click(screen.getByRole('button', { name: /1\/2/ })); + expect(await screen.findByText(/Annotations are stored in full-resolution/)).toBeInTheDocument(); + }); + + it('shows a blocking conflict when the destination holds different uploaded data', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse(ZARR_INFO); + if (url.includes('/api/zarr/preflight')) + return jsonResponse({ + exists: true, + existing: { child_count: 4, external: false, sample_name: 'existing', n_images: 4 }, + }); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + await typePathAndInspect(user); + + expect(await screen.findByText(/and holds 4 uploaded images/)).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /Load volume/ })).toBeDisabled(); + // Not an "external" zarr, so no "Replace it" shortcut is offered. + expect(screen.queryByRole('button', { name: 'Replace it' })).not.toBeInTheDocument(); + }); + + it('offers "Replace it" for a conflict with a previously loaded external zarr, and registers on click', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string, opts?: RequestInit) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse(ZARR_INFO); + if (url.includes('/api/zarr/preflight')) + return jsonResponse({ + exists: true, + existing: { child_count: 1, external: true, sample_name: 'volume', n_images: null }, + }); + if (url.includes('/api/zarr/register')) { + const body = JSON.parse(opts!.body as string); + expect(body.on_conflict).toBe('replace'); + return jsonResponse({ tiled_path: 'browse/volume' }); + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + await typePathAndInspect(user); + + expect(await screen.findByText(/previously loaded Zarr/)).toBeInTheDocument(); + // Load volume stays enabled for an external conflict. + expect(screen.getByRole('button', { name: /Load volume/ })).not.toBeDisabled(); + + await user.click(screen.getByRole('button', { name: 'Replace it' })); + expect(await screen.findByText(/Loaded 100 slices/)).toBeInTheDocument(); + }); + + it('registers a new volume (Load volume), shows success, and Annotate/Browse call back with the level path', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string, opts?: RequestInit) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse(ZARR_INFO); + if (url.includes('/api/zarr/preflight')) return jsonResponse({ exists: false }); + if (url.includes('/api/zarr/register')) { + const body = JSON.parse(opts!.body as string); + expect(body).toMatchObject({ + path: '/data/volume.zarr', + container_path: 'browse', + on_conflict: 'fail', + server_uri: 'http://tiled.example', + }); + return jsonResponse({ tiled_path: 'browse/volume' }); + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + const { props } = renderLoader(); + await typePathAndInspect(user); + await screen.findByText('volume.zarr'); + + await user.click(screen.getByRole('button', { name: /Load volume/ })); + expect(await screen.findByText(/Loaded 100 slices — no data was copied\./)).toBeInTheDocument(); + + await user.click(screen.getByRole('button', { name: 'Annotate' })); + expect(props.onAnnotate).toHaveBeenCalledWith('browse/volume/0'); + + await user.click(screen.getByRole('button', { name: 'Browse' })); + expect(props.onBrowse).toHaveBeenCalledWith('browse', 1); + }); + + it('shows an error message when register fails', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse(ZARR_INFO); + if (url.includes('/api/zarr/preflight')) return jsonResponse({ exists: false }); + if (url.includes('/api/zarr/register')) return jsonResponse({ detail: 'disk full' }, false, 500); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + await typePathAndInspect(user); + await screen.findByText('volume.zarr'); + + await user.click(screen.getByRole('button', { name: /Load volume/ })); + expect(await screen.findByText('disk full')).toBeInTheDocument(); + }); + + describe('server-side directory browser', () => { + function mockBrowseFetch(entriesByRel: Record) { + return vi.fn(async (url: string) => { + if (url.includes('/api/local/root')) return jsonResponse({ root: '/data/raw' }); + if (url.includes('/api/local/list')) { + const rel = new URL(url, 'http://x').searchParams.get('rel') ?? ''; + return jsonResponse(entriesByRel[rel] ?? []); + } + throw new Error(`unexpected fetch: ${url}`); + }); + } + + it('opens the browser, fetches the root, and lists directories (files filtered out)', async () => { + (global.fetch as ReturnType).mockImplementation( + mockBrowseFetch({ + '': [ + { name: 'scratch', path: 'scratch', is_dir: true }, + { name: 'readme.txt', path: 'readme.txt', is_dir: false }, + ], + }), + ); + const user = userEvent.setup(); + renderLoader(); + + await user.click(screen.getByRole('button', { name: /Browse…/ })); + expect(await screen.findByText('/data/raw')).toBeInTheDocument(); + expect(screen.getByText('scratch')).toBeInTheDocument(); + expect(screen.queryByText('readme.txt')).not.toBeInTheDocument(); + }); + + it('descends into a plain subfolder, and selecting a .zarr entry fills the path and closes the browser', async () => { + (global.fetch as ReturnType).mockImplementation( + mockBrowseFetch({ + '': [{ name: 'scratch', path: 'scratch', is_dir: true }], + scratch: [{ name: 'ant_m12.zarr', path: 'scratch/ant_m12.zarr', is_dir: true }], + }), + ); + const user = userEvent.setup(); + renderLoader(); + + await user.click(screen.getByRole('button', { name: /Browse…/ })); + await screen.findByText('scratch'); + await user.click(screen.getByText('scratch')); + + expect(await screen.findByText('ant_m12.zarr')).toBeInTheDocument(); + await user.click(screen.getByText('ant_m12.zarr')); + + // Browser closes and the path field now has the full absolute path — + // no more guessing a container-visible path blind. + expect(screen.queryByText('ant_m12.zarr')).not.toBeInTheDocument(); + expect(screen.getByPlaceholderText('/absolute/path/to/volume.zarr')).toHaveValue( + '/data/raw/scratch/ant_m12.zarr', + ); + }); + + it('"Use this folder" selects the currently-browsed directory even without a .zarr suffix', async () => { + (global.fetch as ReturnType).mockImplementation( + mockBrowseFetch({ '': [{ name: 'my_volume', path: 'my_volume', is_dir: true }] }), + ); + const user = userEvent.setup(); + renderLoader(); + + await user.click(screen.getByRole('button', { name: /Browse…/ })); + await user.click(await screen.findByText('my_volume')); + await screen.findByText(/Use this folder/); + await user.click(screen.getByText(/Use this folder/)); + + expect(screen.getByPlaceholderText('/absolute/path/to/volume.zarr')).toHaveValue('/data/raw/my_volume'); + }); + + it('shows an error if the root or a directory listing fails', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/local/root')) return jsonResponse({ detail: 'no access' }, false, 500); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + + await user.click(screen.getByRole('button', { name: /Browse…/ })); + expect(await screen.findByText('no access')).toBeInTheDocument(); + }); + + it('closes the browser when Browse… is clicked again', async () => { + (global.fetch as ReturnType).mockImplementation(mockBrowseFetch({ '': [] })); + const user = userEvent.setup(); + renderLoader(); + + await user.click(screen.getByRole('button', { name: /Browse…/ })); + await screen.findByText('/data/raw'); + await user.click(screen.getByRole('button', { name: /Browse…/ })); + expect(screen.queryByText('/data/raw')).not.toBeInTheDocument(); + }); + + it('an empty directory shows actionable LOCAL_SOURCE_DIR guidance, not a bare "no sub-folders" dead end', async () => { + (global.fetch as ReturnType).mockImplementation(mockBrowseFetch({ '': [] })); + const user = userEvent.setup(); + renderLoader(); + + await user.click(screen.getByRole('button', { name: /Browse…/ })); + expect(await screen.findByText('No sub-folders here.')).toBeInTheDocument(); + expect(screen.getByText(/LOCAL_SOURCE_DIR=\/path\/to\/your\/data/)).toBeInTheDocument(); + }); + + it('overriding the root (e.g. anywhere on disk when run via start_all.sh) browses from there instead', async () => { + const listCalls: { root: string; rel: string }[] = []; + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/local/root')) return jsonResponse({ root: '/data' }); + if (url.includes('/api/local/list')) { + const u = new URL(url, 'http://x'); + const root = u.searchParams.get('root') ?? ''; + const rel = u.searchParams.get('rel') ?? ''; + listCalls.push({ root, rel }); + if (root === '/Users/me/tomo' && rel === '') { + return jsonResponse([{ name: 'ant_m12.zarr', path: 'ant_m12.zarr', is_dir: true }]); + } + return jsonResponse([]); + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + + await user.click(screen.getByRole('button', { name: /Browse…/ })); + await screen.findByText('/data'); + + const rootField = screen.getByPlaceholderText('Root to browse from'); + await user.clear(rootField); + await user.type(rootField, '/Users/me/tomo'); + await user.click(screen.getByRole('button', { name: 'Go' })); + + expect(await screen.findByText('ant_m12.zarr')).toBeInTheDocument(); + expect(await screen.findByText('/Users/me/tomo')).toBeInTheDocument(); + expect(listCalls).toContainEqual({ root: '/Users/me/tomo', rel: '' }); + + await user.click(screen.getByText('ant_m12.zarr')); + expect(screen.getByPlaceholderText('/absolute/path/to/volume.zarr')).toHaveValue( + '/Users/me/tomo/ant_m12.zarr', + ); + }); + + it('sends the granted root on every subsequent list call (breadcrumbs, descending)', async () => { + const listCalls: { root: string; rel: string }[] = []; + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/local/root')) return jsonResponse({ root: '/data' }); + if (url.includes('/api/local/list')) { + const u = new URL(url, 'http://x'); + const root = u.searchParams.get('root') ?? ''; + const rel = u.searchParams.get('rel') ?? ''; + listCalls.push({ root, rel }); + if (rel === '') return jsonResponse([{ name: 'sub', path: 'sub', is_dir: true }]); + return jsonResponse([]); + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + + await user.click(screen.getByRole('button', { name: /Browse…/ })); + await user.click(await screen.findByText('sub')); + + expect(listCalls).toContainEqual({ root: '/data', rel: 'sub' }); + }); + }); + + it('a bare-array store (single level, empty level.path) builds the Tiled path without a trailing slash', async () => { + const BARE_ARRAY_INFO = { + name: 'plain.zarr', + path: '/data/plain.zarr', + levels: [ + { path: '', shape: [6, 10, 12], dtype: 'uint16', n_slices: 6, height: 10, width: 12, downsample: [1, 1, 1] }, + ], + full_shape: [6, 10, 12], + dtype: 'uint16', + voxel_size: null, + voxel_unit: null, + }; + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/zarr/inspect')) return jsonResponse(BARE_ARRAY_INFO); + if (url.includes('/api/zarr/preflight')) return jsonResponse({ exists: false }); + if (url.includes('/api/zarr/register')) return jsonResponse({ tiled_path: 'browse/plain' }); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + const { props } = renderLoader(); + await typePathAndInspect(user, '/data/plain.zarr'); + await screen.findByText('plain.zarr'); + + await user.click(screen.getByRole('button', { name: /Load volume/ })); + await screen.findByText(/Loaded 6 slices — no data was copied\./); + + await user.click(screen.getByRole('button', { name: 'Annotate' })); + expect(props.onAnnotate).toHaveBeenCalledWith('browse/plain'); + }); + + describe('scan folder for Zarr volumes', () => { + async function openBrowserAt(user: ReturnType, root = '/data') { + await user.click(screen.getByRole('button', { name: /Browse…/ })); + await screen.findByText(root); + } + + it('scans the currently-browsed folder and reports registered/skipped/errors', async () => { + const scanCalls: unknown[] = []; + (global.fetch as ReturnType).mockImplementation(async (url: string, init?: RequestInit) => { + if (url.includes('/api/local/root')) return jsonResponse({ root: '/data' }); + if (url.includes('/api/local/list')) return jsonResponse([]); + if (url.includes('/api/scan-datasets')) { + scanCalls.push(JSON.parse(String(init?.body))); + return jsonResponse({ + scanned: 3, + registered: [{ name: 'a.zarr', key: 'a', tiled_path: 'browse/a' }], + skipped: ['b'], + shadowed: [], + errors: [{ name: 'c.zarr', error: 'boom' }], + }); + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + const { props } = renderLoader({ serverUri: 'http://tiled.example' }); + + await openBrowserAt(user); + await user.click(screen.getByRole('button', { name: /Scan folder for datasets/ })); + + expect(await screen.findByText(/Scanned 3 — registered 1 new, skipped 1 already present, 1 failed\./)) + .toBeInTheDocument(); + expect(screen.getByText('a.zarr')).toBeInTheDocument(); + expect(screen.getByText(/c\.zarr/)).toBeInTheDocument(); + expect(screen.getByText(/boom/)).toBeInTheDocument(); + expect(scanCalls).toEqual([ + { scan_root: '/data', container_path: 'browse', server_uri: 'http://tiled.example' }, + ]); + + await user.click(screen.getByRole('button', { name: 'Go to Browse' })); + expect(props.onBrowse).toHaveBeenCalledWith('browse', 1); + }); + + it('scans a descended-into subfolder using its full path, not just the root', async () => { + const scanCalls: unknown[] = []; + (global.fetch as ReturnType).mockImplementation(async (url: string, init?: RequestInit) => { + if (url.includes('/api/local/root')) return jsonResponse({ root: '/data' }); + if (url.includes('/api/local/list')) { + const rel = new URL(url, 'http://x').searchParams.get('rel') ?? ''; + if (rel === '') return jsonResponse([{ name: 'scratch', path: 'scratch', is_dir: true }]); + return jsonResponse([]); + } + if (url.includes('/api/scan-datasets')) { + scanCalls.push(JSON.parse(String(init?.body))); + return jsonResponse({ scanned: 0, registered: [], skipped: [], shadowed: [], errors: [] }); + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + + await openBrowserAt(user); + await user.click(await screen.findByText('scratch')); + await user.click(await screen.findByRole('button', { name: /Scan folder for datasets \(scratch\)/ })); + + await screen.findByText(/Scanned 0/); + expect(scanCalls).toEqual([ + { scan_root: '/data/scratch', container_path: 'browse', server_uri: 'http://tiled.example' }, + ]); + }); + + it('shows an error if the scan request fails', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/local/root')) return jsonResponse({ root: '/data' }); + if (url.includes('/api/local/list')) return jsonResponse([]); + if (url.includes('/api/scan-datasets')) return jsonResponse({ detail: 'no such directory' }, false, 404); + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + + await openBrowserAt(user); + await user.click(screen.getByRole('button', { name: /Scan folder for datasets/ })); + expect(await screen.findByText('no such directory')).toBeInTheDocument(); + }); + + it('does not show "Go to Browse" when nothing new was registered', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/local/root')) return jsonResponse({ root: '/data' }); + if (url.includes('/api/local/list')) return jsonResponse([]); + if (url.includes('/api/scan-datasets')) { + return jsonResponse({ scanned: 1, registered: [], skipped: ['already-there'], shadowed: [], errors: [] }); + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + + await openBrowserAt(user); + await user.click(screen.getByRole('button', { name: /Scan folder for datasets/ })); + await screen.findByText(/Scanned 1/); + expect(screen.queryByRole('button', { name: 'Go to Browse' })).not.toBeInTheDocument(); + }); + + it('shows a shadowed collision distinctly, and retrying it merges into the existing summary', async () => { + const scanCalls: unknown[] = []; + (global.fetch as ReturnType).mockImplementation(async (url: string, init?: RequestInit) => { + if (url.includes('/api/local/root')) return jsonResponse({ root: '/data' }); + if (url.includes('/api/local/list')) return jsonResponse([]); + if (url.includes('/api/scan-datasets')) { + const body = JSON.parse(String(init?.body)); + scanCalls.push(body); + if (body.renames) { + return jsonResponse({ + scanned: 1, + registered: [{ name: 'stack_a', key: 'stack_a_images', tiled_path: 'browse/stack_a_images' }], + skipped: [], + shadowed: [], + errors: [], + }); + } + return jsonResponse({ + scanned: 1, + registered: [], + skipped: [], + shadowed: [ + { name: 'stack_a', key: 'stack_a', existing_kind: 'zarr', suggested_key: 'stack_a_images' }, + ], + errors: [], + }); + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderLoader(); + + await openBrowserAt(user); + await user.click(screen.getByRole('button', { name: /Scan folder for datasets/ })); + + expect(await screen.findByText(/Scanned 1 — registered 0 new, skipped 0 already present, 1 shadowed\./)) + .toBeInTheDocument(); + expect(screen.getByText('stack_a')).toBeInTheDocument(); + expect(screen.getByText(/already registered as a/)).toBeInTheDocument(); + expect(screen.getByText('zarr')).toBeInTheDocument(); + const input = screen.getByDisplayValue('stack_a_images'); + + await user.click(screen.getByRole('button', { name: 'Register as this' })); + + // The shadowed entry is gone once resolved, replaced by a registered + // one (rendered by name, so still "stack_a" — merged in, not dangling). + await waitFor(() => { + expect(screen.queryByRole('button', { name: 'Register as this' })).not.toBeInTheDocument(); + }); + expect(screen.getByText(/registered 1 new/)).toBeInTheDocument(); + expect(scanCalls).toContainEqual( + expect.objectContaining({ renames: { stack_a: 'stack_a_images' } }), + ); + expect(input).toBeDefined(); + }); + }); +}); diff --git a/frontend/src/components/Ingest/ZarrLoader.tsx b/frontend/src/components/Ingest/ZarrLoader.tsx new file mode 100644 index 0000000..1d9030d --- /dev/null +++ b/frontend/src/components/Ingest/ZarrLoader.tsx @@ -0,0 +1,711 @@ +/** + * ZarrLoader — load an on-disk Zarr volume by path, without uploading anything. + * + * The dropzone next to this streams files through the browser; a tomography + * volume is far too large for that (the reference data is 7-56 GB). Tiled can + * read a Zarr store in place, so loading one is really just registering its path: + * nothing is copied, and slices are read lazily as you annotate. + * + * Flow: paste a server-side path -> Inspect (shows the pyramid) -> pick a level + * -> Load. Inspect is separate from Load so you can see what was found before + * anything is written to the catalog. + */ +import { useCallback, useState } from 'react'; +import { Stack, Warning, CheckCircle, MagnifyingGlass, Folder, FolderOpen } from '@phosphor-icons/react'; +import { API_BASE } from '@/config'; + +/** One resolution level of a multiscale volume. */ +interface ZarrLevel { + path: string; + shape: [number, number, number]; + dtype: string; + n_slices: number; + height: number; + width: number; + /** Full-res voxels spanned per voxel here, as [z, y, x]. */ + downsample: [number, number, number]; +} + +interface ZarrInfo { + name: string; + path: string; + levels: ZarrLevel[]; + full_shape: [number, number, number]; + dtype: string; + voxel_size: number[] | null; + voxel_unit: string | null; +} + +/** What already occupies the destination key, when there is a collision. */ +interface ExistingInfo { + child_count: number; + external: boolean; + sample_name: string; + n_images: number | null; +} + +interface ZarrLoaderProps { + serverUri: string; + onBrowse?: (containerPath: string, sampleCount: number) => void; + onAnnotate?: (tiledPath: string) => void; +} + +/** Human-readable byte-ish size of a level, for a sense of scale. */ +function describeLevel(level: ZarrLevel): string { + const [z, y, x] = level.shape; + return `${z} × ${y} × ${x}`; +} + +interface DirEntry { + name: string; + path: string; + is_dir: boolean; +} + +/** Response of POST /api/scan-datasets — bulk-registers every Zarr store AND + * folder of image slices found directly under a directory, for pointing a + * mounted folder of already-reconstructed volumes at Tiled in one action + * instead of one at a time (or one per data kind). */ +interface ScanResult { + scanned: number; + registered: { name: string; key: string; tiled_path: string }[]; + skipped: string[]; + // A candidate whose natural key collides with an UNRELATED, different-kind + // registration sharing the same stem name (e.g. a raw image folder and its + // own already-registered Zarr reconstruction) — nothing registered, so it + // doesn't silently masquerade as "already present". Retry with a different + // key (see `renames`) to keep both. + shadowed: { name: string; key: string; existing_kind: string; suggested_key: string }[]; + errors: { name: string; error: string }[]; +} + +/** Join an absolute root with a root-relative path (as returned by /api/local/list). */ +function joinPath(root: string, rel: string): string { + if (!rel) return root; + return `${root.replace(/\/$/, '')}/${rel}`; +} + +export default function ZarrLoader({ serverUri, onBrowse, onAnnotate }: ZarrLoaderProps) { + const [path, setPath] = useState(''); + const [containerPath, setContainerPath] = useState('browse'); + const [description, setDescription] = useState(''); + const [info, setInfo] = useState(null); + const [levelIdx, setLevelIdx] = useState(0); + const [busy, setBusy] = useState<'inspect' | 'register' | null>(null); + const [error, setError] = useState(null); + const [existing, setExisting] = useState(null); + const [loaded, setLoaded] = useState<{ tiledPath: string; nSlices: number } | null>(null); + + // Server-side directory browser — the path field takes an absolute path on + // the server's own filesystem (nothing is uploaded), which is easy to get + // wrong by typing a path from your own machine when that's NOT the same + // filesystem the server sees (e.g. inside Docker, only what's actually + // bind-mounted is visible). Browsing what the server can really see avoids + // that guesswork, mirroring the same root+relative-listing pattern the + // Connect page's "Local Folder" mode already uses — including letting the + // root itself be overridden: running via start_all.sh (no container), the + // backend is a native process with the same filesystem access as the rest + // of your machine, so granting a different root here searches anywhere + // you'd expect, exactly like "Local Folder" already does. Running via + // Docker, the SAME override mechanism naturally can't escape the + // container's own sandbox — an unmounted host path just comes back empty, + // which is the correct, honest reflection of what's actually mounted + // (see LOCAL_SOURCE_DIR in the deployment docs for making a real directory + // visible there instead of fighting this). + const [browsing, setBrowsing] = useState(false); + const [defaultRoot, setDefaultRoot] = useState(null); + const [rootInput, setRootInput] = useState(''); + const [browseRoot, setBrowseRoot] = useState(null); + const [browseRel, setBrowseRel] = useState(''); + const [browseEntries, setBrowseEntries] = useState([]); + const [browseError, setBrowseError] = useState(null); + const [browseLoading, setBrowseLoading] = useState(false); + const [scanning, setScanning] = useState(false); + const [scanResult, setScanResult] = useState(null); + const [scanError, setScanError] = useState(null); + const [renameInputs, setRenameInputs] = useState>({}); + + const listDir = useCallback(async (root: string, rel: string) => { + setBrowseLoading(true); + setBrowseError(null); + try { + const res = await fetch( + `${API_BASE}/api/local/list?root=${encodeURIComponent(root)}&rel=${encodeURIComponent(rel)}`, + ); + if (!res.ok) throw new Error(await readError(res)); + const entries: DirEntry[] = await res.json(); + setBrowseEntries(entries.filter((e) => e.is_dir)); + setBrowseRoot(root); + setBrowseRel(rel); + } catch (err) { + setBrowseError(err instanceof Error ? err.message : String(err)); + } finally { + setBrowseLoading(false); + } + }, []); + + const openBrowser = useCallback(async () => { + setBrowsing(true); + let root = defaultRoot; + if (root === null) { + try { + const res = await fetch(`${API_BASE}/api/local/root`); + if (!res.ok) throw new Error(await readError(res)); + ({ root } = await res.json()); + setDefaultRoot(root); + setRootInput(root ?? ''); + } catch (err) { + setBrowseError(err instanceof Error ? err.message : String(err)); + return; + } + } + await listDir(browseRoot ?? root ?? '', browseRoot !== null ? browseRel : ''); + }, [defaultRoot, browseRoot, browseRel, listDir]); + + const grantRoot = () => { + const root = rootInput.trim(); + if (root) void listDir(root, ''); + }; + + const chooseDir = (entry: DirEntry) => { + if (browseRoot === null) return; + setPath(joinPath(browseRoot, entry.path)); + setBrowsing(false); + }; + + const runScan = useCallback( + async (renames?: Record) => { + if (browseRoot === null) return null; + const res = await fetch(`${API_BASE}/api/scan-datasets`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + scan_root: joinPath(browseRoot, browseRel), + container_path: containerPath, + server_uri: serverUri, + ...(renames ? { renames } : {}), + }), + }); + if (!res.ok) throw new Error(await readError(res)); + return (await res.json()) as ScanResult; + }, + [browseRoot, browseRel, containerPath, serverUri], + ); + + const scanFolder = useCallback(async () => { + setScanning(true); + setScanError(null); + setScanResult(null); + try { + setScanResult(await runScan()); + } catch (err) { + setScanError(err instanceof Error ? err.message : String(err)); + } finally { + setScanning(false); + } + }, [runScan]); + + /** Retry one shadowed candidate under an alternate key, merging the result + * into the existing summary rather than replacing it (the rest of the + * scan's findings are still valid and shouldn't disappear from view). */ + const retryShadowed = useCallback( + async (name: string, altKey: string) => { + setScanning(true); + setScanError(null); + try { + const retryResult = await runScan({ [name]: altKey }); + if (!retryResult) return; + setScanResult((prev) => { + const base = prev ?? { scanned: 0, registered: [], skipped: [], shadowed: [], errors: [] }; + return { + scanned: base.scanned, + registered: [...base.registered, ...retryResult.registered], + skipped: base.skipped, + shadowed: base.shadowed.filter((s) => s.name !== name), + errors: [...base.errors, ...retryResult.errors], + }; + }); + } catch (err) { + setScanError(err instanceof Error ? err.message : String(err)); + } finally { + setScanning(false); + } + }, + [runScan], + ); + + /** Pull the FastAPI `detail` message out of an error response. */ + const readError = async (res: Response): Promise => { + try { + const body = await res.json(); + if (typeof body?.detail === 'string') return body.detail; + return JSON.stringify(body); + } catch { + return `Request failed (${res.status})`; + } + }; + + const inspect = useCallback(async () => { + setBusy('inspect'); + setError(null); + setInfo(null); + setExisting(null); + setLoaded(null); + try { + const res = await fetch(`${API_BASE}/api/zarr/inspect?path=${encodeURIComponent(path)}`); + if (!res.ok) throw new Error(await readError(res)); + const data: ZarrInfo = await res.json(); + setInfo(data); + setLevelIdx(0); + + // Warn about a name collision before the user commits to loading. + const pf = await fetch(`${API_BASE}/api/zarr/preflight`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ path, container_path: containerPath, server_uri: serverUri }), + }); + if (pf.ok) { + const conflict = await pf.json(); + setExisting(conflict.exists ? conflict.existing : null); + } + } catch (err) { + setError(err instanceof Error ? err.message : String(err)); + } finally { + setBusy(null); + } + }, [path, containerPath, serverUri]); + + const register = useCallback( + async (onConflict: 'fail' | 'replace') => { + if (!info) return; + setBusy('register'); + setError(null); + try { + const res = await fetch(`${API_BASE}/api/zarr/register`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + path: info.path, + container_path: containerPath, + description, + on_conflict: onConflict, + server_uri: serverUri, + }), + }); + if (!res.ok) throw new Error(await readError(res)); + const data = await res.json(); + const level = info.levels[levelIdx]; + // A bare-array store (no OME-NGFF multiscale group) registers as a + // single leaf array node directly at the container key — its level + // has no sub-path to descend into (level.path === ''), so appending + // '/' + '' would produce a wrong, trailing-slash Tiled path. + setLoaded({ + tiledPath: level.path ? `${data.tiled_path}/${level.path}` : data.tiled_path, + nSlices: info.levels[0].n_slices, + }); + setExisting(null); + } catch (err) { + setError(err instanceof Error ? err.message : String(err)); + } finally { + setBusy(null); + } + }, + [info, containerPath, description, serverUri, levelIdx], + ); + + const level = info?.levels[levelIdx]; + const zFactor = level ? level.downsample[0] : 1; + const coarse = levelIdx > 0; + + return ( +
+

+ Point at a .zarr directory on the server. Nothing is + copied — Tiled reads it in place, so even a 50 GB volume loads in seconds. +

+ +
+ setPath(e.target.value)} + onKeyDown={(e) => { if (e.key === 'Enter' && path.trim()) inspect(); }} + placeholder="/absolute/path/to/volume.zarr" + className="flex-1 min-w-0 rounded-md border border-white/15 bg-white/5 px-2 py-1.5 text-sm text-white placeholder:text-white/30" + /> + + +
+ + {browsing && ( +
+
+ setRootInput(e.target.value)} + onKeyDown={(e) => { if (e.key === 'Enter') grantRoot(); }} + placeholder="Root to browse from" + className="flex-1 min-w-0 rounded-md border border-white/15 bg-white/5 px-2 py-1 font-mono text-xs text-white placeholder:text-white/30" + /> + +
+

+ Only what this server can actually see is listed — running locally + (start_all.sh), that's anywhere on + your machine; in Docker, only what's bind-mounted (e.g. via{' '} + LOCAL_SOURCE_DIR). +

+ + {browseRoot !== null && ( +
+ + {browseRel.split('/').filter(Boolean).map((seg, i, arr) => { + const target = arr.slice(0, i + 1).join('/'); + return ( + + / + + + ); + })} +
+ )} + + {browseError && ( +
+ + {browseError} +
+ )} + + {!browseError && ( +
+ {browseLoading && ( +
Loading…
+ )} + {!browseLoading && browseEntries.length === 0 && ( +
+

No sub-folders here.

+

+ If you expected your own data to show up: unlike the image dropzone + above (which uploads through the browser), a Zarr volume has to + already be visible to this server — nothing gets copied. If + you're running this in Docker, that means bind-mounting it in first, + then restarting: +

+
+                    LOCAL_SOURCE_DIR=/path/to/your/data docker compose -f docker-compose.full.yml up -d
+                  
+

+ (or set LOCAL_SOURCE_DIR in a{' '} + .env file at the repo root — + see the Production deployment docs). Not using Docker? Grant a + different root above instead — the server can see anywhere on that + machine. +

+
+ )} + {!browseLoading && browseEntries.map((entry) => { + const isZarr = entry.name.toLowerCase().endsWith('.zarr'); + return ( +
+ + {isZarr && ( + + .zarr + + )} +
+ ); + })} +
+ )} + + {browseRoot !== null && !browseError && ( +
+ + +
+ )} + + {scanError && ( +
+ + {scanError} +
+ )} + + {scanResult && ( +
+
+ + + Scanned {scanResult.scanned} — registered {scanResult.registered.length} new, + skipped {scanResult.skipped.length} already present + {scanResult.shadowed.length > 0 ? `, ${scanResult.shadowed.length} shadowed` : ''} + {scanResult.errors.length > 0 ? `, ${scanResult.errors.length} failed` : ''}. + +
+ {scanResult.registered.length > 0 && ( +
    + {scanResult.registered.map((r) => ( +
  • {r.name}
  • + ))} +
+ )} + {scanResult.shadowed.length > 0 && ( +
+ {scanResult.shadowed.map((s) => ( +
+

+ {s.name} — same name already registered as a{' '} + {s.existing_kind} dataset. Register this one too, under + a different name: +

+
+ setRenameInputs((prev) => ({ ...prev, [s.name]: e.target.value }))} + className="flex-1 min-w-0 rounded-md border border-white/15 bg-white/5 px-2 py-1 font-mono text-xs text-white" + /> + +
+
+ ))} +
+ )} + {scanResult.errors.length > 0 && ( +
    + {scanResult.errors.map((e) => ( +
  • + {e.name}: {e.error} +
  • + ))} +
+ )} + {scanResult.registered.length > 0 && ( + + )} +
+ )} +
+ )} + + {error && ( +
+ + {error} +
+ )} + + {info && ( +
+
+ {info.name} + + {info.full_shape[0]} × {info.full_shape[1]} × {info.full_shape[2]} · {info.dtype} + {info.voxel_size ? ` · ${info.voxel_size[0]} ${info.voxel_unit ?? ''}/voxel` : ''} + +
+ +
+ +
+ {info.levels.map((lv, i) => ( + + ))} +
+ {coarse && level && ( + // Both consequences of annotating a coarse level, stated plainly: + // strokes are stored as if full-res, and z is subsampled too. +

+ Annotations are stored in full-resolution coordinates, so strokes drawn here carry + less precision than they appear to. This level also has {level.n_slices} slices, so + it addresses roughly every {zFactor.toFixed(zFactor % 1 ? 1 : 0)} + th full-resolution slice. +

+ )} +
+ +
+
+ + setContainerPath(e.target.value)} + className="w-full rounded-md border border-white/15 bg-white/5 px-2 py-1 text-xs text-white" + /> +
+
+ + setDescription(e.target.value)} + placeholder="tomography, sand" + className="w-full rounded-md border border-white/15 bg-white/5 px-2 py-1 text-xs text-white placeholder:text-white/30" + /> +
+
+ + {existing && ( +
+
+ + + {info.name.replace(/\.zarr$/, '')} already exists in{' '} + {containerPath} + {existing.external + ? ' as a previously loaded Zarr. Replacing it only updates the catalog entry.' + : ` and holds ${existing.child_count} uploaded images. That is different data — + replacing it would delete those files, so choose another destination.`} + +
+ {existing.external && ( + + )} +
+ )} + + {!loaded && ( + + )} + + {loaded && ( +
+
+ + Loaded {loaded.nSlices} slices — no data was copied. +
+
+ + +
+
+ )} +
+ )} +
+ ); +} diff --git a/frontend/src/components/annotate/AnnotationCanvas/OverlaysLayer.test.tsx b/frontend/src/components/annotate/AnnotationCanvas/OverlaysLayer.test.tsx new file mode 100644 index 0000000..a7478a5 --- /dev/null +++ b/frontend/src/components/annotate/AnnotationCanvas/OverlaysLayer.test.tsx @@ -0,0 +1,217 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render } from '@testing-library/react'; +import { Stage } from 'react-konva'; +import OverlaysLayer, { type OverlaysLayerProps } from './OverlaysLayer'; +import type { ManifoldPoint } from '@/lib/featureManifold'; + +/** + * Shallow smoke tests: jsdom has no real 2D canvas, so Konva (which every + * Stage/Layer/shape ultimately talks to) needs `HTMLCanvasElement.getContext` + * stubbed before anything renders — otherwise Konva's own canvas setup throws. + * The stub is a Proxy over a handful of explicitly-typed methods/properties + * (fillStyle, drawImage, save/restore, path building, gradients, …) that falls + * back to a `vi.fn()` for anything unlisted, so an internal Konva call this + * suite didn't anticipate returns a harmless mock instead of `undefined is not + * a function`. + */ +function createMockContext(canvas: HTMLCanvasElement) { + const base: Record = { + canvas, + fillStyle: '#000000', + strokeStyle: '#000000', + lineWidth: 1, + lineCap: 'butt', + lineJoin: 'miter', + globalAlpha: 1, + globalCompositeOperation: 'source-over', + font: '10px sans-serif', + textAlign: 'start', + textBaseline: 'alphabetic', + imageSmoothingEnabled: true, + filter: 'none', + fillRect: vi.fn(), + clearRect: vi.fn(), + strokeRect: vi.fn(), + drawImage: vi.fn(), + putImageData: vi.fn(), + getImageData: vi.fn(() => ({ data: new Uint8ClampedArray(4), width: 1, height: 1 })), + createImageData: vi.fn((w: number, h: number) => ({ + data: new Uint8ClampedArray(Math.max(1, w) * Math.max(1, h) * 4), + width: w, + height: h, + })), + save: vi.fn(), + restore: vi.fn(), + translate: vi.fn(), + scale: vi.fn(), + rotate: vi.fn(), + transform: vi.fn(), + setTransform: vi.fn(), + resetTransform: vi.fn(), + beginPath: vi.fn(), + closePath: vi.fn(), + moveTo: vi.fn(), + lineTo: vi.fn(), + bezierCurveTo: vi.fn(), + quadraticCurveTo: vi.fn(), + arc: vi.fn(), + arcTo: vi.fn(), + ellipse: vi.fn(), + rect: vi.fn(), + fill: vi.fn(), + stroke: vi.fn(), + clip: vi.fn(), + isPointInPath: vi.fn(() => false), + measureText: vi.fn(() => ({ width: 0 })), + fillText: vi.fn(), + strokeText: vi.fn(), + createLinearGradient: vi.fn(() => ({ addColorStop: vi.fn() })), + createRadialGradient: vi.fn(() => ({ addColorStop: vi.fn() })), + createPattern: vi.fn(() => ({})), + setLineDash: vi.fn(), + getLineDash: vi.fn(() => []), + }; + return new Proxy(base, { + get(target, prop) { + if (prop in target) return (target as Record)[prop as string]; + return vi.fn(); + }, + set(target, prop, value) { + (target as Record)[prop as string] = value; + return true; + }, + }); +} + +beforeEach(() => { + HTMLCanvasElement.prototype.getContext = vi.fn(function (this: HTMLCanvasElement) { + return createMockContext(this) as unknown as CanvasRenderingContext2D; + }) as unknown as typeof HTMLCanvasElement.prototype.getContext; + HTMLCanvasElement.prototype.toDataURL = vi.fn(() => 'data:,'); + vi.stubGlobal( + 'ResizeObserver', + class { + observe() {} + unobserve() {} + disconnect() {} + }, + ); +}); + +afterEach(() => { + cleanup(); + vi.restoreAllMocks(); + vi.unstubAllGlobals(); +}); + +function baseProps(overrides: Partial = {}): OverlaysLayerProps { + return { + width: 256, + height: 256, + showFeatures: false, + featureChannelUrl: null, + showProba: false, + probaOverlayUrl: null, + probaOpacity: 0.5, + showPredictions: false, + clfCommitUrl: null, + clfStatusUrl: null, + predictionsOpacity: 0.5, + classColorById: new Map(), + predictionClassVisible: {}, + showPredictionMulti: true, + showPredictionAbstain: true, + showManifold: false, + manifoldHeatmapUrl: null, + manifoldHeatmapOpacity: 0.45, + manifoldShowHeatmap: true, + manifoldMarkers: [], + manifoldShowMarkers: true, + manifoldBoxSize: 64, + ...overrides, + }; +} + +/** Every Konva shape/image ultimately needs a Stage ancestor. */ +function renderInStage(props: OverlaysLayerProps) { + return render( + + + , + ); +} + +describe('OverlaysLayer', () => { + it('renders no Konva layer (and no canvas) when every overlay is off', () => { + const { container } = renderInStage(baseProps()); + // OverlaysLayer renders nothing with every flag off — no means no + // canvas gets created at all, which is a real, assertable fact here. + expect(container.querySelectorAll('canvas').length).toBe(0); + }); + + it('renders without throwing when every overlay flag is on but no data has loaded yet', () => { + // urls are set but jsdom's /Image never fires onload, so the feature/ + // proba/prediction overlays stay un-rendered — this only proves the pending + // state doesn't crash. The manifold Layer, unlike those, is gated on + // `showManifold` alone (its heatmap/marker children are conditional inside + // it), so it mounts as one empty canvas even before anything decodes. + const { container } = renderInStage( + baseProps({ + showFeatures: true, + featureChannelUrl: 'blob:feature', + showProba: true, + probaOverlayUrl: 'blob:proba', + showPredictions: true, + clfCommitUrl: 'blob:commit', + clfStatusUrl: 'blob:status', + showManifold: true, + manifoldHeatmapUrl: 'blob:manifold', + manifoldShowMarkers: false, + }), + ); + expect(container.querySelectorAll('canvas').length).toBe(1); + }); + + it('renders manifold marker rects immediately (they need no image decode)', () => { + const { container } = renderInStage( + baseProps({ + showManifold: true, + manifoldShowHeatmap: false, + manifoldShowMarkers: true, + manifoldMarkers: [ + { x: 10, y: 10, cluster: 0 }, + { x: 50, y: 60, cluster: 1 }, + ] satisfies ManifoldPoint[], + }), + ); + // The markers are drawn on a non-listening Layer, which Konva backs with + // exactly one scene canvas (no hit canvas since nothing needs hit-testing). + expect(container.querySelectorAll('canvas').length).toBe(1); + }); + + it('adding an active layer increases the number of canvases Konva maintains', () => { + const off = renderInStage(baseProps()); + const offCount = off.container.querySelectorAll('canvas').length; + cleanup(); + + const on = renderInStage( + baseProps({ + showManifold: true, + manifoldShowMarkers: true, + manifoldMarkers: [{ x: 1, y: 1, cluster: 0 }] satisfies ManifoldPoint[], + }), + ); + const onCount = on.container.querySelectorAll('canvas').length; + expect(onCount).toBeGreaterThan(offCount); + }); + + it('does not throw across width/height prop changes (re-render)', () => { + const { rerender, container } = renderInStage(baseProps({ width: 100, height: 100 })); + rerender( + + + , + ); + expect(container.querySelectorAll('canvas').length).toBeGreaterThan(0); + }); +}); diff --git a/frontend/src/components/annotate/AnnotationCanvas/OverlaysLayer.tsx b/frontend/src/components/annotate/AnnotationCanvas/OverlaysLayer.tsx new file mode 100644 index 0000000..f012cef --- /dev/null +++ b/frontend/src/components/annotate/AnnotationCanvas/OverlaysLayer.tsx @@ -0,0 +1,216 @@ +/** + * OverlaysLayer — iPred proba / conformal-prediction / manifold-suggest overlays. + * + * Extracted from AnnotationCanvas so its PNG-decode work (fetch → ImageData → + * recolor) stays isolated from the main canvas's per-frame render path. Each + * overlay is gated by the matching `layerVisibilityStore` group, already wired + * up in LayersPanel — this component only needs to render when told to. + */ +import { useEffect, useRef, useState } from 'react'; +import { Layer, Image as KonvaImage, Rect } from 'react-konva'; +import { colorizeConformalOverlay, loadLabelPng } from '@/lib/pixelClf'; +import { colorizeManifoldHeatmap } from '@/lib/featureManifold'; +import type { ManifoldPoint } from '@/lib/featureManifold'; +import { manifoldMarkerRect } from '@/lib/featureManifold'; + +/** Loads a blob: URL into an element; null while loading/absent. */ +function useImageFromUrl(url: string | null): HTMLImageElement | null { + const [img, setImg] = useState(null); + useEffect(() => { + if (!url) { + setImg(null); + return; + } + let cancelled = false; + const el = new Image(); + el.onload = () => { + if (!cancelled) setImg(el); + }; + el.src = url; + return () => { + cancelled = true; + }; + }, [url]); + return img; +} + +export interface OverlaysLayerProps { + width: number; + height: number; + imageClip?: { clipX?: number; clipY?: number; clipWidth?: number; clipHeight?: number }; + + showFeatures: boolean; + featureChannelUrl: string | null; + + showProba: boolean; + probaOverlayUrl: string | null; + probaOpacity: number; + + showPredictions: boolean; + clfCommitUrl: string | null; + clfStatusUrl: string | null; + predictionsOpacity: number; + classColorById: Map; + predictionClassVisible: Record; + showPredictionMulti: boolean; + showPredictionAbstain: boolean; + + showManifold: boolean; + manifoldHeatmapUrl: string | null; + manifoldHeatmapOpacity: number; + manifoldShowHeatmap: boolean; + manifoldMarkers: ManifoldPoint[]; + manifoldShowMarkers: boolean; + manifoldBoxSize: number; +} + +/** Fetches + decodes the conformal commit/status PNGs into a colorized canvas. */ +function useConformalOverlayCanvas( + commitUrl: string | null, + statusUrl: string | null, + colorById: Map, + classVisible: Record, + showMulti: boolean, + showAbstain: boolean, +): HTMLCanvasElement | null { + const [canvas, setCanvas] = useState(null); + useEffect(() => { + if (!commitUrl || !statusUrl) { + setCanvas(null); + return; + } + let cancelled = false; + (async () => { + try { + const [commit, status] = await Promise.all([loadLabelPng(commitUrl), loadLabelPng(statusUrl)]); + if (cancelled) return; + const out = colorizeConformalOverlay(commit.data, status.data, commit.width, commit.height, colorById, { + classVisible: (cid) => classVisible[cid] !== false, + showMulti, + showAbstain, + }); + if (!cancelled) setCanvas(out); + } catch { + if (!cancelled) setCanvas(null); + } + })(); + return () => { + cancelled = true; + }; + }, [commitUrl, statusUrl, colorById, classVisible, showMulti, showAbstain]); + return canvas; +} + +/** Fetches + decodes the manifold coverage PNG into a colorized canvas. */ +function useManifoldHeatmapCanvas(url: string | null, opacity: number): HTMLCanvasElement | null { + const [canvas, setCanvas] = useState(null); + useEffect(() => { + if (!url) { + setCanvas(null); + return; + } + let cancelled = false; + (async () => { + try { + const { data, width, height } = await loadLabelPng(url); + if (cancelled) return; + setCanvas(colorizeManifoldHeatmap(data, width, height, opacity)); + } catch { + if (!cancelled) setCanvas(null); + } + })(); + return () => { + cancelled = true; + }; + }, [url, opacity]); + return canvas; +} + +export default function OverlaysLayer({ + width, + height, + imageClip, + showFeatures, + featureChannelUrl, + showProba, + probaOverlayUrl, + probaOpacity, + showPredictions, + clfCommitUrl, + clfStatusUrl, + predictionsOpacity, + classColorById, + predictionClassVisible, + showPredictionMulti, + showPredictionAbstain, + showManifold, + manifoldHeatmapUrl, + manifoldHeatmapOpacity, + manifoldShowHeatmap, + manifoldMarkers, + manifoldShowMarkers, + manifoldBoxSize, +}: OverlaysLayerProps) { + const featureImg = useImageFromUrl(showFeatures ? featureChannelUrl : null); + const probaImg = useImageFromUrl(showProba ? probaOverlayUrl : null); + const conformalCanvas = useConformalOverlayCanvas( + showPredictions ? clfCommitUrl : null, + showPredictions ? clfStatusUrl : null, + classColorById, + predictionClassVisible, + showPredictionMulti, + showPredictionAbstain, + ); + const manifoldCanvas = useManifoldHeatmapCanvas( + showManifold && manifoldShowHeatmap ? manifoldHeatmapUrl : null, + manifoldHeatmapOpacity, + ); + const featureRef = useRef(null); + const probaRef = useRef(null); + const predRef = useRef(null); + const manifoldRef = useRef(null); + + return ( + <> + {showFeatures && featureImg && ( + + + + )} + {showProba && probaImg && ( + + + + )} + {showPredictions && conformalCanvas && ( + + + + )} + {showManifold && ( + + {manifoldShowHeatmap && manifoldCanvas && ( + + )} + {manifoldShowMarkers && + manifoldMarkers.map((pt, i) => { + const r = manifoldMarkerRect(pt, { side: manifoldBoxSize, width, height }); + return ( + + ); + })} + + )} + + ); +} diff --git a/frontend/src/components/annotate/AnnotationCanvas/ShapesLayer.test.tsx b/frontend/src/components/annotate/AnnotationCanvas/ShapesLayer.test.tsx new file mode 100644 index 0000000..e775caf --- /dev/null +++ b/frontend/src/components/annotate/AnnotationCanvas/ShapesLayer.test.tsx @@ -0,0 +1,248 @@ +import { createRef } from 'react'; +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render } from '@testing-library/react'; +import { Stage } from 'react-konva'; +import type Konva from 'konva'; +import ShapesLayer, { type ShapesLayerProps } from './ShapesLayer'; +import type { Shape } from '@/stores/annotationStore'; +import type { AnnotationClass } from '@/stores/classStore'; + +/** + * Shallow smoke tests: jsdom has no real 2D canvas, so Konva (which every + * Stage/Layer/shape ultimately talks to) needs `HTMLCanvasElement.getContext` + * stubbed before anything renders — otherwise Konva's own canvas setup throws. + * The stub is a Proxy over explicitly-typed methods/properties (fillStyle, + * drawImage, path building, `fill('evenodd')` used by the polygon-with-holes + * sceneFunc, …) that falls back to a `vi.fn()` for anything unlisted, so a + * Konva call this suite didn't anticipate returns a harmless mock instead of + * throwing "undefined is not a function". + */ +function createMockContext(canvas: HTMLCanvasElement) { + const base: Record = { + canvas, + fillStyle: '#000000', + strokeStyle: '#000000', + lineWidth: 1, + lineCap: 'butt', + lineJoin: 'miter', + globalAlpha: 1, + globalCompositeOperation: 'source-over', + font: '10px sans-serif', + textAlign: 'start', + textBaseline: 'alphabetic', + imageSmoothingEnabled: true, + filter: 'none', + fillRect: vi.fn(), + clearRect: vi.fn(), + strokeRect: vi.fn(), + drawImage: vi.fn(), + putImageData: vi.fn(), + getImageData: vi.fn(() => ({ data: new Uint8ClampedArray(4), width: 1, height: 1 })), + createImageData: vi.fn((w: number, h: number) => ({ + data: new Uint8ClampedArray(Math.max(1, w) * Math.max(1, h) * 4), + width: w, + height: h, + })), + save: vi.fn(), + restore: vi.fn(), + translate: vi.fn(), + scale: vi.fn(), + rotate: vi.fn(), + transform: vi.fn(), + setTransform: vi.fn(), + resetTransform: vi.fn(), + beginPath: vi.fn(), + closePath: vi.fn(), + moveTo: vi.fn(), + lineTo: vi.fn(), + bezierCurveTo: vi.fn(), + quadraticCurveTo: vi.fn(), + arc: vi.fn(), + arcTo: vi.fn(), + ellipse: vi.fn(), + rect: vi.fn(), + fill: vi.fn(), + stroke: vi.fn(), + clip: vi.fn(), + isPointInPath: vi.fn(() => false), + measureText: vi.fn(() => ({ width: 0 })), + fillText: vi.fn(), + strokeText: vi.fn(), + createLinearGradient: vi.fn(() => ({ addColorStop: vi.fn() })), + createRadialGradient: vi.fn(() => ({ addColorStop: vi.fn() })), + createPattern: vi.fn(() => ({})), + setLineDash: vi.fn(), + getLineDash: vi.fn(() => []), + }; + return new Proxy(base, { + get(target, prop) { + if (prop in target) return (target as Record)[prop as string]; + return vi.fn(); + }, + set(target, prop, value) { + (target as Record)[prop as string] = value; + return true; + }, + }); +} + +beforeEach(() => { + HTMLCanvasElement.prototype.getContext = vi.fn(function (this: HTMLCanvasElement) { + return createMockContext(this) as unknown as CanvasRenderingContext2D; + }) as unknown as typeof HTMLCanvasElement.prototype.getContext; + HTMLCanvasElement.prototype.toDataURL = vi.fn(() => 'data:,'); + vi.stubGlobal( + 'ResizeObserver', + class { + observe() {} + unobserve() {} + disconnect() {} + }, + ); +}); + +afterEach(() => { + cleanup(); + vi.restoreAllMocks(); + vi.unstubAllGlobals(); +}); + +const classA: AnnotationClass = { classId: 1, label: 'A', color: '#ff0000', isVisible: true }; +const classB: AnnotationClass = { classId: 2, label: 'B', color: '#00ff00', isVisible: true }; +const hiddenClass: AnnotationClass = { classId: 3, label: 'Hidden', color: '#0000ff', isVisible: false }; + +const polygonShape: Shape = { id: 'poly-1', classId: 1, kind: 'polygon', points: [0, 0, 10, 0, 10, 10, 0, 10] }; +const polygonWithHoles: Shape = { + id: 'poly-2', + classId: 1, + kind: 'polygon', + points: [0, 0, 20, 0, 20, 20, 0, 20], + holes: [[5, 5, 15, 5, 15, 15, 5, 15]], +}; +const rectShape: Shape = { id: 'rect-1', classId: 2, kind: 'rectangle', x: 1, y: 1, w: 5, h: 5 }; +const ellipseShape: Shape = { id: 'ellipse-1', classId: 2, kind: 'ellipse', cx: 10, cy: 10, rx: 4, ry: 3 }; +const brushShape: Shape = { + id: 'brush-1', + classId: 1, + kind: 'brush', + strokes: [ + { points: [0, 0, 5, 5, 10, 0], radius: 2, mode: 'paint' }, + { points: [2, 2, 6, 6], radius: 1, mode: 'erase' }, + ], +}; +const predictedShape: Shape = { id: 'rect-predicted', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2, origin: 'predicted' }; +const hiddenClassShape: Shape = { id: 'rect-hidden', classId: 3, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }; +const shapeWithErase: Shape = { + id: 'rect-erased', + classId: 2, + kind: 'rectangle', + x: 0, + y: 0, + w: 8, + h: 8, + erased: [{ points: [1, 1, 2, 2], radius: 1 }], +}; + +function baseProps(overrides: Partial = {}): ShapesLayerProps { + return { + layerRef: createRef(), + shapes: [], + classMap: new Map([[1, classA], [2, classB], [3, hiddenClass]]), + fillOpacity: 0.5, + scaleX: 1, + selectedShapeIds: [], + activeBrushShapeId: null, + activeClassId: null, + imageWidth: 256, + imageHeight: 256, + originVisible: { human: true, predicted: true }, + ...overrides, + }; +} + +function renderInStage(props: ShapesLayerProps) { + return render( + + + , + ); +} + +describe('ShapesLayer', () => { + it('renders an empty layer without crashing when there are no shapes', () => { + const { container } = renderInStage(baseProps()); + // The Layer itself still mounts (it always exists so the parent's ref + + // caching logic has something to grab), giving exactly one scene canvas. + expect(container.querySelectorAll('canvas').length).toBe(1); + }); + + it('renders without crashing across every shape kind (polygon, polygon-with-holes, rectangle, ellipse, brush)', () => { + const { container } = renderInStage( + baseProps({ + shapes: [polygonShape, polygonWithHoles, rectShape, ellipseShape, brushShape], + }), + ); + expect(container.querySelectorAll('canvas').length).toBe(1); + }); + + it('renders a shape with erase carve-outs without crashing', () => { + const { container } = renderInStage(baseProps({ shapes: [shapeWithErase] })); + expect(container.querySelectorAll('canvas').length).toBe(1); + }); + + it('filters out shapes belonging to a hidden class', () => { + // Nothing to assert on the canvas pixels themselves (coarse mock), but the + // filter runs in plain JS before any Konva node is built, so at minimum it + // must not throw when the only shape present is on a hidden class. + expect(() => renderInStage(baseProps({ shapes: [hiddenClassShape] }))).not.toThrow(); + }); + + it('filters out shapes whose origin is toggled off', () => { + expect(() => + renderInStage( + baseProps({ + shapes: [predictedShape], + originVisible: { human: true, predicted: false }, + }), + ), + ).not.toThrow(); + }); + + it('recolors the active brush instance to the active class without crashing', () => { + const brushInProgress: Shape = { ...brushShape, id: 'active-brush' }; + expect(() => + renderInStage( + baseProps({ + shapes: [brushInProgress], + activeBrushShapeId: 'active-brush', + activeClassId: 2, + }), + ), + ).not.toThrow(); + }); + + it('renders selected shapes (thicker stroke path) without crashing', () => { + expect(() => + renderInStage(baseProps({ shapes: [rectShape, ellipseShape], selectedShapeIds: ['rect-1'] })), + ).not.toThrow(); + }); + + it('does not throw when scaleX changes (stroke width depends on zoom)', () => { + const props = baseProps({ shapes: [rectShape] }); + const { rerender, container } = renderInStage(props); + rerender( + + + , + ); + expect(container.querySelectorAll('canvas').length).toBe(1); + }); + + it('forwards layerRef to the underlying Konva.Layer', () => { + const layerRef = createRef(); + renderInStage(baseProps({ layerRef })); + expect(layerRef.current).not.toBeNull(); + // Real Konva.Layer instance, not a DOM node — spot-check its own API surface. + expect(typeof layerRef.current?.getCanvas).toBe('function'); + }); +}); diff --git a/frontend/src/components/annotate/AnnotationCanvas/ShapesLayer.tsx b/frontend/src/components/annotate/AnnotationCanvas/ShapesLayer.tsx new file mode 100644 index 0000000..f19c793 --- /dev/null +++ b/frontend/src/components/annotate/AnnotationCanvas/ShapesLayer.tsx @@ -0,0 +1,224 @@ +/** + * ShapesLayer — the cached, non-interactive layer of committed annotations. + * + * Extracted from AnnotationCanvas and wrapped in `React.memo` for one reason: the + * canvas takes brightness / contrast / levels / gamma / blur as props, so every + * tick of a display slider re-rendered it and reconciled every Konva node on this + * layer — despite shapes depending on none of those values. On a slice with many + * annotations that dominated slider latency. Keeping this layer's props limited to + * what it actually draws means those ticks now skip the entire shape tree. + * + * If you add a prop here, make sure it genuinely affects the drawn shapes; + * anything that changes per-frame (a pan offset, a pointer position) would + * reintroduce exactly the problem this exists to solve. + */ +import { memo } from 'react'; +import { Layer, Line, Rect, Ellipse, Group, Shape as KonvaShape } from 'react-konva'; +import type Konva from 'konva'; +import type { Shape, EraseStroke } from '@/stores/annotationStore'; +import type { AnnotationClass } from '@/stores/classStore'; +import { isShapeOriginVisible } from '@/stores/layerVisibilityStore'; + +/** Dash pattern marking a shape as iPred-predicted (vs. hand-drawn, solid). */ +function dashFor(shape: Shape, strokeW: number): number[] | undefined { + return shape.origin === 'predicted' ? [strokeW * 4, strokeW * 3] : undefined; +} + +/** Build an even-odd path (outer + hole rings) on a Konva context. */ +function buildRingsPath(ctx: Konva.Context, rings: number[][]): void { + ctx.beginPath(); + for (const r of rings) { + if (r.length < 6) continue; + ctx.moveTo(r[0], r[1]); + for (let i = 2; i < r.length; i += 2) ctx.lineTo(r[i], r[i + 1]); + ctx.closePath(); + } +} + +/** Erase carve-outs rendered destination-out over the shape they belong to. */ +function renderErased(erased?: EraseStroke[]) { + return (erased ?? []).map((st, i) => ( + + )); +} + +export interface ShapesLayerProps { + /** Konva layer ref — the parent owns caching/recaching of this layer. */ + layerRef: React.Ref; + shapes: Shape[]; + /** Class lookup for color + visibility (O(1) per shape). */ + classMap: Map; + fillOpacity: number; + /** Current zoom, used only to keep stroke width constant on screen. */ + scaleX: number; + selectedShapeIds: string[]; + /** The in-progress brush instance, recolored to the active class. */ + activeBrushShapeId: string | null; + activeClassId: number | null; + /** Image frame, for clipping and for giving hole-shapes a real self-rect. */ + imageWidth: number; + imageHeight: number; + /** Predicted/human sub-visibility (see layerVisibilityStore). */ + originVisible: { human: boolean; predicted: boolean }; +} + +function ShapesLayerImpl({ + layerRef, + shapes, + classMap, + fillOpacity, + scaleX, + selectedShapeIds, + activeBrushShapeId, + activeClassId, + imageWidth, + imageHeight, + originVisible, +}: ShapesLayerProps) { + const colorForClass = (classId: number) => classMap.get(classId)?.color ?? '#ff0000'; + + /** Render a polygon that may have holes via an even-odd fill (outer path minus + * hole subpaths). Even-odd — not destination-out — so a hole reveals whatever + * is *beneath* it (e.g. another class) instead of erasing it off the layer. */ + const renderPolygonWithHoles = ( + points: number[], + holes: number[][], + color: string, + strokeW: number, + dash: number[] | undefined, + ) => ( + { + buildRingsPath(ctx, [points, ...holes]); + const raw = (ctx as unknown as { _context: CanvasRenderingContext2D })._context; + raw.fillStyle = color; + raw.fill('evenodd'); + ctx.strokeShape(node); + }} + /> + ); + + /** Render a committed shape (any kind), with the active brush instance recolored + * to the active class and erase strokes carved out. */ + const renderShape = (shape: Shape) => { + const color = + shape.id === activeBrushShapeId && activeClassId !== null + ? colorForClass(activeClassId) + : colorForClass(shape.classId); + const isSelected = selectedShapeIds.includes(shape.id); + const strokeW = (isSelected ? 2 : 1) / scaleX; + // Predicted (uncommitted-by-a-human) shapes get a dashed outline so they read + // as provisional at a glance — see layerVisibilityStore's origin toggle. + const dash = dashFor(shape, strokeW); + + if (shape.kind === 'polygon') { + return ( + + {shape.holes?.length + ? renderPolygonWithHoles(shape.points, shape.holes, color, strokeW, dash) + : ( + + )} + {renderErased(shape.erased)} + + ); + } + if (shape.kind === 'rectangle') { + return ( + + + {renderErased(shape.erased)} + + ); + } + if (shape.kind === 'ellipse') { + return ( + + + {renderErased(shape.erased)} + + ); + } + if (shape.kind === 'brush') { + return ( + + {shape.strokes.map((stroke, i) => ( + + ))} + + ); + } + return null; + }; + + return ( + + {shapes + .filter((s) => classMap.get(s.classId)?.isVisible !== false) + .filter((s) => isShapeOriginVisible(originVisible, s.origin)) + .map(renderShape)} + + ); +} + +export default memo(ShapesLayerImpl); diff --git a/frontend/src/components/annotate/AnnotationCanvas/index.test.tsx b/frontend/src/components/annotate/AnnotationCanvas/index.test.tsx new file mode 100644 index 0000000..dc3ae64 --- /dev/null +++ b/frontend/src/components/annotate/AnnotationCanvas/index.test.tsx @@ -0,0 +1,283 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render } from '@testing-library/react'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import AnnotationCanvas from './index'; +import { useDatasetStore, type ImageMeta } from '@/stores/datasetStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { useToolStore } from '@/stores/toolStore'; +import { useClassStore, type AnnotationClass } from '@/stores/classStore'; +import { useLayerVisibilityStore } from '@/stores/layerVisibilityStore'; +import { useClipboardStore } from '@/stores/clipboardStore'; + +/** + * Shallow smoke tests for the 3636-line react-konva Stage that is the Annotate + * page's whole editing surface. These deliberately do NOT drive tool + * interactions (drawing, panning, brush strokes, SAM, livewire, …) — only that + * the component mounts and re-renders across its main branches (no dataset + * loaded / dataset + shapes loaded / different active tools / preview mode) + * without throwing. + * + * jsdom has no real 2D canvas, so Konva (which every Stage/Layer/shape + * ultimately talks to) needs `HTMLCanvasElement.getContext` stubbed before + * anything renders — otherwise Konva's own canvas setup throws. The stub is a + * Proxy over explicitly-typed methods/properties (the ones this component's + * own code calls directly — willReadFrequently getImageData/putImageData for + * the histogram and threshold-overlay repaint — plus everything Konva itself + * needs) that falls back to a `vi.fn()` for anything unlisted. + */ +function createMockContext(canvas: HTMLCanvasElement) { + const base: Record = { + canvas, + fillStyle: '#000000', + strokeStyle: '#000000', + lineWidth: 1, + lineCap: 'butt', + lineJoin: 'miter', + globalAlpha: 1, + globalCompositeOperation: 'source-over', + font: '10px sans-serif', + textAlign: 'start', + textBaseline: 'alphabetic', + imageSmoothingEnabled: true, + filter: 'none', + fillRect: vi.fn(), + clearRect: vi.fn(), + strokeRect: vi.fn(), + drawImage: vi.fn(), + putImageData: vi.fn(), + getImageData: vi.fn(() => ({ data: new Uint8ClampedArray(4), width: 1, height: 1 })), + createImageData: vi.fn((w: number, h: number) => ({ + data: new Uint8ClampedArray(Math.max(1, w) * Math.max(1, h) * 4), + width: w, + height: h, + })), + save: vi.fn(), + restore: vi.fn(), + translate: vi.fn(), + scale: vi.fn(), + rotate: vi.fn(), + transform: vi.fn(), + setTransform: vi.fn(), + resetTransform: vi.fn(), + beginPath: vi.fn(), + closePath: vi.fn(), + moveTo: vi.fn(), + lineTo: vi.fn(), + bezierCurveTo: vi.fn(), + quadraticCurveTo: vi.fn(), + arc: vi.fn(), + arcTo: vi.fn(), + ellipse: vi.fn(), + rect: vi.fn(), + fill: vi.fn(), + stroke: vi.fn(), + clip: vi.fn(), + isPointInPath: vi.fn(() => false), + measureText: vi.fn(() => ({ width: 0 })), + fillText: vi.fn(), + strokeText: vi.fn(), + createLinearGradient: vi.fn(() => ({ addColorStop: vi.fn() })), + createRadialGradient: vi.fn(() => ({ addColorStop: vi.fn() })), + createPattern: vi.fn(() => ({})), + setLineDash: vi.fn(), + getLineDash: vi.fn(() => []), + }; + return new Proxy(base, { + get(target, prop) { + if (prop in target) return (target as Record)[prop as string]; + return vi.fn(); + }, + set(target, prop, value) { + (target as Record)[prop as string] = value; + return true; + }, + }); +} + +// Snapshot the zustand stores' initial state once at module load so every test +// can restore a clean slate, per this session's established convention (these +// four stores have no built-in `reset`; annotationStore's own `reset()` is used +// where available). +const initialToolState = useToolStore.getState(); +const initialLayerVisibilityState = useLayerVisibilityStore.getState(); +const initialClipboardState = useClipboardStore.getState(); + +beforeEach(() => { + HTMLCanvasElement.prototype.getContext = vi.fn(function (this: HTMLCanvasElement) { + return createMockContext(this) as unknown as CanvasRenderingContext2D; + }) as unknown as typeof HTMLCanvasElement.prototype.getContext; + HTMLCanvasElement.prototype.toDataURL = vi.fn(() => 'data:,'); + vi.stubGlobal( + 'ResizeObserver', + class { + observe() {} + unobserve() {} + disconnect() {} + }, + ); + // useImageSlice's queryFn calls fetch when a dataset is loaded; keep it from + // ever hitting the network (jsdom's .src load never resolves either way, + // so this only avoids an unhandled real request, never actual image data). + global.fetch = vi.fn(() => Promise.reject(new Error('network disabled in tests'))); + + useDatasetStore.getState().reset(); + useAnnotationStore.getState().reset(); + useClassStore.setState({ classes: [] }); + useToolStore.setState(initialToolState, true); + useLayerVisibilityStore.setState(initialLayerVisibilityState, true); + useClipboardStore.setState(initialClipboardState, true); +}); + +afterEach(() => { + cleanup(); + vi.restoreAllMocks(); + vi.unstubAllGlobals(); +}); + +const CLASS_A: AnnotationClass = { classId: 1, label: 'A', color: '#ff0000', isVisible: true }; +const CLASS_B: AnnotationClass = { classId: 2, label: 'B', color: '#00ff00', isVisible: true }; + +const META: ImageMeta = { + nSlices: 5, + height: 256, + width: 256, + dtype: 'uint8', + isRgb: false, + valueRange: [0, 255], +}; + +function baseProps(overrides: Partial> = {}) { + return { + brightness: 0, + contrast: 0, + levelsLo: 0, + levelsHi: 255, + activeClassId: null, + activeBrushShapeId: null, + onNewBrushInstance: vi.fn(), + ...overrides, + }; +} + +/** Loads a dataset (source/kind/meta) into datasetStore, as ConnectPage/BrowsePage would. */ +function loadDataset() { + useDatasetStore.getState().setDataset('local', 'sample.tif', null, META); +} + +function renderCanvas(overrides: Partial> = {}) { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render( + + + , + ); +} + +describe('AnnotationCanvas', () => { + it('renders without crashing when no dataset is loaded', () => { + const { container } = renderCanvas(); + // The image layer always mounts (even empty), so at least one canvas exists. + expect(container.querySelectorAll('canvas').length).toBeGreaterThan(0); + }); + + it('renders without crashing once a dataset (meta) is loaded, with no shapes', () => { + loadDataset(); + const { container } = renderCanvas(); + expect(container.querySelectorAll('canvas').length).toBeGreaterThan(0); + }); + + it('renders without crashing with committed shapes present on the current slice', () => { + loadDataset(); + useClassStore.setState({ classes: [CLASS_A, CLASS_B] }); + useAnnotationStore.getState().replaceClassShapesOnSlice('local:sample.tif', 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 10, y: 10, w: 20, h: 20 }, + { id: 's2', classId: 2, kind: 'ellipse', cx: 100, cy: 100, rx: 15, ry: 10 }, + ]); + const { container } = renderCanvas({ activeClassId: 1 }); + expect(container.querySelectorAll('canvas').length).toBeGreaterThan(0); + }); + + it('renders without crashing for each drawing tool (pan/select/polygon/rectangle/ellipse/brush/eraser)', () => { + loadDataset(); + useClassStore.setState({ classes: [CLASS_A] }); + const tools = ['pan', 'select', 'polygon', 'magnetic', 'rectangle', 'ellipse', 'brush', 'eraser', 'threshold', 'sampler', 'fill'] as const; + for (const tool of tools) { + useToolStore.setState({ tool }); + expect(() => { + const { unmount } = renderCanvas({ activeClassId: 1 }); + unmount(); + }, `tool=${tool}`).not.toThrow(); + } + }); + + it('renders the select tool with a selection without crashing (Transformer + toolbar overlay)', () => { + loadDataset(); + useClassStore.setState({ classes: [CLASS_A] }); + useAnnotationStore.getState().replaceClassShapesOnSlice('local:sample.tif', 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 10, y: 10, w: 20, h: 20 }, + ]); + useToolStore.setState({ tool: 'select', selectedShapeIds: ['s1'] }); + expect(() => renderCanvas({ activeClassId: 1 })).not.toThrow(); + }); + + it('renders in read-only preview mode (previewShapes/previewClasses) without crashing', () => { + loadDataset(); + const { container } = renderCanvas({ + previewShapes: [{ id: 'p1', classId: 9, kind: 'rectangle', x: 0, y: 0, w: 5, h: 5 }], + previewClasses: [{ classId: 9, label: 'Preview', color: '#123456', isVisible: true }], + }); + expect(container.querySelectorAll('canvas').length).toBeGreaterThan(0); + }); + + it('renders with iPred overlay props set (proba/predictions/manifold) without crashing', () => { + loadDataset(); + const { container } = renderCanvas({ + probaOverlayUrl: 'blob:proba', + clfCommitUrl: 'blob:commit', + clfStatusUrl: 'blob:status', + predictionClassColorById: new Map([[1, '#ff0000']]), + manifoldHeatmapUrl: 'blob:manifold', + manifoldMarkers: [{ x: 5, y: 5, cluster: 0 }], + }); + expect(container.querySelectorAll('canvas').length).toBeGreaterThan(0); + }); + + it('renders across a display-adjustment prop sweep (brightness/contrast/levels/colormap/gamma/blur/upscale) without crashing', () => { + loadDataset(); + expect(() => + renderCanvas({ + brightness: 30, + contrast: -20, + levelsLo: 10, + levelsHi: 240, + colormap: 'viridis', + gamma: 1.4, + clahe: true, + sharpen: true, + blur: 2, + upscale: 2, + }), + ).not.toThrow(); + }); + + it('calls onHistogram only after a slice image has actually decoded (never synchronously with no dataset)', () => { + const onHistogram = vi.fn(); + renderCanvas({ onHistogram }); + // No dataset loaded => no image => the histogram effect bails out early. + expect(onHistogram).not.toHaveBeenCalled(); + }); + + it('re-renders across a focusRegion prop change without crashing', () => { + loadDataset(); + const { rerender, container } = renderCanvas({ focusRegion: null }); + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + rerender( + + + , + ); + expect(container.querySelectorAll('canvas').length).toBeGreaterThan(0); + }); +}); diff --git a/frontend/src/components/annotate/AnnotationCanvas/index.tsx b/frontend/src/components/annotate/AnnotationCanvas/index.tsx index 366b0e4..48a105b 100644 --- a/frontend/src/components/annotate/AnnotationCanvas/index.tsx +++ b/frontend/src/components/annotate/AnnotationCanvas/index.tsx @@ -20,24 +20,38 @@ import { Trash } from '@phosphor-icons/react'; import type Konva from 'konva'; import { v4 as uuidv4 } from 'uuid'; import { useDatasetStore } from '@/stores/datasetStore'; -import { useAnnotationStore, type Shape, type PolygonShape, type BrushStroke, type EraseStroke } from '@/stores/annotationStore'; +import { useAnnotationStore, type Shape, type PolygonShape, type BrushStroke } from '@/stores/annotationStore'; import { useToolStore } from '@/stores/toolStore'; import { useClassStore, type AnnotationClass } from '@/stores/classStore'; -import { toImage, normalizeRect, normalizeEllipse } from '@/lib/geometry'; +import { useLayerVisibilityStore } from '@/stores/layerVisibilityStore'; +import { toImage, normalizeRect, normalizeEllipse, shapeBBox, unionBBox, bboxNear, bboxIntersects, type BBox } from '@/lib/geometry'; import { buildCostMap, dijkstra, tracePath, imageToGrid, simplifyPath, type CostMap } from '@/lib/livewire'; import { useImageSlice } from '@/hooks/useImageSlice'; import { buildSourceKey } from '@/lib/sourceKey'; import { buildField, magicSelect, maskToPolygons, maskToPolygonsWithHoles, type GrayField } from '@/lib/magicwand'; import { useSam } from '@/hooks/useSam'; import { renderAdjusted, renderPreprocessOnly } from '@/lib/sam/adjust'; -import { gridFor, fullResGridFor, rasterizeShapes, rasterizeUnion } from '@/lib/rasterize'; -import { keepComponentsAtPoints } from '@/lib/morphology'; -import { unionShapesToPolygons, unionShapesToMultiPolygon, eraseStampToMultiPolygon, subtractFromShape } from '@/lib/polybool'; +import { gridFor, fullResGridFor, rasterizeShapes, rasterizeUnion, stampStroke } from '@/lib/rasterize'; +import { keepComponentsAtPoints, dilate, erode, removeSmallComponents } from '@/lib/morphology'; +import { unionShapesToPolygons, unionShapesToMultiPolygon, unionShapesChecked, eraseStampToMultiPolygon, subtractFromShape, regionsToMultiPolygon } from '@/lib/polybool'; import { computeRegionOps, type RegionOp } from '@/lib/regionOps'; -import { clipShapesToOthers, hasOtherClass } from '@/lib/clipToClasses'; +import { clipShapesToOthers, clipShapesToOthersMask, hasOtherClass } from '@/lib/clipToClasses'; import { mergeNewWithSameClass, expandSameClassOverlap } from '@/lib/mergeSameClass'; import { useClipboardStore } from '@/stores/clipboardStore'; import { colormapTables, type ColormapName } from '@/lib/colormaps'; +import ShapesLayer from './ShapesLayer'; +import OverlaysLayer from './OverlaysLayer'; +import type { ManifoldPoint } from '@/lib/featureManifold'; +import { displayAffineFor, displayBandToBase, baseToDisplay, displayToBase } from '@/lib/displayTransform'; +import { + backgroundRing, fitWithBlurSweep, combineSigma, BLUR_CANDIDATES, MAX_RING_WIDTH, + fitBand, sampleHistograms, type SweepResult, +} from '@/lib/thresholdFit'; +import { + buildChannels, fitProjection, projectToScore, CHANNEL_NAMES, +} from '@/lib/featureChannels'; +import { applyGaussianBlurGray } from '@/lib/blur'; +import { time } from '@/lib/perf'; // macOS labels the Alt key "Option" (⌥). e.altKey is true for it either way, // so only the on-screen label needs to differ. @@ -46,6 +60,42 @@ const REMOVE_KEY_LABEL = IS_MAC ? 'Option' : 'Alt'; /** Shared stable empty shape list — see `storeShapes` for why identity matters. */ const EMPTY_SHAPES: Shape[] = []; +const EMPTY_MANIFOLD_MARKERS: ManifoldPoint[] = []; +const EMPTY_CLASS_COLOR_BY_ID = new Map(); + +/** Outcome of a Sampler lasso: the fitted band plus what it took to get there. */ +export interface SamplerFit extends SweepResult { + /** The band converted into the DISPLAYED space the threshold knobs use. */ + displayLo: number; + displayHi: number; + /** True when the display transform squashes the fitted range to one value, so + * the band cannot be expressed at the current brightness/contrast/levels. */ + collapsed: boolean; + /** Band and blur in force before the fit, so the UI can offer Revert. */ + previousBand: [number, number]; + previousBlur: number; + /** Blur the fit settled on (current combined with the swept extra). */ + appliedBlur: number; + /** Which signal the brush now gates on: raw displayed intensity, or a fitted + * combination of intensity and texture channels. */ + mode?: 'intensity' | 'projected'; + /** Skill intensity alone managed, so the gain from projecting is visible. */ + intensitySkill?: number; + /** Contribution of each channel to the projection, for the readout. */ + weights?: Array<{ name: string; weight: number }>; +} + +/** How far (image px) the cursor travels before the sampler lasso drops a new + * live-wire anchor. Larger = fewer dijkstra runs but looser snapping. */ +const SAMPLER_ANCHOR_STEP = 40; + +/** Smallest island (grid cells) a threshold stroke keeps after regularization. + * Below this a region is thresholding noise, not a feature. */ +const THRESHOLD_MIN_REGION = 12; + +/** Upscale budget: the working resolution is clamped back to 1x past this many + * pixels so a large slice at 4x can't exhaust memory (~64 MP ≈ 256 MB of RGBA). */ +const MAX_WORKING_PIXELS = 64e6; interface AnnotationCanvasProps { brightness: number; @@ -60,8 +110,18 @@ interface AnnotationCanvasProps { * tools' baked view but NOT the exported pixels). */ clahe?: boolean; sharpen?: boolean; + /** Gaussian pre-blur sigma in image pixels (0 = off). */ + blur?: number; + /** Working-resolution multiplier (1, 2, 4): resamples the slice for the display + * base and every tool field so small features get more pixels. Annotation + * coordinates stay in NATIVE image pixels — they just gain sub-pixel precision. */ + upscale?: number; /** Emits the current slice's 256-bin luminance histogram when it loads. */ onHistogram?: (bins: number[]) => void; + /** Result of a Sampler lasso fit (null when it could not fit), for the Toolbar. */ + onSamplerFit?: (fit: SamplerFit | null) => void; + /** Lets the Sampler apply the blur sigma it chose. */ + onBlurChange?: (sigma: number) => void; activeClassId: number | null; activeBrushShapeId: string | null; onNewBrushInstance: (id: string) => void; @@ -72,6 +132,23 @@ interface AnnotationCanvasProps { /** Zoom to + highlight this image-coord region (e.g. from an Insights QA flag). * `nonce` changes to re-trigger the same region. */ focusRegion?: { x: number; y: number; w: number; h: number; nonce: number } | null; + + // iPred overlays (additive; each gated by its layerVisibilityStore group). + /** Selected feature-channel PNG (grayscale), from useFeatureChannels. */ + featureChannelUrl?: string | null; + /** Colorized softmax probability PNG for the active class, from usePixelClassifier. */ + probaOverlayUrl?: string | null; + /** Conformal commit (argmax classId) + status (abstain/singleton/multi) PNGs. */ + clfCommitUrl?: string | null; + clfStatusUrl?: string | null; + /** classId → hex color, for the conformal overlay and manifold markers. */ + predictionClassColorById?: Map; + manifoldHeatmapUrl?: string | null; + manifoldHeatmapOpacity?: number; + manifoldShowHeatmap?: boolean; + manifoldMarkers?: ManifoldPoint[]; + manifoldShowMarkers?: boolean; + manifoldBoxSize?: number; } // ---- Point-in-shape hit testing (used by the eraser to pick a target) ---- @@ -122,6 +199,15 @@ function withAlpha(hex: string, a: number): string { return `rgba(${(n >> 16) & 255}, ${(n >> 8) & 255}, ${n & 255}, ${a})`; } +/** hex (#rgb or #rrggbb) → [r,g,b]. Falls back to mid-grey if it isn't a hex color. */ +function hexToRgb(hex: string): [number, number, number] { + let h = hex.replace('#', ''); + if (h.length === 3) h = h.split('').map((c) => c + c).join(''); + if (!/^[0-9a-fA-F]{6}$/.test(h)) return [128, 128, 128]; + const n = parseInt(h, 16); + return [(n >> 16) & 255, (n >> 8) & 255, n & 255]; +} + /** Shortest distance from point (px,py) to segment AB. */ function distToSegment(px: number, py: number, ax: number, ay: number, bx: number, by: number): number { const dx = bx - ax, dy = by - ay; @@ -298,13 +384,6 @@ function interiorPoint(shape: Shape): { x: number; y: number } | null { return null; } -interface BBox { x: number; y: number; w: number; h: number; } - -/** True if two AABBs overlap. */ -function bboxIntersects(a: BBox, b: BBox): boolean { - return a.x < b.x + b.w && a.x + a.w > b.x && a.y < b.y + b.h && a.y + a.h > b.y; -} - /** True if segments AB and CD intersect. */ function segIntersects( ax: number, ay: number, bx: number, by: number, @@ -400,13 +479,28 @@ export default function AnnotationCanvas({ gamma = 1, clahe = false, sharpen = false, + blur = 0, + upscale = 1, onHistogram, + onSamplerFit, + onBlurChange, activeClassId, activeBrushShapeId, onNewBrushInstance, previewShapes = null, previewClasses = null, focusRegion = null, + featureChannelUrl = null, + probaOverlayUrl = null, + clfCommitUrl = null, + clfStatusUrl = null, + predictionClassColorById, + manifoldHeatmapUrl = null, + manifoldHeatmapOpacity = 0.45, + manifoldShowHeatmap = true, + manifoldMarkers = EMPTY_MANIFOLD_MARKERS, + manifoldShowMarkers = true, + manifoldBoxSize = 64, }: AnnotationCanvasProps) { const containerRef = useRef(null); const imageRef = useRef(null); @@ -419,6 +513,7 @@ export default function AnnotationCanvas({ // Imperative refs for lag-free brush drawing const draftStrokeLayerRef = useRef(null); const draftLineRef = useRef(null); + const thresholdPreviewRef = useRef(null); const brushCursorLayerRef = useRef(null); const brushCursorRef = useRef(null); @@ -433,7 +528,54 @@ export default function AnnotationCanvas({ eraseTargetKind?: 'brush' | 'vector'; } | null>(null); - const { kind, source, serverUri, meta, currentSlice, renderOpts } = useDatasetStore(); + // Threshold-brush stroke, buffered exactly like `draftStrokeRef`: a binary mask + // over the threshold field's grid that accumulates only in-band pixels, plus the + // offscreen canvas mirroring it for the live preview. Flushed on mouseup. + const thresholdStrokeRef = useRef<{ + mode: 'paint' | 'erase'; + /** Grid geometry of the field this stroke was started against. */ + gw: number; gh: number; scale: number; + mask: Uint8Array; + /** In-band gate for the whole slice (same grid) — the stroke can't leave it. */ + gate: Uint8Array; + /** Last pointer position in image coords, for segment stamping. */ + last: { x: number; y: number }; + canvas: HTMLCanvasElement; + imageData: ImageData; + } | null>(null); + + // Sampler lasso: a loop drawn to teach the Threshold Brush what to select. + // Buffered in a ref like the brush so dragging causes no re-renders, and purely + // a measurement gesture — it commits no shape and touches no history. + // + // It snaps to edges using the same live-wire machinery as the magnetic tool: + // an anchor is dropped every `SAMPLER_ANCHOR_STEP` image pixels and the segment + // between anchors is the least-cost path, so the sample follows the feature's + // real boundary instead of a shaky hand-drawn one. A tighter sample means a + // cleaner positive histogram, which is what the whole fit rests on. + const samplerLassoRef = useRef<{ + /** Snapped points locked in so far (flat [x,y,…] image coords). */ + committed: number[]; + /** Where the current live-wire segment starts. */ + anchor: { x: number; y: number }; + /** First point, so the loop can be closed back to it. */ + start: { x: number; y: number }; + cm: CostMap | null; + /** Dijkstra predecessor map from `anchor`; null when snapping is unavailable. */ + prev: Int32Array | null; + } | null>(null); + + const { kind, source, serverUri, meta, currentSlice, renderOpts, denoise } = useDatasetStore(); + const layerGroups = useLayerVisibilityStore((s) => s.groups); + // The Denoise layer toggle mutes the configured method without discarding it, + // so switching it back on doesn't lose the user's method/strength choice. + const effectiveDenoise = layerGroups.denoise ? denoise : null; + const probaOpacity = useLayerVisibilityStore((s) => s.probaOpacity); + const predictionsOpacity = useLayerVisibilityStore((s) => s.predictionsOpacity); + const predictionClassVisible = useLayerVisibilityStore((s) => s.predictionClassVisible); + const showPredictionMulti = useLayerVisibilityStore((s) => s.showPredictionMulti); + const showPredictionAbstain = useLayerVisibilityStore((s) => s.showPredictionAbstain); + const annotationOriginVisible = useLayerVisibilityStore((s) => s.annotationOriginVisible); const sourceKey = source && kind ? buildSourceKey(kind as 'tiled' | 'local', source, serverUri) : null; @@ -458,6 +600,7 @@ export default function AnnotationCanvas({ const magicEdgeStop = useToolStore((s) => s.magicEdgeStop); const magicEngine = useToolStore((s) => s.magicEngine); const setMagicEngine = useToolStore((s) => s.setMagicEngine); + const setTool = useToolStore((s) => s.setTool); const samDetail = useToolStore((s) => s.samDetail); const samThreshold = useToolStore((s) => s.samThreshold); const samAvoidLabeled = useToolStore((s) => s.samAvoidLabeled); @@ -466,6 +609,13 @@ export default function AnnotationCanvas({ const clipToOtherClasses = useToolStore((s) => s.clipToOtherClasses); const mergeOverlappingSameClass = useToolStore((s) => s.mergeOverlappingSameClass); const fillThreshold = useToolStore((s) => s.fillThreshold); + // NOTE: thresholdLo/thresholdHi are deliberately NOT subscribed — dragging the + // band would then re-render this entire component every frame. They are read + // via getState() at use sites and drive the overlay through a store + // subscription (see the overlay block below). + const thresholdOverlay = useToolStore((s) => s.thresholdOverlay); + const thresholdSampleWidth = useToolStore((s) => s.thresholdSampleWidth); + const setThresholdBand = useToolStore((s) => s.setThresholdBand); const { classes } = useClassStore(); const [stageSize, setStageSize] = useState({ width: 800, height: 600 }); @@ -516,7 +666,12 @@ export default function AnnotationCanvas({ // (expensive) layer re-cache so we don't rebuild the shapes bitmap mid-stroke. const [isDrawing, setIsDrawing] = useState(false); - const { data: sliceUrl } = useImageSlice(source, kind, currentSlice, renderOpts, serverUri); + // Denoising is applied server-side, before normalization — so the PNG this + // returns is already denoised, and everything downstream that samples it + // (displayBase, the Threshold Brush field, the Sampler's fit, magic wand, + // livewire) sees the denoised data with no extra plumbing. That is the point: + // an intensity tool should act on the image the user is actually looking at. + const { data: sliceUrl } = useImageSlice(source, kind, currentSlice, renderOpts, serverUri, effectiveDenoise); const [imageEl, setImageEl] = useState(null); useEffect(() => { @@ -543,14 +698,21 @@ export default function AnnotationCanvas({ // becomes the Konva image base; brightness/contrast/levels/gamma/colormap still // apply on top via the GPU SVG filter. Cache-key fragment so encodes/fields // refresh when toggled. - const preprocess = useMemo(() => ({ clahe, sharpen }), [clahe, sharpen]); - const preprocessKey = `${clahe ? 1 : 0}${sharpen ? 1 : 0}`; + const preprocess = useMemo(() => ({ clahe, sharpen, blur }), [clahe, sharpen, blur]); + const preprocessKey = `${clahe ? 1 : 0}${sharpen ? 1 : 0}b${blur}`; + // Working resolution, clamped so a huge slice can't blow up memory at 4x + // (each level costs 4x the pixels for the base canvas AND every tool field). + const workScale = useMemo(() => { + if (!meta) return 1; + const u = Math.max(1, upscale); + return meta.width * meta.height * u * u > MAX_WORKING_PIXELS ? 1 : u; + }, [meta, upscale]); const displayBase = useMemo(() => { if (!imageEl || !meta) return imageEl; - return (clahe || sharpen) - ? renderPreprocessOnly(imageEl, meta.width, meta.height, preprocess) + return (clahe || sharpen || blur > 0 || workScale > 1) + ? renderPreprocessOnly(imageEl, meta.width, meta.height, preprocess, workScale) : imageEl; - }, [imageEl, meta, clahe, sharpen, preprocess]); + }, [imageEl, meta, clahe, sharpen, blur, workScale, preprocess]); // SAM sees the preprocessed + brightness/contrast/levels-adjusted image // (windowing a low-contrast slice greatly helps), so the encode is keyed on @@ -558,7 +720,9 @@ export default function AnnotationCanvas({ const samEncodeKey = imageEl && meta ? `${sourceKey}|${currentSlice}|b${brightness}|c${contrast}|l${levelsLo}-${levelsHi}|pp${preprocessKey}` : null; - /** Lazily render the display-adjusted slice (preprocess → brightness/contrast/levels) that SAM encodes. */ + /** Lazily render the display-adjusted slice (preprocess → brightness/contrast/levels) + * that SAM encodes. Deliberately kept at 1x: SAM resizes its input to 1024² + * internally, so an upscaled source would only cost memory. */ const makeSamSource = useCallback( () => renderAdjusted(imageEl!, meta!.width, meta!.height, brightness, contrast, levelsLo, levelsHi, preprocess), [imageEl, meta, brightness, contrast, levelsLo, levelsHi, preprocess], @@ -582,17 +746,13 @@ export default function AnnotationCanvas({ // sync with what's displayed. const displayFilterId = 'display-adjust-' + useId().replace(/[^a-zA-Z0-9]/g, ''); const displayAffine = useMemo(() => { - const b255 = brightness * 255; - const adjust = Math.pow((contrast + 100) / 100, 2); - const range = Math.max(1, levelsHi - levelsLo); - const A = adjust * (255 / range); - const Bconst = (255 / range) * (adjust * b255 + 127.5 * (1 - adjust)) - (255 * levelsLo) / range; + const { slope, intercept } = displayAffineFor(brightness, contrast, levelsLo, levelsHi); // The filter is a no-op only when brightness/contrast/levels AND gamma AND // colormap are all identity — otherwise it must stay applied. const identity = brightness === 0 && contrast === 0 && levelsLo <= 0 && levelsHi >= 255 && gamma === 1 && colormap === 'gray'; - return { slope: A, intercept: Bconst / 255, identity }; + return { slope, intercept, identity }; }, [brightness, contrast, levelsLo, levelsHi, gamma, colormap]); // Colormap LUT (per-channel tableValues) for the display filter; null = gray. @@ -718,20 +878,49 @@ export default function AnnotationCanvas({ // Cache the shapes layer for uniform opacity compositing — skipped mid-stroke // so the cache rebuild doesn't stall painting. + // + // The cache pixel ratio is quantized to powers of two rather than tracking zoom + // continuously. Re-caching rasterizes the whole layer, and keying it on the raw + // scale meant every wheel tick paid for that. Rounding UP to the next bucket + // means the cache is never coarser than the old `pixelRatio = scale` (at 1.5x it + // now caches at 2x), so shapes are equally or more crisp while most zoom steps + // reuse the existing cache outright. + const cachePixelRatio = useMemo(() => { + const scale = Math.min(Math.max(transform.scaleX, 1), 4); + return Math.min(4, Math.pow(2, Math.ceil(Math.log2(scale)))); + }, [transform.scaleX]); + useEffect(() => { const layer = shapesLayerRef.current; if (!layer) return; if (isDrawing) return; - layer.clearCache(); - if (displayShapes.length > 0) { - const pr = Math.min(Math.max(transform.scaleX, 1), 4); - layer.cache({ pixelRatio: pr }); - } - layer.batchDraw(); - }, [displayShapes, renderClasses, fillOpacity, meta, transform.scaleX, isDrawing]); + time('layer-cache', () => { + layer.clearCache(); + if (displayShapes.length > 0) { + layer.cache({ pixelRatio: cachePixelRatio }); + } + layer.batchDraw(); + }); + }, [displayShapes, renderClasses, fillOpacity, meta, cachePixelRatio, isDrawing]); + + // Class lookups by id. These run per shape, on both shape layers, on every + // render — a linear scan there is O(shapes × classes) for what should be O(1). + const classMap = useMemo( + () => new Map(classes.map((c) => [c.classId, c])), + [classes], + ); + const renderClassMap = useMemo( + () => (renderClasses === classes ? classMap : new Map(renderClasses.map((c) => [c.classId, c]))), + [renderClasses, classes, classMap], + ); const colorForClass = (classId: number) => - renderClasses.find((c) => c.classId === classId)?.color ?? '#ff0000'; + renderClassMap.get(classId)?.color ?? '#ff0000'; + /** Visibility of a shape's class (defaults to visible for unknown classes). */ + const isShapeVisible = useCallback( + (s: Shape) => classMap.get(s.classId)?.isVisible !== false, + [classMap], + ); /** Commit new shapes, clipping them against other classes when the toggle is on * (neighbor classes act as a hard boundary), and merging with overlapping @@ -742,12 +931,12 @@ export default function AnnotationCanvas({ const computeCommittedSlice = useCallback((newShapes: Shape[], sliceShapes: Shape[]): Shape[] => { let toAdd = newShapes; if (clipToOtherClasses && meta && newShapes.some((s) => hasOtherClass(sliceShapes, s.classId))) { - toAdd = clipShapesToOthers(newShapes, sliceShapes, meta.width, meta.height); + toAdd = time('clip', () => clipShapesToOthers(newShapes, sliceShapes, meta.width, meta.height, workScale)); } // Auto-merge with overlapping same-class shapes (replaces those + the new shape with // one unioned polygon). Runs after clipping so other-class bounds still hold. if (mergeOverlappingSameClass && meta && toAdd.length) { - const { add, removeIds } = mergeNewWithSameClass(toAdd, sliceShapes, meta.width, meta.height); + const { add, removeIds } = time('merge', () => mergeNewWithSameClass(toAdd, sliceShapes, meta.width, meta.height, workScale)); if (removeIds.length) { const kept = sliceShapes.filter((s) => !removeIds.includes(s.id)); return [...kept, ...add]; @@ -755,19 +944,213 @@ export default function AnnotationCanvas({ toAdd = add; } return [...sliceShapes, ...toAdd]; - }, [clipToOtherClasses, mergeOverlappingSameClass, meta]); + }, [clipToOtherClasses, mergeOverlappingSameClass, meta, workScale]); const commitShapes = useCallback((newShapes: Shape[]) => { if (!sourceKey || newShapes.length === 0) return; const slice = useDatasetStore.getState().currentSlice; const sliceShapes = useAnnotationStore.getState().byImage[sourceKey]?.[String(slice)] ?? []; - setShapes(sourceKey, slice, computeCommittedSlice(newShapes, sliceShapes)); + const next = time('commit', () => computeCommittedSlice(newShapes, sliceShapes)); + setShapes(sourceKey, slice, next); }, [sourceKey, computeCommittedSlice, setShapes]); const activeColor = activeClassId !== null ? colorForClass(activeClassId) : '#4090ff'; - const showBrushCursor = (underlyingTool === 'brush' || underlyingTool === 'eraser') && !!meta && !isPreviewing; + const showBrushCursor = + (underlyingTool === 'brush' || underlyingTool === 'eraser' || underlyingTool === 'threshold') && + !!meta && !isPreviewing; + + /** + * Fit the Threshold Brush band from a lassoed example and apply it. + * + * Everything inside the loop is a positive example; a ring just outside it + * supplies negatives, so the result means "fill this and not what it touches" + * rather than "fill everything this bright". Work happens on a crop around the + * lasso, which keeps the blur sweep cheap on a full-resolution slice. + */ + const runSamplerFit = (path: number[]): void => { + if (!meta || path.length < 6) { onSamplerFit?.(null); return; } + const field = ensureThresholdField(); + if (!field) { onSamplerFit?.(null); return; } + const { gw, gh, scale } = field; + + // Everything below works on a CROP around the lasso, never the whole slice. + // Rasterizing and especially dilating full-grid is what made a large sample + // crawl: dilation costs one pass over its grid per cell of ring width, so on + // a 2560² field that is billions of operations for a region a few hundred + // pixels across. Cropping first makes the cost scale with the sample. + let minGX = Infinity, minGY = Infinity, maxGX = -Infinity, maxGY = -Infinity; + for (let i = 0; i + 1 < path.length; i += 2) { + const gx = path[i] / scale; + const gy = path[i + 1] / scale; + if (gx < minGX) minGX = gx; + if (gx > maxGX) maxGX = gx; + if (gy < minGY) minGY = gy; + if (gy > maxGY) maxGY = gy; + } + if (!Number.isFinite(minGX) || maxGX < minGX) { onSamplerFit?.(null); return; } + + // Margin: the ring, plus slack so the widest blur kernel is not distorted by + // the crop edge. Ring width is estimated from the lasso's area in cells. + const approxArea = Math.max(1, (maxGX - minGX) * (maxGY - minGY)); + const ringWidth = Math.min(MAX_RING_WIDTH, Math.max(4, Math.round(Math.sqrt(approxArea) / 2))); + const margin = ringWidth + Math.ceil(3 * Math.max(...BLUR_CANDIDATES)) + 2; + + const x0 = Math.max(0, Math.floor(minGX) - margin); + const y0 = Math.max(0, Math.floor(minGY) - margin); + const x1 = Math.min(gw - 1, Math.ceil(maxGX) + margin); + const y1 = Math.min(gh - 1, Math.ceil(maxGY) + margin); + const cw = x1 - x0 + 1; + const ch = y1 - y0 + 1; + if (cw < 2 || ch < 2) { onSamplerFit?.(null); return; } + + // Rasterize the lasso directly into crop coordinates by shifting it into the + // crop's frame (image units), so no full-size mask is ever allocated. + const shifted = path.map((v, i) => (i % 2 === 0 ? v - x0 * scale : v - y0 * scale)); + const lassoShape: Shape = { id: 'sampler', classId: 0, kind: 'polygon', points: shifted }; + const cropPos = rasterizeShapes([lassoShape], cw, ch, scale); + const cropNeg = backgroundRing(cropPos, cw, ch, dilate, ringWidth); + + const cropField = new Float32Array(cw * ch); + for (let y = 0; y < ch; y++) { + const src = (y0 + y) * gw + x0; + const dst = y * cw; + for (let x = 0; x < cw; x++) cropField[dst + x] = field.gray[src + x]; + } + + const result = fitWithBlurSweep(cropField, cw, ch, cropPos, cropNeg, applyGaussianBlurGray); + if (!result.ok) { onSamplerFit?.(null); return; } + + // Phase 2: also try a texture-aware score. Intensity alone cannot separate + // materials that share a grey range but differ in grain, nor a material whose + // brightness drifts across the slice. Fitting a projection over intensity + + // band-pass + local-std + local-mean-ratio, then running the SAME band fitter + // on the projected score, handles both — and because the score is graded by + // the same skill number, the two options are directly comparable. + const channels = buildChannels(cropField, cw, ch); + const projection = fitProjection(channels, cropPos, cropNeg); + let projected: SweepResult | null = null; + if (projection) { + const score = projectToScore(channels, projection); + const { posHist, negHist } = sampleHistograms(score, cropPos, cropNeg); + const fit = fitBand(posHist, negHist); + if (fit.ok) projected = { ...fit, extraSigma: 0 }; + } + + // Only switch to the projected score when it is meaningfully better. Equal + // results should stay on plain intensity: it is the mode the display sliders + // steer and the histogram picker describes, so it is the one to prefer. + const useProjection = + projection !== null && projected !== null && projected.skill > result.skill + 0.05; + + if (useProjection && projected && projection) { + projectionRef.current = { projection, key: baseKey }; + scoreFieldRef.current = null; // rebuilt lazily for the whole slice + setThresholdBand(projected.lo, projected.hi); + onSamplerFit?.({ + ...projected, + displayLo: projected.lo, + displayHi: projected.hi, + collapsed: false, + previousBand: [useToolStore.getState().thresholdLo, useToolStore.getState().thresholdHi], + previousBlur: blur, + appliedBlur: blur, + mode: 'projected', + intensitySkill: result.skill, + weights: CHANNEL_NAMES.map((name, i) => ({ name, weight: projection.weights[i] })), + }); + setTool('threshold'); + return; + } + + // Plain intensity won — drop any projection left from an earlier sample so the + // brush goes back to gating on what the display shows. + projectionRef.current = null; + scoreFieldRef.current = null; + + // The fit lives in the field's BASE space; the stored band is authored in + // DISPLAYED space. Rather than enumerate the ways that conversion can fail, + // do it and check it round-trips: convert to display, round as the store + // will, convert back, and require the original base band. That catches both + // failure modes at once — + // * collapse (contrast squashing the range onto one displayed value), and + // * saturation, where an endpoint lands on 0 or 255 and `displayBandToBase` + // correctly reopens it to ∓Infinity — which would silently select far + // MORE than was fitted. + const displayLo = Math.round(baseToDisplay(result.lo, displayAffine, gamma)); + const displayHi = Math.round(baseToDisplay(result.hi, displayAffine, gamma)); + const roundTrip = displayBandToBase(displayLo, displayHi, displayAffine, gamma); + // Storing the band as integers costs up to half a displayed level. How much + // that is in BASE units depends on the local slope, which gamma makes vary + // along the range — so measure it at the band edges rather than deriving it + // from the affine part alone. + const levelWidth = (d: number): number => { + const a = displayToBase(d - 0.5, displayAffine, gamma); + const b = displayToBase(d + 0.5, displayAffine, gamma); + return a === null || b === null ? Infinity : Math.abs(b - a); + }; + const tolerance = Math.max(1.5, 1.5 * Math.max(levelWidth(displayLo), levelWidth(displayHi))); + const collapsed = + !Number.isFinite(roundTrip.lo) || + !Number.isFinite(roundTrip.hi) || + Math.abs(roundTrip.lo - result.lo) > tolerance || + Math.abs(roundTrip.hi - result.hi) > tolerance; + + onSamplerFit?.({ + ...result, + displayLo, + displayHi, + collapsed, + previousBand: [useToolStore.getState().thresholdLo, useToolStore.getState().thresholdHi], + previousBlur: blur, + appliedBlur: combineSigma(blur, result.extraSigma), + }); + + if (collapsed) return; // leave the band alone; the UI explains why + setThresholdBand(displayLo, displayHi); + if (result.extraSigma > 0) onBlurChange?.(combineSigma(blur, result.extraSigma)); + // Sampling is a means, not an end: hand the user straight back to the brush + // the band was just fitted for. + setTool('threshold'); + }; + + /** Snapped path from the current anchor to `to`, or a straight line if the + * live-wire is unavailable (no edge map yet, or an off-grid point). */ + const samplerTraceTo = (to: { x: number; y: number }): number[] => { + const st = samplerLassoRef.current; + if (!st) return []; + if (!st.cm || !st.prev) return [to.x, to.y]; + try { + // `.slice(2)` drops the anchor itself, which is already committed. + return tracePath(st.cm, st.prev, imageToGrid(st.cm, to.x, to.y)).slice(2); + } catch { + return [to.x, to.y]; + } + }; + + /** Re-seed the live-wire at `at` so subsequent segments trace from there. */ + const samplerReseed = (at: { x: number; y: number }): void => { + const st = samplerLassoRef.current; + if (!st) return; + st.anchor = at; + st.prev = st.cm ? dijkstra(st.cm, imageToGrid(st.cm, at.x, at.y)) : null; + }; + + /** Close the sampler lasso, run the fit, and clear the transient state. */ + const finishSamplerLasso = (): void => { + const st = samplerLassoRef.current; + samplerLassoRef.current = null; + if (draftLineRef.current) { + draftLineRef.current.visible(false); + draftLineRef.current.closed(false); + draftStrokeLayerRef.current?.batchDraw(); + } + if (!st) return; + // Close the loop along the edge too, rather than cutting straight across it. + const path = [...st.committed, ...samplerTraceTo(st.start)]; + runSamplerFit(path); + }; /** Current pointer position mapped from stage/display coords to image pixels. */ const getPointerImagePos = () => { @@ -786,7 +1169,7 @@ export default function AnnotationCanvas({ if (magneticCostRef.current && magneticBuiltForRef.current === base) { return magneticCostRef.current; } - const cm = buildCostMap(base, meta.width, meta.height); + const cm = time('cost-map', () => buildCostMap(base, meta.width, meta.height, 512, workScale)); magneticCostRef.current = cm; magneticBuiltForRef.current = base; return cm; @@ -815,16 +1198,233 @@ export default function AnnotationCanvas({ // what the user sees, exactly like SAM. Cache key includes the display settings. const magicFieldRef = useRef(null); const magicFieldForRef = useRef(null); + /** Cache-key fragment for fields built from the display-adjusted slice (wand, SAM): + * brightness/contrast/levels are part of what those tools see. */ + const adjustedKey = `${sourceKey}|${currentSlice}|b${brightness}|c${contrast}|l${levelsLo}-${levelsHi}|pp${preprocessKey}|u${workScale}`; + /** Cache-key fragment for fields built from the PREPROCESSED base (threshold brush + * + its overlay, matching the histogram). Deliberately excludes brightness/ + * contrast/levels so those sliders never invalidate a field. */ + const baseKey = `${sourceKey}|${currentSlice}|pp${preprocessKey}|u${workScale}`; const ensureMagicField = useCallback((): GrayField | null => { if (!imageEl || !meta) return null; - const key = `${sourceKey}|${currentSlice}|b${brightness}|c${contrast}|l${levelsLo}-${levelsHi}|pp${preprocessKey}`; - if (magicFieldRef.current && magicFieldForRef.current === key) return magicFieldRef.current; - const src = renderAdjusted(imageEl, meta.width, meta.height, brightness, contrast, levelsLo, levelsHi, preprocess); - const f = buildField(src, meta.width, meta.height); + if (magicFieldRef.current && magicFieldForRef.current === adjustedKey) return magicFieldRef.current; + const src = renderAdjusted(imageEl, meta.width, meta.height, brightness, contrast, levelsLo, levelsHi, preprocess, workScale); + const f = time('magic-field', () => buildField(src, meta.width, meta.height, 1600, workScale)); magicFieldRef.current = f; - magicFieldForRef.current = key; + magicFieldForRef.current = adjustedKey; return f; - }, [imageEl, meta, sourceKey, currentSlice, brightness, contrast, levelsLo, levelsHi, preprocessKey, preprocess]); + }, [imageEl, meta, brightness, contrast, levelsLo, levelsHi, preprocess, workScale, adjustedKey]); + + // Threshold-brush field. Built at FULL working resolution (no 1600-px cap) — a + // brush needs per-pixel accuracy where a click-based wand can afford a coarser grid. + // + // Source is the PREPROCESSED base, not the brightness/contrast/levels-adjusted + // render the wand and SAM build. That is NOT a change in semantics: the band is + // still evaluated against displayed intensity, so the sliders steer the brush. + // The equivalence is moved rather than lost — `bandInBaseSpace` maps the band + // backwards through the display transform (monotonic, hence exactly the same + // pixel set; see lib/displayTransform.ts), which turns a per-tick full-image + // re-render into a couple of `Math.pow` calls. Only blur/CLAHE/sharpen, the + // slice, and the working scale invalidate this field. + // Phase 2: when a sample shows texture beats brightness, the brush gates on a + // fitted projection of several channels instead of raw intensity. Held in refs + // (not state) so activating it does not re-render the canvas. + const projectionRef = useRef<{ + projection: { weights: number[]; centers: number[]; scales: number[] }; + key: string; + } | null>(null); + /** Whole-slice projected score, built lazily from `projectionRef`. */ + const scoreFieldRef = useRef(null); + + const thresholdFieldRef = useRef(null); + const thresholdFieldForRef = useRef(null); + const ensureThresholdField = useCallback((): GrayField | null => { + if (!imageEl || !meta) return null; + if (thresholdFieldRef.current && thresholdFieldForRef.current === baseKey) return thresholdFieldRef.current; + // No gradient: the brush gates purely on intensity, and a Sobel pass over a + // full-resolution (possibly 4x) grid would be pure waste. + const f = time('threshold-field', () => buildField(displayBase ?? imageEl, meta.width, meta.height, Math.max(meta.width, meta.height), workScale, false)); + thresholdFieldRef.current = f; + thresholdFieldForRef.current = baseKey; + return f; + }, [imageEl, meta, displayBase, workScale, baseKey]); + + // Release the cached tool fields when the slice or sample changes. Each holds a + // Float32 gray plane (plus a gradient, for the wand) sized to the working + // resolution — up to tens of MB at 2x/4x — and without this they stay resident + // until the next use happens to replace them, which may be never. + useEffect(() => { + return () => { + magicFieldRef.current = null; + magicFieldForRef.current = null; + thresholdFieldRef.current = null; + thresholdFieldForRef.current = null; + overlayFieldRef.current = null; + overlayFieldForRef.current = null; + gateRef.current = null; + magneticCostRef.current = null; + magneticBuiltForRef.current = null; + }; + }, [sourceKey, currentSlice]); + + /** + * The field the Threshold Brush gates on. + * + * Normally the intensity field. Once a sample shows that texture separates the + * feature better, this becomes the fitted projection of all channels, computed + * across the whole slice and cached — the brush, the overlay and the band all + * read it, so they cannot disagree about what is selected. + */ + const ensureGateField = useCallback((): GrayField | null => { + const base = ensureThresholdField(); + const active = projectionRef.current; + if (!base || !active) return base; + // A projection is only valid for the field it was fitted on; a slice or + // preprocessing change invalidates it rather than silently misapplying it. + if (active.key !== baseKey) { + projectionRef.current = null; + scoreFieldRef.current = null; + return base; + } + if (scoreFieldRef.current) return scoreFieldRef.current; + const channels = buildChannels(base.gray, base.gw, base.gh); + const gray = projectToScore(channels, active.projection); + const score: GrayField = { gw: base.gw, gh: base.gh, scale: base.scale, gray }; + scoreFieldRef.current = score; + return score; + }, [ensureThresholdField, baseKey]); + + /** True while the brush is gating on a fitted score rather than brightness. */ + const usingProjection = (): boolean => + projectionRef.current !== null && projectionRef.current.key === baseKey; + + /** The threshold band, mapped from the DISPLAYED intensities the user authored + * it in into the base space the cached fields live in. Brightness/contrast/ + * levels/gamma therefore steer the brush exactly as they steer the image — + * without any of them invalidating a field or touching a pixel. */ + const bandInBaseSpace = useCallback(() => { + const { thresholdLo, thresholdHi } = useToolStore.getState(); + // With a projection active the gate field is a fitted score, not displayed + // brightness, so inverting the display transform would be meaningless — the + // band is already in the score's own units. (A consequence worth knowing: + // in that mode the display sliders no longer steer the brush.) + if (projectionRef.current && projectionRef.current.key === baseKey) { + return { lo: thresholdLo, hi: thresholdHi }; + } + return displayBandToBase(thresholdLo, thresholdHi, displayAffine, gamma); + }, [displayAffine, gamma, baseKey]); + + /** Binary gate over a field: 1 where the pixel reads as in-band on screen. This + * is what the brush may paint. Reads the band from the store rather than a + * subscribed value — see the overlay below for why the canvas deliberately does + * NOT re-render on band changes. */ + // Cached across strokes: the gate only changes when the band, the display + // transform, or the underlying field changes — not per mousedown, where + // rebuilding it meant a fresh multi-megabyte pass at the start of every stroke. + const gateRef = useRef<{ key: string; field: GrayField; gate: Uint8Array } | null>(null); + const buildThresholdGate = useCallback((field: GrayField): Uint8Array => { + const { lo, hi } = bandInBaseSpace(); + const key = `${baseKey}|${lo}|${hi}`; + const cached = gateRef.current; + if (cached && cached.key === key && cached.field === field) return cached.gate; + + const gate = new Uint8Array(field.gw * field.gh); + const { gray } = field; + for (let i = 0; i < gate.length; i++) { + const v = gray[i]; + if (v >= lo && v <= hi) gate[i] = 1; + } + gateRef.current = { key, field, gate }; + return gate; + }, [bandInBaseSpace, baseKey]); + + // ---- Red in-band overlay (ImageJ's threshold display) ------------------- + // Shows exactly which pixels the brush may paint. Three things keep dragging + // the band (or any display slider) smooth, all of which cost real frames when + // done the obvious way: + // 1. Its own SMALL field (≤1024 px, no upscale, no gradient) instead of the + // brush's full-resolution one — activating the tool or nudging brightness + // must not trigger a full-res renderAdjusted + buildField. + // 2. Repainted IMPERATIVELY into a persistent canvas — routing it through + // React state re-rendered this whole (very large) component per frame. + // 3. Driven by a store subscription, so the canvas never re-renders on a band + // change at all. That's why thresholdLo/Hi are not subscribed above. + const overlayFieldRef = useRef(null); + const overlayFieldForRef = useRef(null); + const overlayCanvasRef = useRef(null); + const overlayImageDataRef = useRef(null); + const overlayImageRef = useRef(null); + const showThresholdOverlay = tool === 'threshold' && thresholdOverlay && !isPreviewing; + + const ensureOverlayField = useCallback((): GrayField | null => { + if (!imageEl || !meta) return null; + // With a projection active the overlay must show the PROJECTED selection, or + // it would advertise a different set of pixels than the brush paints. The + // gate field is already cached, so reuse it rather than fitting a second one. + if (projectionRef.current && projectionRef.current.key === baseKey) { + return ensureGateField(); + } + if (overlayFieldRef.current && overlayFieldForRef.current === baseKey) return overlayFieldRef.current; + // Same preprocessed source as the brush field (see above), so the overlay shows + // exactly what the brush would paint — including as the display sliders move, + // since both read the band through `bandInBaseSpace`. + const f = time('overlay-field', () => buildField(displayBase ?? imageEl, meta.width, meta.height, 1024, 1, false)); + overlayFieldRef.current = f; + overlayFieldForRef.current = baseKey; + return f; + }, [imageEl, meta, displayBase, baseKey, ensureGateField]); + + /** Repaint the overlay canvas for the current band and push it to Konva. */ + const paintThresholdOverlay = useCallback(() => time('overlay-paint', () => { + const node = overlayImageRef.current; + if (!node) return; + const field = ensureOverlayField(); + if (!field) return; + const { gw, gh, gray } = field; + let canvas = overlayCanvasRef.current; + if (!canvas || canvas.width !== gw || canvas.height !== gh) { + canvas = document.createElement('canvas'); + canvas.width = gw; + canvas.height = gh; + overlayCanvasRef.current = canvas; + overlayImageDataRef.current = null; + } + const ctx = canvas.getContext('2d'); + if (!ctx) return; + const { lo, hi } = bandInBaseSpace(); + // Reuse the pixel buffer across repaints — a band drag repaints every frame + // and a fresh multi-MB ImageData per frame is pure GC churn. + let img = overlayImageDataRef.current; + if (!img) { + img = ctx.createImageData(gw, gh); + overlayImageDataRef.current = img; + } + const d = img.data; + for (let i = 0; i < gray.length; i++) { + const v = gray[i]; + const o = i * 4; + if (v < lo || v > hi) { d[o + 3] = 0; continue; } + d[o] = 239; d[o + 1] = 68; d[o + 2] = 68; d[o + 3] = 255; // red-500 + } + ctx.putImageData(img, 0, 0); + node.image(canvas); + node.getLayer()?.batchDraw(); + }), [ensureOverlayField, bandInBaseSpace]); + + useEffect(() => { + if (!showThresholdOverlay) return; + // Coalesce bursts (a band drag commits once per frame) into one repaint. + let pending = 0; + const schedule = () => { + if (pending) return; + pending = requestAnimationFrame(() => { pending = 0; paintThresholdOverlay(); }); + }; + schedule(); + const unsub = useToolStore.subscribe((s, prev) => { + if (s.thresholdLo !== prev.thresholdLo || s.thresholdHi !== prev.thresholdHi) schedule(); + }); + return () => { unsub(); if (pending) cancelAnimationFrame(pending); }; + }, [showThresholdOverlay, paintThresholdOverlay]); // Auto negative ("not") prompts for SAM: interior points of nearby other-class // regions, so a new selection won't bleed into already-labeled areas. Anchored @@ -841,8 +1441,7 @@ export default function AnnotationCanvas({ : null; if (!ref) return []; - const isClassVisible = (s: Shape) => - classes.find((c) => c.classId === s.classId)?.isVisible !== false; + const isClassVisible = isShapeVisible; const cands: Array<{ x: number; y: number }> = []; for (const shape of storeShapes) { @@ -924,6 +1523,19 @@ export default function AnnotationCanvas({ const field = ensureMagicField(); if (!field) { setMagicPreview([]); return; } setMagicLoading(true); + // Fill should stop at a pixel already claimed by a different annotated + // class (Peter's "like a wall" ask) — rasterize the other classes onto + // the same grid `field` uses and pass it as a flood-time barrier, so the + // preview itself is bounded, not just the post-commit clip (`clipToOtherClasses` + // already trims the committed shape, but the flood could wander arbitrarily + // far across a neighbor first). + const blocked = + isFill && activeClassId !== null + ? (() => { + const others = storeShapes.filter((s) => s.classId !== activeClassId); + return others.length > 0 ? rasterizeUnion(others, field.gw, field.gh, field.scale) : undefined; + })() + : undefined; const id = requestAnimationFrame(() => { const polys: number[][] = []; for (const seed of magicSeeds) { @@ -933,6 +1545,7 @@ export default function AnnotationCanvas({ mode: isFill ? 'contiguous' : magicMode, smooth: isFill ? 1 : magicSigma, edgeStop: isFill ? 0 : magicEdgeStop, + blocked, }), ); } @@ -940,7 +1553,7 @@ export default function AnnotationCanvas({ setMagicLoading(false); }); return () => cancelAnimationFrame(id); - }, [tool, magicSeeds, magicBox, autoNegPoints, samDetail, samThreshold, samConnectedOnly, samEncodeKey, makeSamSource, magicEngine, magicTolerance, magicMode, magicSigma, magicEdgeStop, fillThreshold, ensureMagicField, imageEl, meta, sam.ensureEncoded, sam.segment]); // eslint-disable-line react-hooks/exhaustive-deps + }, [tool, magicSeeds, magicBox, autoNegPoints, samDetail, samThreshold, samConnectedOnly, samEncodeKey, makeSamSource, magicEngine, magicTolerance, magicMode, magicSigma, magicEdgeStop, fillThreshold, ensureMagicField, imageEl, meta, sam.ensureEncoded, sam.segment, storeShapes, activeClassId]); // eslint-disable-line react-hooks/exhaustive-deps /** Commit the magic preview polygons as new shapes (one batched undo step). */ const commitMagic = useCallback(() => { @@ -1042,6 +1655,171 @@ export default function AnnotationCanvas({ // eslint-disable-next-line react-hooks/exhaustive-deps }, [draft.magneticSeed, draft.tool, draft.sourceKey, draft.sliceKey, sourceKey, currentSlice, imageEl, displayBase]); + // ---- Threshold brush ---------------------------------------------------- + // Paints only where the displayed intensity is inside the band, so a stroke + // stops dead at a feature boundary. The result can't be expressed as a stroke + + // radius, so it accumulates as a mask and commits as POLYGONS. + + /** Begin a threshold stroke at `pos`, priming the mask, gate, and preview canvas. */ + const startThresholdStroke = (pos: { x: number; y: number }, erase: boolean): void => { + const field = ensureGateField(); + if (!field || !meta) return; + const { gw, gh, scale } = field; + const canvas = document.createElement('canvas'); + canvas.width = gw; + canvas.height = gh; + const ctx = canvas.getContext('2d'); + if (!ctx) return; + thresholdStrokeRef.current = { + mode: erase ? 'erase' : 'paint', + gw, gh, scale, + mask: new Uint8Array(gw * gh), + gate: buildThresholdGate(field), + last: pos, + canvas, + imageData: ctx.createImageData(gw, gh), + }; + if (thresholdPreviewRef.current) { + thresholdPreviewRef.current.image(canvas); + thresholdPreviewRef.current.visible(true); + } + extendThresholdStroke(pos); + }; + + /** Stamp the segment from the last point to `pos` and repaint the preview. + * Only the segment's dirty rect is rewritten — repainting the whole grid every + * mousemove would be tens of millions of writes per frame at a 2x/4x full-res + * grid, which is exactly the size this tool is meant to be used at. */ + const extendThresholdStroke = (pos: { x: number; y: number }): void => { + const st = thresholdStrokeRef.current; + if (!st) return; + const { gw, gh, scale, mask, gate } = st; + const from = st.last; + stampStroke(mask, gw, gh, [from.x, from.y, pos.x, pos.y], brushSize / scale, scale, 1, gate); + st.last = pos; + + // Dirty rect of this segment in grid cells (the stamped capsule's bbox). + const r = Math.max(0.5, brushSize / scale) + 1; + const x0 = Math.max(0, Math.floor(Math.min(from.x, pos.x) / scale - r)); + const x1 = Math.min(gw - 1, Math.ceil(Math.max(from.x, pos.x) / scale + r)); + const y0 = Math.max(0, Math.floor(Math.min(from.y, pos.y) / scale - r)); + const y1 = Math.min(gh - 1, Math.ceil(Math.max(from.y, pos.y) / scale + r)); + if (x1 < x0 || y1 < y0) return; + + // Erase strokes preview white (matching the eraser's draft line); paint + // strokes use the active class color. + const d = st.imageData.data; + const [cr, cg, cb] = st.mode === 'erase' ? [255, 255, 255] : hexToRgb(activeColor); + for (let y = y0; y <= y1; y++) { + for (let x = x0; x <= x1; x++) { + const i = y * gw + x; + if (!mask[i]) continue; + const o = i * 4; + d[o] = cr; d[o + 1] = cg; d[o + 2] = cb; d[o + 3] = 255; + } + } + st.canvas + .getContext('2d') + ?.putImageData(st.imageData, 0, 0, x0, y0, x1 - x0 + 1, y1 - y0 + 1); + thresholdPreviewRef.current?.getLayer()?.batchDraw(); + }; + + /** Flush the buffered threshold stroke: mask → polygons → one undo step. */ + const commitThresholdStroke = (): void => { + const st = thresholdStrokeRef.current; + thresholdStrokeRef.current = null; + if (thresholdPreviewRef.current) { + thresholdPreviewRef.current.visible(false); + thresholdPreviewRef.current.image(undefined); + thresholdPreviewRef.current.getLayer()?.batchDraw(); + } + if (!st || !sourceKey || !meta || activeClassId === null) return; + + const { gw, gh, scale } = st; + let any = false; + for (let i = 0; i < st.mask.length; i++) if (st.mask[i]) { any = true; break; } + if (!any) return; + + // Regularize before vectorizing. A per-pixel gate speckles, and every speck + // becomes its own polygon: that is what made one stroke commit hundreds of + // shapes and dominate clip cost. A morphological opening (erode then dilate) + // drops isolated pixels and pinholes while leaving real regions intact, and + // the small-component pass clears what survives. Better geometry AND a much + // cheaper commit, from the same step. + const mask = removeSmallComponents( + dilate(erode(st.mask, gw, gh, 1), gw, gh, 1), + gw, + gh, + THRESHOLD_MIN_REGION, + ); + let survives = false; + for (let i = 0; i < mask.length; i++) if (mask[i]) { survives = true; break; } + // A thin stroke can be erased entirely by the opening; keep the raw mask + // rather than silently discarding what the user just painted. + const finalMask = survives ? mask : st.mask; + + const regions = maskToPolygonsWithHoles(finalMask, gw, gh, { minRegion: 4, scale }) + .filter((p) => p.points.length >= 6); + if (regions.length === 0) return; + + if (st.mode === 'erase') { + // Subtract the painted region from the shapes it actually overlaps, via the + // same node-preserving boolean path the eraser uses. The raster overlap test + // matters: `subtractFromShape` always rebuilds geometry, so running it on + // untouched shapes would churn their ids and vertices for nothing. + const sliceShapes = useAnnotationStore.getState().byImage[sourceKey]?.[String(currentSlice)] ?? []; + const visible = isShapeVisible; + const inScope = (s: Shape) => eraseAllClasses || s.classId === activeClassId; + // The stroke's own bounds come free from the regions we just vectorized, so + // a shape nowhere near it is rejected before any rasterization happens. + const strokeBounds = unionBBox( + regions.map((p) => ({ id: '', classId: 0, kind: 'polygon' as const, points: p.points })), + ); + const scratch = new Uint8Array(gw * gh); + const overlapsStroke = (s: Shape) => { + if (!bboxNear(shapeBBox(s), strokeBounds, scale + 1)) return false; + scratch.fill(0); + rasterizeShapes([s], gw, gh, scale, scratch); + for (let i = 0; i < finalMask.length; i++) if (finalMask[i] && scratch[i]) return true; + return false; + }; + const stampMP = regionsToMultiPolygon(regions); + if (stampMP.length === 0) return; + + const next: Shape[] = []; + const replacedSelection: string[] = []; + let changed = false; + for (const s of sliceShapes) { + if (!inScope(s) || !visible(s) || !overlapsStroke(s)) { next.push(s); continue; } + const polys = subtractFromShape(s, stampMP, meta.width, meta.height); + if (polys === null) { next.push(s); continue; } + changed = true; + next.push(...polys); + if (selectedShapeIds.includes(s.id)) replacedSelection.push(...polys.map((p) => p.id)); + } + if (changed) { + setShapes(sourceKey, currentSlice, next); + setSelectedShapeIds([ + ...selectedShapeIds.filter((id) => next.some((s) => s.id === id)), + ...replacedSelection, + ]); + } + return; + } + + // Paint: commit as polygons through the shared path, so clip-to-other-classes + // and merge-same-class apply and the whole stroke is a single undo step. + commitShapes( + regions.map((p) => ({ + id: uuidv4(), + classId: activeClassId, + kind: 'polygon' as const, + points: p.points, + ...(p.holes.length ? { holes: p.holes } : {}), + })), + ); + }; + /** Flush the buffered draft stroke to the Zustand store (one write per stroke). */ const commitDraftStroke = () => { const draft = draftStrokeRef.current; @@ -1061,8 +1839,7 @@ export default function AnnotationCanvas({ if (mode === 'erase') { if (!meta) return; const sliceShapes = useAnnotationStore.getState().byImage[sourceKey]?.[String(currentSlice)] ?? []; - const visible = (s: Shape) => - classes.find((c) => c.classId === s.classId)?.isVisible !== false; + const visible = isShapeVisible; const inScope = (s: Shape) => eraseAllClasses || s.classId === activeClassId; const strokeHits = (s: Shape) => { // Radius-aware so grazing a shape's edge (disk overlaps, center outside) @@ -1086,7 +1863,7 @@ export default function AnnotationCanvas({ // Bake the erase directly into polygon geometry for every target — the // vertices always match the visible shape, splits produce independent // polygons, and undo is one clean shape-replacement step. - const { gw, gh, scale } = fullResGridFor(meta.width, meta.height); + const { gw, gh, scale } = fullResGridFor(meta.width, meta.height, workScale); // The erase stamp as polygon geometry, subtracted from each shape via a true // boolean difference so untouched vertices are PRESERVED (only the cut edge // gets new points). Falls back per-shape to a rasterize→re-vectorize round-trip @@ -1139,7 +1916,7 @@ export default function AnnotationCanvas({ // Re-vectorize at full resolution so the round-trip is ~idempotent: existing // merged/clipped regions keep their shape instead of eroding or shifting a // little each time a stroke is added. - const { gw, gh, scale } = fullResGridFor(meta.width, meta.height); + const { gw, gh, scale } = fullResGridFor(meta.width, meta.height, workScale); const prospective = { ...brush, strokes: [...brush.strokes, { points: finalPoints, radius, mode: 'paint' as const }] }; const mine = rasterizeShapes([prospective], gw, gh, scale); let clipChanged = false; @@ -1147,6 +1924,13 @@ export default function AnnotationCanvas({ // Clip detection (fast mask test): does the brush overlap other classes? // The actual clip is applied below via boolean difference so the brush tiles // flush against the neighbor (a mask carve left a ~1px unlabeled seam). + // Bounds of the brush including this stroke. Used ONLY for the same-class + // merge test below — the clip detection deliberately does not pre-filter + // (see the note in lib/clipToClasses.ts about the reverted filter). + const mineBounds = shapeBBox(prospective); + const nearBrush = (s: Shape) => bboxNear(shapeBBox(s), mineBounds, scale + 1); + const overlapScratch = new Uint8Array(gw * gh); + if (clipToOtherClasses) { const others = sliceShapes.filter((s) => s.classId !== brush.classId); if (others.length > 0) { @@ -1159,8 +1943,10 @@ export default function AnnotationCanvas({ const mergeTargets = mergeOverlappingSameClass ? sliceShapes.filter((s) => { if (s.id === shapeId || s.classId !== brush.classId) return false; - const sm = rasterizeShapes([s], gw, gh, scale); - for (let i = 0; i < mine.length; i++) if (mine[i] && sm[i]) return true; + if (!nearBrush(s)) return false; + overlapScratch.fill(0); + rasterizeShapes([s], gw, gh, scale, overlapScratch); + for (let i = 0; i < mine.length; i++) if (mine[i] && overlapScratch[i]) return true; return false; }) : []; @@ -1176,12 +1962,30 @@ export default function AnnotationCanvas({ // Clip via boolean difference so the brush tiles flush against other // classes (no ~1px unlabeled seam); existing shapes are untouched. + // + // `clipChanged` is decided by a mask test, so we KNOW there is overlap to + // remove here. Every failure below must therefore fall back to the mask + // clip rather than keeping the shape as-is: returning `bp` unchanged (the + // old behaviour) leaves the brush overlapping its neighbour while the + // toggle claims otherwise, which looks exactly like clipping being off. if (clipChanged) { - const otherMP = unionShapesToMultiPolygon( - sliceShapes.filter((s) => s.classId !== brush.classId), meta.width, meta.height, - ); - if (otherMP.length) { - brushPolys = brushPolys.flatMap((bp) => subtractFromShape(bp, otherMP, meta.width, meta.height) ?? [bp]); + const others = sliceShapes.filter((s) => s.classId !== brush.classId); + const { mp: otherMP, ok } = unionShapesChecked(others, meta.width, meta.height); + if (ok && otherMP.length) { + // Boolean path per polygon; anything it can't do is batched into a + // single mask clip (that helper caches its masks per call, so one + // call for many shapes is far cheaper than one call each). + const kept: Shape[] = []; + const failed: Shape[] = []; + for (const bp of brushPolys) { + const sub = subtractFromShape(bp, otherMP, meta.width, meta.height); + if (sub === null) failed.push(bp); else kept.push(...sub); + } + brushPolys = failed.length + ? [...kept, ...clipShapesToOthersMask(failed, sliceShapes, meta.width, meta.height, workScale)] + : kept; + } else { + brushPolys = clipShapesToOthersMask(brushPolys, sliceShapes, meta.width, meta.height, workScale); } } @@ -1362,6 +2166,45 @@ export default function AnnotationCanvas({ draftLineRef.current.visible(true); draftStrokeLayerRef.current?.batchDraw(); } + } else if (tool === 'sampler') { + // Magnetic lasso: seed the live-wire here, then snap each dragged segment + // to the strongest edge between anchors. No store writes and no undo entry + // — this measures, it does not annotate. + setIsDrawing(true); + const cm = ensureCostMap(); + samplerLassoRef.current = { + committed: [pos.x, pos.y], + anchor: pos, + start: pos, + cm, + prev: cm ? dijkstra(cm, imageToGrid(cm, pos.x, pos.y)) : null, + }; + if (draftLineRef.current) { + draftLineRef.current.stroke('#38bdf8'); + draftLineRef.current.strokeWidth(2 / transform.scaleX); + draftLineRef.current.points([pos.x, pos.y]); + draftLineRef.current.closed(false); + draftLineRef.current.visible(true); + draftStrokeLayerRef.current?.batchDraw(); + } + } else if (tool === 'threshold') { + // Shift-click samples instead of painting: re-center the band on the pixel + // under the cursor, keeping the configured sample width. + if (e.evt.shiftKey) { + // Sample from the GATE field so the picked value is in the same units as + // the band — with a projection active those are score units, not greys. + const field = ensureGateField(); + if (field) { + const gx = Math.max(0, Math.min(field.gw - 1, Math.floor(pos.x / field.scale))); + const gy = Math.max(0, Math.min(field.gh - 1, Math.floor(pos.y / field.scale))); + const v = field.gray[gy * field.gw + gx]; + const half = thresholdSampleWidth / 2; + setThresholdBand(v - half, v + half); + } + return; + } + setIsDrawing(true); + startThresholdStroke(pos, e.evt.altKey); } else if (tool === 'eraser') { setIsDrawing(true); const sliceShapes = byImage[sourceKey]?.[String(currentSlice)] ?? []; @@ -1371,8 +2214,7 @@ export default function AnnotationCanvas({ // cursor we still start the stroke (so the preview shows and the user can // drag ONTO a shape) — the target is resolved from the whole stroke on // commit (see commitDraftStroke). - const visible = (s: Shape) => - classes.find((c) => c.classId === s.classId)?.isVisible !== false; + const visible = isShapeVisible; // "Erase all classes" ignores the active-class restriction (any visible shape). const inScope = (s: Shape) => eraseAllClasses || s.classId === activeClassId; let target: Shape | undefined; @@ -1447,6 +2289,33 @@ export default function AnnotationCanvas({ return; } + // Sampler lasso: preview the snapped segment from the anchor to the cursor, + // dropping a new anchor once the cursor has travelled far enough. Anchoring + // periodically (rather than per pixel) is what keeps the live-wire honest — + // one dijkstra per anchor, and the locked-in path stops re-flowing behind you. + if (tool === 'sampler' && e.evt.buttons === 1 && samplerLassoRef.current) { + const st = samplerLassoRef.current; + const traced = samplerTraceTo(pos); + if (draftLineRef.current) { + draftLineRef.current.points([...st.committed, ...traced]); + draftStrokeLayerRef.current?.batchDraw(); + } + const dx = pos.x - st.anchor.x; + const dy = pos.y - st.anchor.y; + if (dx * dx + dy * dy >= SAMPLER_ANCHOR_STEP * SAMPLER_ANCHOR_STEP) { + st.committed = [...st.committed, ...traced]; + samplerReseed(pos); + } + return; + } + + // Threshold brush: stamp into the in-band mask and repaint the preview. Like + // the brush, this writes nothing to the store until mouseup. + if (tool === 'threshold' && e.evt.buttons === 1 && thresholdStrokeRef.current) { + extendThresholdStroke(pos); + return; + } + // Buffer brush/eraser points — no store writes here if ((tool === 'brush' || tool === 'eraser') && e.evt.buttons === 1 && draftStrokeRef.current) { draftStrokeRef.current.points.push(pos.x, pos.y); @@ -1471,7 +2340,7 @@ export default function AnnotationCanvas({ const additive = marqueeShiftRef.current; if (rect && rect.w > 3 && rect.h > 3 && sourceKey) { const ids = storeShapes - .filter((s) => classes.find((c) => c.classId === s.classId)?.isVisible !== false) + .filter(isShapeVisible) .filter((s) => shapeIntersectsRect(s, rect)) .map((s) => s.id); setSelectedShapeIds(additive ? Array.from(new Set([...selectedShapeIds, ...ids])) : ids); @@ -1503,6 +2372,8 @@ export default function AnnotationCanvas({ if (isDrawing) { commitDraftStroke(); + commitThresholdStroke(); + finishSamplerLasso(); setIsDrawing(false); } if (!sourceKey || !meta || activeClassId === null) return; @@ -1565,6 +2436,8 @@ export default function AnnotationCanvas({ // Commit any in-progress stroke if (isDrawing) { commitDraftStroke(); + commitThresholdStroke(); + finishSamplerLasso(); setIsDrawing(false); } setDragStart(null); @@ -1595,19 +2468,44 @@ export default function AnnotationCanvas({ }, [showBrushCursor]); // eslint-disable-line react-hooks/exhaustive-deps /** Zoom in/out toward the cursor, keeping the point under the pointer fixed. */ + // Zoom is applied at most once per animation frame. A trackpad or free-spin + // wheel emits events far faster than the display refreshes, and each one used to + // trigger a full React render (and potentially a layer re-cache). Accumulating + // into a ref and flushing on rAF collapses a burst into a single update, with + // the same final scale and focal point. + const pendingZoomRef = useRef<{ steps: number; pointer: { x: number; y: number } } | null>(null); + const zoomRafRef = useRef(null); + useEffect(() => () => { if (zoomRafRef.current != null) cancelAnimationFrame(zoomRafRef.current); }, []); + const handleWheel = (e: Konva.KonvaEventObject) => { e.evt.preventDefault(); const stage = stageRef.current; if (!stage) return; - const scaleBy = 1.1; - const oldScale = transform.scaleX; - const pointer = stage.getPointerPosition()!; - const newScale = e.evt.deltaY < 0 ? oldScale * scaleBy : oldScale / scaleBy; - const mousePointTo = { x: (pointer.x - transform.x) / oldScale, y: (pointer.y - transform.y) / oldScale }; - setTransform({ - scaleX: newScale, scaleY: newScale, - x: pointer.x - mousePointTo.x * newScale, - y: pointer.y - mousePointTo.y * newScale, + const pointer = stage.getPointerPosition(); + if (!pointer) return; + + const pending = pendingZoomRef.current; + const steps = (pending?.steps ?? 0) + (e.evt.deltaY < 0 ? 1 : -1); + pendingZoomRef.current = { steps, pointer }; + if (zoomRafRef.current != null) return; + + zoomRafRef.current = requestAnimationFrame(() => { + zoomRafRef.current = null; + const job = pendingZoomRef.current; + pendingZoomRef.current = null; + if (!job || job.steps === 0) return; + setTransform((t) => { + const newScale = t.scaleX * Math.pow(1.1, job.steps); + const mousePointTo = { + x: (job.pointer.x - t.x) / t.scaleX, + y: (job.pointer.y - t.y) / t.scaleY, + }; + return { + scaleX: newScale, scaleY: newScale, + x: job.pointer.x - mousePointTo.x * newScale, + y: job.pointer.y - mousePointTo.y * newScale, + }; + }); }); }; @@ -1641,138 +2539,6 @@ export default function AnnotationCanvas({ return null; }; - /** Destination-out lines that carve erase strokes out of a vector shape. */ - const renderErased = (erased?: EraseStroke[]) => - (erased ?? []).map((st, i) => ( - - )); - - /** Render a polygon that may have holes via an even-odd fill (outer path minus - * hole subpaths). Even-odd — not destination-out — so a hole reveals whatever - * is *beneath* it (e.g. another class) instead of erasing it off the layer. */ - const renderPolygonWithHoles = (points: number[], holes: number[][], color: string, strokeW: number) => ( - { - buildRingsPath(ctx, [points, ...holes]); - const raw = (ctx as unknown as { _context: CanvasRenderingContext2D })._context; - raw.fillStyle = color; - raw.fill('evenodd'); - ctx.strokeShape(node); - }} - /> - ); - - /** Render a committed shape (any kind) on the cached display layer, with the - * active brush instance recolored to the active class and erase strokes carved out. */ - const renderShape = (shape: Shape) => { - const color = - shape.id === activeBrushShapeId && activeClassId !== null - ? colorForClass(activeClassId) - : colorForClass(shape.classId); - const isSelected = selectedShapeIds.includes(shape.id); - const strokeW = (isSelected ? 2 : 1) / transform.scaleX; - - if (shape.kind === 'polygon') { - return ( - - {shape.holes?.length - ? renderPolygonWithHoles(shape.points, shape.holes, color, strokeW) - : ( - - )} - {renderErased(shape.erased)} - - ); - } - if (shape.kind === 'rectangle') { - return ( - - - {renderErased(shape.erased)} - - ); - } - if (shape.kind === 'ellipse') { - return ( - - - {renderErased(shape.erased)} - - ); - } - if (shape.kind === 'brush') { - return ( - - {shape.strokes.map((stroke, i) => { - if (stroke.mode === 'erase') { - return ( - - ); - } - return ( - - ); - })} - - ); - } - return null; - }; // ----- Interactive selection / move / vertex-edit (select tool only) ----- @@ -1794,7 +2560,7 @@ export default function AnnotationCanvas({ const pos = getPointerImagePos(); if (pos && activeClassId !== null) { const isVisible = (s: Shape) => - classes.find((c) => c.classId === s.classId)?.isVisible !== false; + isShapeVisible(s); let topActive: string | null = null; for (const s of storeShapes) { // Later in draw order = rendered on top → keep the last (topmost) match. @@ -2087,6 +2853,23 @@ export default function AnnotationCanvas({ }; const showInteractive = tool === 'select' && !isPreviewing; + + // Visible shapes with the single-selected one moved LAST, so its vertex handles + // (incl. amber hole vertices) sit above any shape enclosed in its holes. Memoized + // and built with a single partition rather than a comparator sort — this ran on + // every render of the select tool, and a sort whose only job is to hoist one + // element doesn't need to compare every pair. + const interactiveShapes = useMemo(() => { + if (!showInteractive) return EMPTY_SHAPES; + const rest: Shape[] = []; + let selected: Shape | null = null; + for (const s of storeShapes) { + if (!isShapeVisible(s)) continue; + if (s.id === selectedId) selected = s; + else rest.push(s); + } + return selected ? [...rest, selected] : rest; + }, [showInteractive, storeShapes, isShapeVisible, selectedId]); // The single selected shape (resize/transform only applies to one). const selectedShape = sourceKey && selectedId ? storeShapes.find((s) => s.id === selectedId) ?? null @@ -2226,7 +3009,7 @@ export default function AnnotationCanvas({ if (mod && k === 'a') { // Select all shapes on the slice, scoped to the active class or all classes. e.preventDefault(); - const visible = (s: Shape) => classes.find((c) => c.classId === s.classId)?.isVisible !== false; + const visible = isShapeVisible; const ids = storeShapes .filter((s) => (selectScope === 'all' || s.classId === activeClassId) && visible(s)) .map((s) => s.id); @@ -2247,7 +3030,7 @@ export default function AnnotationCanvas({
- {imageEl && meta && ( + {layerGroups.image && imageEl && meta && ( + {/* Threshold band overlay (ImageJ-style): every pixel the threshold brush + is currently allowed to paint, in translucent red. Sits above the image + but below the annotations so existing regions stay readable. */} + {showThresholdOverlay && meta && ( + + + + )} + + {/* iPred overlays: probability heatmap / conformal predictions / manifold + suggestions. Sits above the image, below annotations — see OverlaysLayer. */} + {meta && ( + + )} + {/* Layer 1: committed shapes — cached + opacity applied once at the layer so overlapping same-class shapes render a uniform class color. - Clipped to the image frame so strokes never render past the edges. */} - - {displayShapes - .filter((s) => renderClasses.find((c) => c.classId === s.classId)?.isVisible !== false) - .map(renderShape)} - + Clipped to the image frame so strokes never render past the edges. + Memoized (see ShapesLayer) so display-slider ticks don't reconcile it. */} + {meta && layerGroups.annotations && ( + + )} {/* Interactive layer: hit targets for select/move/edit (select tool only). */} {showInteractive && ( - {storeShapes - .filter((s) => classes.find((c) => c.classId === s.classId)?.isVisible !== false) - // Render the single-selected shape LAST so its vertex handles (incl. - // amber hole vertices) sit above any shape enclosed in its holes. - .sort((a, b) => (a.id === selectedId ? 1 : 0) - (b.id === selectedId ? 1 : 0)) - .map(renderInteractive)} + {interactiveShapes.map(renderInteractive)} + {/* Threshold-brush stroke preview: the in-band mask painted so far, + drawn at image size (its canvas is at the working resolution). */} + {meta && ( + + )} {/* Layer 4: brush/eraser size cursor preview (position updated imperatively) */} diff --git a/frontend/src/components/annotate/ClassManager/index.test.tsx b/frontend/src/components/annotate/ClassManager/index.test.tsx new file mode 100644 index 0000000..f3fdfe6 --- /dev/null +++ b/frontend/src/components/annotate/ClassManager/index.test.tsx @@ -0,0 +1,145 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import ClassManager from './index'; +import { useClassStore } from '@/stores/classStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { useReferenceGuideStore } from '@/stores/referenceGuideStore'; +import { useSettingsStore } from '@/stores/settingsStore'; +import { useDatasetStore } from '@/stores/datasetStore'; + +beforeEach(() => { + useClassStore.setState({ classes: [] }); + useAnnotationStore.getState().reset(); + useReferenceGuideStore.getState().clear(); + useSettingsStore.setState({ colorblindMode: false }); + useDatasetStore.setState({ source: null, kind: null, serverUri: null } as any); + vi.spyOn(window, 'confirm').mockReturnValue(true); + vi.spyOn(window, 'alert').mockImplementation(() => {}); +}); + +afterEach(() => { + cleanup(); + vi.restoreAllMocks(); +}); + +describe('ClassManager', () => { + it('shows an empty-state message with no classes', () => { + render(); + expect(screen.getByText(/No classes yet/)).toBeInTheDocument(); + }); + + it('shows default quick-add suggestions when there is no guide', () => { + render(); + expect(screen.getByRole('button', { name: /air/ })).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /substrate/ })).toBeInTheDocument(); + }); + + it('quick-add creates and activates a class', async () => { + const onActivate = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole('button', { name: 'air' })); + expect(useClassStore.getState().classes).toHaveLength(1); + expect(useClassStore.getState().classes[0].label).toBe('air'); + expect(onActivate).toHaveBeenCalledWith(useClassStore.getState().classes[0].classId); + }); + + it('quick-add re-activates an existing class of the same label instead of duplicating', async () => { + useClassStore.getState().addClass('air', '#111111'); + const onActivate = vi.fn(); + const user = userEvent.setup(); + render(); + // Already exists, so it's no longer offered as a quick-add suggestion — + // the row itself is the only way to reactivate it. + expect(screen.queryByRole('button', { name: 'air' })).not.toBeInTheDocument(); + await user.click(screen.getByRole('option', { name: 'air' })); + expect(useClassStore.getState().classes).toHaveLength(1); + }); + + it('adds a class via the add form', async () => { + const onActivate = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('Add class')); + await user.type(screen.getByPlaceholderText('Class label'), 'Cell'); + await user.click(screen.getByRole('button', { name: 'Add' })); + expect(useClassStore.getState().classes.map((c) => c.label)).toEqual(['Cell']); + }); + + it('rejects a duplicate label with an alert and does not add it', async () => { + useClassStore.getState().addClass('Cell', '#111111'); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('Add class')); + await user.type(screen.getByPlaceholderText('Class label'), 'Cell'); + await user.click(screen.getByRole('button', { name: 'Add' })); + expect(window.alert).toHaveBeenCalledWith('A class with that label already exists.'); + expect(useClassStore.getState().classes).toHaveLength(1); + }); + + it('toggles class visibility', async () => { + useClassStore.getState().addClass('Cell', '#111111'); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('Hide class')); + expect(useClassStore.getState().classes[0].isVisible).toBe(false); + }); + + it('renames a class through the inline editor', async () => { + useClassStore.getState().addClass('Cell', '#111111'); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('Edit class label and color')); + const input = screen.getByDisplayValue('Cell'); + await user.clear(input); + await user.type(input, 'Nucleus{Enter}'); + expect(useClassStore.getState().classes[0].label).toBe('Nucleus'); + }); + + it('deletes a class after confirmation', async () => { + useClassStore.getState().addClass('Cell', '#111111'); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('Delete class and its annotations')); + expect(window.confirm).toHaveBeenCalled(); + expect(useClassStore.getState().classes).toHaveLength(0); + }); + + it('cancelling the confirm dialog keeps the class', async () => { + vi.spyOn(window, 'confirm').mockReturnValue(false); + useClassStore.getState().addClass('Cell', '#111111'); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('Delete class and its annotations')); + expect(useClassStore.getState().classes).toHaveLength(1); + }); + + it('duplicates a class and its shapes into a new class', async () => { + useDatasetStore.setState({ source: 'sample.tif', kind: 'local', serverUri: null } as any); + const classId = useClassStore.getState().addClass('Cell', '#111111'); + useAnnotationStore.getState().replaceClassShapesOnSlice('local:sample.tif', 0, classId, [ + { id: 's1', classId, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + const onActivate = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('Duplicate class and its annotations')); + const classes = useClassStore.getState().classes; + expect(classes).toHaveLength(2); + expect(classes[1].label).toBe('Cell copy'); + const newId = classes[1].classId; + expect(useAnnotationStore.getState().byImage['local:sample.tif']['0'].some((s) => s.classId === newId)).toBe(true); + }); + + it('shows guide-defined suggestions instead of the generic defaults when a guide is loaded', () => { + useReferenceGuideStore.getState().setGuide( + [{ label: 'Pore', color: '#123456', description: '', exampleCrops: [] }], + '', + 'local:sample.tif', + ); + render(); + expect(screen.getByRole('button', { name: /Pore/ })).toBeInTheDocument(); + expect(screen.queryByRole('button', { name: 'air' })).not.toBeInTheDocument(); + }); +}); diff --git a/frontend/src/components/annotate/ClassManager/index.tsx b/frontend/src/components/annotate/ClassManager/index.tsx index 936e5b9..286e321 100644 --- a/frontend/src/components/annotate/ClassManager/index.tsx +++ b/frontend/src/components/annotate/ClassManager/index.tsx @@ -16,6 +16,7 @@ import { getClassPalette } from '@/lib/classColors'; import { buildSourceKey } from '@/lib/sourceKey'; import { cn } from '@/lib/utils'; import { Copy } from '@phosphor-icons/react'; +import CollapsibleSection from '@/components/common/CollapsibleSection'; /** Counts shapes assigned to a class within a single sample (*sourceKey*), across its * slices. Scoped to the current sample so the delete prompt never counts (or deletes) @@ -316,9 +317,9 @@ export default function ClassManager({ activeClassId, onActivate, onClassDeleted ); return ( -
-
- Classes + -
- + } + >
-
+ ); } diff --git a/frontend/src/components/annotate/DenoiseBakeModal.test.tsx b/frontend/src/components/annotate/DenoiseBakeModal.test.tsx new file mode 100644 index 0000000..0ae8a2e --- /dev/null +++ b/frontend/src/components/annotate/DenoiseBakeModal.test.tsx @@ -0,0 +1,195 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import DenoiseBakeModal, { type DenoiseBakeModalProps } from './DenoiseBakeModal'; + +function renderModal(overrides: Partial = {}) { + const props: DenoiseBakeModalProps = { + open: true, + onClose: vi.fn(), + source: 'browse/foo', + serverUri: 'http://tiled.example', + denoise: { method: 'tv', strength: 0.42 }, + nSlices: 10, + methodLabel: 'Total variation', + ...overrides, + }; + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + const invalidateSpy = vi.spyOn(queryClient, 'invalidateQueries'); + const utils = render( + + + + ); + return { ...utils, props, invalidateSpy }; +} + +beforeEach(() => { + global.fetch = vi.fn(); +}); + +afterEach(() => { + cleanup(); + vi.restoreAllMocks(); +}); + +describe('DenoiseBakeModal', () => { + it('renders nothing when open is false', () => { + const { container } = renderModal({ open: false }); + expect(container).toBeEmptyDOMElement(); + }); + + it('shows method label, strength, and slice count; prefills destination', () => { + renderModal({ source: 'browse/foo', methodLabel: 'Total variation', denoise: { method: 'tv', strength: 0.42 }, nSlices: 10 }); + expect(screen.getByText(/Total variation, strength 0\.42/)).toBeInTheDocument(); + expect(screen.getByText(/all 10 slices/)).toBeInTheDocument(); + expect(screen.getByDisplayValue('browse/foo_denoised')).toBeInTheDocument(); + }); + + it('uses singular "slice" when nSlices is 1', () => { + renderModal({ nSlices: 1 }); + expect(screen.getByText(/all 1 slice and written/)).toBeInTheDocument(); + }); + + it('calls onClose from the Cancel button and the X button', async () => { + const onClose = vi.fn(); + const user = userEvent.setup(); + renderModal({ onClose }); + await user.click(screen.getByRole('button', { name: 'Cancel' })); + expect(onClose).toHaveBeenCalledTimes(1); + await user.click(screen.getByLabelText('Close')); + expect(onClose).toHaveBeenCalledTimes(2); + }); + + it('disables Save copy when destination is blank', async () => { + const user = userEvent.setup(); + renderModal(); + const input = screen.getByDisplayValue('browse/foo_denoised'); + await user.clear(input); + expect(screen.getByRole('button', { name: 'Save copy' })).toBeDisabled(); + }); + + it('resets destination/description when reopened for a different source', () => { + const { rerender, props } = renderModal({ source: 'browse/foo' }); + rerender( + + + + ); + expect(screen.getByDisplayValue('browse/bar_denoised')).toBeInTheDocument(); + }); + + it('starts a bake job, polls status, shows progress, then completion and invalidates the browse query', async () => { + let statusCall = 0; + (global.fetch as ReturnType).mockImplementation(async (url: string, opts?: RequestInit) => { + if (url.includes('/api/denoise/bake')) { + expect(opts?.method).toBe('POST'); + const body = JSON.parse(opts!.body as string); + expect(body).toEqual({ + source: 'browse/foo', + server_uri: 'http://tiled.example', + method: 'tv', + strength: 0.42, + target_path: 'browse/foo_denoised', + description: '', + }); + return { ok: true, json: async () => ({ job_id: 'job-1' }) }; + } + if (url.includes('/api/export/status/job-1')) { + statusCall++; + if (statusCall === 1) { + return { + ok: true, + json: async () => ({ state: 'running', phase: 'Denoising', done: 3, total: 10, error: null, result: null }), + }; + } + return { + ok: true, + json: async () => ({ state: 'done', phase: 'Done', done: 10, total: 10, error: null, result: { path: 'browse/foo_denoised', n_slices: 10 } }), + }; + } + throw new Error(`unexpected fetch: ${url}`); + }); + + const user = userEvent.setup(); + const { invalidateSpy } = renderModal(); + await user.click(screen.getByRole('button', { name: 'Save copy' })); + + await waitFor(() => { + expect(screen.getByText(/Denoising — 3\/10/)).toBeInTheDocument(); + }); + + await waitFor(() => { + expect(screen.getByText('Saved 10 slices to browse/foo_denoised.')).toBeInTheDocument(); + }, { timeout: 3000 }); + + expect(invalidateSpy).toHaveBeenCalledWith({ queryKey: ['browse'] }); + }, 10000); + + it('shows a cancelled message when the result reports cancelled', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/denoise/bake')) { + return { ok: true, json: async () => ({ job_id: 'job-2' }) }; + } + if (url.includes('/api/export/status/job-2')) { + return { + ok: true, + json: async () => ({ state: 'done', phase: 'Done', done: 4, total: 10, error: null, result: { cancelled: true } }), + }; + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('button', { name: 'Save copy' })); + await waitFor(() => { + expect(screen.getByText('Stopped — the partial copy was discarded.')).toBeInTheDocument(); + }, { timeout: 3000 }); + }, 10000); + + it('shows an error message when the bake request fails', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/denoise/bake')) { + return { ok: false, status: 500, json: async () => ({ detail: 'boom' }) }; + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('button', { name: 'Save copy' })); + await waitFor(() => { + expect(screen.getByText('boom')).toBeInTheDocument(); + }); + }); + + it('calls the cancel endpoint when Stop is clicked while running', async () => { + (global.fetch as ReturnType).mockImplementation(async (url: string) => { + if (url.includes('/api/denoise/bake')) { + return { ok: true, json: async () => ({ job_id: 'job-3' }) }; + } + if (url.includes('/api/export/status/job-3')) { + // Never resolves to done so it stays in "running" while we click Stop. + return { + ok: true, + json: async () => ({ state: 'running', phase: 'Denoising', done: 1, total: 10, error: null, result: null }), + }; + } + if (url.includes('/api/export/cancel/job-3')) { + return { ok: true, json: async () => ({}) }; + } + throw new Error(`unexpected fetch: ${url}`); + }); + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('button', { name: 'Save copy' })); + await waitFor(() => screen.getByRole('button', { name: 'Stop' })); + await user.click(screen.getByRole('button', { name: 'Stop' })); + await waitFor(() => { + expect(global.fetch).toHaveBeenCalledWith( + expect.stringContaining('/api/export/cancel/job-3'), + expect.objectContaining({ method: 'POST' }) + ); + }); + }); +}); diff --git a/frontend/src/components/annotate/DenoiseBakeModal.tsx b/frontend/src/components/annotate/DenoiseBakeModal.tsx new file mode 100644 index 0000000..80a48a7 --- /dev/null +++ b/frontend/src/components/annotate/DenoiseBakeModal.tsx @@ -0,0 +1,216 @@ +/** + * DenoiseBakeModal — write the current denoise settings out as a new dataset. + * + * The Annotate preview is deliberately non-destructive: it changes what you see + * and what the intensity tools act on, but exports still use the original + * pixels. This is the other half — it produces a first-class dataset you can + * open, annotate and export, sitting next to its source in Browse. + * + * The job runs on the shared export-job registry, so it reports progress and + * cancels through the same routes everything else long-running here uses. + */ +import { useCallback, useEffect, useState } from 'react'; +import { useQueryClient } from '@tanstack/react-query'; +import { X, Warning, CheckCircle, CircleNotch } from '@phosphor-icons/react'; +import { API_BASE } from '@/config'; +import type { DenoiseOpts } from '@/stores/datasetStore'; + +interface JobState { + state: 'pending' | 'running' | 'done' | 'error'; + phase: string; + done: number; + total: number; + error: string | null; + result: { path?: string; n_slices?: number; cancelled?: boolean } | null; +} + +export interface DenoiseBakeModalProps { + open: boolean; + onClose: () => void; + source: string; + serverUri: string | null; + denoise: DenoiseOpts; + /** Slice count, so the estimate is honest about how long this will take. */ + nSlices: number; + /** Method label for display (e.g. "Total variation"). */ + methodLabel: string; +} + +/** `browse/foo` -> `browse/foo_denoised`, matching the backend default. */ +function defaultTarget(source: string): string { + return `${source.replace(/\/+$/, '')}_denoised`; +} + +async function readError(res: Response): Promise { + try { + const body = await res.json(); + if (typeof body?.detail === 'string') return body.detail; + return JSON.stringify(body); + } catch { + return `Request failed (${res.status})`; + } +} + +export default function DenoiseBakeModal({ + open, onClose, source, serverUri, denoise, nSlices, methodLabel, +}: DenoiseBakeModalProps) { + const queryClient = useQueryClient(); + const [target, setTarget] = useState(() => defaultTarget(source)); + const [description, setDescription] = useState(''); + const [job, setJob] = useState(null); + const [jobId, setJobId] = useState(null); + const [error, setError] = useState(null); + + useEffect(() => { + if (open) { + setTarget(defaultTarget(source)); + setJob(null); + setJobId(null); + setError(null); + } + }, [open, source]); + + const start = useCallback(async () => { + setError(null); + setJob({ state: 'pending', phase: 'Starting', done: 0, total: nSlices, error: null, result: null }); + try { + const res = await fetch(`${API_BASE}/api/denoise/bake`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + source, + server_uri: serverUri, + method: denoise.method, + strength: denoise.strength, + target_path: target, + description, + }), + }); + if (!res.ok) throw new Error(await readError(res)); + const { job_id: id } = await res.json(); + setJobId(id); + + for (;;) { + await new Promise((r) => setTimeout(r, 1000)); + const s = await fetch(`${API_BASE}/api/export/status/${id}`); + if (!s.ok) throw new Error(await readError(s)); + const status: JobState = await s.json(); + setJob(status); + if (status.state === 'error') throw new Error(status.error || 'Bake failed'); + if (status.state === 'done') break; + } + // A new dataset exists — let Browse pick it up without a reload. + await queryClient.invalidateQueries({ queryKey: ['browse'] }); + } catch (err) { + setError(err instanceof Error ? err.message : String(err)); + setJob(null); + } + }, [source, serverUri, denoise, target, description, nSlices, queryClient]); + + const cancel = useCallback(async () => { + if (!jobId) return; + // Cooperative: the job stops at its next slice boundary and discards the + // partial dataset, rather than leaving something that looks complete. + await fetch(`${API_BASE}/api/export/cancel/${jobId}`, { method: 'POST' }); + }, [jobId]); + + if (!open) return null; + + const running = job !== null && job.state !== 'done'; + const finished = job?.state === 'done'; + const percent = job && job.total > 0 ? Math.round((job.done / job.total) * 100) : 0; + + return ( +
+
+
+
+

Save denoised copy

+

+ {methodLabel}, strength {denoise.strength.toFixed(2)} — applied to all{' '} + {nSlices} slice{nSlices === 1 ? '' : 's'} and written as a new dataset. +

+
+ +
+ + {!running && !finished && ( + <> + + setTarget(e.target.value)} + className="mb-3 w-full rounded border border-sky-800 bg-sky-900/50 px-2 py-1 text-xs" + /> + + setDescription(e.target.value)} + placeholder="e.g. denoised, tv" + className="mb-4 w-full rounded border border-sky-800 bg-sky-900/50 px-2 py-1 text-xs" + /> +

+ The original dataset is not modified. This writes a copy — on a large + volume it can take a while, and it can be cancelled. +

+
+ + +
+ + )} + + {running && ( +
+

+ + {job.phase} — {job.done}/{job.total} +

+
+
+
+
+ +
+
+ )} + + {finished && ( +
+

+ + {job.result?.cancelled + ? 'Stopped — the partial copy was discarded.' + : `Saved ${job.result?.n_slices ?? job.done} slices to ${job.result?.path ?? target}.`} +

+
+ +
+
+ )} + + {error && ( +

+ + {error} +

+ )} +
+
+ ); +} diff --git a/frontend/src/components/annotate/DenoisePanel/index.test.tsx b/frontend/src/components/annotate/DenoisePanel/index.test.tsx new file mode 100644 index 0000000..102e0b9 --- /dev/null +++ b/frontend/src/components/annotate/DenoisePanel/index.test.tsx @@ -0,0 +1,292 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import DenoisePanel, { type DenoiseMethodInfo } from './index'; +import type { DenoiseOpts } from '@/stores/datasetStore'; + +const METHODS: DenoiseMethodInfo[] = [ + { method: 'none', label: 'None', cost: 'cheap', description: '', available: true, z_radius: 0 }, + { method: 'gaussian', label: 'Gaussian', cost: 'cheap', description: 'Simple blur.', available: true, z_radius: 0 }, + { method: 'bilateral', label: 'Bilateral', cost: 'slow', description: 'Edge-preserving.', available: true, z_radius: 0 }, + { method: 'nlm', label: 'NLM', cost: 'moderate', description: 'Non-local means.', available: false, z_radius: 0 }, +]; + +function renderWithClient(ui: React.ReactElement) { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render({ui}); +} + +function baseDenoise(overrides: Partial = {}): DenoiseOpts { + return { method: 'none', strength: 0.5, ...overrides } as DenoiseOpts; +} + +beforeEach(() => { + vi.stubGlobal( + 'fetch', + vi.fn(async (url: string) => { + if (String(url).includes('/api/denoise/methods')) { + return { ok: true, json: async () => ({ methods: METHODS }) } as Response; + } + if (String(url).includes('/api/denoise/auto')) { + return { ok: true, json: async () => ({ strength: 0.42, noise_sigma: 0.01 }) } as Response; + } + return { ok: false, json: async () => ({}) } as Response; + }), + ); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('DenoisePanel', () => { + it('loads and lists the available methods from the server', async () => { + const onChange = vi.fn(); + renderWithClient( + , + ); + await waitFor(() => { + expect(screen.getByRole('option', { name: 'None' })).toBeInTheDocument(); + }); + expect(screen.getByRole('option', { name: 'Gaussian' })).toBeInTheDocument(); + expect(screen.getByRole('option', { name: /Bilateral \(slow\)/ })).toBeInTheDocument(); + expect(screen.getByRole('option', { name: /NLM — unavailable/ })).toBeDisabled(); + }); + + it('does not show the strength slider or extra controls when method is none', async () => { + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByRole('option', { name: 'None' })).toBeInTheDocument()); + expect(screen.queryByText('Strength')).not.toBeInTheDocument(); + expect(screen.queryByRole('button', { name: /Auto/ })).not.toBeInTheDocument(); + }); + + it('shows the strength slider, description, and Auto button when a method is active', async () => { + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByText('Simple blur.')).toBeInTheDocument()); + expect(screen.getByText('Strength')).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /Auto/ })).toBeInTheDocument(); + }); + + it('calls onChange with the new method when the select changes', async () => { + const onChange = vi.fn(); + const user = userEvent.setup(); + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByRole('option', { name: 'Gaussian' })).toBeInTheDocument()); + await user.selectOptions(screen.getByRole('combobox'), 'gaussian'); + expect(onChange).toHaveBeenCalledWith({ method: 'gaussian' }); + }); + + it('shows the "Save denoised copy" button only when onBake is provided', async () => { + const onBake = vi.fn(); + const user = userEvent.setup(); + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByText('Save denoised copy…')).toBeInTheDocument()); + await user.click(screen.getByText('Save denoised copy…')); + expect(onBake).toHaveBeenCalledOnce(); + }); + + it('omits "Save denoised copy" when onBake is not provided', async () => { + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByText('Strength')).toBeInTheDocument()); + expect(screen.queryByText('Save denoised copy…')).not.toBeInTheDocument(); + }); + + it('shows the busy indicator only while active and busy', async () => { + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByText('filtering…')).toBeInTheDocument()); + }); + + it('applyAuto calls onChange with the measured strength', async () => { + const onChange = vi.fn(); + const user = userEvent.setup(); + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByRole('button', { name: /Auto/ })).toBeInTheDocument()); + await user.click(screen.getByRole('button', { name: /Auto/ })); + await waitFor(() => expect(onChange).toHaveBeenCalledWith({ strength: 0.42 })); + const [url] = (fetch as any).mock.calls.find(([u]: [string]) => String(u).includes('/api/denoise/auto')); + expect(url).toContain('slice_index=3'); + expect(url).toContain('method=gaussian'); + }); + + it('shows a clean-slice warning when Auto measures zero strength', async () => { + (fetch as any).mockImplementation(async (url: string) => { + if (String(url).includes('/api/denoise/methods')) { + return { ok: true, json: async () => ({ methods: METHODS }) }; + } + if (String(url).includes('/api/denoise/auto')) { + return { ok: true, json: async () => ({ strength: 0, noise_sigma: 0.002 }) }; + } + return { ok: false, json: async () => ({}) }; + }); + const user = userEvent.setup(); + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByRole('button', { name: /Auto/ })).toBeInTheDocument()); + await user.click(screen.getByRole('button', { name: /Auto/ })); + expect(await screen.findByText(/This slice looks clean/)).toBeInTheDocument(); + }); + + it('shows an error message when the Auto request fails', async () => { + (fetch as any).mockImplementation(async (url: string) => { + if (String(url).includes('/api/denoise/methods')) { + return { ok: true, json: async () => ({ methods: METHODS }) }; + } + if (String(url).includes('/api/denoise/auto')) { + return { ok: false, json: async () => ({}) }; + } + return { ok: false, json: async () => ({}) }; + }); + const user = userEvent.setup(); + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByRole('button', { name: /Auto/ })).toBeInTheDocument()); + await user.click(screen.getByRole('button', { name: /Auto/ })); + expect(await screen.findByText('Could not measure this slice')).toBeInTheDocument(); + }); + + it('disables the Auto button when there is no source', async () => { + renderWithClient( + , + ); + await waitFor(() => expect(screen.getByRole('button', { name: /Auto/ })).toBeInTheDocument()); + expect(screen.getByRole('button', { name: /Auto/ })).toBeDisabled(); + }); + + it('clears a stale auto error when the method changes', async () => { + (fetch as any).mockImplementation(async (url: string) => { + if (String(url).includes('/api/denoise/methods')) { + return { ok: true, json: async () => ({ methods: METHODS }) }; + } + if (String(url).includes('/api/denoise/auto')) { + return { ok: false, json: async () => ({}) }; + } + return { ok: false, json: async () => ({}) }; + }); + const user = userEvent.setup(); + const { rerender } = renderWithClient( + , + ); + await waitFor(() => expect(screen.getByRole('button', { name: /Auto/ })).toBeInTheDocument()); + await user.click(screen.getByRole('button', { name: /Auto/ })); + expect(await screen.findByText('Could not measure this slice')).toBeInTheDocument(); + + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + rerender( + + + , + ); + expect(screen.queryByText('Could not measure this slice')).not.toBeInTheDocument(); + }); +}); diff --git a/frontend/src/components/annotate/DenoisePanel/index.tsx b/frontend/src/components/annotate/DenoisePanel/index.tsx new file mode 100644 index 0000000..6e782ed --- /dev/null +++ b/frontend/src/components/annotate/DenoisePanel/index.tsx @@ -0,0 +1,185 @@ +/** + * DenoisePanel — server-side denoising of the slice, before it is normalized. + * + * Unlike the sliders next to it, this is NOT a display filter. Denoising runs on + * the raw slice in its own intensity units (noise statistics do not survive a + * trip through the 8-bit display range), so it changes what every intensity + * tool sees — the Threshold Brush, the Sampler's fit, magic wand, livewire — not + * just what is drawn. It still does not change exported pixels; turning a + * denoised view into data is what "Save denoised copy" is for. + * + * The slider is debounced, and results are cached server-side per + * (slice, method, strength) — measured 1.61s -> 0.006s on a repeat — so + * revisiting a setting is free. Two refinements are NOT here yet and are worth + * knowing about: + * + * - **Side-by-side comparison.** Denoising is judged by what it *removes*, + * which the current on/off switch hides: flipping between two images makes + * you compare from memory. A wipe overlay needs a second image layer in + * `AnnotationCanvas`. + * - **Crop while dragging.** `buildSliceUrl` already supports `denoise_crop` + * (measured: bilateral on a 2560² slice, 7.3s full vs 0.49s at 512px), but + * nothing requests it yet, so a slow method is slow on every commit. + */ +import { useCallback, useEffect, useState } from 'react'; +import { useQuery } from '@tanstack/react-query'; +import { MagicWand, Warning } from '@phosphor-icons/react'; +import DebouncedSlider from '@/components/common/DebouncedSlider'; +import { API_BASE } from '@/config'; +import type { DenoiseOpts } from '@/stores/datasetStore'; + +/** One entry of `GET /api/denoise/methods`. */ +export interface DenoiseMethodInfo { + method: string; + label: string; + cost: 'cheap' | 'moderate' | 'slow'; + description: string; + available: boolean; + z_radius: number; +} + +export interface DenoisePanelProps { + denoise: DenoiseOpts; + onChange: (opts: Partial) => void; + /** Dataset identity, for the "Auto" suggestion (which reads this slice). */ + source: string | null; + kind: string | null; + serverUri: string | null; + sliceIndex: number; + /** Opens the bake modal — turning the preview into a real dataset. */ + onBake?: () => void; + /** True while the slice request for the current settings is in flight. */ + busy?: boolean; +} + +export default function DenoisePanel({ + denoise, + onChange, + source, + kind, + serverUri, + sliceIndex, + onBake, + busy, +}: DenoisePanelProps) { + const [autoError, setAutoError] = useState(null); + const [autoBusy, setAutoBusy] = useState(false); + + const { data: methods = [] } = useQuery({ + queryKey: ['denoise-methods'], + queryFn: async () => { + const res = await fetch(`${API_BASE}/api/denoise/methods`); + if (!res.ok) throw new Error('Failed to load denoise methods'); + return (await res.json()).methods; + }, + staleTime: Infinity, // capability of the server, not of the dataset + }); + + useEffect(() => setAutoError(null), [denoise.method]); + + const active = denoise.method !== 'none'; + const current = methods.find((m) => m.method === denoise.method); + + const applyAuto = useCallback(async () => { + if (!source || !kind) return; + setAutoBusy(true); + setAutoError(null); + try { + const params = new URLSearchParams({ + source, kind, slice_index: String(sliceIndex), method: denoise.method, + }); + if (serverUri) params.set('server_uri', serverUri); + const res = await fetch(`${API_BASE}/api/denoise/auto?${params}`); + if (!res.ok) throw new Error('Could not measure this slice'); + const { strength, noise_sigma: sigma } = await res.json(); + onChange({ strength }); + if (strength === 0) { + // Saying "0.0" alone reads as a failure. It is a finding: this slice is + // already clean, and smoothing it would only cost detail. + setAutoError( + `This slice looks clean (noise ≈ ${(sigma * 100).toFixed(2)}% of its range) — ` + + 'denoising it would mostly remove detail.' + ); + } + } catch (err) { + setAutoError(err instanceof Error ? err.message : String(err)); + } finally { + setAutoBusy(false); + } + }, [source, kind, serverUri, sliceIndex, denoise.method, onChange]); + + return ( +
+
+ + {busy && active && filtering…} +
+ + + + {current && current.method !== 'none' && ( +

{current.description}

+ )} + + {active && ( + <> + onChange({ strength: v })} + /> + +
+ + {onBake && ( + + )} +
+ + {autoError && ( +

+ + {autoError} +

+ )} + +

+ Affects the intensity tools too, not just the display. Exports still use + the original pixels. +

+ + )} +
+ ); +} diff --git a/frontend/src/components/annotate/DisplayControls/index.test.tsx b/frontend/src/components/annotate/DisplayControls/index.test.tsx new file mode 100644 index 0000000..dc241ea --- /dev/null +++ b/frontend/src/components/annotate/DisplayControls/index.test.tsx @@ -0,0 +1,104 @@ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import DisplayControls, { type DisplayControlsProps } from './index'; + +afterEach(() => { + cleanup(); +}); + +function baseProps(overrides: Partial = {}): DisplayControlsProps { + return { + brightness: 0, + contrast: 0, + onBrightnessChange: vi.fn(), + onContrastChange: vi.fn(), + onReset: vi.fn(), + histogramBins: null, + levelsLo: 0, + levelsHi: 255, + onLevelsChange: vi.fn(), + onLevelsReset: vi.fn(), + colormap: 'gray', + gamma: 1, + onColormapChange: vi.fn(), + onGammaChange: vi.fn(), + clahe: false, + sharpen: false, + onClaheChange: vi.fn(), + onSharpenChange: vi.fn(), + blur: 0, + onBlurChange: vi.fn(), + upscale: 1, + onUpscaleChange: vi.fn(), + ...overrides, + }; +} + +describe('DisplayControls', () => { + it('renders the reset button and calls onReset', async () => { + const onReset = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('Reset brightness and contrast')); + expect(onReset).toHaveBeenCalledOnce(); + }); + + it('toggles CLAHE and Sharpen checkboxes', async () => { + const onClaheChange = vi.fn(); + const onSharpenChange = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByLabelText('CLAHE')); + expect(onClaheChange).toHaveBeenCalledWith(true); + await user.click(screen.getByLabelText('Sharpen')); + expect(onSharpenChange).toHaveBeenCalledWith(true); + }); + + it('marks the current colormap as pressed', () => { + render(); + expect(screen.getByTitle('viridis')).toHaveAttribute('aria-pressed', 'true'); + expect(screen.getByTitle('gray')).toHaveAttribute('aria-pressed', 'false'); + }); + + it('clicking a colormap swatch calls onColormapChange', async () => { + const onColormapChange = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByTitle('viridis')); + expect(onColormapChange).toHaveBeenCalledWith('viridis'); + }); + + it('working-resolution radio group reflects the current upscale', () => { + render(); + const radios = screen.getAllByRole('radio'); + expect(radios[0]).toHaveAttribute('aria-checked', 'false'); + expect(radios[1]).toHaveAttribute('aria-checked', 'true'); + }); + + it('disables an upscale option above maxUpscale', () => { + render(); + const radios = screen.getAllByRole('radio'); + expect(radios[0]).not.toBeDisabled(); + expect(radios[1]).toBeDisabled(); + expect(radios[2]).toBeDisabled(); + }); + + it('clicking an enabled upscale option calls onUpscaleChange', async () => { + const onUpscaleChange = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getAllByRole('radio')[1]); + expect(onUpscaleChange).toHaveBeenCalledWith(2); + }); + + it('renders the denoise slot when provided', () => { + render( })} />); + expect(screen.getByTestId('denoise-slot')).toBeInTheDocument(); + }); + + it('omits the denoise slot area when not provided', () => { + render(); + expect(screen.queryByTestId('denoise-slot')).not.toBeInTheDocument(); + }); +}); diff --git a/frontend/src/components/annotate/DisplayControls/index.tsx b/frontend/src/components/annotate/DisplayControls/index.tsx index 45a2ffb..db461e3 100644 --- a/frontend/src/components/annotate/DisplayControls/index.tsx +++ b/frontend/src/components/annotate/DisplayControls/index.tsx @@ -6,6 +6,7 @@ import { ArrowCounterClockwise } from '@phosphor-icons/react'; import DebouncedSlider from '@/components/common/DebouncedSlider'; import HistogramControl from '@/components/annotate/HistogramControl'; +import CollapsibleSection from '@/components/common/CollapsibleSection'; import { COLORMAP_NAMES, colormapGradient, type ColormapName } from '@/lib/colormaps'; export interface DisplayControlsProps { @@ -30,8 +31,24 @@ export interface DisplayControlsProps { sharpen: boolean; onClaheChange: (v: boolean) => void; onSharpenChange: (v: boolean) => void; + /** Gaussian pre-blur sigma in image pixels (0 = off) — denoises what the + * intensity-driven tools see, as well as the display. */ + blur: number; + onBlurChange: (v: number) => void; + /** Working resolution multiplier (1, 2, 4) for the drawing tools. */ + upscale: number; + onUpscaleChange: (v: number) => void; + /** Highest upscale this slice can afford before the guard clamps it. */ + maxUpscale?: number; + /** Slot rendered above the pre-blur slider, for server-side denoising. + * A slot rather than props so this component stays free of data-fetching + * concerns — denoising needs the dataset identity and a network round trip, + * neither of which any other control here does. */ + denoiseSlot?: React.ReactNode; } +const UPSCALES = [1, 2, 4] as const; + /** Renders the brightness/contrast sliders + levels histogram with reset buttons. */ export default function DisplayControls({ brightness, @@ -52,11 +69,17 @@ export default function DisplayControls({ sharpen, onClaheChange, onSharpenChange, + blur, + onBlurChange, + upscale, + onUpscaleChange, + maxUpscale = 4, + denoiseSlot, }: DisplayControlsProps) { return ( -
-
- Display + -
- + } + >
- {/* Nonlinear enhancers (display-only; can combine: CLAHE → Sharpen). */} + {/* Server-side denoising, applied to the RAW slice before normalization — + unlike the pre-blur below, which is a client-side filter on the already + 8-bit display image. Placed first because it is the upstream stage. */} + {denoiseSlot && ( +
{denoiseSlot}
+ )} + + {/* Gaussian pre-blur — denoises so threshold/wand/fill see coherent regions. */} +
+ (v === 0 ? 'off' : v.toFixed(2))} + min={0} + max={5} + step={0.25} + value={blur} + onChange={onBlurChange} + // Each commit re-blurs the whole slice (and invalidates every tool + // field), so commit on pause rather than on every tick. + debounceMs={200} + /> +
+ + {/* Nonlinear enhancers (display-only; can combine: Blur → CLAHE → Sharpen). */}
-
+ + {/* Working resolution — resamples the slice for the drawing tools so small + features get more pixels to annotate against. Coordinates stay native. */} +
+ Working resolution +
+ {UPSCALES.map((u) => { + const tooBig = u > maxUpscale; + return ( + + ); + })} +
+

+ Resamples the slice for the drawing tools so smaller features can be annotated. + Exported pixels and annotation coordinates are unchanged. +

+
+ ); } diff --git a/frontend/src/components/annotate/DownloadModal.test.tsx b/frontend/src/components/annotate/DownloadModal.test.tsx new file mode 100644 index 0000000..2cc9371 --- /dev/null +++ b/frontend/src/components/annotate/DownloadModal.test.tsx @@ -0,0 +1,275 @@ +/** + * DownloadModal — scope/format selection, export payload building + * (including predicted_slices pointer logic), and export/mask-sync job wiring. + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import { MemoryRouter } from 'react-router'; +import DownloadModal from './DownloadModal'; +import { useDatasetStore } from '@/stores/datasetStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { usePredictedRasterStore } from '@/stores/predictedRasterStore'; +import { useClassStore } from '@/stores/classStore'; +import { useRatingStore } from '@/stores/ratingStore'; +import { useSettingsStore } from '@/stores/settingsStore'; + +function renderModal(onClose = vi.fn()) { + return render( + + + , + ); +} + +/** A fetch mock whose POST resolves to a synchronous "done" response (no job_id), + * so useExportJob treats it as immediately done — avoids polling in tests. */ +function mockFetchDone(result: Record = {}) { + return vi.fn(async () => ({ + ok: true, + json: async () => result, + text: async () => JSON.stringify(result), + })) as unknown as typeof fetch; +} + +beforeEach(() => { + useDatasetStore.setState({ + kind: 'local', source: 'sample.tif', serverUri: null, currentSlice: 0, + } as any); + useAnnotationStore.getState().reset(); + usePredictedRasterStore.setState({ bySource: {} }); + useClassStore.setState({ classes: [{ classId: 1, label: 'Cell', color: '#111111', isVisible: true }] }); + useRatingStore.setState({ ratings: {} }); + useSettingsStore.setState({ annotatorName: '' }); + vi.spyOn(window, 'alert').mockImplementation(() => {}); +}); + +afterEach(() => { + cleanup(); + vi.restoreAllMocks(); +}); + +describe('DownloadModal', () => { + it('renders scope and format options', () => { + renderModal(); + expect(screen.getByText('Download Dataset')).toBeInTheDocument(); + expect(screen.getByText('Current slice only')).toBeInTheDocument(); + expect(screen.getByText('Current sample only')).toBeInTheDocument(); + expect(screen.getByText('All annotated samples')).toBeInTheDocument(); + expect(screen.getByText(/COCO \(SAM3\)/)).toBeInTheDocument(); + expect(screen.getByText(/DINOv3 \/ Lightly/)).toBeInTheDocument(); + }); + + it('defaults to "current sample" scope with a 1-sample preview when a source is loaded', () => { + renderModal(); + expect(screen.getByRole('radio', { name: /Current sample only/ })).toBeChecked(); + expect(screen.getByText('1 sample will be exported.')).toBeInTheDocument(); + }); + + it('shows "No samples match" for the "all" scope with no annotated samples', async () => { + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('radio', { name: /All annotated samples/ })); + expect(screen.getByText('No samples match.')).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /Export/ })).toBeDisabled(); + }); + + it('counts annotated samples for the "all" scope, applying the star-rating filter', async () => { + useAnnotationStore.getState().setShapes('local:a.tif', 0, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + useAnnotationStore.getState().setShapes('local:b.tif', 0, [ + { id: 's2', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + useRatingStore.setState({ ratings: { 'local:b.tif': 2 } }); + const user = userEvent.setup(); + renderModal(); + + await user.click(screen.getByRole('radio', { name: /All annotated samples/ })); + expect(screen.getByText('2 samples will be exported.')).toBeInTheDocument(); + + await user.click(screen.getByRole('radio', { name: /★★ and above/ })); + expect(screen.getByText('1 sample will be exported.')).toBeInTheDocument(); + + await user.click(screen.getByRole('radio', { name: /★★★ only/ })); + expect(screen.getByText('No samples match.')).toBeInTheDocument(); + }); + + it('POSTs the current-sample export payload with predicted_slices for un-vectorized pointers', async () => { + useAnnotationStore.getState().setShapes('local:sample.tif', 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + usePredictedRasterStore.getState().setPointers('local:sample.tif', { + '0': { runId: 'run-1', classIds: [1] }, + '1': { runId: 'run-2', classIds: [1] }, + }); + useSettingsStore.setState({ annotatorName: ' Ada ' }); + + const fetchMock = mockFetchDone({ dataset_path: '/exports/foo' }); + global.fetch = fetchMock; + + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('button', { name: /^Export$/ })); + + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + const [url, opts] = (fetchMock as any).mock.calls[0]; + expect(url).toBe('/api/export/coco'); + const body = JSON.parse(opts.body); + expect(body.annotator).toBe('Ada'); + expect(body.format).toBe('coco_sam3'); + expect(body.sources).toHaveLength(1); + const src = body.sources[0]; + expect(src.kind).toBe('local'); + expect(src.source).toBe('sample.tif'); + // Slice 1 has real shapes, so its pointer is excluded; slice 0 has no + // shapes, so its pointer is included as an un-vectorized predicted_slice. + expect(src.predicted_slices).toEqual({ '0': { run_id: 'run-1', class_ids: [1] } }); + expect(src.slices).toEqual({ '1': expect.any(Array) }); + }); + + it('scopes predicted_slices and negative_slices to just the current slice for "slice" scope', async () => { + useDatasetStore.setState({ currentSlice: 0 } as any); + usePredictedRasterStore.getState().setPointers('local:sample.tif', { + '0': { runId: 'run-a', classIds: [1] }, + '2': { runId: 'run-b', classIds: [1] }, + }); + useAnnotationStore.getState().toggleNegativeSlice('local:sample.tif', 0); + useAnnotationStore.getState().toggleNegativeSlice('local:sample.tif', 2); + + const fetchMock = mockFetchDone({}); + global.fetch = fetchMock; + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('radio', { name: /Current slice only/ })); + await user.click(screen.getByRole('button', { name: /^Export$/ })); + + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + const body = JSON.parse((fetchMock as any).mock.calls[0][1].body); + const src = body.sources[0]; + expect(src.predicted_slices).toEqual({ '0': { run_id: 'run-a', class_ids: [1] } }); + expect(src.negative_slices).toEqual(['0']); + }); + + it('includes include_polygons only relevant for coco_sam3 and reflects the checkbox', async () => { + const fetchMock = mockFetchDone({}); + global.fetch = fetchMock; + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByLabelText(/Include polygon copy in COCO/)); + await user.click(screen.getByRole('button', { name: /^Export$/ })); + + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + const body = JSON.parse((fetchMock as any).mock.calls[0][1].body); + expect(body.include_polygons).toBe(true); + }); + + it('switches format to lightly_dinov3', async () => { + const fetchMock = mockFetchDone({}); + global.fetch = fetchMock; + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByLabelText(/DINOv3 \/ Lightly/)); + await user.click(screen.getByRole('button', { name: /^Export$/ })); + + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + const body = JSON.parse((fetchMock as any).mock.calls[0][1].body); + expect(body.format).toBe('lightly_dinov3'); + }); + + it('shows the saved-to-Tiled message on success (no zip when the response has no job_id)', async () => { + // A synchronous response (no job_id) is treated as immediately "done" with + // no jobId, so downloadUrl stays null and only the "Saved to" line shows. + global.fetch = mockFetchDone({ dataset_path: '/exports/foo.zip' }); + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('button', { name: /^Export$/ })); + + await screen.findByText(/Saved to/); + expect(screen.getByText('/exports/foo.zip')).toBeInTheDocument(); + expect(screen.queryByRole('link', { name: /Download \.zip/ })).not.toBeInTheDocument(); + }); + + it('shows an error message when the export request fails', async () => { + global.fetch = vi.fn(async () => ({ + ok: false, + status: 500, + text: async () => JSON.stringify({ detail: 'boom' }), + })) as unknown as typeof fetch; + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('button', { name: /^Export$/ })); + + await waitFor(() => expect(screen.getByText(/boom|Request failed/)).toBeInTheDocument()); + }); + + it('alerts and does not call fetch when no sample is loaded for "current" scope', async () => { + useDatasetStore.setState({ kind: null, source: null } as any); + const fetchMock = mockFetchDone({}); + global.fetch = fetchMock; + const user = userEvent.setup(); + renderModal(); + // previewCount is 0 with no source, disabling Export — assert directly instead. + expect(screen.getByRole('button', { name: /^Export$/ })).toBeDisabled(); + expect(fetchMock).not.toHaveBeenCalled(); + }); + + it('disables "Push masks to Tiled" for a local source and enables it for a tiled source', async () => { + renderModal(); + expect(screen.getByRole('button', { name: /Push masks to Tiled/ })).toBeDisabled(); + + cleanup(); + useDatasetStore.setState({ kind: 'tiled', source: 'a', serverUri: 'http://x', currentSlice: 0 } as any); + renderModal(); + expect(screen.getByRole('button', { name: /Push masks to Tiled/ })).not.toBeDisabled(); + }); + + it('POSTs to the mask-sync endpoint, filtering to tiled sources only, and shows a "View in 3D" link', async () => { + useDatasetStore.setState({ kind: 'tiled', source: 'vol.tif', serverUri: 'http://x', currentSlice: 0 } as any); + useAnnotationStore.getState().setShapes('tiled:http://x:vol.tif', 0, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + const fetchMock = mockFetchDone({ + written: [{ container: 'vol.tif', n_slices: 10, updated: 1 }], + }); + global.fetch = fetchMock; + const user = userEvent.setup(); + renderModal(); + await user.click(screen.getByRole('button', { name: /Push masks to Tiled/ })); + + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + const [url, opts] = (fetchMock as any).mock.calls[0]; + expect(url).toBe('/api/masks/to-tiled'); + const body = JSON.parse(opts.body); + expect(body.sources.every((s: any) => s.kind === 'tiled')).toBe(true); + + await screen.findByText(/Masks merged into Tiled/); + const link = screen.getByRole('button', { name: /View in 3D/ }); + expect(link).toBeInTheDocument(); + }); + + it('navigates to the 3D volume view when "View in 3D" is clicked', async () => { + useDatasetStore.setState({ kind: 'tiled', source: 'vol.tif', serverUri: 'http://x', currentSlice: 0 } as any); + useAnnotationStore.getState().setShapes('tiled:http://x:vol.tif', 0, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + global.fetch = mockFetchDone({ + written: [{ container: 'vol.tif', n_slices: 10, updated: 1 }], + }); + const onClose = vi.fn(); + const user = userEvent.setup(); + renderModal(onClose); + await user.click(screen.getByRole('button', { name: /Push masks to Tiled/ })); + await screen.findByText(/Masks merged into Tiled/); + await user.click(screen.getByRole('button', { name: /View in 3D/ })); + expect(onClose).toHaveBeenCalled(); + }); + + it('calls onClose when Cancel/Close is clicked', async () => { + const onClose = vi.fn(); + const user = userEvent.setup(); + renderModal(onClose); + await user.click(screen.getByRole('button', { name: 'Cancel' })); + expect(onClose).toHaveBeenCalled(); + }); +}); diff --git a/frontend/src/components/annotate/DownloadModal.tsx b/frontend/src/components/annotate/DownloadModal.tsx index be6cb50..5b5f8df 100644 --- a/frontend/src/components/annotate/DownloadModal.tsx +++ b/frontend/src/components/annotate/DownloadModal.tsx @@ -2,13 +2,15 @@ * DownloadModal — scope picker + optional star-rating filter, then COCO export. */ import { useMemo, useState } from 'react'; -import { DownloadSimple, X, CheckCircle, WarningCircle } from '@phosphor-icons/react'; +import { useNavigate } from 'react-router'; +import { DownloadSimple, X, CheckCircle, WarningCircle, Cube } from '@phosphor-icons/react'; import { useDatasetStore } from '@/stores/datasetStore'; import { useAnnotationStore } from '@/stores/annotationStore'; +import { usePredictedRasterStore } from '@/stores/predictedRasterStore'; import { useClassStore } from '@/stores/classStore'; import { useRatingStore } from '@/stores/ratingStore'; import { useSettingsStore } from '@/stores/settingsStore'; -import { buildSourceKey } from '@/lib/sourceKey'; +import { buildSourceKey, parseSourceKey } from '@/lib/sourceKey'; import { useExportJob } from '@/hooks/useExportJob'; interface DownloadModalProps { @@ -17,16 +19,6 @@ interface DownloadModalProps { type Scope = 'slice' | 'current' | 'all' | 'stars1' | 'stars2' | 'stars3'; -/** Parse a canonical sourceKey back to { kind, source, serverUri }. */ -function parseSourceKey(sk: string) { - if (sk.startsWith('tiled:')) { - const rest = sk.slice('tiled:'.length); - const sep = rest.indexOf(':'); - return { kind: 'tiled' as const, serverUri: rest.slice(0, sep) || null, source: rest.slice(sep + 1) }; - } - return { kind: 'local' as const, serverUri: null, source: sk.slice('local:'.length) }; -} - const SCOPE_OPTIONS: { value: Scope; label: string; desc: string; stars?: string }[] = [ { value: 'slice', @@ -65,8 +57,10 @@ const SCOPE_OPTIONS: { value: Scope; label: string; desc: string; stars?: string /** Renders the COCO download dialog and drives the export job for the chosen scope. */ export default function DownloadModal({ onClose }: DownloadModalProps) { + const navigate = useNavigate(); const { source, kind, serverUri, currentSlice } = useDatasetStore(); const { byImage, splitBySlice, negativeSlices } = useAnnotationStore(); + const predictedPointers = usePredictedRasterStore((s) => s.bySource); const { classes } = useClassStore(); const ratings = useRatingStore((s) => s.ratings); const annotatorName = useSettingsStore((s) => s.annotatorName); @@ -106,6 +100,19 @@ export default function DownloadModal({ onClose }: DownloadModalProps) { const split_by_slice = scope === 'slice' ? (allSplits[cur] ? { [cur]: allSplits[cur] } : {}) : allSplits; const negative_slices = scope === 'slice' ? allNeg.filter((k) => String(k) === cur) : allNeg; + // Un-vectorized predicted pointers (see predictedRasterStore) for this + // sample — only relevant to mask-sync (build_mask_volumes fetches + // their commit.png directly); COCO export ignores this field today, + // same as before. "Current slice only" keeps just the one pointer + // matching cur, mirroring how slices/negative_slices are scoped above. + const allPredicted = predictedPointers[sk] ?? {}; + const predicted_slices = Object.fromEntries( + Object.entries(allPredicted) + .filter(([k]) => (scope === 'slice' ? k === cur : true)) + .filter(([k]) => !(slices[k]?.length)) + .map(([k, p]) => [k, { run_id: p.runId, class_ids: p.classIds }]), + ); + return [{ kind, source, @@ -113,6 +120,7 @@ export default function DownloadModal({ onClose }: DownloadModalProps) { slices, split_by_slice, negative_slices, + predicted_slices, }]; } @@ -305,13 +313,23 @@ export default function DownloadModal({ onClose }: DownloadModalProps) { )} {status === 'done' && maskResult && ( -
+
- + {maskResult.length === 0 ? 'No masks written (no Tiled sources or no annotated slices).' : <>Masks merged into Tiled: {maskResult.map((w) => `${w.container} (${w.n_slices} slices total, ${w.updated ?? 0} updated)`).join(', ')}.} + {maskResult.length > 0 && ( + + )}
)} {status === 'done' && !maskResult && ( diff --git a/frontend/src/components/annotate/FeatureChannelsPanel/index.test.tsx b/frontend/src/components/annotate/FeatureChannelsPanel/index.test.tsx new file mode 100644 index 0000000..52d622a --- /dev/null +++ b/frontend/src/components/annotate/FeatureChannelsPanel/index.test.tsx @@ -0,0 +1,330 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen, waitFor, within } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import FeatureChannelsPanel from './index'; +import { useIpredStore, DEFAULT_COMPOSITION_ID } from '@/stores/ipredStore'; +import type { FeatureJobInfo } from '@/hooks/useFeatureChannels'; +import { listIpredCompositions, listIpredModules } from '@/lib/ipredApi'; + +vi.mock('@/lib/ipredApi', () => ({ + listIpredCompositions: vi.fn(async () => []), + listIpredModules: vi.fn(async () => []), +})); + +const mockListIpredCompositions = vi.mocked(listIpredCompositions); +const mockListIpredModules = vi.mocked(listIpredModules); + +const initialIpredState = useIpredStore.getState(); + +const compSkimage = { + id: 'comp-skimage', + name: 'Skimage only', + nodes: [{ id: 'n1', module: 'skimage_multiscale' }], + outputs: ['n1'], +}; + +const compSlimsam = { + id: 'comp-skimage-slimsam', + name: 'Skimage + SlimSAM', + nodes: [ + { id: 'n1', module: 'skimage_multiscale' }, + { id: 'n2', module: 'slimsam' }, + ], + outputs: ['n1', 'n2'], +}; + +const compMark25 = { + id: 'comp-skimage-mark25', + name: 'Skimage + Mark25', + nodes: [ + { id: 'n1', module: 'skimage_multiscale' }, + { id: 'n2', module: 'tomojepa_mark25' }, + ], + outputs: ['n1', 'n2'], +}; + +const modSkimage = { + id: 'skimage_multiscale', + name: 'Skimage multiscale', + description: '', + runtime: 'numpy', + ready: true, + accepts_input_from: false, + produces_channels: true, + produces_embedding: false, + params_schema: {}, +}; + +const modSlimsam = { + id: 'slimsam', + name: 'SlimSAM', + description: '', + runtime: 'torch', + ready: true, + accepts_input_from: true, + produces_channels: false, + produces_embedding: true, + params_schema: {}, +}; + +const modMark25NotReady = { + id: 'tomojepa_mark25', + name: 'TomoJEPA Mark25', + description: '', + runtime: 'torch', + ready: false, + accepts_input_from: true, + produces_channels: false, + produces_embedding: true, + params_schema: {}, +}; + +function baseProps() { + return { + job: null as FeatureJobInfo | null, + channelIndex: null as number | null, + computing: false, + error: null as string | null, + onCompute: vi.fn(), + onSelectChannel: vi.fn(), + onCycle: vi.fn(), + onOriginal: vi.fn(), + }; +} + +function makeJob(overrides: Partial = {}): FeatureJobInfo { + return { + jobId: 'job-1', + width: 100, + height: 100, + channels: [ + { index: 0, label: 'edges' }, + { index: 1, label: 'texture' }, + ], + hasSam: false, + ...overrides, + }; +} + +beforeEach(() => { + useIpredStore.setState(initialIpredState, true); + mockListIpredCompositions.mockReset().mockResolvedValue([]); + mockListIpredModules.mockReset().mockResolvedValue([]); +}); + +afterEach(() => { + cleanup(); +}); + +describe('FeatureChannelsPanel', () => { + it('shows a loading placeholder for the recipe select while presets load', () => { + mockListIpredCompositions.mockImplementation(() => new Promise(() => {})); + const props = baseProps(); + render(); + expect(screen.queryByRole('combobox')).not.toBeInTheDocument(); + expect(screen.getByText('Recipe').parentElement?.querySelector('.animate-pulse')).toBeInTheDocument(); + }); + + it('renders preset options once loaded, marking the recommended one and disabling unready ones', async () => { + mockListIpredCompositions.mockResolvedValue([compSkimage, compSlimsam, compMark25]); + mockListIpredModules.mockResolvedValue([modSkimage, modSlimsam, modMark25NotReady]); + useIpredStore.setState({ preferredCompositionId: 'comp-skimage-slimsam' }); + const props = baseProps(); + render(); + + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + const selects = screen.getAllByRole('combobox'); + const recipeSelect = selects[0]; + const options = within(recipeSelect).getAllByRole('option'); + expect(options).toHaveLength(3); + + const recommended = options.find((o) => o.textContent?.includes('★')); + expect(recommended).toBeDefined(); + expect(recommended?.textContent).toContain('Texture-aware (+ SlimSAM)'); + + const mark25Option = options.find((o) => o.getAttribute('value') === 'comp-skimage-mark25'); + expect(mark25Option).toBeDisabled(); + expect(mark25Option?.textContent).toContain('(unavailable)'); + + // Recommended preset hint text is shown for the currently selected preset. + expect( + screen.getByText('Skimage filters + SlimSAM vision-encoder embeddings (PCA-reduced). Recommended default.'), + ).toBeInTheDocument(); + }); + + it('changing the recipe select calls setPreferredCompositionId', async () => { + mockListIpredCompositions.mockResolvedValue([compSkimage, compSlimsam]); + mockListIpredModules.mockResolvedValue([modSkimage, modSlimsam]); + useIpredStore.setState({ preferredCompositionId: 'comp-skimage-slimsam' }); + const user = userEvent.setup(); + const props = baseProps(); + render(); + + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + const recipeSelect = screen.getAllByRole('combobox')[0]; + await user.selectOptions(recipeSelect, 'comp-skimage'); + expect(useIpredStore.getState().preferredCompositionId).toBe('comp-skimage'); + }); + + it('shows a warning and disables Compute when the selected preset needs a missing module', async () => { + mockListIpredCompositions.mockResolvedValue([compMark25]); + mockListIpredModules.mockResolvedValue([modSkimage, modMark25NotReady]); + useIpredStore.setState({ preferredCompositionId: 'comp-skimage-mark25' }); + const props = baseProps(); + render(); + + await waitFor(() => expect(screen.getByText(/isn't installed yet/)).toBeInTheDocument()); + expect(screen.getByText(/isn't installed yet/).closest('div')).toHaveTextContent('TomoJEPA Mark25'); + expect(screen.getByRole('button', { name: /Compute/ })).toBeDisabled(); + }); + + it('falls back to the recommended preset if the stored preference is not ready once presets load', async () => { + mockListIpredCompositions.mockResolvedValue([compSkimage, compSlimsam, compMark25]); + mockListIpredModules.mockResolvedValue([modSkimage, modSlimsam, modMark25NotReady]); + useIpredStore.setState({ preferredCompositionId: 'comp-skimage-mark25' }); + render(); + + await waitFor(() => + expect(useIpredStore.getState().preferredCompositionId).toBe(DEFAULT_COMPOSITION_ID), + ); + }); + + it('surfaces a presets error message when loading compositions/modules fails', async () => { + mockListIpredCompositions.mockRejectedValue(new Error('network down')); + render(); + await waitFor(() => expect(screen.getByText('network down')).toBeInTheDocument()); + }); + + it('clicking Compute calls onCompute', async () => { + mockListIpredCompositions.mockResolvedValue([compSkimage]); + mockListIpredModules.mockResolvedValue([modSkimage]); + useIpredStore.setState({ preferredCompositionId: 'comp-skimage' }); + const user = userEvent.setup(); + const props = baseProps(); + render(); + + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + await user.click(screen.getByRole('button', { name: /Compute/ })); + expect(props.onCompute).toHaveBeenCalledTimes(1); + }); + + it('shows a spinner label and disables Compute + recipe select while computing', async () => { + mockListIpredCompositions.mockResolvedValue([compSkimage]); + mockListIpredModules.mockResolvedValue([modSkimage]); + const props = { ...baseProps(), computing: true }; + render(); + + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + expect(screen.getByText('Computing…')).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /Computing/ })).toBeDisabled(); + expect(screen.getAllByRole('combobox')[0]).toBeDisabled(); + }); + + it('disables the whole panel controls when disabled prop is set', async () => { + mockListIpredCompositions.mockResolvedValue([compSkimage]); + mockListIpredModules.mockResolvedValue([modSkimage]); + const props = { ...baseProps(), disabled: true }; + render(); + + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + expect(screen.getByRole('button', { name: /Compute/ })).toBeDisabled(); + expect(screen.getAllByRole('combobox')[0]).toBeDisabled(); + }); + + it('renders the compute-error message when error is set', async () => { + const props = { ...baseProps(), error: 'Compute failed: boom' }; + render(); + expect(screen.getByText('Compute failed: boom')).toBeInTheDocument(); + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + }); + + it('shows no channel section when there is no job', async () => { + render(); + expect(screen.queryByRole('button', { name: 'Previous channel' })).not.toBeInTheDocument(); + expect(screen.queryByText('Show original')).not.toBeInTheDocument(); + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + }); + + it('shows the header channel count summary including SAM and cache flags', async () => { + const job = makeJob({ hasSam: true, cacheHit: true }); + render(); + expect(screen.getByText('2 ch +SAM · cache')).toBeInTheDocument(); + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + }); + + it('renders channel list, active label, and cycle/select controls once a job is present', async () => { + const job = makeJob(); + const props = { ...baseProps(), job, channelIndex: 1 }; + render(); + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + + expect(screen.getByRole('button', { name: 'Previous channel' })).toBeEnabled(); + expect(screen.getByRole('button', { name: 'Next channel' })).toBeEnabled(); + expect(screen.getByTitle('texture')).toBeInTheDocument(); + + const channelSelect = screen.getByDisplayValue('texture'); + const options = within(channelSelect).getAllByRole('option'); + expect(options.map((o) => o.textContent)).toEqual(['Original', 'edges', 'texture']); + }); + + it('disables cycle buttons when channelIndex is null (Original selected)', async () => { + const job = makeJob(); + const props = { ...baseProps(), job, channelIndex: null }; + render(); + await waitFor(() => expect(screen.getAllByRole('combobox').length).toBeGreaterThan(0)); + expect(screen.getByRole('button', { name: 'Previous channel' })).toBeDisabled(); + expect(screen.getByRole('button', { name: 'Next channel' })).toBeDisabled(); + expect(screen.queryByTitle('texture')).not.toBeInTheDocument(); + }); + + it('clicking next/previous channel calls onCycle with the right delta', async () => { + const job = makeJob(); + const props = { ...baseProps(), job, channelIndex: 0 }; + const user = userEvent.setup(); + render(); + + await user.click(screen.getByRole('button', { name: 'Next channel' })); + expect(props.onCycle).toHaveBeenCalledWith(1); + await user.click(screen.getByRole('button', { name: 'Previous channel' })); + expect(props.onCycle).toHaveBeenCalledWith(-1); + }); + + it('selecting a channel option calls onSelectChannel with the numeric index or null for Original', async () => { + const job = makeJob(); + const props = { ...baseProps(), job, channelIndex: 0 }; + const user = userEvent.setup(); + render(); + + const channelSelect = screen.getByDisplayValue('edges'); + await user.selectOptions(channelSelect, '1'); + expect(props.onSelectChannel).toHaveBeenCalledWith(1); + + await user.selectOptions(channelSelect, ''); + expect(props.onSelectChannel).toHaveBeenCalledWith(null); + }); + + it('clicking "Show original" calls onOriginal', async () => { + const job = makeJob(); + const props = { ...baseProps(), job, channelIndex: 0 }; + const user = userEvent.setup(); + render(); + + await user.click(screen.getByText('Show original')); + expect(props.onOriginal).toHaveBeenCalledTimes(1); + }); + + it('toggles the Advanced disclosure to reveal the composition panel', async () => { + const user = userEvent.setup(); + render(); + + expect(screen.queryByText(/edit feature recipe/)).toBeInTheDocument(); + const toggle = screen.getByText(/Advanced: edit feature recipe/); + expect(toggle.textContent).toContain('▸'); + + await user.click(toggle); + expect(toggle.textContent).toContain('▾'); + + await user.click(toggle); + expect(toggle.textContent).toContain('▸'); + }); +}); diff --git a/frontend/src/components/annotate/FeatureChannelsPanel/index.tsx b/frontend/src/components/annotate/FeatureChannelsPanel/index.tsx new file mode 100644 index 0000000..39345fc --- /dev/null +++ b/frontend/src/components/annotate/FeatureChannelsPanel/index.tsx @@ -0,0 +1,265 @@ +/** + * FeatureChannelsPanel — the Assist-stage sidebar panel. + * + * Compositions (feature-bank recipes) are surfaced as named presets so most + * users never need the underlying node graph. The full graph editor + * (`CompositionPanel`) is still there for anyone who wants a custom recipe — + * it just lives behind a collapsed "Advanced" disclosure instead of its own + * tab, since composing a feature graph is a power-user path, not the default. + * + * Presets whose modules aren't ready (e.g. TomoJEPA checkpoints not present — + * see ipred/models/README.md) are greyed out rather than left to fail at + * Compute time; only SlimSAM ships auto-vendored today, so it's the + * recommended default. + */ +import { useEffect, useMemo, useState } from 'react'; +import { CaretLeft, CaretRight, CircleNotch, Stack, Warning } from '@phosphor-icons/react'; +import type { FeatureJobInfo } from '@/hooks/useFeatureChannels'; +import { useIpredStore } from '@/stores/ipredStore'; +import { listIpredCompositions, listIpredModules, type CompositionDoc, type FeatureModuleInfo } from '@/lib/ipredApi'; +import CompositionPanel from '@/components/CompositionPanel'; +import CollapsibleSection from '@/components/common/CollapsibleSection'; +import { cn } from '@/lib/utils'; + +/** Plain-language labels for the 7 built-in compositions, by id. */ +const PRESET_LABELS: Record = { + 'comp-skimage': 'Fast (skimage only)', + 'comp-skimage-slimsam': 'Texture-aware (+ SlimSAM)', + 'comp-slimsam-clahe': 'Texture-aware, contrast-enhanced', + 'comp-skimage-mark25': 'Deep features (TomoJEPA Mark25)', + 'comp-mark25-clahe': 'Deep features, contrast-enhanced (Mark25)', + 'comp-skimage-mark11': 'Deep features (TomoJEPA Mark11)', + 'comp-mark11-clahe': 'Deep features, contrast-enhanced (Mark11)', +}; + +/** One-line explanation of what each preset trades off, shown under the picker. */ +const PRESET_HINTS: Record = { + 'comp-skimage': 'Multiscale edge/texture filters only. Cheapest, no model download.', + 'comp-skimage-slimsam': 'Skimage filters + SlimSAM vision-encoder embeddings (PCA-reduced). Recommended default.', + 'comp-slimsam-clahe': 'SlimSAM embeddings on a contrast-equalized slice — helps on low-contrast data.', + 'comp-skimage-mark25': 'Skimage filters + TomoJEPA Mark25 embeddings. Requires a private checkpoint.', + 'comp-mark25-clahe': 'TomoJEPA Mark25 on a contrast-equalized slice. Requires a private checkpoint.', + 'comp-skimage-mark11': 'Skimage filters + TomoJEPA Mark11 embeddings. Requires a private checkpoint.', + 'comp-mark11-clahe': 'TomoJEPA Mark11 on a contrast-equalized slice. Requires a private checkpoint.', +}; + +const RECOMMENDED_PRESET_ID = 'comp-skimage-slimsam'; + +export interface FeatureChannelsPanelProps { + job: FeatureJobInfo | null; + channelIndex: number | null; + computing: boolean; + error: string | null; + onCompute: () => void; + onSelectChannel: (index: number | null) => void; + onCycle: (delta: number) => void; + onOriginal: () => void; + disabled?: boolean; +} + +export default function FeatureChannelsPanel({ + job, + channelIndex, + computing, + error, + onCompute, + onSelectChannel, + onCycle, + onOriginal, + disabled = false, +}: FeatureChannelsPanelProps) { + const preferredCompositionId = useIpredStore((s) => s.preferredCompositionId); + const setPreferredCompositionId = useIpredStore((s) => s.setPreferredCompositionId); + const [presets, setPresets] = useState([]); + const [modules, setModules] = useState([]); + const [loadingPresets, setLoadingPresets] = useState(true); + const [presetsError, setPresetsError] = useState(null); + const [showAdvanced, setShowAdvanced] = useState(false); + + useEffect(() => { + let cancelled = false; + Promise.all([listIpredCompositions(), listIpredModules()]) + .then(([comps, mods]) => { + if (cancelled) return; + setPresets(comps); + setModules(mods); + }) + .catch((e) => { + if (!cancelled) setPresetsError(e instanceof Error ? e.message : String(e)); + }) + .finally(() => { + if (!cancelled) setLoadingPresets(false); + }); + return () => { + cancelled = true; + }; + }, []); + + const moduleById = useMemo(() => new Map(modules.map((m) => [m.id, m])), [modules]); + + /** A preset is ready only if every module it wires in is ready (e.g. weights present). */ + const presetReady = useMemo(() => { + const out = new Map(); + for (const p of presets) { + out.set(p.id, p.nodes.every((n) => moduleById.get(n.module)?.ready !== false)); + } + return out; + }, [presets, moduleById]); + + const selected = presets.find((p) => p.id === preferredCompositionId) ?? null; + const selectedReady = selected ? (presetReady.get(selected.id) ?? true) : true; + const missingModule = selected?.nodes.find((n) => moduleById.get(n.module)?.ready === false); + const missingModuleName = missingModule ? (moduleById.get(missingModule.module)?.name ?? missingModule.module) : null; + + // If the preferred preset isn't ready once modules load (e.g. a stale choice from + // a previous session), fall back to the recommended one instead of a dead-end. + useEffect(() => { + if (loadingPresets || presets.length === 0) return; + if (presetReady.get(preferredCompositionId) === false) { + const fallback = presetReady.get(RECOMMENDED_PRESET_ID) ? RECOMMENDED_PRESET_ID : presets.find((p) => presetReady.get(p.id))?.id; + if (fallback && fallback !== preferredCompositionId) setPreferredCompositionId(fallback); + } + // Only re-run when the readiness picture itself changes, not on every keystroke. + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [loadingPresets, presetReady]); + + const n = job?.channels.length ?? 0; + const activeLabel = + channelIndex !== null && job ? (job.channels[channelIndex]?.label ?? `Channel ${channelIndex}`) : null; + + return ( + 0 ? ( + + {n} ch{job?.hasSam ? ' +SAM' : ''} + {job?.cacheHit ? ' · cache' : ''} + + ) : undefined + } + > + ); } diff --git a/frontend/src/components/annotate/PerfOverlay/index.test.tsx b/frontend/src/components/annotate/PerfOverlay/index.test.tsx new file mode 100644 index 0000000..fdd3290 --- /dev/null +++ b/frontend/src/components/annotate/PerfOverlay/index.test.tsx @@ -0,0 +1,63 @@ +import { afterEach, beforeEach, describe, expect, it } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import PerfOverlay from './index'; +import { initPerf, mark, resetPerf } from '@/lib/perf'; + +beforeEach(() => { + window.localStorage.setItem('perf', '1'); + initPerf(); + resetPerf(); +}); + +afterEach(() => { + cleanup(); + window.localStorage.clear(); + initPerf(); + resetPerf(); +}); + +describe('PerfOverlay', () => { + it('shows a placeholder when no samples have been recorded', () => { + render(); + expect(screen.getByText('no samples yet — draw or zoom')).toBeInTheDocument(); + }); + + it('shows the shape count', () => { + render(); + expect(screen.getByText('42 shapes')).toBeInTheDocument(); + }); + + it('renders a row per recorded label with p50/p95/count', () => { + mark('commit', 5); + mark('commit', 7); + render(); + expect(screen.getByText('commit')).toBeInTheDocument(); + const row = screen.getByText('commit').closest('tr')!; + expect(row).toHaveTextContent('2'); // count column + }); + + it('colors a slow p95 red and a fast one green', () => { + mark('clip', 60); + render(); + const row = screen.getByText('clip').closest('tr')!; + const cells = row.querySelectorAll('td'); + expect(cells[2]).toHaveClass('text-red-400'); + }); + + it('reset button clears samples back to the placeholder', async () => { + mark('commit', 5); + const user = userEvent.setup(); + render(); + expect(screen.queryByText('no samples yet — draw or zoom')).not.toBeInTheDocument(); + await user.click(screen.getByTitle('Reset samples')); + expect(await screen.findByText('no samples yet — draw or zoom')).toBeInTheDocument(); + }); + + it('hide button removes the overlay entirely', async () => { + const user = userEvent.setup(); + const { container } = render(); + await user.click(screen.getByTitle('Hide (reload to show again)')); + expect(container).toBeEmptyDOMElement(); + }); +}); diff --git a/frontend/src/components/annotate/PerfOverlay/index.tsx b/frontend/src/components/annotate/PerfOverlay/index.tsx new file mode 100644 index 0000000..950e5b9 --- /dev/null +++ b/frontend/src/components/annotate/PerfOverlay/index.tsx @@ -0,0 +1,87 @@ +/** + * PerfOverlay — dev-only timing HUD for the annotation workspace. + * + * Rendered only when `?perf=1` (see lib/perf.ts). Shows rolling p50/p95 per + * instrumented path plus the current shape count, so an optimization can be + * confirmed on real data — the slices that matter are far larger than anything + * reproducible in a test. + */ +import { useSyncExternalStore, useState } from 'react'; +import { X, ArrowCounterClockwise } from '@phosphor-icons/react'; +import { snapshot, subscribe, getVersion, resetPerf, type PerfStat } from '@/lib/perf'; + +interface PerfOverlayProps { + /** Shapes on the current slice — the main driver of the costs below. */ + shapeCount: number; +} + +/** Colour by how close a p95 is to a dropped frame (16.7ms). */ +function severity(ms: number): string { + if (ms >= 50) return 'text-red-400'; + if (ms >= 16.7) return 'text-amber-300'; + return 'text-emerald-300'; +} + +function Row({ stat }: { stat: PerfStat }) { + return ( + + {stat.label} + {stat.p50.toFixed(1)} + {stat.p95.toFixed(1)} + {stat.count} + + ); +} + +export default function PerfOverlay({ shapeCount }: PerfOverlayProps) { + const [hidden, setHidden] = useState(false); + // Re-render when new samples land (the version counter is the store value). + useSyncExternalStore(subscribe, getVersion, () => 0); + const stats = snapshot(); + + if (hidden) return null; + + return ( +
+
+ perf + {shapeCount} shapes +
+ + +
+
+ {stats.length === 0 ? ( +
no samples yet — draw or zoom
+ ) : ( + + + + + + + + + + + {stats.map((s) => )} + +
pathp50p95n
+ )} +
+ ); +} diff --git a/frontend/src/components/annotate/PixelClassifierPanel/index.test.tsx b/frontend/src/components/annotate/PixelClassifierPanel/index.test.tsx new file mode 100644 index 0000000..b04c4b6 --- /dev/null +++ b/frontend/src/components/annotate/PixelClassifierPanel/index.test.tsx @@ -0,0 +1,542 @@ +/** + * PixelClassifierPanel — fully prop-driven (no internal store/hook access), so + * tests construct prop objects directly rather than mocking usePixelClassifier + * or @/lib/ipredApi. Covers: idle/canTrain gating, busy states during train/ + * multi-slice train/predict, model summary + feature importances, class- + * probability cycling, commit/dismiss, volume-apply flow (apply -> progress -> + * result -> commit/dismiss), predicted-pointer "make editable" banner, push-to- + * Tiled + view-in-3D, train-deep-model hand-off, and error surfacing. + * (Suggest-labels/manifold sampling was pulled out into its own + * SuggestLabelsPanel — see that component's own test file.) + */ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, fireEvent, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import PixelClassifierPanel, { type PixelClassifierPanelProps } from './index'; +import type { ClfTrainResult, ClfPredictCounts } from '@/hooks/usePixelClassifier'; + +afterEach(() => { + cleanup(); +}); + +const BASE_PARAMS = { iterations: 200, depth: 6, learningRate: 0.1, alpha: 0.05 }; + +function makeModel(overrides: Partial = {}): ClfTrainResult { + return { + modelId: 'model-1', + featureId: 'feat-1', + nSamples: 100, + nTrain: 80, + nCal: 20, + classIds: [1, 2], + trainAccuracy: 0.912, + params: BASE_PARAMS, + nTrees: 200, + usesSam: false, + featureImportances: [ + { label: 'intensity', importance: 5.5 }, + { label: 'edges', importance: 2.1 }, + ], + trainerId: 'catboost', + compositionId: 'comp-1', + ...overrides, + }; +} + +function makeCounts(overrides: Partial = {}): ClfPredictCounts { + return { singleton: 700, multi: 200, abstain: 100, ...overrides }; +} + +function baseProps(overrides: Partial = {}): PixelClassifierPanelProps { + return { + hasFeatureJob: true, + canAutoPreprocess: false, + hasShapes: true, + training: false, + predicting: false, + params: BASE_PARAMS, + onParamsChange: vi.fn(), + model: null, + hasPrediction: false, + predictCounts: null, + error: null, + onTrain: vi.fn(), + onPredict: vi.fn(), + onCommit: vi.fn(), + onDismiss: vi.fn(), + annotatedSliceCount: 1, + trainAcrossSlices: false, + onTrainAcrossSlicesChange: vi.fn(), + multiTraining: false, + multiTrainProgress: null, + totalSliceCount: 1, + commitClassIds: [], + onToggleCommitClassId: vi.fn(), + volumeApplying: false, + volumeApplyProgress: null, + volumeApplyResult: null, + onApplyToVolume: vi.fn(), + onCommitVolumeApply: vi.fn(), + onCancelVolumeApply: vi.fn(), + onDismissVolumeApply: vi.fn(), + hasPredictedPointerOnCurrentSlice: false, + vectorizingSlice: false, + onMakeSliceEditable: vi.fn(), + hasAnyPredictedPointers: false, + ...overrides, + }; +} + +describe('PixelClassifierPanel', () => { + describe('idle / train gating', () => { + it('enables Train when features exist and shapes are annotated', () => { + render(); + expect(screen.getByRole('button', { name: /train classifier/i })).toBeEnabled(); + // Predict is disabled with no model yet. + expect(screen.getByRole('button', { name: /^predict$/i })).toBeDisabled(); + }); + + it('disables Train when there are no shapes and no auto-preprocess', () => { + render(); + expect(screen.getByRole('button', { name: /train classifier/i })).toBeDisabled(); + }); + + it('allows training with no feature job when canAutoPreprocess is set, and notes it will compute features', () => { + render( + , + ); + expect(screen.getByRole('button', { name: /train classifier/i })).toBeEnabled(); + expect(screen.getByText(/Train will compute features/i)).toBeInTheDocument(); + }); + + it('calls onTrain when the train button is clicked', async () => { + const user = userEvent.setup(); + const onTrain = vi.fn(); + render(); + await user.click(screen.getByRole('button', { name: /train classifier/i })); + expect(onTrain).toHaveBeenCalledTimes(1); + }); + + it('shows a training busy bar while training', () => { + render(); + expect(screen.getByText(/Training…/i)).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /train classifier/i })).toBeDisabled(); + }); + + it('updates iteration/depth/learning-rate params via onParamsChange', () => { + const onParamsChange = vi.fn(); + render(); + const trees = screen.getByLabelText(/trees/i); + fireEvent.change(trees, { target: { value: '300' } }); + expect(onParamsChange).toHaveBeenCalledTimes(1); + expect(onParamsChange.mock.calls[0][0].iterations).toBe(300); + }); + }); + + describe('train across slices', () => { + it('disables the checkbox when only one slice is annotated', () => { + render(); + expect(screen.getByRole('checkbox', { name: /train across all annotated slices/i })).toBeDisabled(); + }); + + it('enables the checkbox with multiple annotated slices and toggles it', async () => { + const user = userEvent.setup(); + const onTrainAcrossSlicesChange = vi.fn(); + render( + , + ); + const checkbox = screen.getByRole('checkbox', { name: /train across all annotated slices \(3\)/i }); + expect(checkbox).toBeEnabled(); + await user.click(checkbox); + expect(onTrainAcrossSlicesChange).toHaveBeenCalledWith(true); + }); + + it('label switches to "Train across slices" and gates on annotatedSliceCount only when trainAcrossSlices is set', () => { + render( + , + ); + expect(screen.getByRole('button', { name: /train across slices/i })).toBeEnabled(); + }); + + it('shows multi-train job progress bar with done/total', () => { + render( + , + ); + expect(screen.getByText(/Training across slices…/i)).toBeInTheDocument(); + expect(screen.getByText('2/5')).toBeInTheDocument(); + }); + }); + + describe('predict', () => { + it('enables Predict once a model exists and calls onPredict', async () => { + const user = userEvent.setup(); + const onPredict = vi.fn(); + render(); + const btn = screen.getByRole('button', { name: /^predict$/i }); + expect(btn).toBeEnabled(); + await user.click(btn); + expect(onPredict).toHaveBeenCalledTimes(1); + }); + + it('shows a predicting busy bar', () => { + render(); + expect(screen.getByText(/Predicting…/i)).toBeInTheDocument(); + }); + }); + + describe('model summary', () => { + it('renders accuracy, sample counts, trees, classes, and feature importances', () => { + render(); + expect(screen.getByText(/Acc 91\.2%/)).toBeInTheDocument(); + expect(screen.getByText(/train 80 \/ cal 20/)).toBeInTheDocument(); + expect(screen.getByText(/200 trees/)).toBeInTheDocument(); + expect(screen.getByText(/classes 1, 2/)).toBeInTheDocument(); + expect(screen.getByText('intensity')).toBeInTheDocument(); + expect(screen.getByText('edges')).toBeInTheDocument(); + }); + + it('shows the SlimSAM suffix when usesSam is true', () => { + render(); + expect(screen.getByText(/\+SlimSAM/)).toBeInTheDocument(); + }); + }); + + describe('commit classes + apply across volume', () => { + it('toggles commit-class chips via onToggleCommitClassId', async () => { + const user = userEvent.setup(); + const onToggleCommitClassId = vi.fn(); + render( + , + ); + await user.click(screen.getByText('class 1')); + expect(onToggleCommitClassId).toHaveBeenCalledWith(1); + await user.click(screen.getByText('class 2')); + expect(onToggleCommitClassId).toHaveBeenCalledWith(2); + }); + + it('disables Apply across volume with a single slice or no classes selected', () => { + const { rerender } = render( + , + ); + expect(screen.getByRole('button', { name: /apply across volume/i })).toBeDisabled(); + + rerender( + , + ); + expect(screen.getByRole('button', { name: /apply across volume/i })).toBeDisabled(); + }); + + it('enables Apply across volume with multiple slices and at least one selected class, and calls onApplyToVolume', async () => { + const user = userEvent.setup(); + const onApplyToVolume = vi.fn(); + render( + , + ); + const btn = screen.getByRole('button', { name: /apply across volume \(10 slices\)/i }); + expect(btn).toBeEnabled(); + await user.click(btn); + expect(onApplyToVolume).toHaveBeenCalledTimes(1); + }); + + it('shows volume-apply progress and a cancel button while running', async () => { + const user = userEvent.setup(); + const onCancelVolumeApply = vi.fn(); + render( + , + ); + expect(screen.getByText(/Applying across volume…/i)).toBeInTheDocument(); + expect(screen.getByText('3/10')).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: /cancel/i })); + expect(onCancelVolumeApply).toHaveBeenCalledTimes(1); + }); + + it('shows the volume-apply result with commit/dismiss actions', async () => { + const user = userEvent.setup(); + const onCommitVolumeApply = vi.fn(); + const onDismissVolumeApply = vi.fn(); + render( + , + ); + expect(screen.getByText(/Predicted 8 slice\(s\), 2 failed\./)).toBeInTheDocument(); + const commitBtn = screen.getByRole('button', { name: /Commit predicted shapes \(2 classes\)/i }); + expect(commitBtn).toBeEnabled(); + await user.click(commitBtn); + expect(onCommitVolumeApply).toHaveBeenCalledTimes(1); + await user.click(screen.getByRole('button', { name: /^dismiss$/i })); + expect(onDismissVolumeApply).toHaveBeenCalledTimes(1); + }); + + it('disables commit-predicted-shapes when runCount is 0, and shows cancelled prefix', () => { + render( + , + ); + expect(screen.getByText(/^Cancelled — Predicted 0 slice\(s\)\./)).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /Commit predicted shapes \(1 class\)/i })).toBeDisabled(); + }); + }); + + describe('predicted pointer / make editable', () => { + it('shows the "make this slice editable" banner and calls onMakeSliceEditable', async () => { + const user = userEvent.setup(); + const onMakeSliceEditable = vi.fn(); + render( + , + ); + const btn = screen.getByRole('button', { name: /make this slice editable/i }); + await user.click(btn); + expect(onMakeSliceEditable).toHaveBeenCalledTimes(1); + }); + + it('shows "Vectorizing…" and disables the button while vectorizingSlice is true', () => { + render( + , + ); + const btn = screen.getByRole('button', { name: /vectorizing…/i }); + expect(btn).toBeDisabled(); + }); + }); + + describe('push to Tiled / view in 3D / train deep model', () => { + it('does not render the sync row without onSyncToTiled', () => { + render(); + expect(screen.queryByRole('button', { name: /push to tiled/i })).not.toBeInTheDocument(); + }); + + it('renders push-to-Tiled and view-in-3D and calls their handlers', async () => { + const user = userEvent.setup(); + const onSyncToTiled = vi.fn(); + const onViewIn3D = vi.fn(); + render( + , + ); + await user.click(screen.getByRole('button', { name: /push to tiled/i })); + expect(onSyncToTiled).toHaveBeenCalledTimes(1); + await user.click(screen.getByRole('button', { name: /view in 3d/i })); + expect(onViewIn3D).toHaveBeenCalledTimes(1); + }); + + it('shows the syncing state, success message, and error message', () => { + const { rerender } = render( + , + ); + expect(screen.getByRole('button', { name: /pushing…/i })).toBeDisabled(); + + rerender( + , + ); + expect(screen.getByText(/pushed to tiled/i)).toBeInTheDocument(); + + rerender( + , + ); + expect(screen.getByText('network failed')).toBeInTheDocument(); + }); + + it('renders the sync row when there are only predicted pointers (no real shapes yet)', () => { + render( + , + ); + expect(screen.getByRole('button', { name: /push to tiled/i })).toBeInTheDocument(); + }); + + it('renders "Train a deep model on this" and calls onTrainDeepModel', async () => { + const user = userEvent.setup(); + const onTrainDeepModel = vi.fn(); + render( + , + ); + await user.click(screen.getByRole('button', { name: /train a deep model on this/i })); + expect(onTrainDeepModel).toHaveBeenCalledTimes(1); + }); + }); + + describe('prediction results: coverage headline, class probability, commit/dismiss', () => { + it('shows a "good" coverage headline when confident percentage is high', () => { + render( + , + ); + expect(screen.getByText(/90% of pixels are confidently labeled at 95% coverage\./)).toBeInTheDocument(); + }); + + it('shows a "warn" coverage headline when confident percentage is low', () => { + render( + , + ); + expect(screen.getByText(/Only 10% confidently labeled/)).toBeInTheDocument(); + }); + + it('renders the class-probability cycler and calls onCycleProbaClass / onProbaThresholdChange', async () => { + const user = userEvent.setup(); + const onCycleProbaClass = vi.fn(); + const onProbaThresholdChange = vi.fn(); + render( + (id === 1 ? 'Cell' : 'Background'), + onCycleProbaClass, + onProbaThresholdChange, + })} + />, + ); + expect(screen.getByTitle('Cell')).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: /next class/i })); + expect(onCycleProbaClass).toHaveBeenCalledWith(1); + await user.click(screen.getByRole('button', { name: /previous class/i })); + expect(onCycleProbaClass).toHaveBeenCalledWith(-1); + expect(onProbaThresholdChange).not.toHaveBeenCalled(); + }); + + it('does not render the class-probability cycler without the cycle/threshold callbacks', () => { + render( + , + ); + expect(screen.queryByText(/class probability/i)).not.toBeInTheDocument(); + }); + + it('renders commit and dismiss buttons, using a custom commitLabel, and calls their handlers', async () => { + const user = userEvent.setup(); + const onCommit = vi.fn(); + const onDismiss = vi.fn(); + render( + , + ); + await user.click(screen.getByRole('button', { name: /commit my classes/i })); + expect(onCommit).toHaveBeenCalledTimes(1); + await user.click(screen.getByRole('button', { name: /^dismiss$/i })); + expect(onDismiss).toHaveBeenCalledTimes(1); + }); + }); + + describe('error state', () => { + it('renders the top-level error message', () => { + render(); + expect(screen.getByText('Feature bank missing.')).toBeInTheDocument(); + }); + }); +}); diff --git a/frontend/src/components/annotate/PixelClassifierPanel/index.tsx b/frontend/src/components/annotate/PixelClassifierPanel/index.tsx new file mode 100644 index 0000000..2b5b3d6 --- /dev/null +++ b/frontend/src/components/annotate/PixelClassifierPanel/index.tsx @@ -0,0 +1,639 @@ +/** + * PixelClassifierPanel — train / conformal-predict / commit, with conformal + * coverage surfaced as headline guidance (not a legend entry). + * + * Suggest-Labels (manifold coverage) used to be folded in here; it now has + * its own section (`SuggestLabelsPanel`), positioned right after feature-bank + * setup instead of behind classifier training, which it never depended on. + */ +import { Brain, CaretLeft, CaretRight, CircleNotch, Cube, TreeStructure, WarningCircle } from '@phosphor-icons/react'; +import type { ClfParams, ClfPredictCounts, ClfTrainResult } from '@/hooks/usePixelClassifier'; +import { cn } from '@/lib/utils'; +import CollapsibleSection from '@/components/common/CollapsibleSection'; + +/** Shared progress-bar treatment (see useExportJob / DownloadModal). */ +function BusyBar({ label }: { label: string }) { + return ( +
+
+ {label} +
+
+
+
+
+ ); +} + +/** Real done/total progress bar for batch jobs (see DownloadModal's job bar). */ +function JobProgressBar({ label, done, total }: { label: string; done: number; total: number }) { + const pct = total > 0 ? Math.round((done / total) * 100) : 0; + return ( +
+
+ {label} + {total > 0 && {done}/{total}} +
+
+
+
+
+ ); +} + +/** Coverage guidance derived from conformal set counts — the headline, not a legend line. */ +function coverageHeadline(counts: ClfPredictCounts, alpha: number): { text: string; tone: 'good' | 'warn' } { + const total = counts.singleton + counts.multi + counts.abstain; + if (total === 0) return { text: 'No predicted pixels yet.', tone: 'warn' }; + const confidentPct = Math.round((counts.singleton / total) * 100); + const alphaPct = Math.round(alpha * 100); + if (confidentPct >= 70) { + return { + text: `${confidentPct}% of pixels are confidently labeled at ${100 - alphaPct}% coverage.`, + tone: 'good', + }; + } + return { + text: `Only ${confidentPct}% confidently labeled — most pixels need more training data or a lower α.`, + tone: 'warn', + }; +} + +export interface PixelClassifierPanelProps { + hasFeatureJob: boolean; + canAutoPreprocess?: boolean; + hasShapes: boolean; + training: boolean; + predicting: boolean; + params: ClfParams; + onParamsChange: (p: ClfParams) => void; + model: ClfTrainResult | null; + hasPrediction: boolean; + predictCounts: ClfPredictCounts | null; + probaClassIndex?: number; + activeProbaClassId?: number | null; + activeProbaThreshold?: number; + classLabelForId?: (classId: number) => string; + onCycleProbaClass?: (delta: number) => void; + onProbaThresholdChange?: (threshold: number) => void; + error: string | null; + onTrain: () => void; + onPredict: () => void; + onCommit: () => void; + onDismiss: () => void; + commitLabel?: string; + featureSetupId?: string | null; + trainerId?: string | null; + + // Multi-slice train: pool labeled pixels across every annotated slice of the sample. + annotatedSliceCount: number; + trainAcrossSlices: boolean; + onTrainAcrossSlicesChange: (v: boolean) => void; + multiTraining: boolean; + multiTrainProgress: { done: number; total: number } | null; + + // Apply a trained model across many slices at once, then commit selected classes. + totalSliceCount: number; + commitClassIds: number[]; + onToggleCommitClassId: (classId: number) => void; + volumeApplying: boolean; + volumeApplyProgress: { done: number; total: number } | null; + volumeApplyResult: { runCount: number; errorCount: number; cancelled: boolean } | null; + onApplyToVolume: () => void; + onCommitVolumeApply: () => void; + onCancelVolumeApply: () => void; + onDismissVolumeApply: () => void; + + // The slice on screen has an un-vectorized predicted region (a + // predictedRasterStore pointer, set by "Commit predicted shapes" — see + // AnnotatePage's handleCommitVolumeApply) that's only ever shown as a + // raster overlay until explicitly turned into editable Shape[]. + hasPredictedPointerOnCurrentSlice: boolean; + vectorizingSlice: boolean; + onMakeSliceEditable: () => void; + // True when this sample has ANY committed-but-not-yet-vectorized predicted + // pointer, on any slice — annotatedSliceCount alone (real shapes only, + // correctly so for multi-slice training) would otherwise hide "Push to + // Tiled" entirely for a sample whose only committed content is still + // pointers, even though there's real predicted content to push. + hasAnyPredictedPointers: boolean; + + // Hand-off to the deep-training tab: the iPred annotation work already done + // on this sample becomes the training set, pre-selected there. + onTrainDeepModel?: () => void; + + // Push this sample's current shapes (every slice, both origins) to Tiled's + // __masks container and jump straight to the 3D view's Fast + // (iPred) layer — the direct path that skips having to separately + // discover the Export modal's "Sync masks to Tiled" action first. + // Split into two independent actions — a combined "push + navigate" used + // to force you onto the 3D page (and into "Build Volume" if the dataset's + // pyramid didn't exist yet) with no way back to Annotate short of the + // browser's own back button. Pushing to Tiled no longer navigates at all; + // viewing in 3D no longer requires a push to have just happened. + onSyncToTiled?: () => void; + syncingToTiled?: boolean; + syncToTiledError?: string | null; + syncedToTiled?: boolean; + onViewIn3D?: () => void; +} + +/** Sidebar controls for CatBoost + split-conformal sets, with suggest-labels folded in. */ +export default function PixelClassifierPanel({ + hasFeatureJob, + canAutoPreprocess = false, + hasShapes, + training, + predicting, + params, + onParamsChange, + model, + hasPrediction, + predictCounts, + probaClassIndex = 0, + activeProbaClassId = null, + activeProbaThreshold = 0.5, + classLabelForId, + onCycleProbaClass, + onProbaThresholdChange, + error, + onTrain, + onPredict, + onCommit, + onDismiss, + commitLabel = 'Commit singletons', + featureSetupId = null, + trainerId = null, + annotatedSliceCount, + trainAcrossSlices, + onTrainAcrossSlicesChange, + multiTraining, + multiTrainProgress, + totalSliceCount, + commitClassIds, + onToggleCommitClassId, + volumeApplying, + volumeApplyProgress, + volumeApplyResult, + onApplyToVolume, + onCommitVolumeApply, + onCancelVolumeApply, + onDismissVolumeApply, + hasPredictedPointerOnCurrentSlice, + vectorizingSlice, + onMakeSliceEditable, + hasAnyPredictedPointers, + onTrainDeepModel, + onSyncToTiled, + syncingToTiled = false, + syncToTiledError = null, + syncedToTiled = false, + onViewIn3D, +}: PixelClassifierPanelProps) { + const busy = training || predicting || multiTraining || volumeApplying; + const canTrain = + (trainAcrossSlices ? annotatedSliceCount > 0 : (hasFeatureJob || canAutoPreprocess) && hasShapes) && + !busy; + // ensureFeatureBank() (inside predict()) computes a bank for the current slice + // when one isn't ready yet, so predicting doesn't require hasFeatureJob up front — + // only that a model exists to apply. + const canPredict = !!model && !busy; + const maxImp = model?.featureImportances[0]?.importance ?? 1; + const alphaPct = Math.round(params.alpha * 100); + const nClasses = model?.classIds.length ?? 0; + const threshPct = Math.round(activeProbaThreshold * 100); + const probaLabel = + activeProbaClassId !== null + ? (classLabelForId?.(activeProbaClassId) ?? `class ${activeProbaClassId}`) + : '—'; + const headline = predictCounts ? coverageHeadline(predictCounts, params.alpha) : null; + + return ( + + {trainerId ?? 'catboost'} · conformal + + } + > +

+ Setup:{' '} + + {featureSetupId ?? 'none — pick a recipe under Features'} + + {!hasFeatureJob && canAutoPreprocess ? ( + · Train will compute features + ) : null} +

+ +
+ + + +
+ + + + + + + {training && } + {multiTraining && ( + + )} + + + {predicting && } + + {model && ( +
+

+ Acc {(model.trainAccuracy * 100).toFixed(1)}% · train {model.nTrain.toLocaleString()} / + cal {model.nCal.toLocaleString()} · {model.nTrees} trees · classes{' '} + {model.classIds.join(', ')} + {model.usesSam ? ' · +SlimSAM' : ''} +

+ + {model.featureImportances.length > 0 && ( +
+ Feature importance + {model.featureImportances.slice(0, 8).map((fi) => ( +
+ {fi.label} +
+
+
+ {fi.importance.toFixed(1)} +
+ ))} +
+ )} +
+ )} + + {model && ( +
+ Classes to commit +
+ {model.classIds.map((cid) => { + const on = commitClassIds.includes(cid); + return ( + + ); + })} +
+ + + + {volumeApplying && ( +
+ + +
+ )} + + {volumeApplyResult && !volumeApplying && ( +
+

+ {volumeApplyResult.cancelled ? 'Cancelled — ' : ''} + Predicted {volumeApplyResult.runCount} slice(s) + {volumeApplyResult.errorCount > 0 ? `, ${volumeApplyResult.errorCount} failed` : ''}. +

+
+ + +
+
+ )} + + {hasPredictedPointerOnCurrentSlice && ( +
+

This slice's predicted regions are shown but not yet editable.

+ +
+ )} + + {onSyncToTiled && (annotatedSliceCount > 0 || hasAnyPredictedPointers) && ( +
+
+ + {onViewIn3D && ( + + )} +
+ {syncedToTiled && !syncingToTiled && ( +

Pushed to Tiled.

+ )} + {syncToTiledError && ( +

+ + {syncToTiledError} +

+ )} +
+ )} + + {onTrainDeepModel && annotatedSliceCount > 0 && ( + + )} +
+ )} + + {hasPrediction && headline && ( +
+ {headline.text} + {predictCounts && ( +

+ singleton {predictCounts.singleton.toLocaleString()} · multi{' '} + {predictCounts.multi.toLocaleString()} · abstain {predictCounts.abstain.toLocaleString()} +

+ )} +
+ )} + + {hasPrediction && nClasses > 0 && onCycleProbaClass && onProbaThresholdChange && ( +
+
+ Class probability + + {probaClassIndex + 1}/{nClasses} + +
+
+ + + {probaLabel} + + +
+ +
+ )} + + {hasPrediction && ( +
+ + +
+ )} + + {error &&

{error}

} + + ); +} diff --git a/frontend/src/components/annotate/SaveModal.test.tsx b/frontend/src/components/annotate/SaveModal.test.tsx new file mode 100644 index 0000000..89f9e1d --- /dev/null +++ b/frontend/src/components/annotate/SaveModal.test.tsx @@ -0,0 +1,144 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen, waitFor } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import SaveModal, { type SaveModalProps } from './SaveModal'; +import { useSettingsStore } from '@/stores/settingsStore'; +import type { SaveDraftPayload } from '@/hooks/useSave'; + +type OnSaveArg = Parameters[0]; + +const basePayload: SaveDraftPayload = { + classes: [], + slices: {}, + split_by_slice: {}, + negative_slices: [], +}; + +function baseProps(overrides: Partial[0]> = {}) { + return { + sourceKey: 'local:sample.tif', + payload: basePayload, + shapeCount: 3, + classCount: 2, + isSaving: false, + onSave: vi.fn(async () => {}), + onClose: vi.fn(), + ...overrides, + }; +} + +beforeEach(() => { + useSettingsStore.setState({ annotatorName: '' }); + localStorage.clear(); + global.fetch = vi.fn(async () => ({ + ok: true, + blob: async () => new Blob(['fake-png'], { type: 'image/png' }), + })) as unknown as typeof fetch; + // jsdom doesn't implement these — stub them so the preview pipeline doesn't throw. + global.URL.createObjectURL = vi.fn(() => 'blob:preview-url'); + global.URL.revokeObjectURL = vi.fn(); +}); + +afterEach(() => { + cleanup(); + vi.restoreAllMocks(); +}); + +describe('SaveModal', () => { + it('shows a loading state then the fetched preview thumbnail', async () => { + render(); + expect(screen.getByText(/Generating preview/)).toBeInTheDocument(); + await waitFor(() => { + expect(screen.getByAltText('Annotation preview')).toBeInTheDocument(); + }); + expect(screen.getByAltText('Annotation preview')).toHaveAttribute('src', 'blob:preview-url'); + }); + + it('shows preview-unavailable when the fetch fails', async () => { + global.fetch = vi.fn(async () => ({ ok: false })) as unknown as typeof fetch; + render(); + await waitFor(() => { + expect(screen.getByText('Preview unavailable')).toBeInTheDocument(); + }); + }); + + it('displays shape and class counts, singular vs plural', async () => { + render(); + await waitFor(() => screen.getByAltText('Annotation preview')); + expect(screen.getByText('1 shape · 1 class')).toBeInTheDocument(); + }); + + it('displays plural shape/class counts', async () => { + render(); + await waitFor(() => screen.getByAltText('Annotation preview')); + expect(screen.getByText('3 shapes · 2 classes')).toBeInTheDocument(); + }); + + it('prefills annotator name from the settings store', async () => { + useSettingsStore.setState({ annotatorName: 'Ada' }); + render(); + expect(screen.getByLabelText('Who annotated this')).toHaveValue('Ada'); + await waitFor(() => screen.getByAltText('Annotation preview')); + }); + + it('falls back to the legacy localStorage key when the store is empty', async () => { + localStorage.setItem('sam3_annotator_name', 'Legacy Name'); + render(); + expect(screen.getByLabelText('Who annotated this')).toHaveValue('Legacy Name'); + await waitFor(() => screen.getByAltText('Annotation preview')); + }); + + it('calls onClose when Cancel is clicked', async () => { + const onClose = vi.fn(); + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole('button', { name: 'Cancel' })); + expect(onClose).toHaveBeenCalledOnce(); + }); + + it('calls onClose when the X button is clicked', async () => { + const onClose = vi.fn(); + const user = userEvent.setup(); + render(); + const buttons = screen.getAllByRole('button'); + // The X close button is the first button in the header (no accessible name). + await user.click(buttons[0]); + expect(onClose).toHaveBeenCalledOnce(); + }); + + it('submits with trimmed annotator name, notes, and base64 thumbnail; persists name to the store', async () => { + const onSave = vi.fn(async (_opts: OnSaveArg) => {}); + const user = userEvent.setup(); + render(); + await waitFor(() => screen.getByAltText('Annotation preview')); + + await user.type(screen.getByLabelText('Who annotated this'), ' Ada Lovelace '); + await user.type(screen.getByLabelText('Notes'), ' looks good '); + await user.click(screen.getByRole('button', { name: /Save version/ })); + + await waitFor(() => expect(onSave).toHaveBeenCalledTimes(1)); + const arg = onSave.mock.calls[0]![0]; + expect(arg.annotatedBy).toBe('Ada Lovelace'); + expect(arg.notes).toBe('looks good'); + expect(typeof arg.thumbnailBase64).toBe('string'); + expect(useSettingsStore.getState().annotatorName).toBe('Ada Lovelace'); + }); + + it('submits without a thumbnail when the preview fetch failed', async () => { + global.fetch = vi.fn(async () => ({ ok: false })) as unknown as typeof fetch; + const onSave = vi.fn(async (_opts: OnSaveArg) => {}); + const user = userEvent.setup(); + render(); + await waitFor(() => screen.getByText('Preview unavailable')); + await user.click(screen.getByRole('button', { name: /Save version/ })); + await waitFor(() => expect(onSave).toHaveBeenCalledTimes(1)); + expect(onSave.mock.calls[0]![0].thumbnailBase64).toBeUndefined(); + }); + + it('disables Cancel/X/Submit and shows Saving state while isSaving is true', async () => { + render(); + expect(screen.getByRole('button', { name: 'Cancel' })).toBeDisabled(); + expect(screen.getByRole('button', { name: /Saving/ })).toBeDisabled(); + await waitFor(() => screen.getByAltText('Annotation preview')); + }); +}); diff --git a/frontend/src/components/annotate/SliceNavigator/index.test.tsx b/frontend/src/components/annotate/SliceNavigator/index.test.tsx new file mode 100644 index 0000000..30b27a8 --- /dev/null +++ b/frontend/src/components/annotate/SliceNavigator/index.test.tsx @@ -0,0 +1,93 @@ +import { afterEach, beforeEach, describe, expect, it } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import SliceNavigator from './index'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { useDatasetStore } from '@/stores/datasetStore'; + +const SOURCE_KEY = 'local:sample.tif'; + +function setLoadedDataset(overrides: Partial> = {}) { + useDatasetStore.setState({ + meta: { + nSlices: 10, height: 32, width: 32, dtype: 'uint8', isRgb: false, valueRange: [0, 255], + }, + currentSlice: 3, + source: 'sample.tif', + kind: 'local', + serverUri: null, + ...overrides, + } as any); +} + +beforeEach(() => { + useAnnotationStore.getState().reset(); +}); + +afterEach(() => { + cleanup(); +}); + +describe('SliceNavigator', () => { + it('shows a placeholder message when no dataset is loaded', () => { + useDatasetStore.setState({ meta: null, source: null, kind: null } as any); + render(); + expect(screen.getByText('No dataset loaded.')).toBeInTheDocument(); + }); + + it('shows the current slice / total in the header', () => { + setLoadedDataset(); + render(); + expect(screen.getByText('Slice 4 / 10')).toBeInTheDocument(); + }); + + it('prev/next buttons step the slice and clamp at the ends', async () => { + setLoadedDataset({ currentSlice: 0 } as any); + const user = userEvent.setup(); + render(); + expect(screen.getByLabelText('Previous slice')).toBeDisabled(); + + await user.click(screen.getByLabelText('Next slice')); + expect(useDatasetStore.getState().currentSlice).toBe(1); + }); + + it('next is disabled on the last slice', () => { + setLoadedDataset({ currentSlice: 9 } as any); + render(); + expect(screen.getByLabelText('Next slice')).toBeDisabled(); + }); + + it('shows a jump-to-annotated dropdown only when slices are annotated', () => { + setLoadedDataset(); + render(); + expect(screen.queryByLabelText('Jump to annotated slice')).not.toBeInTheDocument(); + + cleanup(); + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE_KEY, 5, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + render(); + expect(screen.getByLabelText('Jump to annotated slice')).toBeInTheDocument(); + }); + + it('marks the current slice as a negative example', async () => { + setLoadedDataset(); + const user = userEvent.setup(); + render(); + const toggle = screen.getByRole('button', { name: /Mark as negative/ }); + await user.click(toggle); + expect(useAnnotationStore.getState().negativeSlices[SOURCE_KEY]).toContain('3'); + expect(await screen.findByRole('button', { name: /Negative example/ })).toHaveAttribute('aria-pressed', 'true'); + }); + + it('shows a coarse-level warning when a non-finest pyramid level is open', () => { + setLoadedDataset({ + meta: { + nSlices: 10, height: 32, width: 32, dtype: 'uint8', isRgb: false, valueRange: [0, 255], + levelKey: 'scale1', levelIndex: 1, levelWidth: 16, levelNSlices: 5, + }, + } as any); + render(); + expect(screen.getByText(/Viewing scale1/)).toBeInTheDocument(); + }); +}); diff --git a/frontend/src/components/annotate/SliceNavigator/index.tsx b/frontend/src/components/annotate/SliceNavigator/index.tsx index 33f5099..2c75b6e 100644 --- a/frontend/src/components/annotate/SliceNavigator/index.tsx +++ b/frontend/src/components/annotate/SliceNavigator/index.tsx @@ -6,6 +6,7 @@ import { useAnnotationStore } from '@/stores/annotationStore'; import { useDatasetStore } from '@/stores/datasetStore'; import { buildSourceKey } from '@/lib/sourceKey'; import DebouncedSlider from '@/components/common/DebouncedSlider'; +import CollapsibleSection from '@/components/common/CollapsibleSection'; /** Renders slice navigation controls bound to the dataset and annotation stores. */ export default function SliceNavigator() { @@ -34,11 +35,22 @@ export default function SliceNavigator() { /** Steps forward one slice (clamped at the last slice); writes to the dataset store. */ const next = () => setSlice(Math.min(n - 1, currentSlice + 1)); + // For a multiscale volume the slider is in FULL-RESOLUTION slice indices even + // when a coarse level is displayed, so say which level is actually on screen — + // otherwise "slice 345 / 690" over a 345-slice level is quietly confusing. + const onCoarseLevel = !!meta.levelKey && (meta.levelIndex ?? 0) > 0; + return ( -
- - Slice {currentSlice + 1} / {n} - + + {onCoarseLevel && ( + + Viewing {meta.levelKey} ({meta.levelWidth}² px, {meta.levelNSlices} slices) — indices stay + full-resolution. + + )} {isNegative ? 'Negative example' : 'Mark as negative'} -
+ ); } diff --git a/frontend/src/components/annotate/SuggestLabelsPanel/index.test.tsx b/frontend/src/components/annotate/SuggestLabelsPanel/index.test.tsx new file mode 100644 index 0000000..f5f3771 --- /dev/null +++ b/frontend/src/components/annotate/SuggestLabelsPanel/index.test.tsx @@ -0,0 +1,148 @@ +/** + * SuggestLabelsPanel — fully prop-driven, same style as PixelClassifierPanel's + * own test file. Extracted from PixelClassifierPanel's "suggest labels + * (manifold)" tests when the section was pulled out into its own component. + */ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, fireEvent, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import SuggestLabelsPanel, { type SuggestLabelsPanelProps } from './index'; +import type { ManifoldParams } from '@/hooks/useFeatureManifold'; + +afterEach(() => { + cleanup(); +}); + +const BASE_MANIFOLD_PARAMS: ManifoldParams = { k: 24, boxSize: 64 } as ManifoldParams; + +function baseProps(overrides: Partial = {}): SuggestLabelsPanelProps { + return { + hasFeatureJob: true, + manifoldParams: BASE_MANIFOLD_PARAMS, + onManifoldParamsChange: vi.fn(), + manifoldSampling: false, + manifoldHasSample: false, + manifoldShowHeatmap: false, + onManifoldShowHeatmapChange: vi.fn(), + manifoldShowMarkers: false, + onManifoldShowMarkersChange: vi.fn(), + manifoldHeatmapOpacity: 0.5, + onManifoldHeatmapOpacityChange: vi.fn(), + manifoldMeta: null, + manifoldError: null, + onManifoldSample: vi.fn(), + onManifoldDismiss: vi.fn(), + manifoldRoiShapeCount: 0, + canCaptureManifoldRoi: false, + onCaptureManifoldRoi: vi.fn(), + onClearManifoldRoi: vi.fn(), + ...overrides, + }; +} + +describe('SuggestLabelsPanel', () => { + it('is open by default (not buried like when it lived inside PixelClassifierPanel)', () => { + render(); + expect(screen.getByText(/suggest regions to label/i)).toBeInTheDocument(); + }); + + it('is disabled (non-clickable) when there is no feature job', async () => { + const user = userEvent.setup(); + render(); + const header = screen.getByRole('button', { name: /suggest labels/i }); + expect(header).toBeDisabled(); + await user.click(header); + // Still expanded by default; the header itself being disabled is what matters. + expect(screen.getByRole('button', { name: /suggest regions to label/i })).toBeDisabled(); + }); + + it('calls onManifoldSample and shows a sampling busy bar', async () => { + const user = userEvent.setup(); + const onManifoldSample = vi.fn(); + const { rerender } = render(); + const sampleBtn = screen.getByRole('button', { name: /suggest regions to label/i }); + await user.click(sampleBtn); + expect(onManifoldSample).toHaveBeenCalledTimes(1); + + rerender(); + expect(screen.getByText(/sampling manifold coverage…/i)).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /sampling…/i })).toBeDisabled(); + }); + + it('shows manifold meta summary once sampled', () => { + render( + , + ); + expect(screen.getByText(/12 boxes from 5,000 px · 87% variance explained/)).toBeInTheDocument(); + }); + + it('shows heatmap/marker toggles and opacity slider once a sample exists, and dismiss clears it', async () => { + const user = userEvent.setup(); + const onManifoldShowHeatmapChange = vi.fn(); + const onManifoldShowMarkersChange = vi.fn(); + const onManifoldDismiss = vi.fn(); + render( + , + ); + await user.click(screen.getByRole('checkbox', { name: /show coverage heatmap/i })); + expect(onManifoldShowHeatmapChange).toHaveBeenCalledWith(true); + await user.click(screen.getByRole('checkbox', { name: /show suggested boxes/i })); + expect(onManifoldShowMarkersChange).toHaveBeenCalledWith(true); + await user.click(screen.getByRole('button', { name: /dismiss suggestions/i })); + expect(onManifoldDismiss).toHaveBeenCalledTimes(1); + }); + + it('restrict-to-selection button is disabled without a capturable ROI, and enabled+clearable with one', async () => { + const user = userEvent.setup(); + const onCaptureManifoldRoi = vi.fn(); + const onClearManifoldRoi = vi.fn(); + const { rerender } = render( + , + ); + expect(screen.getByRole('button', { name: /restrict to selection \(0\)/i })).toBeDisabled(); + expect(screen.queryByRole('button', { name: /^clear$/i })).not.toBeInTheDocument(); + + rerender( + , + ); + const restrictBtn = screen.getByRole('button', { name: /restrict to selection \(2\)/i }); + expect(restrictBtn).toBeEnabled(); + await user.click(restrictBtn); + expect(onCaptureManifoldRoi).toHaveBeenCalledTimes(1); + await user.click(screen.getByRole('button', { name: /^clear$/i })); + expect(onClearManifoldRoi).toHaveBeenCalledTimes(1); + }); + + it('shows a manifold error message', () => { + render(); + expect(screen.getByText('Sampling failed.')).toBeInTheDocument(); + }); + + it('updates K boxes / box size params', () => { + const onManifoldParamsChange = vi.fn(); + render(); + // A fully-controlled numeric input whose value prop never changes across + // this render — user.type() would accumulate keystrokes against jsdom's + // own uncommitted DOM value instead. A single fireEvent.change avoids that. + fireEvent.change(screen.getByLabelText(/k boxes/i), { target: { value: '50' } }); + expect(onManifoldParamsChange).toHaveBeenLastCalledWith({ k: 50, boxSize: 64 }); + }); +}); diff --git a/frontend/src/components/annotate/SuggestLabelsPanel/index.tsx b/frontend/src/components/annotate/SuggestLabelsPanel/index.tsx new file mode 100644 index 0000000..3f7507a --- /dev/null +++ b/frontend/src/components/annotate/SuggestLabelsPanel/index.tsx @@ -0,0 +1,216 @@ +/** + * SuggestLabelsPanel — manifold-coverage sampling ("Suggest regions to + * label"), as its own sidebar section. + * + * Previously folded into PixelClassifierPanel, near the bottom of the + * sidebar — but useFeatureManifold only ever depends on a feature bank + * (featureJobId), never on a trained classifier, so it was buried behind + * classifier training for no functional reason. Pulled out to stand next to + * FeatureChannelsPanel instead: the earliest point it actually has what it + * needs, and well before a user is asked to train anything. + */ +import { CircleNotch, Compass } from '@phosphor-icons/react'; +import type { ManifoldParams } from '@/hooks/useFeatureManifold'; +import { cn } from '@/lib/utils'; +import CollapsibleSection from '@/components/common/CollapsibleSection'; + +/** Shared progress-bar treatment (see useExportJob / DownloadModal). */ +function BusyBar({ label }: { label: string }) { + return ( +
+
+ {label} +
+
+
+
+
+ ); +} + +export interface SuggestLabelsPanelProps { + hasFeatureJob: boolean; + manifoldParams: ManifoldParams; + onManifoldParamsChange: (p: ManifoldParams) => void; + manifoldSampling: boolean; + manifoldHasSample: boolean; + manifoldShowHeatmap: boolean; + onManifoldShowHeatmapChange: (v: boolean) => void; + manifoldShowMarkers: boolean; + onManifoldShowMarkersChange: (v: boolean) => void; + manifoldHeatmapOpacity: number; + onManifoldHeatmapOpacityChange: (v: number) => void; + manifoldMeta: { nPicked: number; nSubsample: number; explainedVariance: number } | null; + manifoldError: string | null; + onManifoldSample: () => void; + onManifoldDismiss: () => void; + manifoldRoiShapeCount: number; + canCaptureManifoldRoi: boolean; + onCaptureManifoldRoi: () => void; + onClearManifoldRoi: () => void; +} + +/** Sidebar section for manifold-coverage label suggestions. */ +export default function SuggestLabelsPanel({ + hasFeatureJob, + manifoldParams, + onManifoldParamsChange, + manifoldSampling, + manifoldHasSample, + manifoldShowHeatmap, + onManifoldShowHeatmapChange, + manifoldShowMarkers, + onManifoldShowMarkersChange, + manifoldHeatmapOpacity, + onManifoldHeatmapOpacityChange, + manifoldMeta, + manifoldError, + onManifoldSample, + onManifoldDismiss, + manifoldRoiShapeCount, + canCaptureManifoldRoi, + onCaptureManifoldRoi, + onClearManifoldRoi, +}: SuggestLabelsPanelProps) { + return ( + } + disabled={!hasFeatureJob} + > +
+
+ + +
+ +
+ + {manifoldRoiShapeCount > 0 && ( + + )} +
+ + + {manifoldSampling && } + + {manifoldMeta && ( +

+ {manifoldMeta.nPicked} boxes from {manifoldMeta.nSubsample.toLocaleString()} px ·{' '} + {(manifoldMeta.explainedVariance * 100).toFixed(0)}% variance explained +

+ )} + + {manifoldHasSample && ( +
+ + + + +
+ )} + + {manifoldError && ( +

{manifoldError}

+ )} +
+
+ ); +} diff --git a/frontend/src/components/annotate/Toolbar/index.test.tsx b/frontend/src/components/annotate/Toolbar/index.test.tsx new file mode 100644 index 0000000..3e63983 --- /dev/null +++ b/frontend/src/components/annotate/Toolbar/index.test.tsx @@ -0,0 +1,365 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, fireEvent, render, screen, within } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import Toolbar from './index'; +import { useToolStore } from '@/stores/toolStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import * as editHistory from '@/hooks/editHistory'; + +// Avoid spinning up the real SAM worker (unavailable in jsdom) — the magic-tool +// panel only needs a stable, controllable status/support surface. +vi.mock('@/hooks/useSam', () => ({ + useSam: vi.fn(() => ({ + status: 'idle', + error: null, + backend: null, + webgpu: false, + supported: true, + ensureEncoded: vi.fn(), + segment: vi.fn(), + })), +})); + +vi.mock('@/hooks/editHistory', async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + undo: vi.fn(), + redo: vi.fn(), + }; +}); + +const initialToolState = useToolStore.getState(); + +beforeEach(() => { + useToolStore.setState(initialToolState, true); + useAnnotationStore.getState().reset(); + useAnnotationStore.temporal.getState().clear(); + vi.mocked(editHistory.undo).mockClear(); + vi.mocked(editHistory.redo).mockClear(); +}); + +afterEach(() => { + cleanup(); +}); + +describe('Toolbar', () => { + it('renders every tool as a radio button, defaulting to Pan active', () => { + render(); + const group = screen.getByRole('radiogroup', { name: 'Drawing tools' }); + const radios = within(group).getAllByRole('radio'); + expect(radios).toHaveLength(11); + expect(screen.getByRole('radio', { name: /Pan/ })).toHaveAttribute('aria-checked', 'true'); + expect(screen.getByRole('radio', { name: /Brush/ })).toHaveAttribute('aria-checked', 'false'); + }); + + it('clicking a tool button selects it in the store and updates aria-checked', async () => { + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole('radio', { name: /Brush \(b\)/ })); + expect(useToolStore.getState().tool).toBe('brush'); + expect(screen.getByRole('radio', { name: /Brush/ })).toHaveAttribute('aria-checked', 'true'); + expect(screen.getByRole('radio', { name: /Pan/ })).toHaveAttribute('aria-checked', 'false'); + }); + + it('reflects the active tool from the store even when set externally', () => { + useToolStore.setState({ tool: 'polygon' }); + render(); + expect(screen.getByRole('radio', { name: /Polygon/ })).toHaveAttribute('aria-checked', 'true'); + }); + + it('disables every tool button and shows the no-class hint when disabled', () => { + render(); + expect(screen.getByText('Add a class above to start annotating.')).toBeInTheDocument(); + const group = screen.getByRole('radiogroup', { name: 'Drawing tools' }); + within(group).getAllByRole('radio').forEach((btn) => { + expect(btn).toBeDisabled(); + // Disabled tools are never shown as active, even the current tool. + expect(btn).toHaveAttribute('aria-checked', 'false'); + }); + }); + + it('does not show the no-class hint when not disabled', () => { + render(); + expect(screen.queryByText('Add a class above to start annotating.')).not.toBeInTheDocument(); + }); + + describe('undo/redo', () => { + it('disables undo and redo when the history is empty', () => { + render(); + expect(screen.getByRole('button', { name: 'Undo' })).toBeDisabled(); + expect(screen.getByRole('button', { name: 'Redo' })).toBeDisabled(); + }); + + it('enables undo when there is past history and calls editHistory.undo on click', async () => { + useAnnotationStore.temporal.setState({ pastStates: [{} as any] }); + const user = userEvent.setup(); + render(); + const undoBtn = screen.getByRole('button', { name: 'Undo' }); + expect(undoBtn).toBeEnabled(); + await user.click(undoBtn); + expect(editHistory.undo).toHaveBeenCalledTimes(1); + }); + + it('enables redo when there is future history and calls editHistory.redo on click', async () => { + useAnnotationStore.temporal.setState({ futureStates: [{} as any] }); + const user = userEvent.setup(); + render(); + const redoBtn = screen.getByRole('button', { name: 'Redo' }); + expect(redoBtn).toBeEnabled(); + await user.click(redoBtn); + expect(editHistory.redo).toHaveBeenCalledTimes(1); + }); + }); + + describe('brush / eraser / threshold panel', () => { + it('shows the brush radius controls for the brush tool', () => { + useToolStore.setState({ tool: 'brush' }); + render(); + expect(screen.getByLabelText('Brush radius (px)')).toBeInTheDocument(); + }); + + it('does not show brush radius controls for tools without a brush', () => { + useToolStore.setState({ tool: 'polygon' }); + render(); + expect(screen.queryByLabelText('Brush radius')).not.toBeInTheDocument(); + }); + + it('changing the brush radius number input snaps and updates the store', () => { + useToolStore.setState({ tool: 'brush', brushSize: 10 }); + render(); + const input = screen.getByLabelText('Brush radius (px)') as HTMLInputElement; + fireEvent.change(input, { target: { value: '25' } }); + expect(useToolStore.getState().brushSize).toBe(25); + }); + + it('shows the erase-scope radiogroup only for the eraser tool', () => { + useToolStore.setState({ tool: 'eraser' }); + render(); + expect(screen.getByRole('radiogroup', { name: 'Erase scope' })).toBeInTheDocument(); + }); + + it('toggles erase scope between class and all classes', async () => { + useToolStore.setState({ tool: 'eraser', eraseAllClasses: false }); + const user = userEvent.setup(); + render(); + const allRadio = screen.getByRole('radio', { name: 'Erase all classes' }); + expect(allRadio).not.toBeChecked(); + await user.click(allRadio); + expect(useToolStore.getState().eraseAllClasses).toBe(true); + }); + }); + + describe('threshold / sampler panel', () => { + it('shows the threshold band controls for the threshold tool', () => { + useToolStore.setState({ tool: 'threshold' }); + render(); + expect(screen.getByText('Set band from a region')).toBeInTheDocument(); + expect(screen.getByText('Show in-range overlay')).toBeInTheDocument(); + }); + + it('shows the threshold panel for the sampler tool too, in sampling mode', () => { + useToolStore.setState({ tool: 'sampler' }); + render(); + expect(screen.getByText('Sampling — draw a loop')).toBeInTheDocument(); + }); + + it('toggles into sampler mode and back via the sample button', async () => { + useToolStore.setState({ tool: 'threshold' }); + const user = userEvent.setup(); + render(); + await user.click(screen.getByText('Set band from a region')); + expect(useToolStore.getState().tool).toBe('sampler'); + await user.click(screen.getByText('Sampling — draw a loop')); + expect(useToolStore.getState().tool).toBe('threshold'); + }); + + it('toggles the threshold overlay checkbox', async () => { + useToolStore.setState({ tool: 'threshold', thresholdOverlay: true }); + const user = userEvent.setup(); + render(); + const checkbox = screen.getByRole('checkbox', { name: 'Show in-range overlay' }); + expect(checkbox).toBeChecked(); + await user.click(checkbox); + expect(useToolStore.getState().thresholdOverlay).toBe(false); + }); + + it('shows a Sampler result readout when samplerFit is provided', () => { + useToolStore.setState({ tool: 'threshold' }); + render( + , + ); + expect(screen.getByText(/match 90%/)).toBeInTheDocument(); + }); + + it('shows the collapsed-band warning when the fit collapsed', () => { + useToolStore.setState({ tool: 'threshold' }); + render(); + expect(screen.getByText('Band not applied')).toBeInTheDocument(); + }); + + it('offers "View band in 3D" for a plain (non-projected) fit and sends the native lo/hi', async () => { + const user = userEvent.setup(); + const onSendBandTo3D = vi.fn(); + useToolStore.setState({ tool: 'threshold' }); + render( + , + ); + const button = screen.getByRole('button', { name: /view band in 3d/i }); + await user.click(button); + // Native (lo/hi), not the displayed range shown in the readout above it. + expect(onSendBandTo3D).toHaveBeenCalledWith(10, 200); + }); + + it('does not offer "View band in 3D" for a texture-projected fit', () => { + useToolStore.setState({ tool: 'threshold' }); + render( + , + ); + expect(screen.queryByRole('button', { name: /view band in 3d/i })).not.toBeInTheDocument(); + }); + }); + + describe('select panel', () => { + it('shows the select-scope radiogroup for the select tool', () => { + useToolStore.setState({ tool: 'select' }); + render(); + expect(screen.getByRole('radiogroup', { name: 'Select scope' })).toBeInTheDocument(); + }); + + it('toggles select scope between class and all', async () => { + useToolStore.setState({ tool: 'select', selectScope: 'all' }); + const user = userEvent.setup(); + render(); + const classRadio = screen.getByRole('radio', { name: 'Select this class' }); + expect(classRadio).not.toBeChecked(); + await user.click(classRadio); + expect(useToolStore.getState().selectScope).toBe('class'); + }); + }); + + describe('fill panel', () => { + it('shows the fill threshold slider for the fill tool', () => { + useToolStore.setState({ tool: 'fill', fillThreshold: 0.1 }); + render(); + expect(screen.getByText(/Fill threshold/)).toBeInTheDocument(); + expect(screen.getByText(/10%/)).toBeInTheDocument(); + }); + }); + + describe('magic panel', () => { + it('shows the SAM engine controls by default', () => { + useToolStore.setState({ tool: 'magic', magicEngine: 'sam' }); + render(); + expect(screen.getByRole('button', { name: 'Smart (AI)' })).toBeInTheDocument(); + expect(screen.getByText('Detail', { selector: 'label' })).toBeInTheDocument(); + expect(screen.getByText('Avoid other-class regions')).toBeInTheDocument(); + }); + + it('switches to the classic engine and shows its controls', async () => { + useToolStore.setState({ tool: 'magic', magicEngine: 'sam' }); + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole('button', { name: 'Classic' })); + expect(useToolStore.getState().magicEngine).toBe('classic'); + expect(await screen.findByRole('button', { name: 'Connected' })).toBeInTheDocument(); + expect(screen.getByRole('button', { name: 'All similar' })).toBeInTheDocument(); + }); + + it('shows the edge-stop slider only in contiguous classic mode', () => { + useToolStore.setState({ tool: 'magic', magicEngine: 'classic', magicMode: 'contiguous' }); + render(); + expect(screen.getByRole('slider', { name: 'Edge stop' })).toBeInTheDocument(); + }); + + it('hides the edge-stop slider in global classic mode', () => { + useToolStore.setState({ tool: 'magic', magicEngine: 'classic', magicMode: 'global' }); + render(); + expect(screen.queryByRole('slider', { name: 'Edge stop' })).not.toBeInTheDocument(); + }); + + it('selecting a SAM detail level updates the store', async () => { + useToolStore.setState({ tool: 'magic', magicEngine: 'sam', samDetail: 'auto' }); + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole('button', { name: 'fine' })); + expect(useToolStore.getState().samDetail).toBe('fine'); + }); + + it('toggles the avoid-labeled and connected-only SAM checkboxes', async () => { + useToolStore.setState({ + tool: 'magic', magicEngine: 'sam', samAvoidLabeled: true, samConnectedOnly: true, + }); + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole('checkbox', { name: /Avoid other-class regions/ })); + expect(useToolStore.getState().samAvoidLabeled).toBe(false); + await user.click(screen.getByRole('checkbox', { name: /Connected regions only/ })); + expect(useToolStore.getState().samConnectedOnly).toBe(false); + }); + }); + + describe('global toggles', () => { + it('toggles "Clip to other classes"', async () => { + useToolStore.setState({ clipToOtherClasses: true }); + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole('checkbox', { name: 'Clip to other classes' })); + expect(useToolStore.getState().clipToOtherClasses).toBe(false); + }); + + it('toggles "Merge overlapping same class"', async () => { + useToolStore.setState({ mergeOverlappingSameClass: false }); + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole('checkbox', { name: 'Merge overlapping same class' })); + expect(useToolStore.getState().mergeOverlappingSameClass).toBe(true); + }); + }); +}); diff --git a/frontend/src/components/annotate/Toolbar/index.tsx b/frontend/src/components/annotate/Toolbar/index.tsx index 9f897c8..98c5074 100644 --- a/frontend/src/components/annotate/Toolbar/index.tsx +++ b/frontend/src/components/annotate/Toolbar/index.tsx @@ -1,17 +1,35 @@ /** * Toolbar — tool selector (radiogroup), brush size, opacity, undo/redo. - * Keybinds: p=polygon, l=ellipse, e=rectangle, r=eraser, b=brush, f=fill, s=select, - * g=magic, m=magnetic, Space=pan (hold), x=next slice, t=fit to screen, Ctrl/Cmd+Z=undo + * Keybinds: p=polygon, l=ellipse, e=rectangle, r=eraser, b=brush, h=threshold brush, + * f=fill, s=select, g=magic, m=magnetic, Space=pan (hold), x=next slice, + * t=fit to screen, Ctrl/Cmd+Z=undo */ -import { Hand, Cursor, Polygon, MagnetStraight, MagicWand, Rectangle, Circle, PaintBrush, PaintBucket, Eraser, ArrowBendUpLeft, ArrowBendUpRight } from '@phosphor-icons/react'; +import { Hand, Cursor, Polygon, MagnetStraight, MagicWand, Rectangle, Circle, PaintBrush, Drop, Eyedropper, PaintBucket, Eraser, ArrowBendUpLeft, ArrowBendUpRight, ArrowCounterClockwise, Cube } from '@phosphor-icons/react'; +import { useMemo } from 'react'; import { useStore } from 'zustand'; import { useToolStore, type Tool } from '@/stores/toolStore'; import { useAnnotationStore } from '@/stores/annotationStore'; import * as editHistory from '@/hooks/editHistory'; import { cn } from '@/lib/utils'; import DebouncedSlider from '@/components/common/DebouncedSlider'; +import HistogramControl from '@/components/annotate/HistogramControl'; +import { otsuThreshold } from '@/lib/magicwand'; +import { displayAffineFor, remapHistogramToDisplay } from '@/lib/displayTransform'; +import { describeFit } from '@/lib/thresholdFit'; + +import type { SamplerFit } from '@/components/annotate/AnnotationCanvas'; import { useSam } from '@/hooks/useSam'; +/** Channel names in the user's terms, for the projection readout. */ +const CHANNEL_LABELS: Record = { + intensity: 'brightness', + dogFine: 'fine texture', + dogCoarse: 'coarse texture', + localStd: 'graininess', + meanRatio: 'local contrast', +}; + + // macOS labels the Alt key "Option" (⌥); the key name only differs on screen. const IS_MAC = typeof navigator !== 'undefined' && /mac/i.test(navigator.userAgent); const REMOVE_KEY_LABEL = IS_MAC ? 'Option' : 'Alt'; @@ -76,15 +94,128 @@ function ToolButton({ tool, label, icon, keybind, activeTool, disabled, onSelect ); } +/** Readout for one Sampler fit: what it chose, how well it did, and an undo. */ +function SamplerResult({ + fit, + onRevert, + onSendBandTo3D, +}: { + fit: SamplerFit; + onRevert?: () => void; + /** Isolate this band in the 3D transfer function (native lo/hi, 0–255). */ + onSendBandTo3D?: (lo: number, hi: number) => void; +}) { + const { label, quality } = describeFit(fit); + const tone = + quality === 'good' ? 'text-emerald-600' : quality === 'fair' ? 'text-amber-600' : 'text-red-600'; + + // The band could not be expressed at the current display settings — applying it + // would have selected nothing, so nothing was applied. + if (fit.collapsed) { + return ( +
+ Band not applied + + Your brightness/contrast/levels squash the fitted range ({fit.lo}–{fit.hi}) into a + single displayed value, so no band can express it. Reset Levels (or lower Contrast) + and sample again. + +
+ ); + } + + const projected = fit.mode === 'projected'; + // Which channels the projection actually leaned on — the reason it beat plain + // brightness, in the user's terms rather than as a weight vector. + const topChannels = (fit.weights ?? []) + .map((w) => ({ ...w, mag: Math.abs(w.weight) })) + .sort((a, b) => b.mag - a.mag) + .filter((w) => w.mag > 0.15) + .slice(0, 2) + .map((w) => CHANNEL_LABELS[w.name] ?? w.name); + + return ( +
+
+ + Band {fit.displayLo}–{fit.displayHi} + + {onRevert && ( + + )} +
+ {label} + + match {(fit.dice * 100).toFixed(0)}% · skill {(fit.skill * 100).toFixed(0)}% ·{' '} + covers {(fit.coverage * 100).toFixed(0)}% + {fit.extraSigma > 0 && ` · blur ${fit.appliedBlur.toFixed(2)}`} + + {!projected && onSendBandTo3D && ( + // Only a plain-intensity fit is expressible in the 3D viewer, which + // has just the raw scalar per voxel — a texture-projected band has + // no equivalent there. + + )} + {projected && ( + // The gate is no longer brightness, which changes how the rest of the + // panel behaves — say so rather than letting it be discovered. + + Using a texture-aware score + {topChannels.length > 0 && ` (mostly ${topChannels.join(' + ')})`} — brightness alone + scored {((fit.intensitySkill ?? 0) * 100).toFixed(0)}%. The band below now applies to + that score, so the Display sliders no longer steer this brush. Sample a plain region + to go back to brightness. + + )} +
+ ); +} + interface ToolbarProps { /** When true, drawing tools are greyed out (e.g. no class defined yet). */ disabled?: boolean; + /** Latest Sampler lasso result, shown as a quality readout. */ + samplerFit?: SamplerFit | null; + /** Restore the band/blur that were in force before the last fit. */ + onRevertSamplerFit?: () => void; + /** Isolate the last fitted band in the 3D transfer function (native lo/hi, 0–255). */ + onSendBandTo3D?: (lo: number, hi: number) => void; + /** 256-bin luminance histogram of the current slice — drives the threshold band + * picker. Owned by AnnotatePage (the canvas emits it); null before load. It is + * sampled from the PREPROCESSED base, so it must be remapped through the display + * transform to line up with the band (which is authored in displayed space). */ + histogramBins?: number[] | null; + /** Live brightness/contrast/levels/gamma, used for exactly that remap. */ + display?: { brightness: number; contrast: number; levelsLo: number; levelsHi: number; gamma: number }; + /** Working resolution multiplier — sets the sub-pixel brush radius floor. */ + upscale?: number; } /** Renders the tool radiogroup, undo/redo, and the active tool's parameter controls. */ -export default function Toolbar({ disabled = false }: ToolbarProps) { +export default function Toolbar({ + disabled = false, histogramBins = null, display, upscale = 1, + samplerFit = null, onRevertSamplerFit, onSendBandTo3D, +}: ToolbarProps) { const { tool, setTool, brushSize, setBrushSize, fillThreshold, setFillThreshold, + thresholdLo, thresholdHi, setThresholdBand, + thresholdOverlay, setThresholdOverlay, magicTolerance, setMagicTolerance, magicMode, setMagicMode, magicSigma, setMagicSigma, magicEdgeStop, setMagicEdgeStop, magicEngine, setMagicEngine, samDetail, setSamDetail, samThreshold, setSamThreshold, @@ -96,6 +227,19 @@ export default function Toolbar({ disabled = false }: ToolbarProps) { selectScope, setSelectScope, } = useToolStore(); const sam = useSam(tool === 'magic' && magicEngine === 'sam'); + // At an upscaled working resolution the brush can go sub-pixel — a 2x grid + // resolves a 0.5 px radius, which is the whole point of upscaling. + const minRadius = 1 / Math.max(1, upscale); + const snapRadius = (n: number) => Math.round(n / minRadius) * minRadius; + + // The band is authored in DISPLAYED intensity, but the histogram is sampled from + // the preprocessed base — remap it so the plot under the knobs shows the same + // image the user is looking at (and the same one the band cuts). + const bandHistogram = useMemo(() => { + if (!histogramBins || !display) return histogramBins; + const affine = displayAffineFor(display.brightness, display.contrast, display.levelsLo, display.levelsHi); + return remapHistogramToDisplay(histogramBins, affine, display.gamma); + }, [histogramBins, display]); // Undo/redo route through editHistory so a class deletion replays alongside its region // change; canUndo/canRedo still reflect the (1:1) zundo stack. const canUndo = useStore(useAnnotationStore.temporal, (s) => s.pastStates.length > 0); @@ -110,6 +254,7 @@ export default function Toolbar({ disabled = false }: ToolbarProps) { { tool: 'rectangle', label: 'Rect', icon: , keybind: 'e' }, { tool: 'ellipse', label: 'Ellipse', icon: , keybind: 'l' }, { tool: 'brush', label: 'Brush', icon: , keybind: 'b' }, + { tool: 'threshold', label: 'Thresh', icon: , keybind: 'h' }, { tool: 'fill', label: 'Fill', icon: , keybind: 'f' }, { tool: 'eraser', label: 'Eraser', icon: , keybind: 'r' }, ]; @@ -178,27 +323,29 @@ export default function Toolbar({ disabled = false }: ToolbarProps) { ⌘/Ctrl+Z undo
- {(tool === 'brush' || tool === 'eraser') && ( + {(tool === 'brush' || tool === 'eraser' || tool === 'threshold') && (
setBrushSize(Math.max(1, Number(e.target.value)))} + onChange={(e) => setBrushSize(Math.max(minRadius, Number(e.target.value)))} className="flex-1 min-w-0" aria-label="Brush radius" /> { const n = Number(e.target.value); - if (Number.isFinite(n)) setBrushSize(Math.min(500, Math.max(1, Math.round(n)))); + if (Number.isFinite(n)) setBrushSize(Math.min(500, Math.max(minRadius, snapRadius(n)))); }} className="w-16 flex-shrink-0 border border-gray-200 rounded px-1 py-0.5 text-xs text-right tabular-nums" aria-label="Brush radius (px)" @@ -223,6 +370,83 @@ export default function Toolbar({ disabled = false }: ToolbarProps) {
)} + {(tool === 'threshold' || tool === 'sampler') && ( +
+ {/* Sampling is a mode of this tool, not a tool of its own: its only + output is this panel's band (and blur), so it belongs here. */} +
+ +

+ Lasso one example of the feature. The band is fitted to match inside it and + avoid the ring just outside — which also highlights similar features elsewhere. + Nothing is annotated. +

+ {samplerFit && ( + + )} +
+ + setThresholdBand(0, 255)} + label="Threshold" + accent="red" + actions={ + + } + /> + +

+ Paints only where intensity falls inside the band, so a stroke stops at the + feature boundary. Shift-click samples the pixel under the cursor;{' '} + {REMOVE_KEY_LABEL}-drag erases. For noisy scans raise Blur in + Display first. +

+
+ )} + {tool === 'select' && (
diff --git a/frontend/src/components/annotate/VersionHistoryModal.test.tsx b/frontend/src/components/annotate/VersionHistoryModal.test.tsx new file mode 100644 index 0000000..cd75792 --- /dev/null +++ b/frontend/src/components/annotate/VersionHistoryModal.test.tsx @@ -0,0 +1,110 @@ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, fireEvent, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import VersionHistoryModal from './VersionHistoryModal'; +import type { VersionMeta } from '@/hooks/useSave'; + +const versions: VersionMeta[] = [ + { version: 1, saved_at: '2024-01-01T12:00:00Z', shape_count: 1, class_count: 1, annotated_by: null, notes: null }, + { + version: 2, + saved_at: '2024-01-02T12:00:00Z', + shape_count: 5, + class_count: 2, + annotated_by: 'Ada', + notes: 'second pass', + }, +]; + +afterEach(() => { + cleanup(); +}); + +describe('VersionHistoryModal', () => { + it('shows an empty-state message with no versions', () => { + render( + + ); + expect(screen.getByText(/No saved versions yet/)).toBeInTheDocument(); + }); + + it('renders versions newest-first with shape/class counts and pluralization', () => { + render( + + ); + const rows = screen.getAllByText(/^v\d/); + expect(rows[0]).toHaveTextContent('v2'); + expect(rows[1]).toHaveTextContent('v1'); + expect(screen.getByText(/5 shapes/)).toBeInTheDocument(); + expect(screen.getByText(/2 classes/)).toBeInTheDocument(); + expect(screen.getByText(/1 shape\b/)).toBeInTheDocument(); + expect(screen.getByText(/1 class\b/)).toBeInTheDocument(); + }); + + it('shows annotator and notes only when present', () => { + render( + + ); + expect(screen.getByText(/by Ada/)).toBeInTheDocument(); + expect(screen.getByText('second pass')).toBeInTheDocument(); + }); + + it('does not render thumbnails when sourceKey is null', () => { + render( + + ); + expect(screen.queryByAltText(/thumbnail/)).not.toBeInTheDocument(); + }); + + it('renders a thumbnail image per version when sourceKey is set', () => { + render( + + ); + expect(screen.getByAltText('Version 1 thumbnail')).toBeInTheDocument(); + expect(screen.getByAltText('Version 2 thumbnail')).toBeInTheDocument(); + }); + + it('falls back to a placeholder icon when a thumbnail fails to load', () => { + render( + + ); + const img = screen.getByAltText('Version 2 thumbnail'); + fireEvent.error(img); + expect(screen.queryByAltText('Version 2 thumbnail')).not.toBeInTheDocument(); + }); + + it('calls onPreview with the version number and does not close the modal', async () => { + const onPreview = vi.fn(); + const onClose = vi.fn(); + const user = userEvent.setup(); + render( + + ); + await user.click(screen.getAllByRole('button', { name: /Preview/ })[0]); + expect(onPreview).toHaveBeenCalledWith(2); + expect(onClose).not.toHaveBeenCalled(); + }); + + it('calls onRestore with the version number and then onClose', async () => { + const onRestore = vi.fn(); + const onClose = vi.fn(); + const user = userEvent.setup(); + render( + + ); + await user.click(screen.getAllByRole('button', { name: /Restore/ })[0]); + expect(onRestore).toHaveBeenCalledWith(2); + expect(onClose).toHaveBeenCalledOnce(); + }); + + it('calls onClose from the header close button', async () => { + const onClose = vi.fn(); + const user = userEvent.setup(); + render( + + ); + const headerButtons = screen.getAllByRole('button').filter((b) => !/Preview|Restore/.test(b.textContent ?? '')); + await user.click(headerButtons[0]); + expect(onClose).toHaveBeenCalledOnce(); + }); +}); diff --git a/frontend/src/components/annotate/VersionPreviewBar.test.tsx b/frontend/src/components/annotate/VersionPreviewBar.test.tsx new file mode 100644 index 0000000..72a6148 --- /dev/null +++ b/frontend/src/components/annotate/VersionPreviewBar.test.tsx @@ -0,0 +1,162 @@ +import { afterEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import VersionPreviewBar from './VersionPreviewBar'; +import type { VersionMeta } from '@/hooks/useSave'; + +afterEach(() => { + cleanup(); +}); + +function makeVersions(): VersionMeta[] { + return [ + { version: 1, saved_at: '2024-01-01T00:00:00Z', shape_count: 1, class_count: 1 }, + { version: 2, saved_at: '2024-01-02T00:00:00Z', shape_count: 3, class_count: 2 }, + { version: 3, saved_at: '2024-01-03T00:00:00Z', shape_count: 5, class_count: 2 }, + ]; +} + +describe('VersionPreviewBar', () => { + it('renders nothing when there are no versions', () => { + const { container } = render( + , + ); + expect(container).toBeEmptyDOMElement(); + }); + + it('shows the previewed version number, shape count, and (latest) tag on the max version', () => { + render( + , + ); + expect(screen.getByText(/Previewing v3/)).toBeInTheDocument(); + expect(screen.getByText('(latest)')).toBeInTheDocument(); + expect(screen.getByText('5 shapes')).toBeInTheDocument(); + }); + + it('omits the (latest) tag and uses singular "shape" for a single-shape non-latest version', () => { + render( + , + ); + expect(screen.queryByText('(latest)')).not.toBeInTheDocument(); + expect(screen.getByText('1 shape')).toBeInTheDocument(); + }); + + it('sorts versions oldest-first regardless of input order for min/max labels', () => { + const shuffled = [makeVersions()[2], makeVersions()[0], makeVersions()[1]]; + render( + , + ); + expect(screen.getByText('v1')).toBeInTheDocument(); + expect(screen.getByText('v3')).toBeInTheDocument(); + }); + + it('falls back to the latest version meta when current does not match any version', () => { + render( + , + ); + // meta falls back to sorted[last] (v3), so the label reads v3, not the unmatched `current`. + // isLatest compares `current` (99) to max (3), so the "(latest)" tag is NOT shown + // even though the displayed meta is the latest version's — a slight label mismatch. + expect(screen.getByText(/Previewing v3/)).toBeInTheDocument(); + expect(screen.queryByText('(latest)')).not.toBeInTheDocument(); + expect(screen.getByText('5 shapes')).toBeInTheDocument(); + }); + + it('shows the loading indicator only when loading is true', () => { + const { rerender } = render( + , + ); + expect(screen.getByText('loading…')).toBeInTheDocument(); + rerender( + , + ); + expect(screen.queryByText('loading…')).not.toBeInTheDocument(); + }); + + it('calls onRestore with the current version when Restore is clicked', async () => { + const onRestore = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByRole('button', { name: /Restore this version/ })); + expect(onRestore).toHaveBeenCalledWith(2); + }); + + it('calls onExit when Exit is clicked', async () => { + const onExit = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByRole('button', { name: 'Exit' })); + expect(onExit).toHaveBeenCalledOnce(); + }); + + it('steps to the older version and disables the older button at the min', async () => { + const onChange = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByRole('button', { name: 'Older version' })); + expect(onChange).toHaveBeenCalledWith(1); + + render( + , + ); + expect(screen.getAllByRole('button', { name: 'Older version' })[1]).toBeDisabled(); + }); + + it('steps to the newer version and disables the newer button at the max', async () => { + const onChange = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByRole('button', { name: 'Newer version' })); + expect(onChange).toHaveBeenCalledWith(3); + + render( + , + ); + expect(screen.getAllByRole('button', { name: 'Newer version' })[1]).toBeDisabled(); + }); + + it('does not step past the range even if step is somehow invoked at the boundary', async () => { + const onChange = vi.fn(); + const user = userEvent.setup(); + render( + , + ); + const olderBtn = screen.getByRole('button', { name: 'Older version' }); + expect(olderBtn).toBeDisabled(); + await user.click(olderBtn); // disabled — should not fire + expect(onChange).not.toHaveBeenCalled(); + }); + + it('renders a formatted date from the ISO timestamp', () => { + render( + , + ); + const expected = new Date('2024-01-02T00:00:00Z').toLocaleString(undefined, { + month: 'short', day: 'numeric', hour: '2-digit', minute: '2-digit', + }); + expect(screen.getByText(expected)).toBeInTheDocument(); + }); + + it('falls back to the raw string when the date cannot be formatted', () => { + const versions = [{ version: 1, saved_at: 'not-a-date-###', shape_count: 0, class_count: 0 }]; + const original = Date.prototype.toLocaleString; + // Force toLocaleString to throw to exercise the catch path. + Date.prototype.toLocaleString = () => { throw new Error('boom'); }; + try { + render( + , + ); + expect(screen.getByText('not-a-date-###')).toBeInTheDocument(); + } finally { + Date.prototype.toLocaleString = original; + } + }); +}); diff --git a/frontend/src/components/common/CollapsibleSection.test.tsx b/frontend/src/components/common/CollapsibleSection.test.tsx new file mode 100644 index 0000000..41506e9 --- /dev/null +++ b/frontend/src/components/common/CollapsibleSection.test.tsx @@ -0,0 +1,50 @@ +/** + * CollapsibleSection — open/close toggle, defaultOpen, headerRight slot. + */ +import { describe, expect, it } from 'vitest'; +import { render, screen } from '@testing-library/react'; +import userEvent from '@testing-library/user-event'; +import CollapsibleSection from './CollapsibleSection'; + +describe('CollapsibleSection', () => { + it('renders children by default (defaultOpen defaults to true)', () => { + render( + +

body content

+
, + ); + expect(screen.getByText('body content')).toBeInTheDocument(); + expect(screen.getByRole('button', { name: /layers/i })).toHaveAttribute('aria-expanded', 'true'); + }); + + it('hides children when defaultOpen is false, and toggling shows them', async () => { + const user = userEvent.setup(); + render( + +

hidden body

+
, + ); + expect(screen.queryByText('hidden body')).not.toBeInTheDocument(); + const header = screen.getByRole('button', { name: /classifier/i }); + expect(header).toHaveAttribute('aria-expanded', 'false'); + + await user.click(header); + expect(screen.getByText('hidden body')).toBeInTheDocument(); + expect(header).toHaveAttribute('aria-expanded', 'true'); + + await user.click(header); + expect(screen.queryByText('hidden body')).not.toBeInTheDocument(); + }); + + it('renders headerRight content regardless of open state', async () => { + const user = userEvent.setup(); + render( + +}> +

rows

+
, + ); + expect(screen.getByLabelText('Add class')).toBeInTheDocument(); + await user.click(screen.getByRole('button', { name: /classes/i })); + expect(screen.getByLabelText('Add class')).toBeInTheDocument(); + }); +}); diff --git a/frontend/src/components/common/CollapsibleSection.tsx b/frontend/src/components/common/CollapsibleSection.tsx new file mode 100644 index 0000000..2ca801c --- /dev/null +++ b/frontend/src/components/common/CollapsibleSection.tsx @@ -0,0 +1,58 @@ +/** + * CollapsibleSection — shared expand/collapse header for sidebar subsections + * (Classes, Layers, Display, Slice, Cross-slice, Measure, Features, Classifier). + * + * Reproduces the house header convention already used ad hoc by those panels + * (`text-xs font-semibold uppercase tracking-wide text-gray-500` + a small + * leading icon) and adds a Phosphor caret toggle in front of it, replacing two + * divergent one-off disclosure patterns that existed before this component. + */ +import { useState, type ReactNode } from 'react'; +import { CaretDown, CaretRight } from '@phosphor-icons/react'; +import { cn } from '@/lib/utils'; + +export interface CollapsibleSectionProps { + title: string; + /** Small leading icon, e.g. `` — matches each section's existing icon. */ + icon?: ReactNode; + /** Whether the section starts open. Not persisted — a fresh mount always uses this. */ + defaultOpen?: boolean; + /** Extra control(s) on the header's right edge, e.g. ClassManager's "+ Add" button. */ + headerRight?: ReactNode; + disabled?: boolean; + children: ReactNode; +} + +/** Sidebar section wrapper: click the header to show/hide `children`. */ +export default function CollapsibleSection({ + title, + icon, + defaultOpen = true, + headerRight, + disabled = false, + children, +}: CollapsibleSectionProps) { + const [open, setOpen] = useState(defaultOpen); + return ( +
+
+ + {headerRight} +
+ {open &&
{children}
} +
+ ); +} diff --git a/frontend/src/components/train/CapabilityBanner.tsx b/frontend/src/components/train/CapabilityBanner.tsx new file mode 100644 index 0000000..29128d4 --- /dev/null +++ b/frontend/src/components/train/CapabilityBanner.tsx @@ -0,0 +1,45 @@ +/** + * CapabilityBanner — Train-tab readiness status. Shows install guidance when + * torch/dlsia aren't available (e.g. the Docker image, which intentionally + * excludes them) instead of letting every downstream action fail opaquely. + */ +import { CheckCircle, WarningCircle } from '@phosphor-icons/react'; +import type { TrainCapability } from '@/hooks/useTrainCapability'; + +interface CapabilityBannerProps { + capability: TrainCapability; +} + +export default function CapabilityBanner({ capability }: CapabilityBannerProps) { + if (!capability.torch_available) { + return ( +
+ +
+

Training is unavailable on this server.

+

+ torch isn't installed. Run this app via start_all.sh with{' '} + INSTALL_ML=1 (default on Apple Silicon) on a machine where it + can install ML dependencies — the Docker image intentionally omits them. +

+
+
+ ); + } + + return ( +
+ + + torch {capability.torch_version} · {capability.device ?? 'cpu'} + + + dlsia (TUNet):{' '} + + {capability.dlsia.available ? 'available' : 'not installed'} + + + {capability.busy && A training/inference job is currently running.} +
+ ); +} diff --git a/frontend/src/components/train/HyperparamsPanel.tsx b/frontend/src/components/train/HyperparamsPanel.tsx new file mode 100644 index 0000000..dfcb9c8 --- /dev/null +++ b/frontend/src/components/train/HyperparamsPanel.tsx @@ -0,0 +1,188 @@ +/** + * HyperparamsPanel — training hyperparameters for the dlsia TUNet family + * (epochs, lr, batch size, image size, flip augment, depth/base_channels/ + * growth_rate). Collapsed behind
— sensible defaults are supplied + * by the backend schema, so most users never open it. + * + * DINOv3 LoRA is deferred (see Phase 5.5), so there is only one family here — + * no model-family switch, no LoRA rank/alpha fields. + */ +import { IMAGE_SIZE_CONSTRAINTS, validateImageSize } from '@/lib/trainConstraints'; + +export interface HyperparamsState { + epochs: number; + lr: number; + batch_size: number; + image_size: number; + flip_augment: boolean; + /** Cut native-resolution `image_size` windows out of each slice instead of + * shrinking whole slices to `image_size`. */ + tiling: boolean; + depth: number; + base_channels: number; + growth_rate: number; +} + +interface HyperparamsPanelProps { + values: HyperparamsState; + onChange: (updates: Partial) => void; + runName: string; + onRunNameChange: (name: string) => void; + /** Measure the largest batch size this config fits (see backend/batch_probe.py). */ + onEstimateBatch: () => void; + /** Request cancellation of the running probe (cooperative — see /api/train/cancel). */ + onCancelEstimate: () => void; + /** True while the probe job runs — it holds the device, so it can't overlap. */ + estimatingBatch: boolean; + /** Progress/result line for the probe, shown under the field. */ + batchEstimateNote: string | null; + /** Why Estimate can't run right now (shown as the button's tooltip), or null when it can. */ + estimateDisabledReason: string | null; +} + +const inputClass = + 'w-full rounded-md border border-slate-600 bg-slate-900/60 px-2.5 py-1.5 text-sm text-slate-200 focus:border-sky-500 focus:outline-none'; +const labelClass = 'block text-xs text-slate-400 mb-1'; + +export default function HyperparamsPanel({ + values, onChange, runName, onRunNameChange, + onEstimateBatch, onCancelEstimate, estimatingBatch, batchEstimateNote, estimateDisabledReason, +}: HyperparamsPanelProps) { + const sizeLimits = IMAGE_SIZE_CONSTRAINTS.dlsia_tunet; + const sizeError = validateImageSize( + 'dlsia_tunet', values.image_size, values.tiling ? 'Patch size' : 'Image size', + ); + + return ( +
+ + Hyperparameters (advanced) + +
+
+ + onRunNameChange(e.target.value)} + placeholder="auto-generated" className={inputClass} + /> +
+
+ + onChange({ epochs: Number(e.target.value) })} className={inputClass} + /> +
+
+ + onChange({ lr: Number(e.target.value) })} className={inputClass} + /> +
+
+
+ + {/* Measures the real ceiling by running actual training steps, so it + needs the device to itself and can't run during another job. */} +
+ + {estimatingBatch && ( + + )} +
+
+ onChange({ batch_size: Number(e.target.value) })} className={inputClass} + /> + {batchEstimateNote && ( +

{batchEstimateNote}

+ )} +
+
+ {/* Same number, two meanings: the tile window when tiling, or the square + the whole slice is squashed into when not. */} + + onChange({ image_size: Number(e.target.value) })} + className={`${inputClass} ${sizeError ? 'border-red-500' : ''}`} + /> +

+ {sizeError ?? `${sizeLimits.min}–${sizeLimits.max}, in steps of ${sizeLimits.multipleOf}`} +

+
+
+ +
+ +
+ +

+ {values.tiling + ? `Cuts ${values.image_size}px patches at full resolution with 25% overlap, then blends the + predictions back together. Keeps fine detail on images larger than the patch size.` + : `Shrinks each whole image to ${values.image_size}px before training — faster, but detail on + large images is lost before the model sees it.`} +

+
+ +
+ + onChange({ depth: Number(e.target.value) })} className={inputClass} + /> +
+
+ + onChange({ base_channels: Number(e.target.value) })} className={inputClass} + /> +
+
+ + onChange({ growth_rate: Number(e.target.value) })} className={inputClass} + /> +
+
+
+ ); +} diff --git a/frontend/src/components/train/InferencePanel.tsx b/frontend/src/components/train/InferencePanel.tsx new file mode 100644 index 0000000..a542132 --- /dev/null +++ b/frontend/src/components/train/InferencePanel.tsx @@ -0,0 +1,304 @@ +/** + * InferencePanel — run a saved fine-tuned run over the current sample, with + * an overlay preview, then import the predictions as editable annotations + * and/or push them into Tiled as masks. + */ +import { useEffect, useMemo, useState } from 'react'; +import { useNavigate } from 'react-router'; +import { Cube } from '@phosphor-icons/react'; +import { API_BASE } from '@/config'; +import { useExportJob } from '@/hooks/useExportJob'; +import { buildSliceUrl } from '@/hooks/useImageSlice'; +import { useDatasetStore } from '@/stores/datasetStore'; +import type { RunClass } from '@/lib/importPredictions'; +import JobProgressBar from './JobProgressBar'; + +const MAX_SLICE_INDICES = 2000; + +interface InferencePanelProps { + selectedRunId: string | null; + hasOpenSample: boolean; + isTiledSource: boolean; + source: string | null; + serverUri: string | null; + currentSlice: number; + nSlices: number; + baseImageUrl: string | null; + onImportPredictions: (runClasses: RunClass[], slices: Record) => void; +} + +type Scope = 'current' | 'range' | 'all'; + +export default function InferencePanel({ + selectedRunId, hasOpenSample, isTiledSource, source, serverUri, currentSlice, nSlices, baseImageUrl, onImportPredictions, +}: InferencePanelProps) { + const navigate = useNavigate(); + const [scope, setScope] = useState('current'); + const [rangeStart, setRangeStart] = useState(0); + const [rangeEnd, setRangeEnd] = useState(Math.max(0, nSlices - 1)); + const [previewSlice, setPreviewSlice] = useState(null); + const [opacity, setOpacity] = useState(0.7); + // Set right after "Import as annotations" is clicked; cleared on a new run + // so a stale count from a previous job never lingers. + const [importedCount, setImportedCount] = useState(null); + + // Fixed keys: same reasoning as TrainPage's jobs — once started, an infer job + // is bound to its own job_id independent of anything selected afterward. + const { state: job, startJob, reset } = useExportJob('train:infer'); + const { state: writeJob, startJob: startWriteJob } = useExportJob('train:infer-write'); + + useEffect(() => { + setRangeEnd(Math.max(0, nSlices - 1)); + }, [nSlices]); + + const sliceIndices = (): number[] => { + if (scope === 'current') return [currentSlice]; + if (scope === 'range') { + const lo = Math.max(0, Math.min(rangeStart, rangeEnd)); + const hi = Math.min(nSlices - 1, Math.max(rangeStart, rangeEnd)); + return Array.from({ length: hi - lo + 1 }, (_, i) => lo + i); + } + return Array.from({ length: Math.min(nSlices, MAX_SLICE_INDICES) }, (_, i) => i); + }; + + const handleRunInference = () => { + if (!selectedRunId || !source) return; + reset(); + setImportedCount(null); + void startJob('/api/train/infer', { + run_id: selectedRunId, + kind: isTiledSource ? 'tiled' : 'local', + source, + server_uri: serverUri, + slice_indices: sliceIndices(), + }); + }; + + const handleCancelInference = () => { + // Shared registry — same cancel route every background job in this app uses. + if (job.jobId) void fetch(`${API_BASE}/api/export/cancel/${job.jobId}`, { method: 'POST' }); + }; + + const previewSlices: number[] = Array.isArray(job.result?.preview_slices) + ? (job.result!.preview_slices as number[]) + : []; + const activePreviewSlice = previewSlice ?? previewSlices[0] ?? null; + // Keyed on the job, deliberately NOT on selectedRunId: the overlay is cached + // server-side against this job_id (which run produced it is already baked in), + // and selectedRunId is page-local state that resets to null when the Train tab + // remounts — requiring it here made a finished job's overlay vanish on the way + // back to the tab, leaving just the bare base image under a "done" progress bar. + // + // Gated on `previewSlices.length > 0`, NOT `job.status === 'done'` (#15): + // infer_jobs.py now publishes `result.preview_slices` incrementally as each + // slice finishes, so a still-`running` job already has real, servable + // preview PNGs for whatever's completed so far — waiting for `done` meant + // switching slices mid-job showed nothing yet (no preview existed at all + // until the whole job finished), which read as "it stopped predicting" + // even though the job was running fine server-side. + const previewUrl = activePreviewSlice != null && job.jobId + ? `${API_BASE}/api/train/infer/preview/${job.jobId}/${activePreviewSlice}` + : null; + + // The image UNDER the overlay must follow the preview-slice slider, not the + // slice open in the app. `baseImageUrl` is bound to the open slice, so on a + // multi-slice job the overlay advanced while the picture beneath it stayed + // put — predictions from slice N drawn over the pixels of slice 0. Rebuilt + // here from activePreviewSlice; the prop stays as the fallback for the + // single-slice case and before a job has produced any preview. + const { kind, renderOpts } = useDatasetStore(); + const basePreviewUrl = useMemo(() => { + if (activePreviewSlice == null || activePreviewSlice === currentSlice) return baseImageUrl; + if (!source || !kind) return baseImageUrl; + return buildSliceUrl(source, kind, activePreviewSlice, renderOpts, serverUri); + }, [activePreviewSlice, currentSlice, source, kind, renderOpts, serverUri, baseImageUrl]); + + const totalShapes = typeof job.result?.n_shapes === 'number' ? job.result.n_shapes : 0; + + // Gates the "start a new job" controls only. Progress/results for a job + // already running or finished must stay visible on their own — e.g. after + // navigating away and back, TrainPage's selectedRunId resets to null before + // the run list settles, but a still-running job keeps going server-side and + // this hook already reattached to it (see useExportJob's persistKey). Gating + // the whole section on selectedRunId used to hide that reattached job behind + // "select a run", making a still-running (or already-finished) job look like + // it vanished. + const canStartNewJob = hasOpenSample && !!selectedRunId; + const hasJob = job.status !== 'idle'; + + return ( +
+

Inference

+ {!hasOpenSample && !hasJob && ( +

Open a sample in Browse/Annotate to run inference on it.

+ )} + {hasOpenSample && !selectedRunId && !hasJob && ( +

Select a saved run above to enable inference.

+ )} + + {canStartNewJob && ( + <> +
+ {([ + ['current', `Current slice (${currentSlice})`], + ['range', 'Slice range'], + ['all', `All slices (${nSlices})`], + ] as const).map(([value, label]) => ( + + ))} + {scope === 'range' && ( +
+ setRangeStart(Number(e.target.value))} + className="w-16 rounded border border-slate-600 bg-slate-900/60 px-1.5 py-1" + /> + to + setRangeEnd(Number(e.target.value))} + className="w-16 rounded border border-slate-600 bg-slate-900/60 px-1.5 py-1" + /> +
+ )} +
+ + + + )} + + {hasJob && ( + <> +
+
+ {job.status === 'running' && ( + + )} +
+ + {(job.status === 'done' || (job.status === 'running' && previewSlices.length > 0)) && ( +
+

+ {totalShapes} predicted region{totalShapes !== 1 ? 's' : ''} + {job.status === 'running' + ? ` so far (${previewSlices.length}/${job.total} slices predicted)` + : ''} + {job.result?.cancelled === true ? ' (cancelled — partial result)' : ''}. +

+ {previewSlices.length > 0 && basePreviewUrl && ( + <> +
+ + {previewUrl && ( + + )} +
+
+ Opacity + setOpacity(Number(e.target.value))} className="flex-1 accent-sky-500" + /> +
+ {previewSlices.length > 1 && ( +
+ Slice + setPreviewSlice(previewSlices[Number(e.target.value)])} + className="flex-1 accent-sky-500" + /> + {activePreviewSlice} +
+ )} + + )} + {/* Acting on the result (importing shapes, pushing to Tiled) waits for + `done` — the preview above is fine to show mid-job, but these two + operate on `job.result.slices`/all label PNGs, so doing them against + a still-growing partial result would silently drop whatever hasn't + predicted yet instead of erroring, which is worse than just waiting. */} + {hasOpenSample && job.status === 'done' && ( +
+ + {/* Not gated on the live `isTiledSource` prop: that reflects + whatever sample is CURRENTLY open in the global dataset + store, not the sample this job actually ran against — a + long-running job (e.g. hundreds of slices) can easily + outlive the user switching samples/tabs and coming back, + at which point `isTiledSource` no longer describes this + job at all and could wrongly hide the button for a job + that really did run on a Tiled source. The write route + is already correctly bound to the job's own source + server-side (infer_jobs.py's cached entry) and reports a + clear error if it truly isn't Tiled — surfaced below via + writeJob's own error state — so there's no need to + duplicate that check here against state that can go + stale. */} + +
+ )} + {importedCount !== null && ( +

+ Imported {importedCount} region{importedCount === 1 ? '' : 's'} as annotations — they're on the + Annotate tab and autosaved. +

+ )} + {writeJob.status === 'done' && ( +
+

+ Masks saved to Tiled + {typeof writeJob.result?.n_slices === 'number' ? ` (${writeJob.result.n_slices} slices)` : ''} — + load them in Annotate anytime with "Load saved masks". +

+ +
+ )} + +
+ )} + + )} +
+ ); +} diff --git a/frontend/src/components/train/JobProgressBar.tsx b/frontend/src/components/train/JobProgressBar.tsx new file mode 100644 index 0000000..5bedda7 --- /dev/null +++ b/frontend/src/components/train/JobProgressBar.tsx @@ -0,0 +1,48 @@ +/** + * JobProgressBar — phase/progress bar + scrolling log pane + error banner, + * shared by the training and inference sections of the Train tab. Styling + * matches DownloadModal's export-job progress display. + */ +import { WarningCircle } from '@phosphor-icons/react'; +import type { ExportJobState } from '@/hooks/useExportJob'; + +interface JobProgressBarProps { + job: ExportJobState; + /** Unit label for the done/total counter (e.g. "batches", "slices"). */ + unit?: string; +} + +export default function JobProgressBar({ job, unit = 'steps' }: JobProgressBarProps) { + if (job.status === 'idle') return null; + const pct = job.total > 0 ? Math.round((job.done / job.total) * 100) : 0; + + return ( +
+ {(job.status === 'running' || job.status === 'done') && ( + <> +
+ {job.phase || 'working'}… + {job.total > 0 && {job.done}/{job.total} {unit}} +
+
+
+
+ {job.log.length > 0 && ( +
+ {job.log.slice(-16).map((line, i) =>
{line}
)} +
+ )} + + )} + {job.status === 'error' && ( +
+ + {job.error ?? 'Job failed.'} +
+ )} +
+ ); +} diff --git a/frontend/src/components/train/RunsPanel.tsx b/frontend/src/components/train/RunsPanel.tsx new file mode 100644 index 0000000..11a6836 --- /dev/null +++ b/frontend/src/components/train/RunsPanel.tsx @@ -0,0 +1,94 @@ +/** + * RunsPanel — saved fine-tune runs (dlsia_tunet + the dlsia_denoiser family); + * pick one to run inference with (see InferencePanel), or permanently delete + * one. + */ +import { Trash } from '@phosphor-icons/react'; +import type { TrainRun } from '@/hooks/useTrainRuns'; +import { useTrainCapability } from '@/hooks/useTrainCapability'; +import { denoiseMethodLabel } from '@/lib/trainDenoiseOption'; + +interface RunsPanelProps { + runs: TrainRun[]; + selectedRunId: string | null; + onSelectRun: (runId: string) => void; + onDeleteRun: (runId: string) => void; +} + +export const FAMILY_LABELS: Record = { + dlsia_tunet: 'dlsia TUNet', + // Architecture-agnostic: this family covers both a dlsia TUNet and a plain + // convolutional autoencoder, and more than the two original schemes. + dlsia_denoiser: 'Denoiser (self-supervised)', +}; + +export default function RunsPanel({ runs, selectedRunId, onSelectRun, onDeleteRun }: RunsPanelProps) { + // Only to turn a stored method id ("tv") into its label ("Total variation"); + // the query is shared and cached across the whole Train tab. + const { capability } = useTrainCapability(); + + const handleDelete = (run: TrainRun) => { + // Includes run_id, not just family + timestamp: two runs of the same family + // trained close together can otherwise look identical in this dialog, with + // nothing to tell the user which one they're actually about to delete. + const label = `${FAMILY_LABELS[run.model_family] ?? run.model_family} — ${new Date(run.created_at).toLocaleString()} (${run.run_id})`; + if (!window.confirm(`Permanently delete this run?\n\n${label}\n\nThis removes its saved weights and cannot be undone.`)) return; + onDeleteRun(run.run_id); + }; + + return ( +
+

Saved runs

+ {runs.length === 0 ? ( +

No runs saved yet — train a model above to create one.

+ ) : ( +
+ {runs.map((run) => ( + + ))} +
+ )} +
+ ); +} diff --git a/frontend/src/components/train/TrainDenoiseToggle.tsx b/frontend/src/components/train/TrainDenoiseToggle.tsx new file mode 100644 index 0000000..f83f6a7 --- /dev/null +++ b/frontend/src/components/train/TrainDenoiseToggle.tsx @@ -0,0 +1,73 @@ +/** + * TrainDenoiseToggle — the "Train on denoised input" opt-in. + * + * It reads the setting straight from `datasetStore`'s denoise state — the same + * one Phase 2's Annotate-tab `DenoisePanel` reads/writes — so the label names + * the exact filter the user tuned there rather than offering a second, + * independent copy of the controls that could disagree with what they were + * actually looking at. + * + * Deliberately opt-in and default-off: unlike every other denoise control in + * the app this one changes what the model LEARNS, and a run trained on + * filtered pixels is not interchangeable with one trained on raw pixels. + * Whether it's on is owned by the host (it's part of that form's submit + * state); this component owns only the presentation and the "can it be used + * at all" guard — the payload builder in `lib/trainDenoiseOption` re-checks + * the guard, so a stale checked box can't leak an unusable method into a + * request. + */ +import { useDatasetStore } from '@/stores/datasetStore'; +import { useTrainCapability } from '@/hooks/useTrainCapability'; +import { trainDenoiseBlockedReason, trainDenoiseSummary } from '@/lib/trainDenoiseOption'; + +interface TrainDenoiseToggleProps { + checked: boolean; + onChange: (next: boolean) => void; + /** Host-level busy state (a job in flight); the guard disables independently. */ + disabled?: boolean; +} + +export default function TrainDenoiseToggle({ + checked, onChange, disabled = false, +}: TrainDenoiseToggleProps) { + const denoise = useDatasetStore((s) => s.denoise); + const { capability } = useTrainCapability(); + + const blocked = trainDenoiseBlockedReason(denoise.method); + const summary = trainDenoiseSummary(denoise, capability.denoise.methods); + + return ( +
+ + {blocked ? ( +

{blocked}

+ ) : ( + checked && ( +

+ Recorded on the run — inference will reapply it automatically. +

+ ) + )} +
+ ); +} diff --git a/frontend/src/components/train/TrainingDataPanel.tsx b/frontend/src/components/train/TrainingDataPanel.tsx new file mode 100644 index 0000000..2aa07e5 --- /dev/null +++ b/frontend/src/components/train/TrainingDataPanel.tsx @@ -0,0 +1,42 @@ +/** + * TrainingDataPanel — checkbox list of this session's annotated samples to + * train on (see gatherTrainingSources.ts for the session-scoping rationale). + */ +import type { TrainingCandidate } from '@/lib/gatherTrainingSources'; + +interface TrainingDataPanelProps { + candidates: TrainingCandidate[]; + selected: Set; + onToggle: (sourceKey: string) => void; +} + +export default function TrainingDataPanel({ candidates, selected, onToggle }: TrainingDataPanelProps) { + return ( +
+
+

Training data

+ {selected.size} of {candidates.length} selected +
+ {candidates.length === 0 ? ( +

+ No annotated samples yet this session. Annotate a few slices in the Annotate tab, then come back here. +

+ ) : ( +
+ {candidates.map((c) => ( + + ))} +
+ )} +
+ ); +} diff --git a/frontend/src/components/volume/BuildVolumePanel.tsx b/frontend/src/components/volume/BuildVolumePanel.tsx new file mode 100644 index 0000000..97d4411 --- /dev/null +++ b/frontend/src/components/volume/BuildVolumePanel.tsx @@ -0,0 +1,178 @@ +/** + * BuildVolumePanel — turn a per-slice dataset into a renderable 3-D volume. + * + * A dataset ingested as individual 2-D slices cannot be streamed as a volume; + * the 3-D view needs a multiscale pyramid. Building one is a one-click job here + * rather than a form, because the slices are already in Tiled — asking for a + * source directory would mean asking for data the app already holds. + * + * Only the downsampled levels are written. Full resolution stays in the existing + * per-slice nodes: the renderer only ever uploads a level that fits a GPU 3-D + * texture, so a full-resolution copy would be cost with no benefit. + */ +import { useCallback, useState } from 'react'; +import { useQuery, useQueryClient } from '@tanstack/react-query'; +import { Cube, Warning, CircleNotch } from '@phosphor-icons/react'; +import { API_BASE } from '@/config'; + +interface PyramidLevel { + path: string; + factor: [number, number, number]; + shape: [number, number, number]; +} + +interface BuildInfo { + full_shape: [number, number, number]; + dtype: string; + pyramid_plan: PyramidLevel[]; + slices_to_read: number; + already_small: boolean; +} + +interface JobState { + state: 'pending' | 'running' | 'done' | 'error'; + phase: string; + done: number; + total: number; + error: string | null; +} + +interface BuildVolumePanelProps { + source: string; + serverUri: string | null; + /** Why no volume exists, from `GET /api/volume/resolve`. */ + message: string; +} + +/** Pull the FastAPI `detail` message out of an error response. */ +async function readError(res: Response): Promise { + try { + const body = await res.json(); + if (typeof body?.detail === 'string') return body.detail; + return JSON.stringify(body); + } catch { + return `Request failed (${res.status})`; + } +} + +export default function BuildVolumePanel({ source, serverUri, message }: BuildVolumePanelProps) { + const queryClient = useQueryClient(); + const [job, setJob] = useState(null); + const [error, setError] = useState(null); + + const { data: info, isLoading, error: inspectError } = useQuery({ + queryKey: ['volume-build-inspect', serverUri, source], + queryFn: async () => { + const params = new URLSearchParams({ source, kind: 'tiled' }); + if (serverUri) params.set('server_uri', serverUri); + const res = await fetch(`${API_BASE}/api/volume/build/inspect?${params}`); + if (!res.ok) throw new Error(await readError(res)); + return res.json(); + }, + retry: false, + }); + + const build = useCallback(async () => { + setError(null); + setJob({ state: 'pending', phase: 'Starting', done: 0, total: 1, error: null }); + try { + const res = await fetch(`${API_BASE}/api/volume/build`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ source, kind: 'tiled', server_uri: serverUri }), + }); + if (!res.ok) throw new Error(await readError(res)); + const { job_id: jobId } = await res.json(); + + // Poll rather than stream: this reuses the same export-job registry (and + // the same status route) every other long-running task here already uses. + for (;;) { + await new Promise((resolve) => setTimeout(resolve, 1000)); + const statusRes = await fetch(`${API_BASE}/api/export/status/${jobId}`); + if (!statusRes.ok) throw new Error(await readError(statusRes)); + const status: JobState = await statusRes.json(); + setJob(status); + if (status.state === 'error') throw new Error(status.error || 'Build failed'); + if (status.state === 'done') break; + } + // The volume now exists — re-resolve so the viewer picks it up. + await queryClient.invalidateQueries({ queryKey: ['volume-node'] }); + } catch (err) { + setError(err instanceof Error ? err.message : String(err)); + setJob(null); + } + }, [source, serverUri, queryClient]); + + const running = job !== null && job.state !== 'done'; + const percent = job && job.total > 0 ? Math.round((job.done / job.total) * 100) : 0; + + return ( +
+
+ +

No 3D volume for this dataset yet

+

{message}

+ + {isLoading &&

Checking this dataset…

} + + {inspectError && ( +

+ + {inspectError instanceof Error ? inspectError.message : String(inspectError)} +

+ )} + + {info && !info.already_small && ( + <> +
+

+ Source{' '} + {info.full_shape.join(' × ')} {info.dtype} +

+

+ Will build{' '} + {info.pyramid_plan.map((l) => l.shape.join('×')).join(', ')} +

+

+ Reads {info.slices_to_read} slices once. Full resolution is not copied — + it stays where it is, and the 3D view never loads a level that large. +

+
+ + {!running && ( + + )} + + {running && ( +
+

+ + {job.phase} — {job.done}/{job.total} +

+
+
+
+
+ )} + + )} + + {error && ( +

+ + {error} +

+ )} +
+
+ ); +} diff --git a/frontend/src/components/volume/MaskLayersPanel.tsx b/frontend/src/components/volume/MaskLayersPanel.tsx new file mode 100644 index 0000000..4e4b689 --- /dev/null +++ b/frontend/src/components/volume/MaskLayersPanel.tsx @@ -0,0 +1,381 @@ +/** + * MaskLayersPanel — sidebar overlay for the 3-D view's two mask/annotation + * layers: "Fast (iPred)" and "Deep (dlsia)". The viewer's own `loadMask`/ + * `loadMaskFromArray`/etc. are slot-neutral (0 | 1, no idea what produced a + * mask — see `WebGpuViewerInstance`'s own doc comment); this panel is where + * that meaning ("fast" vs "deep", which Tiled container each comes from) + * actually lives, kept out of the vendored viewer entirely. + * + * Deep is always Tiled-backed (`buildMaskZarrUrl`), pointed at the + * `__masks_deep` container `tiled_mask_sync.write_masks_to_tiled` + * writes (see its `container_suffix` doc) — there is no local equivalent for + * a from-scratch-trained model. Fast defaults to a zero-latency, no-network + * "Live" mode instead: `buildLiveMaskVolume` rasterizes the sample's CURRENT + * shapes (every slice, client-side) into a coarse class-id array and loads + * it via `loadMaskFromArray` — no "Sync masks to Tiled" step required first. + * Fast can still be pointed at the Tiled-backed `__masks` container + * (the precise result, built by the real backend rasterizer) via its own + * toggle. + */ +import { useEffect, useRef, useState } from 'react'; +import { Eye, EyeSlash, CircleNotch } from '@phosphor-icons/react'; +import { buildMaskZarrUrl } from '@/lib/zarrUrl'; +import { buildLiveMaskVolume } from '@/lib/volumeMaskPreview'; +import type { Shape } from '@/stores/annotationStore'; +import type { WebGpuViewerInstance } from './VolumeViewer'; + +type MaskClasses = NonNullable>; + +interface SlotConfig { + slot: 0 | 1; + label: string; + suffix: '' | '_deep'; + /** Only the Fast slot has a no-Tiled-round-trip option — Deep is always a + * from-scratch-trained model's saved output, nothing to rasterize live. */ + canLive: boolean; +} + +const SLOTS: SlotConfig[] = [ + { slot: 0, label: 'Fast (iPred)', suffix: '', canLive: true }, + { slot: 1, label: 'Deep (dlsia)', suffix: '_deep', canLive: false }, +]; + +type SourceMode = 'live' | 'tiled'; + +interface SlotState { + loading: boolean; + error: string | null; + loaded: boolean; + enabled: boolean; + opacity: number; + classes: MaskClasses; + /** Ignored for slots where `canLive` is false. */ + mode: SourceMode; + /** Set once the loading poll (see `waitForClasses`) has been running a + * while — a "hasn't shown up yet" hint, not an error, since a large Tiled + * mask (e.g. a 690-slice dlsia result) can legitimately take a while over + * the network. */ + slowLoad: boolean; + /** True only when `error` came from `waitForClasses` timing out — the + * underlying (fire-and-forget) load may still finish moments later, so + * this is the one error case worth offering a re-check for. Other errors + * (no source open, nothing annotated yet) won't change by re-polling. */ + canRetry: boolean; +} + +const initialSlotState: SlotState = { + loading: false, + error: null, + loaded: false, + enabled: true, + opacity: 0.6, + classes: [], + mode: 'live', + slowLoad: false, + canRetry: false, +}; + +interface MaskLayersPanelProps { + instance: WebGpuViewerInstance | null; + kind: string | null; + source: string | null; + serverUri: string | null; + /** Slot to load automatically once the viewer is ready — the "View in 3D" + * hand-off from Train/Annotate arrives here already knowing which result + * the user just produced, so it shouldn't need a second manual click. */ + autoLoadSlot?: 0 | 1; + /** Current sample's shapes (every slice) for the Fast slot's "Live" mode — + * `byImage[sourceKey]` from `annotationStore`, undefined if none open. */ + liveShapes?: Record; + imageWidth?: number; + imageHeight?: number; + nSlices?: number; +} + +export default function MaskLayersPanel({ + instance, kind, source, serverUri, autoLoadSlot, liveShapes, imageWidth, imageHeight, nSlices, +}: MaskLayersPanelProps) { + const [state, setState] = useState>({ + 0: initialSlotState, + 1: initialSlotState, + }); + // Guards the auto-load effect against StrictMode's double-invoke and + // against re-firing on every re-render once it's already kicked off once + // for this instance. + const autoLoadedFor = useRef(null); + + // A new dataset invalidates every previously-loaded mask — the viewer + // itself remounts on source change (see VolumePage's `key={url}`), so + // there is nothing to explicitly unload here, only local status to reset. + useEffect(() => { + setState({ 0: initialSlotState, 1: initialSlotState }); + autoLoadedFor.current = null; + }, [source, serverUri]); + + useEffect(() => { + if (!instance || autoLoadSlot === undefined || autoLoadedFor.current === instance) return; + autoLoadedFor.current = instance; + const cfg = SLOTS.find((s) => s.slot === autoLoadSlot); + if (!cfg) return; + // A "View in 3D" hand-off means the user just pushed fresh data to + // Tiled (or is coming from a saved dlsia run) — always the Tiled-backed + // result here, regardless of whichever mode the Fast slot's toggle was + // last left on. + setState((s) => ({ ...s, [cfg.slot]: { ...s[cfg.slot], mode: 'tiled' } })); + void load(cfg, 'tiled'); + // eslint-disable-next-line react-hooks/exhaustive-deps -- `load` is redefined every render but stable in effect: it only reads current instance/kind/source/serverUri via closure, matching the effect's own deps. + }, [instance, autoLoadSlot]); + + if (!instance) return null; + + /** Shared tail for both load paths: poll for classes, softening to a + * "still loading" hint rather than a bare spinner once it's taking a + * while, and update state on both success and eventual (rare) failure. */ + const pollForClasses = async (cfg: SlotConfig, noClassesMessage: string | null) => { + setState((s) => ({ ...s, [cfg.slot]: { ...s[cfg.slot], loading: true, slowLoad: false, canRetry: false } })); + try { + const classes = await waitForClasses(instance, cfg.slot, { + onSlow: () => setState((s) => ({ ...s, [cfg.slot]: { ...s[cfg.slot], slowLoad: true } })), + }); + setState((s) => ({ + ...s, + [cfg.slot]: { + ...s[cfg.slot], loading: false, loaded: true, classes, slowLoad: false, canRetry: false, + error: classes.length === 0 ? noClassesMessage : null, + }, + })); + } catch (err) { + setState((s) => ({ + ...s, + [cfg.slot]: { + ...s[cfg.slot], loading: false, slowLoad: false, canRetry: true, + error: err instanceof Error ? err.message : String(err), + }, + })); + } + }; + + /** Re-poll without re-triggering the underlying load — for when + * `waitForClasses`'s own (generous) budget ran out, but the fire-and-forget + * load may well have finished moments later in the background regardless. */ + const checkAgain = (cfg: SlotConfig) => void pollForClasses(cfg, null); + + const load = async (cfg: SlotConfig, modeOverride?: SourceMode) => { + const mode = modeOverride ?? state[cfg.slot].mode; + + if (cfg.canLive && mode === 'live') { + const volume = imageWidth && imageHeight && nSlices + ? buildLiveMaskVolume(liveShapes ?? {}, imageWidth, imageHeight, nSlices) + : null; + if (!volume) { + setState((s) => ({ + ...s, + [cfg.slot]: { ...s[cfg.slot], error: 'Nothing annotated for this sample yet.', loading: false, canRetry: false }, + })); + return; + } + instance.loadMaskFromArray(cfg.slot, volume.data, volume.dims); + await pollForClasses(cfg, null); + return; + } + + const { url, reason } = buildMaskZarrUrl(kind, source, serverUri, cfg.suffix); + if (!url) { + setState((s) => ({ + ...s, + [cfg.slot]: { ...s[cfg.slot], error: reason ?? 'No source open', loading: false, canRetry: false }, + })); + return; + } + instance.loadMask(cfg.slot, url); + await pollForClasses( + cfg, + 'Loaded, but no classes found (every voxel is background, or nothing has been pushed to Tiled for this dataset yet).', + ); + }; + + const remove = (cfg: SlotConfig) => { + instance.removeMask(cfg.slot); + setState((s) => ({ ...s, [cfg.slot]: initialSlotState })); + }; + + const setOpacity = (cfg: SlotConfig, opacity: number) => { + setState((s) => ({ ...s, [cfg.slot]: { ...s[cfg.slot], opacity } })); + for (const cls of state[cfg.slot].classes) { + instance.setMaskClassOpacity(cfg.slot, cls.id, opacity); + } + }; + + const setMode = (cfg: SlotConfig, mode: SourceMode) => { + setState((s) => ({ ...s, [cfg.slot]: { ...s[cfg.slot], mode, error: null } })); + }; + + const toggleClass = (cfg: SlotConfig, classId: number) => { + instance.toggleMaskClassVisible(cfg.slot, classId); + setState((s) => ({ + ...s, + [cfg.slot]: { + ...s[cfg.slot], + classes: s[cfg.slot].classes.map((c) => (c.id === classId ? { ...c, visible: !c.visible } : c)), + }, + })); + }; + + return ( +
+ {SLOTS.map((cfg) => { + const s = state[cfg.slot]; + return ( +
+
+ {cfg.label} + {s.loaded ? ( + + ) : ( + + )} +
+ + {cfg.canLive && !s.loaded && ( +
+ + +
+ )} + + {s.error && ( +
+

{s.error}

+ {s.canRetry && ( + + )} +
+ )} + + {s.loaded && s.classes.length > 0 && ( +
+ +
+ {s.classes.map((cls) => ( + + ))} +
+
+ )} +
+ ); + })} +
+ ); +} + +/** + * `loadMask` mutates viewer-internal state asynchronously with no returned + * promise and no error/progress signal on the public interface (see the + * upstream brief's Step 3 — it's fire-and-forget, matching how the HUD's own + * click handler calls it, and `getMaskClasses` is the only observable outcome: + * `undefined` covers both "still loading" and "failed" indistinguishably). + * Poll rather than assuming synchronous completion — mirrors how + * `InferencePanel`/`useExportJob` poll a job status elsewhere in this app. + * + * Budget is generous (5 minutes) rather than the original 5 seconds: a real + * Tiled-backed "Deep" mask (e.g. a 690-slice dlsia result, freshly written) + * can legitimately take much longer than a few seconds to fetch over the + * network, and a short timeout here doesn't stop the load — it keeps running + * in the vendored viewer regardless — it just makes OUR panel declare defeat + * (and hide the opacity/visibility controls) while the mask goes on to load + * and render successfully moments later, exactly the bug this widening fixes. + * `onSlow` fires once after `slowAfterMs` so the caller can soften the UI + * from a bare spinner into an explicit "still loading" hint, since 5 minutes + * of unexplained silence would look broken even though it's still working. + */ +async function waitForClasses( + instance: WebGpuViewerInstance, + slot: 0 | 1, + { + attempts = 300, + intervalMs = 1000, + slowAfterMs = 5000, + onSlow, + }: { attempts?: number; intervalMs?: number; slowAfterMs?: number; onSlow?: () => void } = {}, +): Promise { + let slowFired = false; + for (let i = 0; i < attempts; i++) { + const classes = instance.getMaskClasses(slot); + if (classes !== undefined) return classes; + if (!slowFired && i * intervalMs >= slowAfterMs) { + slowFired = true; + onSlow?.(); + } + await new Promise((r) => setTimeout(r, intervalMs)); + } + throw new Error( + 'Still waiting for the mask to load after 5 minutes — it may finish shortly on its own ' + + '(click "Check again" below), or check the browser console for a renderer-side error.', + ); +} diff --git a/frontend/src/components/volume/RebuildVolumeControl.tsx b/frontend/src/components/volume/RebuildVolumeControl.tsx new file mode 100644 index 0000000..1337cfc --- /dev/null +++ b/frontend/src/components/volume/RebuildVolumeControl.tsx @@ -0,0 +1,101 @@ +/** + * RebuildVolumeControl — re-trigger "Build 3D volume" for a dataset that + * already has one, so a fidelity-affecting change (e.g. raising + * `tiff_stack_source.TARGET_DIM`) can actually reach volumes built before + * the change. `volume_build.build_volume` already replaces any prior build + * for the same key safely (see its own "re-running must refresh the volume" + * comment) — this only adds the UI path to trigger that for an EXISTING + * volume; `BuildVolumePanel` covers the "none exists yet" case. + * + * Deliberately no staleness detection (comparing what's registered against + * what a fresh build would produce) — a plain always-available button is + * simpler and correct for what is, for now, a one-time fidelity bump. + */ +import { useCallback, useState } from 'react'; +import { ArrowsClockwise, CircleNotch, Warning } from '@phosphor-icons/react'; +import { API_BASE } from '@/config'; + +interface JobState { + state: 'pending' | 'running' | 'done' | 'error'; + phase: string; + done: number; + total: number; + error: string | null; +} + +async function readError(res: Response): Promise { + try { + const body = await res.json(); + if (typeof body?.detail === 'string') return body.detail; + return JSON.stringify(body); + } catch { + return `Request failed (${res.status})`; + } +} + +interface RebuildVolumeControlProps { + source: string; + serverUri: string | null; + /** Called once the rebuild finishes — the caller should bump whatever key + * forces VolumeViewer to remount, since the resolved path doesn't change + * on a rebuild (same key, same location) the way a brand-new build does. */ + onRebuilt: () => void; +} + +export default function RebuildVolumeControl({ source, serverUri, onRebuilt }: RebuildVolumeControlProps) { + const [job, setJob] = useState(null); + const [error, setError] = useState(null); + + const rebuild = useCallback(async () => { + setError(null); + setJob({ state: 'pending', phase: 'Starting', done: 0, total: 1, error: null }); + try { + const res = await fetch(`${API_BASE}/api/volume/build`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ source, kind: 'tiled', server_uri: serverUri }), + }); + if (!res.ok) throw new Error(await readError(res)); + const { job_id: jobId } = await res.json(); + + for (;;) { + await new Promise((resolve) => setTimeout(resolve, 1000)); + const statusRes = await fetch(`${API_BASE}/api/export/status/${jobId}`); + if (!statusRes.ok) throw new Error(await readError(statusRes)); + const status: JobState = await statusRes.json(); + setJob(status); + if (status.state === 'error') throw new Error(status.error || 'Rebuild failed'); + if (status.state === 'done') break; + } + setJob(null); + onRebuilt(); + } catch (err) { + setError(err instanceof Error ? err.message : String(err)); + setJob(null); + } + }, [source, serverUri, onRebuilt]); + + const running = job !== null; + const percent = job && job.total > 0 ? Math.round((job.done / job.total) * 100) : 0; + + return ( +
+ + {error && ( +

+ + {error} +

+ )} +
+ ); +} diff --git a/frontend/src/components/volume/VolumeViewer.tsx b/frontend/src/components/volume/VolumeViewer.tsx new file mode 100644 index 0000000..38ed5ff --- /dev/null +++ b/frontend/src/components/volume/VolumeViewer.tsx @@ -0,0 +1,142 @@ +/** + * VolumeViewer — React wrapper around the vendored WebGPU volume renderer. + * + * The renderer (pinned as a git submodule, see `.gitmodules` and + * `tsconfig.app.json`'s `@zarrviewer/*` alias) is not a component: it is an + * imperative `run(canvas, { zarrUrl, hudMount })` that boots a volume view into + * a canvas and returns a disposable handle. This wrapper owns only the React + * lifecycle around it — create the hosts, boot, dispose. Everything about *how* + * the volume looks belongs upstream. + * + * Adapted from the viewer repo's own `src/WebGpuNative.tsx`, which had already + * solved the mount/dispose shape; kept here rather than imported because it is + * host-app glue, not renderer code. + * + * Auth: chunks are fetched with a plain `fetch()` straight to the Tiled origin. + * That works because anonymous access is read-only-enabled there — see + * `lib/zarrUrl.ts`. Deliberately no token interceptor: the Tiled API key is + * write-scoped and stays server-side. + */ +import { useEffect, useRef } from 'react'; +import { run, type WebGpuViewerInstance } from '@zarrviewer/ome-zarr-viewer'; + +export type { WebGpuViewerInstance }; + +/** Width of the docked HUD column, in px. Mirrored into the renderer's stage. */ +const HUD_WIDTH = 320; + +/** + * Whether the renderer can run here, with a reason when it cannot. + * + * WebGPU is gated on a *secure context*, so over plain http `navigator.gpu` is + * simply undefined. Distinguishing that from "this browser has no WebGPU" + * matters: the first is a deployment fix, the second is not, and collapsing + * them sends people looking in the wrong place. + */ +export function webGpuAvailability(): { ok: boolean; reason: string } { + if (typeof navigator !== 'undefined' && (navigator as Navigator & { gpu?: unknown }).gpu) { + return { ok: true, reason: '' }; + } + if (typeof window !== 'undefined' && !window.isSecureContext) { + return { + ok: false, + reason: + 'WebGPU needs a secure context. Open the app over https, or via localhost / 127.0.0.1 rather than a LAN address.', + }; + } + return { + ok: false, + reason: 'This browser does not support WebGPU. Chrome or Edge 113+ is required for the 3D view.', + }; +} + +interface VolumeViewerProps { + /** Zarr store root, from `buildZarrUrl`. */ + zarrUrl: string; + /** Receives the handle when the renderer boots, and `null` when it is torn down. */ + onReady?: (instance: WebGpuViewerInstance | null) => void; + /** Receives a boot failure (bad store, unsupported codec, no multiscales). */ + onError?: (error: unknown) => void; +} + +export default function VolumeViewer({ zarrUrl, onReady, onError }: VolumeViewerProps) { + const containerRef = useRef(null); + // Held in refs so a changed callback identity never restarts the renderer — + // rebooting a WebGPU device on every parent render would be ruinous. + const onReadyRef = useRef(onReady); + const onErrorRef = useRef(onError); + onReadyRef.current = onReady; + onErrorRef.current = onError; + + useEffect(() => { + const container = containerRef.current; + if (!container || !zarrUrl || !webGpuAvailability().ok) return; + + let cancelled = false; + let handle: WebGpuViewerInstance | null = null; + + // The renderer treats `canvas.parentElement` as its stage and re-fits the + // backing store to that box every frame, so a canvas that shares the row + // with the HUD sidebar sizes itself correctly with no resize plumbing here. + const canvasHost = document.createElement('div'); + canvasHost.style.flex = '1'; + canvasHost.style.minWidth = '0'; + canvasHost.style.position = 'relative'; + + const canvas = document.createElement('canvas'); + canvas.style.width = '100%'; + canvas.style.height = '100%'; + canvas.style.display = 'block'; + canvasHost.appendChild(canvas); + + const hudHost = document.createElement('div'); + hudHost.style.width = `${HUD_WIDTH}px`; + hudHost.style.flexShrink = '0'; + hudHost.style.overflow = 'auto'; + + container.appendChild(canvasHost); + container.appendChild(hudHost); + + run(canvas, { zarrUrl, hudMount: hudHost }) + .then((created) => { + // StrictMode double-invokes effects in dev, so a run can resolve after + // its own cleanup. Disposing here is what stops GPU devices piling up. + if (cancelled) { + created.dispose(); + return; + } + handle = created; + onReadyRef.current?.(created); + }) + .catch((error: unknown) => { + if (cancelled) return; + console.error('VolumeViewer: renderer failed to start:', error); + onErrorRef.current?.(error); + }); + + return () => { + cancelled = true; + try { + handle?.dispose(); + } catch (error) { + console.warn('VolumeViewer: dispose failed:', error); + } + if (handle) onReadyRef.current?.(null); + container.replaceChildren(); + }; + }, [zarrUrl]); + + return ( +
+ ); +} diff --git a/frontend/src/config.ts b/frontend/src/config.ts index 71703a4..2471b6c 100644 --- a/frontend/src/config.ts +++ b/frontend/src/config.ts @@ -1,11 +1,16 @@ /** * Base URL for the backend API. - * - Default is '' (same origin): dev uses Vite's `/api` proxy; the production - * container serves the SPA from FastAPI itself, so `/api` is same-origin too. + * - Default is derived from Vite's `BASE_URL` (itself driven by `VITE_BASE_PATH` + * at build time, see vite.config.ts): '/' when root-hosted (dev's `/api` + * proxy, or the production container serving the SPA same-origin), or the + * trimmed subpath (e.g. `/bl832/seg_studio`) when hosted behind a + * path-stripping reverse proxy, so `fetch(`${API_BASE}/api/...`)` still + * resolves to a path the proxy actually routes. * - Set `VITE_API_BASE` at build time only for split deployments where the API * lives on a different origin (e.g. `https://api.example.com`). */ -export const API_BASE = import.meta.env.VITE_API_BASE?.trim() || ''; +export const API_BASE = + import.meta.env.VITE_API_BASE?.trim() || import.meta.env.BASE_URL.replace(/\/$/, ''); /** * URL of the user documentation site (MkDocs). diff --git a/frontend/src/hooks/buildSliceUrl.denoise.test.ts b/frontend/src/hooks/buildSliceUrl.denoise.test.ts new file mode 100644 index 0000000..6d4d8fd --- /dev/null +++ b/frontend/src/hooks/buildSliceUrl.denoise.test.ts @@ -0,0 +1,78 @@ +import { describe, it, expect } from 'vitest'; +import { buildSliceUrl } from './useImageSlice'; +import type { RenderOpts, DenoiseOpts } from '@/stores/datasetStore'; + +const RENDER: RenderOpts = { + norm: 'global', + scale: 'linear', + vminPct: 1, + vmaxPct: 99, + cmap: 'gray', +}; + +const params = (url: string) => new URL(url, 'http://x').searchParams; + +describe('buildSliceUrl denoise parameters', () => { + it('omits denoise entirely when it is off', () => { + // The un-denoised request must stay byte-identical to what it has always + // been, so it keeps hitting the same backend cache entry. + const off: DenoiseOpts = { method: 'none', strength: 0.7 }; + const withOff = buildSliceUrl('s', 'tiled', 3, RENDER, null, off); + const without = buildSliceUrl('s', 'tiled', 3, RENDER, null); + expect(withOff).toBe(without); + expect(params(withOff).has('denoise_method')).toBe(false); + }); + + it('sends method and strength when active', () => { + const p = params(buildSliceUrl('s', 'tiled', 3, RENDER, null, { method: 'tv', strength: 0.4 })); + expect(p.get('denoise_method')).toBe('tv'); + expect(p.get('denoise_strength')).toBe('0.4'); + }); + + it('sends a crop only alongside an active method', () => { + // Cropping is what keeps slider-dragging interactive: filtering a full + // 2560² slice costs seconds, a 512 crop costs ~0.3s. + const cropped = params( + buildSliceUrl('s', 'tiled', 3, RENDER, null, { method: 'nlm', strength: 0.5 }, 512), + ); + expect(cropped.get('denoise_crop')).toBe('512'); + + const offWithCrop = params( + buildSliceUrl('s', 'tiled', 3, RENDER, null, { method: 'none', strength: 0.5 }, 512), + ); + expect(offWithCrop.has('denoise_crop')).toBe(false); + }); + + it('ignores a zero or negative crop', () => { + const p = params( + buildSliceUrl('s', 'tiled', 3, RENDER, null, { method: 'tv', strength: 0.5 }, 0), + ); + expect(p.has('denoise_crop')).toBe(false); + }); + + it('still carries the render options', () => { + // Denoise is additive: it must not displace normalization, which is applied + // after it on the backend. + const p = params(buildSliceUrl('s', 'tiled', 3, RENDER, null, { method: 'tv', strength: 0.5 })); + expect(p.get('norm')).toBe('global'); + expect(p.get('vmin_pct')).toBe('1'); + expect(p.get('cmap')).toBe('gray'); + }); + + it('keeps the server uri', () => { + const p = params( + buildSliceUrl('s', 'tiled', 3, RENDER, 'http://127.0.0.1:8010', { method: 'tv', strength: 0.5 }), + ); + expect(p.get('server_uri')).toBe('http://127.0.0.1:8010'); + }); +}); + +describe('denoise is not a render option', () => { + it('is absent from RenderOpts, so it cannot reach export payloads', () => { + // Export payloads are built from RenderOpts. If denoise lived there, a + // preview would silently change exported pixels — the bake exists precisely + // so that turning a denoised view into data is an explicit act. + expect(Object.keys(RENDER)).not.toContain('denoise'); + expect(Object.keys(RENDER)).toEqual(['norm', 'scale', 'vminPct', 'vmaxPct', 'cmap']); + }); +}); diff --git a/frontend/src/hooks/editHistory.test.ts b/frontend/src/hooks/editHistory.test.ts new file mode 100644 index 0000000..f82fb08 --- /dev/null +++ b/frontend/src/hooks/editHistory.test.ts @@ -0,0 +1,152 @@ +import { afterEach, beforeEach, describe, expect, it } from 'vitest'; +import { clearHistory, markClassDelete, redo, undo } from './editHistory'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { useClassStore, type AnnotationClass } from '@/stores/classStore'; +import { useDatasetStore } from '@/stores/datasetStore'; + +const SOURCE_KEY = 'local:x.tif'; + +function openSample() { + useDatasetStore.getState().setDataset('local', 'x.tif', null, { + nSlices: 3, + height: 10, + width: 10, + dtype: 'uint8', + isRgb: false, + valueRange: [0, 255], + }); +} + +beforeEach(() => { + useAnnotationStore.getState().reset(); + useAnnotationStore.temporal.getState().clear(); + useClassStore.setState({ classes: [] }); + useDatasetStore.getState().reset(); +}); + +afterEach(() => { + clearHistory(); +}); + +describe('editHistory', () => { + it('undo()/redo() are no-ops when there is no history', () => { + expect(() => undo()).not.toThrow(); + expect(() => redo()).not.toThrow(); + expect(useAnnotationStore.temporal.getState().pastStates).toHaveLength(0); + }); + + it('undo reverts a plain region edit (not tagged as a class delete)', () => { + openSample(); + useAnnotationStore.getState().addShape(SOURCE_KEY, 0, { + id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2, + }); + expect(useAnnotationStore.getState().byImage[SOURCE_KEY]?.['0']).toHaveLength(1); + + undo(); + expect(useAnnotationStore.getState().byImage[SOURCE_KEY]?.['0'] ?? []).toHaveLength(0); + + redo(); + expect(useAnnotationStore.getState().byImage[SOURCE_KEY]?.['0']).toHaveLength(1); + }); + + it('markClassDelete + undo re-inserts the class at its original index when sourceKey matches', () => { + openSample(); + const cls: AnnotationClass = { classId: 5, label: 'Pore', color: '#f00', isVisible: true }; + useClassStore.setState({ classes: [cls] }); + + markClassDelete(SOURCE_KEY, cls, 0); + useClassStore.getState().deleteClass(5); + // The class-delete journal only records on the NEXT tracked annotationStore edit + // (onTrackedEdit fires from annotationStore, not classStore), so we need one. + useAnnotationStore.getState().touchHistory(); + + expect(useClassStore.getState().classes).toHaveLength(0); + + undo(); + expect(useClassStore.getState().classes).toEqual([cls]); + + redo(); + expect(useClassStore.getState().classes).toHaveLength(0); + }); + + it('does not replay a class-delete entry when the current sourceKey differs', () => { + openSample(); + const cls: AnnotationClass = { classId: 5, label: 'Pore', color: '#f00', isVisible: true }; + useClassStore.setState({ classes: [cls] }); + + markClassDelete(SOURCE_KEY, cls, 0); + useClassStore.getState().deleteClass(5); + useAnnotationStore.getState().touchHistory(); + expect(useClassStore.getState().classes).toHaveLength(0); + + // Switch to a different sample before undoing. + useDatasetStore.getState().setDataset('local', 'other.tif', null, { + nSlices: 1, height: 1, width: 1, dtype: 'uint8', isRgb: false, valueRange: [0, 255], + }); + + undo(); + // Region history still undoes (temporal stack is shared), but the class is + // NOT re-inserted because currentSourceKey() no longer matches the entry. + expect(useClassStore.getState().classes).toHaveLength(0); + }); + + it('redo/undo stacks stay ordered across a mix of plain and class-delete edits', () => { + openSample(); + const cls: AnnotationClass = { classId: 5, label: 'Pore', color: '#f00', isVisible: true }; + useClassStore.setState({ classes: [cls] }); + + // 1: plain region edit + useAnnotationStore.getState().addShape(SOURCE_KEY, 0, { + id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 1, h: 1, + }); + // 2: class delete + markClassDelete(SOURCE_KEY, cls, 0); + useClassStore.getState().deleteClass(5); + useAnnotationStore.getState().touchHistory(); + + // Undo class delete first (LIFO) + undo(); + expect(useClassStore.getState().classes).toEqual([cls]); + expect(useAnnotationStore.getState().byImage[SOURCE_KEY]?.['0']).toHaveLength(1); + + // Undo plain region edit + undo(); + expect(useAnnotationStore.getState().byImage[SOURCE_KEY]?.['0'] ?? []).toHaveLength(0); + + // Redo both back in order + redo(); + expect(useAnnotationStore.getState().byImage[SOURCE_KEY]?.['0']).toHaveLength(1); + redo(); + expect(useClassStore.getState().classes).toHaveLength(0); + }); + + it('a new edit after undo clears the redo stack', () => { + openSample(); + useAnnotationStore.getState().addShape(SOURCE_KEY, 0, { + id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 1, h: 1, + }); + undo(); + // A fresh edit (not via redo) should drop the stale redo entry. + useAnnotationStore.getState().addShape(SOURCE_KEY, 0, { + id: 's2', classId: 1, kind: 'rectangle', x: 1, y: 1, w: 1, h: 1, + }); + redo(); // should be a no-op now (no future states) + expect(useAnnotationStore.getState().byImage[SOURCE_KEY]?.['0']).toHaveLength(1); + expect(useAnnotationStore.getState().byImage[SOURCE_KEY]?.['0']?.[0].id).toBe('s2'); + }); + + it('clearHistory wipes the draft, journal, and temporal stack', () => { + openSample(); + useAnnotationStore.getState().addShape(SOURCE_KEY, 0, { + id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 1, h: 1, + }); + expect(useAnnotationStore.temporal.getState().pastStates.length).toBeGreaterThan(0); + + clearHistory(); + expect(useAnnotationStore.temporal.getState().pastStates).toHaveLength(0); + expect(useAnnotationStore.temporal.getState().futureStates).toHaveLength(0); + // undo/redo are now no-ops (journal cleared alongside temporal stack) + expect(() => undo()).not.toThrow(); + expect(() => redo()).not.toThrow(); + }); +}); diff --git a/frontend/src/hooks/useAnnotatedSourceKeys.test.tsx b/frontend/src/hooks/useAnnotatedSourceKeys.test.tsx new file mode 100644 index 0000000..389c12d --- /dev/null +++ b/frontend/src/hooks/useAnnotatedSourceKeys.test.tsx @@ -0,0 +1,57 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, renderHook, waitFor } from '@testing-library/react'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import type { ReactNode } from 'react'; +import { useAnnotatedSourceKeys } from './useAnnotatedSourceKeys'; +import { useAnnotationStore } from '@/stores/annotationStore'; + +function wrapper({ children }: { children: ReactNode }) { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return {children}; +} + +beforeEach(() => { + useAnnotationStore.getState().reset(); + vi.stubGlobal('fetch', vi.fn().mockResolvedValue({ ok: true, json: async () => [] })); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('useAnnotatedSourceKeys', () => { + it('includes in-session sources with at least one shape', () => { + useAnnotationStore.getState().replaceClassShapesOnSlice('local:a.tif', 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + const { result } = renderHook(() => useAnnotatedSourceKeys(), { wrapper }); + expect(result.current.has('local:a.tif')).toBe(true); + }); + + it('excludes a source whose slices are all empty', () => { + useAnnotationStore.getState().replaceClassShapesOnSlice('local:b.tif', 0, 1, []); + const { result } = renderHook(() => useAnnotatedSourceKeys(), { wrapper }); + expect(result.current.has('local:b.tif')).toBe(false); + }); + + it('merges in annotated sources from the backend drafts endpoint', async () => { + (fetch as any).mockResolvedValue({ + ok: true, + json: async () => [ + { source_key: 'local:c.tif', has_annotations: true }, + { source_key: 'local:d.tif', has_annotations: false }, + ], + }); + const { result } = renderHook(() => useAnnotatedSourceKeys(), { wrapper }); + await waitFor(() => expect(result.current.has('local:c.tif')).toBe(true)); + expect(result.current.has('local:d.tif')).toBe(false); + }); + + it('returns an empty set when the drafts fetch fails', async () => { + (fetch as any).mockResolvedValue({ ok: false, status: 500 }); + const { result } = renderHook(() => useAnnotatedSourceKeys(), { wrapper }); + await waitFor(() => expect(fetch).toHaveBeenCalled()); + expect(result.current.size).toBe(0); + }); +}); diff --git a/frontend/src/hooks/useConnectionHealth.test.ts b/frontend/src/hooks/useConnectionHealth.test.ts new file mode 100644 index 0000000..4f42d23 --- /dev/null +++ b/frontend/src/hooks/useConnectionHealth.test.ts @@ -0,0 +1,78 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook, waitFor } from '@testing-library/react'; +import { useConnectionHealth } from './useConnectionHealth'; +import { useConnectionStore } from '@/stores/connectionStore'; + +beforeEach(() => { + useConnectionStore.setState({ + kind: null, + serverUri: null, + browseContainerPath: null, + browseFocusPath: null, + localRoot: null, + localRel: null, + label: null, + sampleCount: null, + status: 'unknown', + }); + vi.stubGlobal('fetch', vi.fn()); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); + vi.useRealTimers(); +}); + +describe('useConnectionHealth', () => { + it('stays unknown and never fetches for a local connection', async () => { + useConnectionStore.setState({ kind: 'local' }); + renderHook(() => useConnectionHealth()); + await Promise.resolve(); + expect(fetch).not.toHaveBeenCalled(); + expect(useConnectionStore.getState().status).toBe('unknown'); + }); + + it('stays unknown when there is no connection at all', async () => { + renderHook(() => useConnectionHealth()); + await Promise.resolve(); + expect(fetch).not.toHaveBeenCalled(); + expect(useConnectionStore.getState().status).toBe('unknown'); + }); + + it('sets status to ok on a successful tiled health check', async () => { + (fetch as any).mockResolvedValue({ ok: true, json: async () => [] }); + useConnectionStore.setState({ kind: 'tiled', serverUri: 'http://tiled.example' }); + renderHook(() => useConnectionHealth()); + await waitFor(() => expect(useConnectionStore.getState().status).toBe('ok')); + expect(fetch).toHaveBeenCalledWith(expect.stringContaining('/api/tiled/list?')); + expect(fetch).toHaveBeenCalledWith(expect.stringContaining('server_uri=http')); + }); + + it('sets status to error on a failed response', async () => { + (fetch as any).mockResolvedValue({ ok: false }); + useConnectionStore.setState({ kind: 'tiled', serverUri: 'http://tiled.example' }); + renderHook(() => useConnectionHealth()); + await waitFor(() => expect(useConnectionStore.getState().status).toBe('error')); + }); + + it('sets status to error when fetch throws', async () => { + (fetch as any).mockRejectedValue(new Error('network down')); + useConnectionStore.setState({ kind: 'tiled', serverUri: 'http://tiled.example' }); + renderHook(() => useConnectionHealth()); + await waitFor(() => expect(useConnectionStore.getState().status).toBe('error')); + }); + + it('resets to unknown when the connection is cleared', async () => { + (fetch as any).mockResolvedValue({ ok: true, json: async () => [] }); + useConnectionStore.setState({ kind: 'tiled', serverUri: 'http://tiled.example' }); + const { rerender } = renderHook(() => useConnectionHealth()); + await waitFor(() => expect(useConnectionStore.getState().status).toBe('ok')); + + act(() => { + useConnectionStore.setState({ kind: null, serverUri: null }); + rerender(); + }); + await waitFor(() => expect(useConnectionStore.getState().status).toBe('unknown')); + }); +}); diff --git a/frontend/src/hooks/useConnectionHealth.ts b/frontend/src/hooks/useConnectionHealth.ts new file mode 100644 index 0000000..dca220d --- /dev/null +++ b/frontend/src/hooks/useConnectionHealth.ts @@ -0,0 +1,62 @@ +/** + * useConnectionHealth — periodic Tiled reachability check, backed by + * connectionStore's `status` field. Mounted once in the app shell so + * connection health is visible from anywhere (HubHeader), not just Browse's + * own local facets poll or Connect's initial probe. + * + * Local connections have no network dependency to check, so status simply + * stays 'unknown' (hidden in the UI) for kind === 'local' or no connection. + */ +import { useEffect, useRef } from 'react'; +import { API_BASE } from '@/config'; +import { useConnectionStore } from '@/stores/connectionStore'; + +const BASE_INTERVAL_MS = 20_000; +const MAX_INTERVAL_MS = 120_000; + +export function useConnectionHealth(): void { + const kind = useConnectionStore((s) => s.kind); + const serverUri = useConnectionStore((s) => s.serverUri); + const setStatus = useConnectionStore((s) => s.setStatus); + const intervalRef = useRef(BASE_INTERVAL_MS); + + useEffect(() => { + if (kind !== 'tiled') { + setStatus('unknown'); + return; + } + + let cancelled = false; + let timer: ReturnType | null = null; + intervalRef.current = BASE_INTERVAL_MS; + + const check = async () => { + try { + const params = new URLSearchParams({ path: '' }); + if (serverUri) params.set('server_uri', serverUri); + const res = await fetch(`${API_BASE}/api/tiled/list?${params}`); + if (cancelled) return; + if (res.ok) { + setStatus('ok'); + intervalRef.current = BASE_INTERVAL_MS; + } else { + setStatus('error'); + intervalRef.current = Math.min(intervalRef.current * 2, MAX_INTERVAL_MS); + } + } catch { + if (cancelled) return; + setStatus('error'); + intervalRef.current = Math.min(intervalRef.current * 2, MAX_INTERVAL_MS); + } finally { + if (!cancelled) timer = setTimeout(check, intervalRef.current); + } + }; + + void check(); + + return () => { + cancelled = true; + if (timer) clearTimeout(timer); + }; + }, [kind, serverUri, setStatus]); +} diff --git a/frontend/src/hooks/useDraftSync.test.tsx b/frontend/src/hooks/useDraftSync.test.tsx new file mode 100644 index 0000000..ad306af --- /dev/null +++ b/frontend/src/hooks/useDraftSync.test.tsx @@ -0,0 +1,94 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook } from '@testing-library/react'; +import { loadDraft, useDraftSync } from './useDraftSync'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { useClassStore } from '@/stores/classStore'; + +beforeEach(() => { + useAnnotationStore.getState().reset(); + useClassStore.setState({ classes: [] }); + vi.stubGlobal('fetch', vi.fn().mockResolvedValue({ ok: true })); + vi.stubGlobal('sendBeacon', undefined); + Object.defineProperty(window.navigator, 'sendBeacon', { value: vi.fn(), configurable: true }); +}); + +afterEach(() => { + cleanup(); // unmount while fetch/sendBeacon are still stubbed, before unstubbing them + vi.unstubAllGlobals(); + vi.useRealTimers(); +}); + +describe('loadDraft', () => { + it('returns null on 404', async () => { + (fetch as any).mockResolvedValue({ status: 404, ok: false }); + expect(await loadDraft('local:x.tif')).toBeNull(); + }); + + it('returns null on any error', async () => { + (fetch as any).mockRejectedValue(new Error('down')); + expect(await loadDraft('local:x.tif')).toBeNull(); + }); + + it('returns the parsed draft on success', async () => { + (fetch as any).mockResolvedValue({ status: 200, ok: true, json: async () => ({ payload: { classes: [] } }) }); + expect(await loadDraft('local:x.tif')).toEqual({ payload: { classes: [] } }); + }); +}); + +describe('useDraftSync', () => { + it('does nothing when sourceKey is null', () => { + vi.useFakeTimers(); + renderHook(() => useDraftSync(null)); + vi.advanceTimersByTime(5000); + expect(fetch).not.toHaveBeenCalled(); + }); + + it('autosaves 1.5s after a store change, with the right payload', async () => { + vi.useFakeTimers({ shouldAdvanceTime: true }); + useClassStore.setState({ classes: [{ classId: 1, label: 'Cell', color: '#f00', isVisible: true }] }); + renderHook(() => useDraftSync('local:x.tif')); + + act(() => { + useAnnotationStore.getState().replaceClassShapesOnSlice('local:x.tif', 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + }); + await vi.advanceTimersByTimeAsync(1500); + + expect(fetch).toHaveBeenCalledTimes(1); + const [url, init] = (fetch as any).mock.calls[0]; + expect(url).toContain('/api/annotations/draft?source_key='); + expect(init.method).toBe('PUT'); + const body = JSON.parse(init.body); + expect(body.classes).toEqual([{ classId: 1, label: 'Cell', color: '#f00', isVisible: true }]); + expect(body.slices['0']).toHaveLength(1); + }); + + it('re-arms the debounce on each subsequent change (no save before the quiet period)', async () => { + vi.useFakeTimers({ shouldAdvanceTime: true }); + renderHook(() => useDraftSync('local:x.tif')); + act(() => { useAnnotationStore.getState().replaceClassShapesOnSlice('local:x.tif', 0, 1, []); }); + await vi.advanceTimersByTimeAsync(1000); + act(() => { useAnnotationStore.getState().replaceClassShapesOnSlice('local:x.tif', 1, 1, []); }); + await vi.advanceTimersByTimeAsync(1000); + expect(fetch).not.toHaveBeenCalled(); // only 1s since the last change, need 1.5s + await vi.advanceTimersByTimeAsync(500); + expect(fetch).toHaveBeenCalledTimes(1); + }); + + it('flushes immediately (bypassing the debounce) when the component unmounts', () => { + vi.useFakeTimers(); + const { unmount } = renderHook(() => useDraftSync('local:x.tif')); + unmount(); + expect(fetch).toHaveBeenCalledTimes(1); + }); + + it('sends a beacon on beforeunload with the latest payload', () => { + renderHook(() => useDraftSync('local:x.tif')); + window.dispatchEvent(new Event('beforeunload')); + expect(navigator.sendBeacon).toHaveBeenCalledWith( + expect.stringContaining('/api/annotations/draft?source_key='), + expect.any(String), + ); + }); +}); diff --git a/frontend/src/hooks/useExportJob.test.tsx b/frontend/src/hooks/useExportJob.test.tsx new file mode 100644 index 0000000..a764c35 --- /dev/null +++ b/frontend/src/hooks/useExportJob.test.tsx @@ -0,0 +1,101 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook, waitFor } from '@testing-library/react'; +import { useExportJob } from './useExportJob'; + +beforeEach(() => { + sessionStorage.clear(); + vi.stubGlobal('fetch', vi.fn()); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('useExportJob', () => { + it('starts idle', () => { + const { result } = renderHook(() => useExportJob()); + expect(result.current.state.status).toBe('idle'); + expect(result.current.downloadUrl).toBeNull(); + }); + + it('a synchronous (no job_id) response is treated as immediately done', async () => { + (fetch as any).mockResolvedValue({ ok: true, json: async () => ({ ok: true, zip_available: false }) }); + const { result } = renderHook(() => useExportJob()); + await act(async () => { await result.current.start({}); }); + expect(result.current.state.status).toBe('done'); + expect(result.current.state.result).toEqual({ ok: true, zip_available: false }); + }); + + it('a job_id response starts polling and reflects progress then done', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => ({ job_id: 'j1' }) }) + .mockResolvedValueOnce({ ok: true, json: async () => ({ state: 'running', phase: 'working', done: 1, total: 4 }) }) + .mockResolvedValueOnce({ ok: true, json: async () => ({ state: 'done', phase: 'done', done: 4, total: 4, result: { zip_available: true } }) }); + + const { result } = renderHook(() => useExportJob()); + await act(async () => { await result.current.start({}); }); + expect(result.current.state.jobId).toBe('j1'); + + await waitFor(() => expect(result.current.state.phase).toBe('working')); + await waitFor(() => expect(result.current.state.status).toBe('done'), { timeout: 3000 }); + expect(result.current.downloadUrl).toContain('/api/export/download/j1'); + }); + + it('a non-ok POST response surfaces a formatted error and does not poll', async () => { + (fetch as any).mockResolvedValue({ ok: false, status: 422, text: async () => JSON.stringify({ detail: [{ msg: 'bad', loc: ['body', 'x'] }] }) }); + const { result } = renderHook(() => useExportJob()); + await act(async () => { await result.current.start({}); }); + expect(result.current.state.status).toBe('error'); + expect(result.current.state.error).toBeTruthy(); + expect(fetch).toHaveBeenCalledTimes(1); // no follow-up poll + }); + + it('a network error during start() surfaces as an error state', async () => { + (fetch as any).mockRejectedValue(new Error('offline')); + const { result } = renderHook(() => useExportJob()); + await act(async () => { await result.current.start({}); }); + expect(result.current.state.status).toBe('error'); + expect(result.current.state.error).toContain('offline'); + }); + + it('a 404 while polling a resumed job clears it and returns to idle', async () => { + sessionStorage.setItem('exportJob:key1', 'stale-job'); + (fetch as any).mockResolvedValue({ status: 404, ok: false }); + const { result } = renderHook(() => useExportJob('key1')); + expect(result.current.state.status).toBe('running'); // optimistic placeholder + await waitFor(() => expect(result.current.state.status).toBe('idle')); + expect(sessionStorage.getItem('exportJob:key1')).toBeNull(); + }); + + it('resumes polling a persisted job id on a fresh mount', async () => { + sessionStorage.setItem('exportJob:key2', 'resumed-job'); + (fetch as any).mockResolvedValue({ ok: true, json: async () => ({ state: 'done', phase: 'done', result: {} }) }); + const { result } = renderHook(() => useExportJob('key2')); + await waitFor(() => expect(result.current.state.status).toBe('done')); + expect(fetch).toHaveBeenCalledWith(expect.stringContaining('/api/export/status/resumed-job')); + }); + + it('reset clears state and any persisted job id', async () => { + sessionStorage.setItem('exportJob:key3', 'j1'); + (fetch as any).mockResolvedValue({ ok: true, json: async () => ({ state: 'running', phase: 'x' }) }); + const { result } = renderHook(() => useExportJob('key3')); + await waitFor(() => expect(result.current.state.phase).toBe('x')); + act(() => result.current.reset()); + expect(result.current.state.status).toBe('idle'); + expect(sessionStorage.getItem('exportJob:key3')).toBeNull(); + }); + + it('startMaskSync/startIpredBatchTrain/startIpredBatchApply hit their own routes', async () => { + (fetch as any).mockResolvedValue({ ok: true, json: async () => ({ ok: true }) }); + const { result } = renderHook(() => useExportJob()); + await act(async () => { await result.current.startMaskSync({}); }); + expect(fetch).toHaveBeenLastCalledWith(expect.stringContaining('/api/masks/to-tiled'), expect.anything()); + + await act(async () => { await result.current.startIpredBatchTrain({}); }); + expect(fetch).toHaveBeenLastCalledWith(expect.stringContaining('/api/ipred/batch/train'), expect.anything()); + + await act(async () => { await result.current.startIpredBatchApply({}); }); + expect(fetch).toHaveBeenLastCalledWith(expect.stringContaining('/api/ipred/batch/apply'), expect.anything()); + }); +}); diff --git a/frontend/src/hooks/useExportJob.ts b/frontend/src/hooks/useExportJob.ts index 392f983..d73356b 100644 --- a/frontend/src/hooks/useExportJob.ts +++ b/frontend/src/hooks/useExportJob.ts @@ -1,7 +1,9 @@ /** - * useExportJob — drive a background COCO/mask export and stream its progress. + * useExportJob — drive a background job (COCO/mask export, iPred batch + * train/apply, or the Train tab's train/infer/batch-probe jobs) and stream + * its progress. * - * POSTs to /api/export/coco (which returns a job_id), then polls + * POSTs to a job-starting route (which returns a job_id), then polls * /api/export/status/{job_id} until done/error, exposing phase, done/total, and * a live log for the UI. When finished, `downloadUrl` points at the .zip so the * browser's save dialog can write it to the user's machine. The server-side @@ -9,6 +11,7 @@ */ import { useCallback, useEffect, useRef, useState } from 'react'; import { API_BASE } from '@/config'; +import { formatApiError } from '@/lib/apiError'; export interface ExportJobState { status: 'idle' | 'running' | 'done' | 'error'; @@ -25,12 +28,32 @@ const IDLE: ExportJobState = { status: 'idle', phase: '', done: 0, total: 0, log: [], result: null, error: null, jobId: null, }; +const storageKeyFor = (persistKey: string) => `exportJob:${persistKey}`; + /** - * Drives a background COCO export: returns the live job `state`, a `start` action, + * Drives a background job: returns the live job `state`, a `start` action, * `reset`, and a `downloadUrl` for the result zip when available. + * + * `persistKey`, when given, survives this component unmounting mid-job — e.g. the + * user switches from Train to Browse and back while training runs in the + * background. Without it, the job keeps running server-side (it's independent of + * any client), but the polling loop lived in this hook's local state, so a fresh + * mount used to come back to a blank "idle" progress bar with no way to tell a + * job was still going. The job id is stashed in sessionStorage on start and + * reattached to the same polling loop on the next mount; callers using the same + * `persistKey` for genuinely different jobs (e.g. a different sample) should + * include whatever varies in the key so they don't reconnect to the wrong one. */ -export function useExportJob() { - const [state, setState] = useState(IDLE); +export function useExportJob(persistKey?: string) { + const storageKey = persistKey ? storageKeyFor(persistKey) : null; + + const [state, setState] = useState(() => { + const savedId = storageKey ? sessionStorage.getItem(storageKey) : null; + // Placeholder until the resume effect's first tick fills in the real + // phase/progress/result — avoids a flash of "idle" for a job that, most of + // the time, is still genuinely running. + return savedId ? { ...IDLE, status: 'running', jobId: savedId } : IDLE; + }); const timer = useRef(null); /** Cancel any pending poll timeout. */ @@ -40,13 +63,26 @@ export function useExportJob() { useEffect(() => clearTimer, []); /** Stop polling and return state to idle. */ - const reset = useCallback(() => { clearTimer(); setState(IDLE); }, []); + const reset = useCallback(() => { + clearTimer(); + if (storageKey) sessionStorage.removeItem(storageKey); + setState(IDLE); + }, [storageKey]); /** Recursively poll /api/export/status/{jobId} every 500ms until done/error. */ const poll = useCallback((jobId: string) => { const tick = async () => { try { const r = await fetch(`${API_BASE}/api/export/status/${jobId}`); + if (r.status === 404) { + // Only reachable when resuming a persisted id: the job registry is + // in-memory and TTL-pruned, so a job from long enough ago (or from + // before a server restart) is gone. That's not a failure worth + // showing — there's simply nothing left to resume. + if (storageKey) sessionStorage.removeItem(storageKey); + setState(IDLE); + return; + } if (!r.ok) { setState((s) => ({ ...s, status: 'error', error: `Status ${r.status}` })); return; } const j = await r.json(); const status: ExportJobState['status'] = @@ -63,6 +99,14 @@ export function useExportJob() { timer.current = window.setTimeout(tick, 500); }; tick(); + }, [storageKey]); + + // Reattach the polling loop to a job that was already running when this hook + // last mounted. Mount-only: a job started via run()/startJob() below already + // calls poll() itself and must not be double-polled by this effect too. + useEffect(() => { + if (state.jobId && state.status === 'running') poll(state.jobId); + // eslint-disable-next-line react-hooks/exhaustive-deps }, []); /** POST a job-starting request to `path`, then poll the shared status route. */ @@ -75,19 +119,27 @@ export function useExportJob() { headers: { 'Content-Type': 'application/json' }, body: JSON.stringify(payload), }); - if (!res.ok) { const msg = await res.text(); setState((s) => ({ ...s, status: 'error', error: msg })); return; } + if (!res.ok) { + // Rejected requests come back as FastAPI validation JSON; show the field + // and message rather than the raw body. + const msg = formatApiError(await res.text(), `Request failed (${res.status}).`); + setState((s) => ({ ...s, status: 'error', error: msg })); + return; + } const data = await res.json(); if (data.job_id) { + if (storageKey) sessionStorage.setItem(storageKey, data.job_id); setState((s) => ({ ...s, jobId: data.job_id })); poll(data.job_id); } else { // Synchronous response (e.g. dry-run) — treat as immediately done. + if (storageKey) sessionStorage.removeItem(storageKey); setState((s) => ({ ...s, status: 'done', result: data })); } } catch (e) { setState((s) => ({ ...s, status: 'error', error: String(e) })); } - }, [poll]); + }, [poll, storageKey]); /** Start a real COCO export. Returns once the job is queued (progress via state). */ const start = useCallback((payload: unknown) => run('/api/export/coco', payload), [run]); @@ -95,10 +147,33 @@ export function useExportJob() { /** Write rasterized masks into Tiled (standalone). Shares the status polling. */ const startMaskSync = useCallback((payload: unknown) => run('/api/masks/to-tiled', payload), [run]); + /** Train one iPred model pooling labeled pixels across multiple slices. */ + const startIpredBatchTrain = useCallback( + (payload: unknown) => run('/api/ipred/batch/train', payload), + [run], + ); + + /** Run iPred inference across a set of slices (e.g. the whole volume). */ + const startIpredBatchApply = useCallback( + (payload: unknown) => run('/api/ipred/batch/apply', payload), + [run], + ); + const downloadUrl = state.jobId && (state.result as { zip_available?: boolean } | null)?.zip_available ? `${API_BASE}/api/export/download/${state.jobId}` : null; - return { state, start, startMaskSync, reset, downloadUrl }; + return { + state, + start, + startMaskSync, + startIpredBatchTrain, + startIpredBatchApply, + /** Start a job at an arbitrary path (e.g. /api/train/start, /api/train/infer) + * that returns a job_id and shares this same status-polling machinery. */ + startJob: run, + reset, + downloadUrl, + }; } diff --git a/frontend/src/hooks/useFeatureChannels.test.ts b/frontend/src/hooks/useFeatureChannels.test.ts new file mode 100644 index 0000000..e19bc16 --- /dev/null +++ b/frontend/src/hooks/useFeatureChannels.test.ts @@ -0,0 +1,293 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook, waitFor } from '@testing-library/react'; +import { useFeatureChannels } from './useFeatureChannels'; +import { useIpredStore, DEFAULT_COMPOSITION_ID } from '@/stores/ipredStore'; +import { useConnectionStore } from '@/stores/connectionStore'; + +const INITIAL_IPRED_STATE = useIpredStore.getState(); +const INITIAL_CONN_STATE = useConnectionStore.getState(); + +function modulesResponse(overrides: Record[] = []) { + return { + ok: true, + json: async () => ({ modules: overrides }), + }; +} + +function preprocessResponse(overrides: Record = {}) { + return { + ok: true, + json: async () => ({ + feature_id: 'feat-1', + project_id: 'proj-1', + setup_id: 'setup-1', + slice_index: 0, + n_channels: 3, + width: 64, + height: 64, + labels: ['a', 'b', 'c'], + cache_hit: false, + ...overrides, + }), + }; +} + +beforeEach(() => { + useIpredStore.setState(INITIAL_IPRED_STATE, true); + useConnectionStore.setState(INITIAL_CONN_STATE, true); + vi.stubGlobal('URL', Object.assign(URL, { + createObjectURL: vi.fn(() => 'blob:mock-url'), + revokeObjectURL: vi.fn(), + })); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +const ARGS = { source: 'x.tif', kind: 'local', sliceIndex: 0, serverUri: null }; + +describe('useFeatureChannels', () => { + it('probes SAM availability on mount and sets samAvailable when slimsam is ready', async () => { + vi.stubGlobal('fetch', vi.fn().mockResolvedValue(modulesResponse([{ id: 'slimsam', ready: true }]))); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(result.current.samAvailable).toBe(true)); + }); + + it('samAvailable stays false when slimsam is absent or not ready', async () => { + vi.stubGlobal('fetch', vi.fn().mockResolvedValue(modulesResponse([{ id: 'slimsam', ready: false }]))); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect((fetch as any).mock.calls.length).toBeGreaterThan(0)); + expect(result.current.samAvailable).toBe(false); + }); + + it('samAvailable stays false when the modules fetch rejects', async () => { + vi.stubGlobal('fetch', vi.fn().mockRejectedValue(new Error('down'))); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect((fetch as any).mock.calls.length).toBeGreaterThan(0)); + expect(result.current.samAvailable).toBe(false); + }); + + it('compute() fails fast with no composition selected', async () => { + useIpredStore.setState({ preferredCompositionId: '' }); + vi.stubGlobal('fetch', vi.fn().mockResolvedValue(modulesResponse())); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await act(async () => { + await result.current.compute(); + }); + expect(result.current.error).toBe('Select a composition first.'); + expect(result.current.job).toBeNull(); + }); + + it('compute() opens a session, preprocesses, and populates job + channelIndex 0', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) // sam probe on mount + .mockResolvedValueOnce({ ok: true, json: async () => ({ session_id: 'sess-1', project_id: 'proj-1' }) }) // openIpredSession + .mockResolvedValueOnce(preprocessResponse()) // ipredPreprocess + .mockResolvedValueOnce({ ok: true, blob: async () => new Blob(['x']) }); // channel fetch triggered by channelIndex effect + vi.stubGlobal('fetch', fetchMock); + + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + + await act(async () => { + await result.current.compute(); + }); + + expect(result.current.job?.jobId).toBe('feat-1'); + expect(result.current.job?.channels).toEqual([ + { index: 0, label: 'a' }, { index: 1, label: 'b' }, { index: 2, label: 'c' }, + ]); + expect(result.current.channelIndex).toBe(0); + expect(useIpredStore.getState().ipredSessionId).toBe('sess-1'); + + await waitFor(() => expect(result.current.channelUrl).toBe('blob:mock-url')); + }); + + it('compute() reuses an existing ipredSessionId instead of opening a new session', async () => { + useIpredStore.setState({ ipredSessionId: 'existing-sess' }); + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) + .mockResolvedValueOnce(preprocessResponse()) + .mockResolvedValueOnce({ ok: true, blob: async () => new Blob(['x']) }); + vi.stubGlobal('fetch', fetchMock); + + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + // Only 3 total calls: sam probe, preprocess, channel fetch — no session POST. + expect(fetchMock).toHaveBeenCalledTimes(3); + expect(result.current.job?.jobId).toBe('feat-1'); + }); + + it('compute() builds synthetic channel labels when the server returns none', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) + .mockResolvedValueOnce({ ok: true, json: async () => ({ session_id: 'sess-1', project_id: 'proj-1' }) }) + .mockResolvedValueOnce(preprocessResponse({ labels: [], n_channels: 2 })) + .mockResolvedValueOnce({ ok: true, blob: async () => new Blob(['x']) }); + vi.stubGlobal('fetch', fetchMock); + + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + expect(result.current.job?.channels).toEqual([ + { index: 0, label: 'channel 0' }, { index: 1, label: 'channel 1' }, + ]); + }); + + it('compute() sets hasSam true when the composition id mentions sam/slimsam/mark', async () => { + useIpredStore.setState({ preferredCompositionId: 'comp-slimsam-x' }); + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) + .mockResolvedValueOnce({ ok: true, json: async () => ({ session_id: 'sess-1', project_id: 'proj-1' }) }) + .mockResolvedValueOnce(preprocessResponse()) + .mockResolvedValueOnce({ ok: true, blob: async () => new Blob(['x']) }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + expect(result.current.job?.hasSam).toBe(true); + }); + + it('compute() sets error and clears job on a failed preprocess', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) + .mockResolvedValueOnce({ ok: true, json: async () => ({ session_id: 'sess-1', project_id: 'proj-1' }) }) + .mockResolvedValueOnce({ ok: false, text: async () => 'boom' }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + expect(result.current.error).toBe('boom'); + expect(result.current.job).toBeNull(); + expect(result.current.computing).toBe(false); + }); + + it('compute() is a silent no-op when there is no source/kind open (guarded before ensureSession)', async () => { + const fetchMock = vi.fn().mockResolvedValue(modulesResponse()); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => + useFeatureChannels({ source: null, kind: null, sliceIndex: 0, serverUri: null })); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + expect(result.current.error).toBeNull(); + expect(result.current.job).toBeNull(); + }); + + it('selectChannel/cycleChannel/clearSelection manipulate channelIndex', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) + .mockResolvedValueOnce({ ok: true, json: async () => ({ session_id: 'sess-1', project_id: 'proj-1' }) }) + .mockResolvedValueOnce(preprocessResponse()) + .mockResolvedValue({ ok: true, blob: async () => new Blob(['x']) }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + expect(result.current.channelIndex).toBe(0); + + act(() => result.current.cycleChannel(1)); + expect(result.current.channelIndex).toBe(1); + + act(() => result.current.cycleChannel(-2)); // wraps around (3 channels) + expect(result.current.channelIndex).toBe(2); + + act(() => result.current.selectChannel(0)); + expect(result.current.channelIndex).toBe(0); + + act(() => result.current.clearSelection()); + expect(result.current.channelIndex).toBeNull(); + }); + + it('cycleChannel is a no-op with no job', () => { + vi.stubGlobal('fetch', vi.fn().mockResolvedValue(modulesResponse())); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + act(() => result.current.cycleChannel(1)); + expect(result.current.channelIndex).toBeNull(); + }); + + it('invalidateJob clears job/channelIndex/channelUrl', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) + .mockResolvedValueOnce({ ok: true, json: async () => ({ session_id: 'sess-1', project_id: 'proj-1' }) }) + .mockResolvedValueOnce(preprocessResponse()) + .mockResolvedValue({ ok: true, blob: async () => new Blob(['x']) }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + await waitFor(() => expect(result.current.channelUrl).toBe('blob:mock-url')); + + act(() => result.current.invalidateJob()); + expect(result.current.job).toBeNull(); + expect(result.current.channelIndex).toBeNull(); + expect(result.current.channelUrl).toBeNull(); + }); + + it('adoptFeatureBank populates job from an externally-produced feature bank', () => { + vi.stubGlobal('fetch', vi.fn().mockResolvedValue(modulesResponse())); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + act(() => { + result.current.adoptFeatureBank({ + featureId: 'ext-1', width: 32, height: 32, labels: ['x', 'y'], setupId: 'setup-sam', + }); + }); + expect(result.current.job).toEqual({ + jobId: 'ext-1', width: 32, height: 32, + channels: [{ index: 0, label: 'x' }, { index: 1, label: 'y' }], + hasSam: true, setupId: 'setup-sam', cacheHit: undefined, + }); + }); + + it('resets job/channelIndex/error and revokes the channel URL when source/kind/slice/serverUri changes', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) + .mockResolvedValueOnce({ ok: true, json: async () => ({ session_id: 'sess-1', project_id: 'proj-1' }) }) + .mockResolvedValueOnce(preprocessResponse()) + .mockResolvedValue({ ok: true, blob: async () => new Blob(['x']) }); + vi.stubGlobal('fetch', fetchMock); + const { result, rerender } = renderHook((props) => useFeatureChannels(props), { initialProps: ARGS }); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + await waitFor(() => expect(result.current.channelUrl).toBe('blob:mock-url')); + + act(() => { rerender({ ...ARGS, sliceIndex: 1 }); }); + expect(result.current.job).toBeNull(); + expect(result.current.channelIndex).toBeNull(); + expect(URL.revokeObjectURL).toHaveBeenCalledWith('blob:mock-url'); + }); + + it('sets an error when the channel PNG fetch itself fails', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(modulesResponse()) + .mockResolvedValueOnce({ ok: true, json: async () => ({ session_id: 'sess-1', project_id: 'proj-1' }) }) + .mockResolvedValueOnce(preprocessResponse()) + .mockResolvedValueOnce({ ok: false, status: 500 }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureChannels(ARGS)); + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(1)); + await act(async () => { + await result.current.compute(); + }); + await waitFor(() => expect(result.current.error).toMatch(/Channel fetch failed/)); + expect(result.current.channelUrl).toBeNull(); + }); +}); diff --git a/frontend/src/hooks/useFeatureChannels.ts b/frontend/src/hooks/useFeatureChannels.ts new file mode 100644 index 0000000..0034360 --- /dev/null +++ b/frontend/src/hooks/useFeatureChannels.ts @@ -0,0 +1,259 @@ +/** + * useFeatureChannels — compute feature bank via ipred and load channel PNGs. + */ +import { useCallback, useEffect, useRef, useState } from 'react'; +import { useConnectionStore } from '@/stores/connectionStore'; +import { useIpredStore } from '@/stores/ipredStore'; +import { ipredChannelUrl, ipredPreprocess, listIpredModules, openIpredSession } from '@/lib/ipredApi'; + +export interface FeatureChannelInfo { + index: number; + label: string; +} + +export interface FeatureJobInfo { + jobId: string; + width: number; + height: number; + channels: FeatureChannelInfo[]; + hasSam: boolean; + setupId?: string; + cacheHit?: boolean; +} + +export interface UseFeatureChannelsArgs { + source: string | null; + kind: string | null; + sliceIndex: number; + serverUri: string | null; +} + +/** + * Owns feature compute + channel blob URL lifecycle via ipred. + * `channelIndex === null` means show the original rendered slice. + */ +export function useFeatureChannels({ + source, + kind, + sliceIndex, + serverUri, +}: UseFeatureChannelsArgs) { + const [job, setJob] = useState(null); + const [channelIndex, setChannelIndex] = useState(null); + const [channelUrl, setChannelUrl] = useState(null); + const [computing, setComputing] = useState(false); + const [error, setError] = useState(null); + const [samAvailable, setSamAvailable] = useState(false); + const channelUrlRef = useRef(null); + + const preferredCompositionId = useIpredStore((s) => s.preferredCompositionId); + const ipredSessionId = useIpredStore((s) => s.ipredSessionId); + const setIpredSession = useIpredStore((s) => s.setIpredSession); + const localRoot = useConnectionStore((s) => s.localRoot); + + useEffect(() => { + let cancelled = false; + listIpredModules() + .then((modules) => { + if (cancelled) return; + const slimsam = modules.find((m) => m.id === 'slimsam'); + setSamAvailable(!!slimsam?.ready); + }) + .catch(() => { + if (!cancelled) setSamAvailable(false); + }); + return () => { + cancelled = true; + }; + }, []); + + const revokeChannelUrl = useCallback(() => { + if (channelUrlRef.current) { + URL.revokeObjectURL(channelUrlRef.current); + channelUrlRef.current = null; + } + setChannelUrl(null); + }, []); + + useEffect(() => { + setJob(null); + setChannelIndex(null); + revokeChannelUrl(); + setError(null); + }, [source, kind, sliceIndex, serverUri, revokeChannelUrl]); + + useEffect(() => { + if (!job || channelIndex === null) { + revokeChannelUrl(); + return; + } + let cancelled = false; + const url = ipredChannelUrl(job.jobId, channelIndex); + (async () => { + try { + const res = await fetch(url); + if (!res.ok) throw new Error(`Channel fetch failed: ${res.status}`); + const blob = await res.blob(); + if (cancelled) return; + const objUrl = URL.createObjectURL(blob); + if (channelUrlRef.current) URL.revokeObjectURL(channelUrlRef.current); + channelUrlRef.current = objUrl; + setChannelUrl(objUrl); + } catch (e) { + if (!cancelled) { + setError(e instanceof Error ? e.message : String(e)); + revokeChannelUrl(); + } + } + })(); + return () => { + cancelled = true; + }; + }, [job, channelIndex, revokeChannelUrl]); + + useEffect( + () => () => { + if (channelUrlRef.current) URL.revokeObjectURL(channelUrlRef.current); + }, + [], + ); + + const ensureSession = useCallback(async (): Promise => { + if (ipredSessionId) return ipredSessionId; + if (!source || !kind) throw new Error('No sample open'); + const session = await openIpredSession({ + kind, + source, + server_uri: serverUri, + root: kind === 'local' ? localRoot : null, + }); + setIpredSession({ + sessionId: session.session_id, + projectId: session.project_id, + }); + return session.session_id; + }, [ipredSessionId, source, kind, serverUri, localRoot, setIpredSession]); + + const compute = useCallback(async () => { + if (!source || !kind || computing) return; + const compositionId = preferredCompositionId; + if (!compositionId) { + setError('Select a composition first.'); + return; + } + setComputing(true); + setError(null); + try { + const sessionId = await ensureSession(); + const data = await ipredPreprocess({ + session_id: sessionId, + composition_id: compositionId, + slice_index: sliceIndex, + }); + const channels: FeatureChannelInfo[] = (data.labels ?? []).map((label, index) => ({ + index, + label, + })); + if (channels.length === 0) { + for (let i = 0; i < data.n_channels; i += 1) { + channels.push({ index: i, label: `channel ${i}` }); + } + } + setJob({ + jobId: data.feature_id, + width: data.width, + height: data.height, + channels, + hasSam: + compositionId.includes('slimsam') || + compositionId.includes('sam') || + compositionId.includes('mark'), + setupId: data.setup_id, + cacheHit: data.cache_hit, + }); + setChannelIndex(0); + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + setJob(null); + setChannelIndex(null); + } finally { + setComputing(false); + } + }, [source, kind, sliceIndex, computing, preferredCompositionId, ensureSession]); + + const selectChannel = useCallback((index: number | null) => { + setChannelIndex(index); + }, []); + + const cycleChannel = useCallback( + (delta: number) => { + if (!job || job.channels.length === 0) return; + setChannelIndex((prev) => { + const cur = prev ?? 0; + const n = job.channels.length; + return (((cur + delta) % n) + n) % n; + }); + }, + [job], + ); + + const clearSelection = useCallback(() => { + setChannelIndex(null); + }, []); + + const invalidateJob = useCallback(() => { + setJob(null); + setChannelIndex(null); + revokeChannelUrl(); + }, [revokeChannelUrl]); + + /** Adopt a feature bank produced elsewhere (e.g. Train auto-preprocess). */ + const adoptFeatureBank = useCallback( + (payload: { + featureId: string; + width: number; + height: number; + labels?: string[]; + nChannels?: number; + setupId?: string; + cacheHit?: boolean; + }) => { + const n = payload.nChannels ?? payload.labels?.length ?? 0; + const channels: FeatureChannelInfo[] = (payload.labels ?? []).map((label, index) => ({ + index, + label, + })); + if (channels.length === 0) { + for (let i = 0; i < n; i += 1) { + channels.push({ index: i, label: `channel ${i}` }); + } + } + setJob({ + jobId: payload.featureId, + width: payload.width, + height: payload.height, + channels, + hasSam: (payload.setupId ?? '').includes('sam'), + setupId: payload.setupId, + cacheHit: payload.cacheHit, + }); + }, + [], + ); + + return { + job, + channelIndex, + channelUrl, + computing, + error, + samAvailable, + preferredCompositionId, + compute, + selectChannel, + cycleChannel, + clearSelection, + invalidateJob, + adoptFeatureBank, + }; +} diff --git a/frontend/src/hooks/useFeatureManifold.test.ts b/frontend/src/hooks/useFeatureManifold.test.ts new file mode 100644 index 0000000..42b4f2d --- /dev/null +++ b/frontend/src/hooks/useFeatureManifold.test.ts @@ -0,0 +1,297 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook, waitFor } from '@testing-library/react'; +import { DEFAULT_MANIFOLD_PARAMS, useFeatureManifold } from './useFeatureManifold'; +import type { Shape } from '@/stores/annotationStore'; + +function sampleOk(overrides: Record = {}) { + return { + ok: true, + json: async () => ({ + sample_id: 'samp-1', + points: [{ x: 1, y: 2, radius: 5 }], + k: 24, + n_picked: 1, + n_subsample: 100, + explained_variance: 0.8, + radius: 5, + box_size: 64, + ...overrides, + }), + }; +} +function heatmapOk() { + return { ok: true, blob: async () => new Blob(['x']) }; +} + +beforeEach(() => { + vi.stubGlobal('URL', Object.assign(URL, { + createObjectURL: vi.fn(() => 'blob:mock-url'), + revokeObjectURL: vi.fn(), + })); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); + vi.useRealTimers(); +}); + +describe('useFeatureManifold', () => { + it('starts with default params and empty results', () => { + vi.stubGlobal('fetch', vi.fn()); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: null })); + expect(result.current.params).toEqual(DEFAULT_MANIFOLD_PARAMS); + expect(result.current.points).toEqual([]); + expect(result.current.heatmapUrl).toBeNull(); + expect(result.current.hasSample).toBe(false); + expect(result.current.meta).toBeNull(); + }); + + it('sample() is a no-op with no featureJobId', async () => { + const fetchMock = vi.fn(); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: null })); + await act(async () => { + await result.current.sample(); + }); + expect(fetchMock).not.toHaveBeenCalled(); + }); + + it('sample() posts feature_id/k/box_size, then fetches the heatmap, populating points/meta/heatmapUrl', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(sampleOk()) + .mockResolvedValueOnce(heatmapOk()); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + + await act(async () => { + await result.current.sample(); + }); + + expect(result.current.points).toEqual([{ x: 1, y: 2, radius: 5 }]); + expect(result.current.heatmapUrl).toBe('blob:mock-url'); + expect(result.current.meta).toEqual({ + k: 24, nPicked: 1, nSubsample: 100, explainedVariance: 0.8, radius: 5, boxSize: 64, + hasMask: undefined, maskPixels: undefined, + }); + expect(result.current.hasSample).toBe(true); + expect(result.current.sampling).toBe(false); + + const [sampleUrl, sampleInit] = fetchMock.mock.calls[0]; + expect(sampleUrl).toBe('/api/ipred/manifold/sample'); + const body = JSON.parse(sampleInit.body); + expect(body).toEqual({ feature_id: 'feat-1', k: 24, box_size: 64 }); + const [heatUrl] = fetchMock.mock.calls[1]; + expect(heatUrl).toBe('/api/ipred/manifold/samp-1/heatmap.png'); + }); + + it('sample() includes shapes in the body when a placement mask is set', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(sampleOk()) + .mockResolvedValueOnce(heatmapOk()); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + const shapes: Shape[] = [{ id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 1, h: 1 }]; + + act(() => result.current.setPlacementMaskFromShapes(shapes)); + expect(result.current.placementMask).toEqual(shapes); + + await act(async () => { + await result.current.sample(); + }); + const body = JSON.parse(fetchMock.mock.calls[0][1].body); + expect(body.shapes).toEqual(shapes); + }); + + it('setPlacementMaskFromShapes([]) clears the mask instead of storing an empty array', () => { + vi.stubGlobal('fetch', vi.fn()); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + act(() => result.current.setPlacementMaskFromShapes([])); + expect(result.current.placementMask).toBeNull(); + }); + + it('clearPlacementMask resets placementMask to null', () => { + vi.stubGlobal('fetch', vi.fn()); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + act(() => result.current.setPlacementMaskFromShapes([ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 1, h: 1 }, + ])); + act(() => result.current.clearPlacementMask()); + expect(result.current.placementMask).toBeNull(); + }); + + it('sample() surfaces a parsed JSON `detail` error message and revokes any prior heatmap', async () => { + const fetchMock = vi.fn().mockResolvedValueOnce({ + ok: false, + text: async () => JSON.stringify({ detail: 'bad k value' }), + }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + await act(async () => { + await result.current.sample(); + }); + expect(result.current.error).toBe('bad k value'); + expect(result.current.heatmapUrl).toBeNull(); + expect(result.current.sampling).toBe(false); + }); + + it('sample() falls back to the raw body text when it is not JSON', async () => { + const fetchMock = vi.fn().mockResolvedValueOnce({ ok: false, status: 500, text: async () => 'plain text error' }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + await act(async () => { + await result.current.sample(); + }); + expect(result.current.error).toBe('plain text error'); + }); + + it('sample() treats a "not found"/"unknown feature" detail specially: calls onFeatureJobExpired and shows a recompute hint', async () => { + const onFeatureJobExpired = vi.fn(); + const fetchMock = vi.fn().mockResolvedValueOnce({ + ok: false, + text: async () => JSON.stringify({ detail: 'unknown feature id' }), + }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1', onFeatureJobExpired })); + await act(async () => { + await result.current.sample(); + }); + expect(onFeatureJobExpired).toHaveBeenCalledTimes(1); + expect(result.current.error).toMatch(/Compute again/); + }); + + it('sample() errors when the heatmap fetch itself fails, after a successful sample POST', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(sampleOk()) + .mockResolvedValueOnce({ ok: false, status: 500 }); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + await act(async () => { + await result.current.sample(); + }); + expect(result.current.error).toBe('Failed to fetch manifold heatmap'); + expect(result.current.heatmapUrl).toBeNull(); + }); + + it('a concurrent sample() call while one is in-flight is ignored', async () => { + let resolveFirst!: (v: unknown) => void; + const fetchMock = vi.fn() + .mockReturnValueOnce(new Promise((resolve) => { resolveFirst = resolve; })) + .mockResolvedValueOnce(heatmapOk()); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + + let p1: Promise; + let p2: Promise; + act(() => { + p1 = result.current.sample(); + p2 = result.current.sample(); + }); + expect(fetchMock).toHaveBeenCalledTimes(1); // second call short-circuited by samplingRef + + await act(async () => { + resolveFirst(sampleOk()); + await Promise.all([p1, p2]); + }); + }); + + it('resets heatmap/points/meta/error/placementMask when featureJobId changes', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(sampleOk()) + .mockResolvedValueOnce(heatmapOk()); + vi.stubGlobal('fetch', fetchMock); + const { result, rerender } = renderHook( + ({ featureJobId }) => useFeatureManifold({ featureJobId }), + { initialProps: { featureJobId: 'feat-1' } }, + ); + await act(async () => { + await result.current.sample(); + }); + expect(result.current.hasSample).toBe(true); + + act(() => { rerender({ featureJobId: 'feat-2' }); }); + expect(result.current.heatmapUrl).toBeNull(); + expect(result.current.points).toEqual([]); + expect(result.current.meta).toBeNull(); + expect(result.current.placementMask).toBeNull(); + expect(URL.revokeObjectURL).toHaveBeenCalledWith('blob:mock-url'); + }); + + it('debounces a re-sample 350ms after k/boxSize changes, but only once a sample already exists', async () => { + vi.useFakeTimers({ shouldAdvanceTime: true }); + const fetchMock = vi.fn() + .mockResolvedValue(sampleOk()) + // interleave heatmap responses; every other call is the heatmap fetch + ; + fetchMock.mockImplementation(async (url: string) => (url.includes('heatmap') ? heatmapOk() : sampleOk())); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + + await act(async () => { + await result.current.sample(); + }); + expect(result.current.hasSample).toBe(true); + fetchMock.mockClear(); + + act(() => { result.current.setParams((p) => ({ ...p, k: 40 })); }); + // Not yet — debounce hasn't elapsed. + await act(async () => { await vi.advanceTimersByTimeAsync(200); }); + expect(fetchMock).not.toHaveBeenCalled(); + + await act(async () => { await vi.advanceTimersByTimeAsync(200); }); + await waitFor(() => expect(fetchMock).toHaveBeenCalled()); + }); + + it('does not auto re-sample on param change before any sample has been taken', async () => { + vi.useFakeTimers({ shouldAdvanceTime: true }); + const fetchMock = vi.fn().mockResolvedValue(sampleOk()); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + + act(() => { result.current.setParams((p) => ({ ...p, k: 40 })); }); + await act(async () => { await vi.advanceTimersByTimeAsync(500); }); + expect(fetchMock).not.toHaveBeenCalled(); + }); + + it('dismiss() clears the heatmap/points/error', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(sampleOk()) + .mockResolvedValueOnce(heatmapOk()); + vi.stubGlobal('fetch', fetchMock); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + await act(async () => { + await result.current.sample(); + }); + expect(result.current.hasSample).toBe(true); + + act(() => result.current.dismiss()); + expect(result.current.heatmapUrl).toBeNull(); + expect(result.current.points).toEqual([]); + expect(result.current.error).toBeNull(); + expect(result.current.hasSample).toBe(false); + }); + + it('revokes the heatmap object URL on unmount', async () => { + const fetchMock = vi.fn() + .mockResolvedValueOnce(sampleOk()) + .mockResolvedValueOnce(heatmapOk()); + vi.stubGlobal('fetch', fetchMock); + const { result, unmount } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + await act(async () => { + await result.current.sample(); + }); + unmount(); + expect(URL.revokeObjectURL).toHaveBeenCalledWith('blob:mock-url'); + }); + + it('setShowHeatmap/setShowMarkers/setHeatmapOpacity update their state', () => { + vi.stubGlobal('fetch', vi.fn()); + const { result } = renderHook(() => useFeatureManifold({ featureJobId: 'feat-1' })); + act(() => result.current.setShowHeatmap(false)); + expect(result.current.showHeatmap).toBe(false); + act(() => result.current.setShowMarkers(false)); + expect(result.current.showMarkers).toBe(false); + act(() => result.current.setHeatmapOpacity(0.9)); + expect(result.current.heatmapOpacity).toBe(0.9); + }); +}); diff --git a/frontend/src/hooks/useFeatureManifold.ts b/frontend/src/hooks/useFeatureManifold.ts new file mode 100644 index 0000000..12cb253 --- /dev/null +++ b/frontend/src/hooks/useFeatureManifold.ts @@ -0,0 +1,207 @@ +/** + * useFeatureManifold — greedy variance-box suggestions + residual heatmap. + */ +import { useCallback, useEffect, useRef, useState } from 'react'; +import { API_BASE } from '@/config'; +import type { ManifoldPoint } from '@/lib/featureManifold'; +import type { Shape } from '@/stores/annotationStore'; + +export interface ManifoldParams { + k: number; + /** Full square box side length in image pixels. */ + boxSize: number; +} + +export const DEFAULT_MANIFOLD_PARAMS: ManifoldParams = { + k: 24, + boxSize: 64, +}; + +export interface UseFeatureManifoldArgs { + featureJobId: string | null; + onFeatureJobExpired?: () => void; +} + +async function readErrorDetail(res: Response): Promise { + const text = await res.text(); + try { + const parsed = JSON.parse(text) as { detail?: unknown }; + if (typeof parsed.detail === 'string') return parsed.detail; + } catch { + /* plain */ + } + return text || `Request failed: ${res.status}`; +} + +export function useFeatureManifold({ + featureJobId, + onFeatureJobExpired, +}: UseFeatureManifoldArgs) { + const [params, setParams] = useState(DEFAULT_MANIFOLD_PARAMS); + const [points, setPoints] = useState([]); + const [heatmapUrl, setHeatmapUrl] = useState(null); + const [showHeatmap, setShowHeatmap] = useState(true); + const [showMarkers, setShowMarkers] = useState(true); + const [heatmapOpacity, setHeatmapOpacity] = useState(0.45); + const [sampling, setSampling] = useState(false); + const [error, setError] = useState(null); + const [placementMask, setPlacementMask] = useState(null); + const [meta, setMeta] = useState<{ + k: number; + nPicked: number; + nSubsample: number; + explainedVariance: number; + radius: number; + boxSize: number; + hasMask?: boolean; + maskPixels?: number; + } | null>(null); + const heatmapUrlRef = useRef(null); + const samplingRef = useRef(false); + const hasSampleRef = useRef(false); + const paramsRef = useRef(params); + paramsRef.current = params; + const placementMaskRef = useRef(placementMask); + placementMaskRef.current = placementMask; + const onExpiredRef = useRef(onFeatureJobExpired); + onExpiredRef.current = onFeatureJobExpired; + + const revoke = useCallback(() => { + if (heatmapUrlRef.current) { + URL.revokeObjectURL(heatmapUrlRef.current); + heatmapUrlRef.current = null; + } + setHeatmapUrl(null); + setPoints([]); + setMeta(null); + hasSampleRef.current = false; + }, []); + + useEffect(() => { + revoke(); + setError(null); + setPlacementMask(null); + }, [featureJobId, revoke]); + + useEffect( + () => () => { + if (heatmapUrlRef.current) URL.revokeObjectURL(heatmapUrlRef.current); + }, + [], + ); + + const sample = useCallback(async () => { + if (!featureJobId || samplingRef.current) return; + samplingRef.current = true; + setSampling(true); + setError(null); + const { k, boxSize } = paramsRef.current; + const maskShapes = placementMaskRef.current; + try { + const res = await fetch(`${API_BASE}/api/ipred/manifold/sample`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + feature_id: featureJobId, + k, + box_size: boxSize, + ...(maskShapes && maskShapes.length > 0 ? { shapes: maskShapes } : {}), + }), + }); + if (!res.ok) { + const detail = await readErrorDetail(res); + if (/not found|unknown feature/i.test(detail)) { + onExpiredRef.current?.(); + throw new Error('Feature bank missing. Click Compute again, then Suggest labels.'); + } + throw new Error(detail); + } + const body = (await res.json()) as { + sample_id: string; + points: ManifoldPoint[]; + k: number; + n_picked?: number; + n_subsample: number; + explained_variance: number; + radius?: number; + box_size?: number; + has_mask?: boolean; + mask_pixels?: number; + }; + const heatRes = await fetch( + `${API_BASE}/api/ipred/manifold/${body.sample_id}/heatmap.png`, + ); + if (!heatRes.ok) throw new Error('Failed to fetch manifold heatmap'); + const blob = await heatRes.blob(); + const url = URL.createObjectURL(blob); + if (heatmapUrlRef.current) URL.revokeObjectURL(heatmapUrlRef.current); + heatmapUrlRef.current = url; + setHeatmapUrl(url); + setPoints(body.points ?? []); + hasSampleRef.current = (body.points?.length ?? 0) > 0 || !!url; + setMeta({ + k: body.k, + nPicked: body.n_picked ?? (body.points?.length ?? 0), + nSubsample: body.n_subsample, + explainedVariance: body.explained_variance, + radius: body.radius ?? body.points?.[0]?.radius ?? 0, + boxSize: body.box_size ?? boxSize, + hasMask: body.has_mask, + maskPixels: body.mask_pixels, + }); + } catch (e) { + revoke(); + setError(e instanceof Error ? e.message : String(e)); + } finally { + samplingRef.current = false; + setSampling(false); + } + }, [featureJobId, revoke]); + + const sampleRef = useRef(sample); + sampleRef.current = sample; + + // Re-place boxes when K / box size / mask change after an existing suggestion. + useEffect(() => { + if (!featureJobId || !hasSampleRef.current) return; + const t = window.setTimeout(() => { + void sampleRef.current(); + }, 350); + return () => window.clearTimeout(t); + }, [featureJobId, params.k, params.boxSize, placementMask]); + + const setPlacementMaskFromShapes = useCallback((shapes: Shape[]) => { + setPlacementMask(shapes.length ? shapes.map((s) => ({ ...s })) : null); + }, []); + + const clearPlacementMask = useCallback(() => { + setPlacementMask(null); + }, []); + + const dismiss = useCallback(() => { + revoke(); + setError(null); + }, [revoke]); + + return { + params, + setParams, + points, + heatmapUrl, + showHeatmap, + setShowHeatmap, + showMarkers, + setShowMarkers, + heatmapOpacity, + setHeatmapOpacity, + sampling, + error, + meta, + sample, + dismiss, + placementMask, + setPlacementMaskFromShapes, + clearPlacementMask, + hasSample: points.length > 0 || !!heatmapUrl, + }; +} diff --git a/frontend/src/hooks/useGuideSync.test.tsx b/frontend/src/hooks/useGuideSync.test.tsx new file mode 100644 index 0000000..d6b1a73 --- /dev/null +++ b/frontend/src/hooks/useGuideSync.test.tsx @@ -0,0 +1,132 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, renderHook, waitFor } from '@testing-library/react'; +import { + generateGuide, loadGuide, saveGuide, useGuideLoad, useGuideSync, +} from './useGuideSync'; +import { useReferenceGuideStore } from '@/stores/referenceGuideStore'; + +beforeEach(() => { + useReferenceGuideStore.getState().clear(); + vi.stubGlobal('fetch', vi.fn()); +}); + +afterEach(() => { + vi.unstubAllGlobals(); + vi.useRealTimers(); +}); + +describe('loadGuide', () => { + it('returns null on a 404 (no guide exists)', async () => { + (fetch as any).mockResolvedValue({ status: 404, ok: false }); + expect(await loadGuide('local:x.tif')).toBeNull(); + }); + + it('returns null on any fetch error (never throws)', async () => { + (fetch as any).mockRejectedValue(new Error('network down')); + expect(await loadGuide('local:x.tif')).toBeNull(); + }); + + it('parses classes/notes from a successful response', async () => { + (fetch as any).mockResolvedValue({ + status: 200, ok: true, + json: async () => ({ guide: { classes: [{ label: 'Cell', color: '#f00', description: '', exampleCrops: [] }], notes: 'hi' } }), + }); + expect(await loadGuide('local:x.tif')).toEqual({ + classes: [{ label: 'Cell', color: '#f00', description: '', exampleCrops: [] }], notes: 'hi', + }); + }); + + it('defaults to empty classes/notes when the guide object is bare', async () => { + (fetch as any).mockResolvedValue({ status: 200, ok: true, json: async () => ({ guide: {} }) }); + expect(await loadGuide('local:x.tif')).toEqual({ classes: [], notes: '' }); + }); +}); + +describe('saveGuide', () => { + it('PUTs the classes/notes and returns true on success', async () => { + (fetch as any).mockResolvedValue({ ok: true }); + const ok = await saveGuide('local:x.tif', [], 'notes'); + expect(ok).toBe(true); + const [url, init] = (fetch as any).mock.calls[0]; + expect(url).toContain('/api/guide?source_key='); + expect(init.method).toBe('PUT'); + expect(JSON.parse(init.body)).toEqual({ classes: [], notes: 'notes' }); + }); + + it('returns false without throwing on a network error', async () => { + (fetch as any).mockRejectedValue(new Error('down')); + expect(await saveGuide('local:x.tif', [], '')).toBe(false); + }); + + it('returns false when the server responds not-ok', async () => { + (fetch as any).mockResolvedValue({ ok: false }); + expect(await saveGuide('local:x.tif', [], '')).toBe(false); + }); +}); + +describe('generateGuide', () => { + it('POSTs the payload and returns the generated guide', async () => { + (fetch as any).mockResolvedValue({ + ok: true, json: async () => ({ classes: [{ label: 'Pore', color: '#000', description: '', exampleCrops: [] }], notes: '' }), + }); + const result = await generateGuide('local:x.tif', { classes: [], slices: {} }); + expect(result.classes).toHaveLength(1); + }); + + it('throws a descriptive error including the response detail on failure', async () => { + (fetch as any).mockResolvedValue({ ok: false, status: 500, text: async () => 'boom' }); + await expect(generateGuide('local:x.tif', { classes: [], slices: {} })).rejects.toThrow(/500.*boom/); + }); +}); + +describe('useGuideLoad', () => { + it('loads the guide into the store for a given sourceKey', async () => { + (fetch as any).mockResolvedValue({ + status: 200, ok: true, json: async () => ({ guide: { classes: [{ label: 'Cell', color: '#f00', description: '', exampleCrops: [] }], notes: 'n' } }), + }); + renderHook(() => useGuideLoad('local:x.tif')); + await waitFor(() => expect(useReferenceGuideStore.getState().loadedFor).toBe('local:x.tif')); + expect(useReferenceGuideStore.getState().entries).toHaveLength(1); + }); + + it('clears the store when sourceKey is null', () => { + useReferenceGuideStore.getState().setGuide([{ label: 'x', color: '#000', description: '', exampleCrops: [] }], 'n', 'local:y.tif'); + renderHook(() => useGuideLoad(null)); + expect(useReferenceGuideStore.getState().entries).toEqual([]); + }); +}); + +describe('useGuideSync', () => { + it('debounces an autosave after the guide has loaded for this source', async () => { + vi.useFakeTimers({ shouldAdvanceTime: true }); + (fetch as any).mockImplementation(async (url: string) => { + if (url.includes('source_key=local%3Ax.tif') && !url.includes('generate')) { + return { status: 200, ok: true, json: async () => ({ guide: { classes: [], notes: '' } }) }; + } + return { ok: true }; + }); + renderHook(() => useGuideSync('local:x.tif')); + + await waitFor(() => expect(useReferenceGuideStore.getState().loadedFor).toBe('local:x.tif')); + (fetch as any).mockClear(); + + act(() => { + useReferenceGuideStore.getState().setGuide([{ label: 'Cell', color: '#f00', description: '', exampleCrops: [] }], '', 'local:x.tif'); + }); + await vi.advanceTimersByTimeAsync(1000); + expect(fetch).toHaveBeenCalledWith( + expect.stringContaining('/api/guide?source_key='), + expect.objectContaining({ method: 'PUT' }), + ); + }); + + it('does not autosave before the guide has finished loading for this source', async () => { + vi.useFakeTimers({ shouldAdvanceTime: true }); + (fetch as any).mockReturnValue(new Promise(() => {})); // never resolves -> loadedFor stays unset + renderHook(() => useGuideSync('local:x.tif')); + await vi.advanceTimersByTimeAsync(2000); + // Only the initial GET attempt (still pending), no PUT. + const putCalls = (fetch as any).mock.calls.filter(([, init]: any) => init?.method === 'PUT'); + expect(putCalls).toHaveLength(0); + }); +}); diff --git a/frontend/src/hooks/useHubSelectedTabs.test.ts b/frontend/src/hooks/useHubSelectedTabs.test.ts new file mode 100644 index 0000000..9988b55 --- /dev/null +++ b/frontend/src/hooks/useHubSelectedTabs.test.ts @@ -0,0 +1,61 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, renderHook } from '@testing-library/react'; +import { useHubSelectedTabs } from './useHubSelectedTabs'; + +beforeEach(() => { + localStorage.clear(); +}); + +afterEach(() => { + vi.restoreAllMocks(); +}); + +describe('useHubSelectedTabs', () => { + it('starts null with nothing stored', () => { + const { result } = renderHook(() => useHubSelectedTabs()); + expect(result.current.selectedPaths).toBeNull(); + }); + + it('loads previously stored paths on mount', () => { + localStorage.setItem('sam3_hub_selected_tab_paths', JSON.stringify(['/connect', '/browse'])); + const { result } = renderHook(() => useHubSelectedTabs()); + expect(result.current.selectedPaths).toEqual(['/connect', '/browse']); + }); + + it('ignores malformed JSON and starts null', () => { + localStorage.setItem('sam3_hub_selected_tab_paths', 'not json'); + const { result } = renderHook(() => useHubSelectedTabs()); + expect(result.current.selectedPaths).toBeNull(); + }); + + it('ignores a non-array stored value', () => { + localStorage.setItem('sam3_hub_selected_tab_paths', JSON.stringify({ a: 1 })); + const { result } = renderHook(() => useHubSelectedTabs()); + expect(result.current.selectedPaths).toBeNull(); + }); + + it('ignores an array with non-string items', () => { + localStorage.setItem('sam3_hub_selected_tab_paths', JSON.stringify([1, 2, 3])); + const { result } = renderHook(() => useHubSelectedTabs()); + expect(result.current.selectedPaths).toBeNull(); + }); + + it('setSelectedPaths writes through to localStorage and updates state', () => { + const { result } = renderHook(() => useHubSelectedTabs()); + act(() => result.current.setSelectedPaths(['/annotate'])); + expect(result.current.selectedPaths).toEqual(['/annotate']); + expect(JSON.parse(localStorage.getItem('sam3_hub_selected_tab_paths')!)).toEqual(['/annotate']); + }); + + it('logs and keeps prior state when localStorage.setItem throws', () => { + const errorSpy = vi.spyOn(console, 'error').mockImplementation(() => {}); + const setItemSpy = vi.spyOn(Storage.prototype, 'setItem').mockImplementation(() => { + throw new Error('quota exceeded'); + }); + const { result } = renderHook(() => useHubSelectedTabs()); + act(() => result.current.setSelectedPaths(['/annotate'])); + expect(result.current.selectedPaths).toBeNull(); + expect(errorSpy).toHaveBeenCalled(); + setItemSpy.mockRestore(); + }); +}); diff --git a/frontend/src/hooks/useImageSlice.gc.test.ts b/frontend/src/hooks/useImageSlice.gc.test.ts new file mode 100644 index 0000000..491d83e --- /dev/null +++ b/frontend/src/hooks/useImageSlice.gc.test.ts @@ -0,0 +1,89 @@ +/** + * The slice cache hands out blob object URLs, which the browser keeps alive until + * explicitly revoked — evicting the query is not enough. This is exactly the kind + * of leak that never shows up in a feature test (everything still works, the tab + * just grows), so it gets its own. + */ +import { describe, it, expect, vi, beforeEach, afterEach } from 'vitest'; +import { QueryClient } from '@tanstack/react-query'; +import { installImageSliceGc } from './useImageSlice'; + +let revoked: string[] = []; + +beforeEach(() => { + revoked = []; + vi.stubGlobal('URL', { + ...URL, + createObjectURL: (b: Blob) => `blob:mock/${(b as unknown as { id?: string }).id ?? 'x'}`, + revokeObjectURL: (u: string) => { revoked.push(u); }, + }); +}); + +afterEach(() => { + vi.unstubAllGlobals(); +}); + +/** Seed the cache with a slice query already holding a blob URL. */ +function seedSlice(client: QueryClient, sliceIndex: number, url: string) { + const key = ['imageSlice', 'src', 'tiled', sliceIndex, {}, null]; + client.setQueryData(key, url); + return key; +} + +describe('installImageSliceGc', () => { + it('revokes a slice URL when its query is removed from the cache', () => { + const client = new QueryClient(); + const stop = installImageSliceGc(client); + const key = seedSlice(client, 0, 'blob:mock/slice0'); + + client.removeQueries({ queryKey: key, exact: true }); + + expect(revoked).toEqual(['blob:mock/slice0']); + stop(); + }); + + it('revokes the previous URL when a slice is refetched into the same key', () => { + const client = new QueryClient(); + const stop = installImageSliceGc(client); + const key = seedSlice(client, 3, 'blob:mock/old'); + + client.setQueryData(key, 'blob:mock/new'); + + expect(revoked).toEqual(['blob:mock/old']); + stop(); + }); + + it('revokes every visited slice as the cache is cleared', () => { + const client = new QueryClient(); + const stop = installImageSliceGc(client); + for (let i = 0; i < 5; i++) seedSlice(client, i, `blob:mock/s${i}`); + + client.clear(); + + expect(revoked.sort()).toEqual( + ['blob:mock/s0', 'blob:mock/s1', 'blob:mock/s2', 'blob:mock/s3', 'blob:mock/s4'], + ); + stop(); + }); + + it('ignores non-slice queries and non-blob data', () => { + const client = new QueryClient(); + const stop = installImageSliceGc(client); + client.setQueryData(['somethingElse', 1], 'blob:mock/not-a-slice'); + client.setQueryData(['imageSlice', 'src', 'tiled', 9, {}, null], { notAString: true }); + + client.clear(); + + expect(revoked).toEqual([]); + stop(); + }); + + it('stops revoking once unsubscribed', () => { + const client = new QueryClient(); + const stop = installImageSliceGc(client); + stop(); + seedSlice(client, 1, 'blob:mock/after-stop'); + client.clear(); + expect(revoked).toEqual([]); + }); +}); diff --git a/frontend/src/hooks/useImageSlice.test.tsx b/frontend/src/hooks/useImageSlice.test.tsx new file mode 100644 index 0000000..2c1f0e4 --- /dev/null +++ b/frontend/src/hooks/useImageSlice.test.tsx @@ -0,0 +1,75 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { cleanup, renderHook, waitFor } from '@testing-library/react'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import type { ReactNode } from 'react'; +import { useImageSlice } from './useImageSlice'; +import type { RenderOpts } from '@/stores/datasetStore'; + +function wrapper({ children }: { children: ReactNode }) { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return {children}; +} + +const RENDER_OPTS: RenderOpts = { norm: 'slice', scale: 'linear', vminPct: 1, vmaxPct: 99, cmap: 'gray' }; + +beforeEach(() => { + vi.stubGlobal('fetch', vi.fn()); + vi.stubGlobal('URL', class extends URL { + static createObjectURL = vi.fn(() => 'blob:fake-url'); + static revokeObjectURL = vi.fn(); + }); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('useImageSlice', () => { + it('is disabled (no fetch) when source is null', () => { + const { result } = renderHook( + () => useImageSlice(null, 'local', 0, RENDER_OPTS, null), + { wrapper }, + ); + expect(result.current.fetchStatus).toBe('idle'); + expect(fetch).not.toHaveBeenCalled(); + }); + + it('is disabled when kind is null', () => { + renderHook(() => useImageSlice('x.tif', null, 0, RENDER_OPTS, null), { wrapper }); + expect(fetch).not.toHaveBeenCalled(); + }); + + it('fetches the slice and resolves to a blob object URL', async () => { + const fakeBlob = new Blob(['data']); + (fetch as any).mockResolvedValue({ ok: true, blob: async () => fakeBlob }); + const { result } = renderHook( + () => useImageSlice('x.tif', 'local', 0, RENDER_OPTS, null), + { wrapper }, + ); + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(result.current.data).toBe('blob:fake-url'); + expect(fetch).toHaveBeenCalledWith(expect.stringContaining('/api/image/slice?')); + }); + + it('surfaces a fetch failure as an error state', async () => { + (fetch as any).mockResolvedValue({ ok: false, status: 500 }); + const { result } = renderHook( + () => useImageSlice('x.tif', 'local', 0, RENDER_OPTS, null), + { wrapper }, + ); + await waitFor(() => expect(result.current.isError).toBe(true)); + }); + + it('includes the denoise method/strength in the request URL when active', async () => { + (fetch as any).mockResolvedValue({ ok: true, blob: async () => new Blob() }); + renderHook( + () => useImageSlice('x.tif', 'local', 0, RENDER_OPTS, null, { method: 'median', strength: 0.4 }), + { wrapper }, + ); + await waitFor(() => expect(fetch).toHaveBeenCalled()); + const url = (fetch as any).mock.calls[0][0]; + expect(url).toContain('denoise_method=median'); + expect(url).toContain('denoise_strength=0.4'); + }); +}); diff --git a/frontend/src/hooks/useImageSlice.ts b/frontend/src/hooks/useImageSlice.ts index 8547aac..e742ff1 100644 --- a/frontend/src/hooks/useImageSlice.ts +++ b/frontend/src/hooks/useImageSlice.ts @@ -1,21 +1,40 @@ /** * useImageSlice — TanStack Query hook to fetch a PNG slice from the backend. + * + * Slices are cached as blob object URLs. Those are NOT garbage collected when the + * query holding them is evicted — an object URL pins its blob until explicitly + * revoked — so `installImageSliceGc` must be wired up once at startup, or every + * slice the user visits stays resident for the lifetime of the tab. */ -import { useQuery } from '@tanstack/react-query'; +import { useQuery, type QueryClient } from '@tanstack/react-query'; import { API_BASE } from '@/config'; -import { RenderOpts } from '@/stores/datasetStore'; +import { RenderOpts, type DenoiseOpts } from '@/stores/datasetStore'; + +const SLICE_QUERY_KEY = 'imageSlice'; export interface SliceResult { url: string; } -/** Build the /api/image/slice URL encoding source, slice index, and render options. */ +/** Build the /api/image/slice URL encoding source, slice index, and render options. + * + * `denoise` is optional and omitted from the URL entirely when off, so the + * un-denoised request stays byte-identical to what it has always been — the + * backend only pays (and caches) the filter cost when a method is actually set. + * + * `crop`, when set, asks the backend to filter and return only a centred square + * of that size at 1:1. Filtering a full slice costs seconds for NLM and TV; + * cropping keeps slider-dragging interactive (measured: 7.3s -> 0.49s for + * bilateral on a 2560² slice). + */ export function buildSliceUrl( source: string, kind: string, sliceIndex: number, renderOpts: RenderOpts, - serverUri: string | null + serverUri: string | null, + denoise?: DenoiseOpts | null, + crop?: number ): string { const params = new URLSearchParams({ source, @@ -28,6 +47,11 @@ export function buildSliceUrl( cmap: renderOpts.cmap, }); if (serverUri) params.set('server_uri', serverUri); + if (denoise && denoise.method !== 'none') { + params.set('denoise_method', denoise.method); + params.set('denoise_strength', String(denoise.strength)); + if (crop && crop > 0) params.set('denoise_crop', String(crop)); + } return `${API_BASE}/api/image/slice?${params.toString()}`; } @@ -40,15 +64,22 @@ export function useImageSlice( kind: string | null, sliceIndex: number, renderOpts: RenderOpts, - serverUri: string | null + serverUri: string | null, + denoise?: DenoiseOpts | null ) { const enabled = Boolean(source && kind); const url = enabled - ? buildSliceUrl(source!, kind!, sliceIndex, renderOpts, serverUri) + ? buildSliceUrl(source!, kind!, sliceIndex, renderOpts, serverUri, denoise) + : null; + + // Denoise is keyed only when active, so switching it off returns to the exact + // cache entry the un-denoised view already had rather than refetching. + const denoiseKey = denoise && denoise.method !== 'none' + ? [denoise.method, denoise.strength] : null; return useQuery({ - queryKey: ['imageSlice', source, kind, sliceIndex, renderOpts, serverUri], + queryKey: [SLICE_QUERY_KEY, source, kind, sliceIndex, renderOpts, serverUri, denoiseKey], queryFn: async () => { const res = await fetch(url!); if (!res.ok) throw new Error(`Slice fetch failed: ${res.status}`); @@ -59,3 +90,43 @@ export function useImageSlice( staleTime: 1000 * 300, }); } + +/** + * Revoke a slice's object URL when its query leaves the cache, so paging through a + * volume doesn't retain every PNG for the session. Without this, the browser holds + * each blob alive indefinitely — a few hundred slices is easily hundreds of MB, and + * the resulting GC pressure shows up as the whole tab getting slower the longer it + * is used. + * + * Also revokes the previous URL when a query's data is replaced (a refetch of the + * same slice), which would otherwise orphan the old blob. + * + * Call once, next to the QueryClient. Returns the unsubscribe function. + */ +export function installImageSliceGc(queryClient: QueryClient): () => void { + const isSliceQuery = (key: readonly unknown[]) => key[0] === SLICE_QUERY_KEY; + const revoke = (value: unknown) => { + if (typeof value === 'string' && value.startsWith('blob:')) URL.revokeObjectURL(value); + }; + // Track the last-seen URL per query so an 'updated' event can revoke the one + // being replaced (the event carries the new state, not the old). + const seen = new Map(); + + return queryClient.getQueryCache().subscribe((event) => { + const { query } = event; + if (!isSliceQuery(query.queryKey)) return; + const hash = query.queryHash; + + if (event.type === 'removed') { + revoke(query.state.data); + seen.delete(hash); + return; + } + const data = query.state.data; + if (typeof data === 'string') { + const prev = seen.get(hash); + if (prev && prev !== data) revoke(prev); + seen.set(hash, data); + } + }); +} diff --git a/frontend/src/hooks/useKeybinds.test.ts b/frontend/src/hooks/useKeybinds.test.ts new file mode 100644 index 0000000..b7e1dae --- /dev/null +++ b/frontend/src/hooks/useKeybinds.test.ts @@ -0,0 +1,199 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook } from '@testing-library/react'; +import { useKeybinds } from './useKeybinds'; +import { useToolStore } from '@/stores/toolStore'; +import { useDatasetStore } from '@/stores/datasetStore'; +import { useClassStore } from '@/stores/classStore'; +import * as editHistory from '@/hooks/editHistory'; + +const INITIAL_TOOL_STATE = useToolStore.getState(); +const INITIAL_DATASET_STATE = useDatasetStore.getState(); + +function openSample(nSlices = 5) { + useDatasetStore.getState().setDataset('local', 'x.tif', null, { + nSlices, + height: 10, + width: 10, + dtype: 'uint8', + isRgb: false, + valueRange: [0, 255], + }); +} + +function keydown(key: string, opts: Partial = {}) { + act(() => { + window.dispatchEvent(new KeyboardEvent('keydown', { key, bubbles: true, cancelable: true, ...opts })); + }); +} +function keyup(key: string, opts: Partial = {}) { + act(() => { + window.dispatchEvent(new KeyboardEvent('keyup', { key, bubbles: true, cancelable: true, ...opts })); + }); +} + +describe('useKeybinds', () => { + let onActivateClass: ReturnType; + let onNewBrushInstance: ReturnType; + let onDeleteSelected: ReturnType; + let onCancelDraft: ReturnType; + + beforeEach(() => { + useToolStore.setState(INITIAL_TOOL_STATE, true); + useDatasetStore.setState(INITIAL_DATASET_STATE, true); + useClassStore.setState({ classes: [] }); + onActivateClass = vi.fn(); + onNewBrushInstance = vi.fn(); + onDeleteSelected = vi.fn(); + onCancelDraft = vi.fn(); + }); + + afterEach(() => { + cleanup(); + }); + + function mount(activeClassId: number | null = null) { + return renderHook(() => + useKeybinds(activeClassId, onActivateClass, onNewBrushInstance, onDeleteSelected, onCancelDraft), + ); + } + + it('switches tools on their letter key', () => { + mount(); + keydown('p'); + expect(useToolStore.getState().tool).toBe('polygon'); + keydown('b'); + expect(useToolStore.getState().tool).toBe('brush'); + keydown('r'); + expect(useToolStore.getState().tool).toBe('eraser'); + keydown('g'); + expect(useToolStore.getState().tool).toBe('magic'); + }); + + it('ignores keys typed into an editable target (input/textarea/select)', () => { + mount(); + const input = document.createElement('input'); + document.body.appendChild(input); + act(() => { + input.dispatchEvent(new KeyboardEvent('keydown', { key: 'p', bubbles: true })); + }); + expect(useToolStore.getState().tool).not.toBe('polygon'); + document.body.removeChild(input); + }); + + it('activates a class 1-9 via digit keys, mapped by position', () => { + useClassStore.setState({ + classes: [ + { classId: 10, label: 'A', color: '#f00', isVisible: true }, + { classId: 20, label: 'B', color: '#0f0', isVisible: true }, + ], + }); + mount(); + keydown('2'); + expect(onActivateClass).toHaveBeenCalledWith(20); + }); + + it('does nothing for a digit beyond the class list length', () => { + useClassStore.setState({ classes: [{ classId: 1, label: 'A', color: '#f00', isVisible: true }] }); + mount(); + keydown('5'); + expect(onActivateClass).not.toHaveBeenCalled(); + }); + + it('holding Space switches to pan and releasing restores the previous tool', () => { + mount(); + act(() => { useToolStore.getState().setTool('brush'); }); + keydown(' '); + expect(useToolStore.getState().tool).toBe('pan'); + expect(useToolStore.getState().panReturnTool).toBe('brush'); + + keyup(' '); + expect(useToolStore.getState().tool).toBe('brush'); + expect(useToolStore.getState().panReturnTool).toBeNull(); + }); + + it('a repeated (auto-repeat) Space keydown does not re-arm the pan-return tool', () => { + mount(); + act(() => { useToolStore.getState().setTool('brush'); }); + keydown(' '); + expect(useToolStore.getState().tool).toBe('pan'); + // Simulate OS key-repeat: tool is already 'pan', repeat=true should just bail out. + keydown(' ', { repeat: true }); + expect(useToolStore.getState().panReturnTool).toBe('brush'); + }); + + it('Ctrl/Cmd+Z calls editHistory.undo(); Ctrl+Shift+Z and Ctrl+Y call redo()', () => { + const undoSpy = vi.spyOn(editHistory, 'undo').mockImplementation(() => {}); + const redoSpy = vi.spyOn(editHistory, 'redo').mockImplementation(() => {}); + mount(); + + keydown('z', { ctrlKey: true }); + expect(undoSpy).toHaveBeenCalledTimes(1); + expect(redoSpy).not.toHaveBeenCalled(); + + keydown('z', { ctrlKey: true, shiftKey: true }); + expect(redoSpy).toHaveBeenCalledTimes(1); + + keydown('y', { ctrlKey: true }); + expect(redoSpy).toHaveBeenCalledTimes(2); + + keydown('z', { metaKey: true }); + expect(undoSpy).toHaveBeenCalledTimes(2); + + undoSpy.mockRestore(); + redoSpy.mockRestore(); + }); + + it('"t" requests a fit-to-screen (bumps fitRequestId)', () => { + mount(); + const before = useToolStore.getState().fitRequestId; + keydown('t'); + expect(useToolStore.getState().fitRequestId).toBe(before + 1); + }); + + it('"x" advances to the next slice, clamped to the last slice', () => { + openSample(3); + mount(); + keydown('x'); + expect(useDatasetStore.getState().currentSlice).toBe(1); + keydown('x'); + keydown('x'); + keydown('x'); // beyond the end, clamps to nSlices-1 + expect(useDatasetStore.getState().currentSlice).toBe(2); + }); + + it('ArrowLeft/ArrowRight navigate slices, clamped at 0', () => { + openSample(3); + mount(); + keydown('ArrowLeft'); // already at 0, clamps + expect(useDatasetStore.getState().currentSlice).toBe(0); + keydown('ArrowRight'); + expect(useDatasetStore.getState().currentSlice).toBe(1); + keydown('ArrowLeft'); + expect(useDatasetStore.getState().currentSlice).toBe(0); + }); + + it('slice navigation is a no-op with no dataset loaded (meta is null)', () => { + mount(); + keydown('x'); + expect(useDatasetStore.getState().currentSlice).toBe(0); + }); + + it('"n" triggers onNewBrushInstance, Delete/Backspace trigger onDeleteSelected, Escape triggers onCancelDraft', () => { + mount(); + keydown('n'); + expect(onNewBrushInstance).toHaveBeenCalledTimes(1); + keydown('Delete'); + expect(onDeleteSelected).toHaveBeenCalledTimes(1); + keydown('Backspace'); + expect(onDeleteSelected).toHaveBeenCalledTimes(2); + keydown('Escape'); + expect(onCancelDraft).toHaveBeenCalledTimes(1); + }); + + it('removes its listeners on unmount', () => { + const { unmount } = mount(); + unmount(); + keydown('p'); + expect(useToolStore.getState().tool).not.toBe('polygon'); + }); +}); diff --git a/frontend/src/hooks/useKeybinds.ts b/frontend/src/hooks/useKeybinds.ts index 321203c..31cff08 100644 --- a/frontend/src/hooks/useKeybinds.ts +++ b/frontend/src/hooks/useKeybinds.ts @@ -1,8 +1,9 @@ /** * useKeybinds — keyboard shortcuts for the annotation workspace. * - * Tools: p=polygon l=ellipse e=rectangle r=eraser b=brush f=fill g=magic - * m=magnetic s=select (each key is a letter in the tool's label) + * Tools: p=polygon l=ellipse e=rectangle r=eraser b=brush h=threshold brush + * k=sampler (fits the threshold band from a lassoed example) + * f=fill g=magic m=magnetic s=select (each key is a letter in the label) * Pan: hold Space (reverts to the previous tool on release) * View: x=next slice (forward) t=fit image to screen arrows=prev/next slice * Edit: Ctrl/Cmd+Z=undo Ctrl/Cmd+Shift+Z / Ctrl+Y=redo @@ -23,6 +24,8 @@ const TOOL_KEYBINDS: Record = { e: 'rectangle', // Rect l: 'ellipse', // eLlipse (e is taken by rect) b: 'brush', + h: 'threshold', // tHreshold brush (t is fit-to-screen) + k: 'sampler', // eyedropper lasso that fits the threshold band f: 'fill', r: 'eraser', // eRaser }; diff --git a/frontend/src/hooks/useMaskOps.test.ts b/frontend/src/hooks/useMaskOps.test.ts new file mode 100644 index 0000000..b39b89f --- /dev/null +++ b/frontend/src/hooks/useMaskOps.test.ts @@ -0,0 +1,91 @@ +import { beforeEach, describe, expect, it } from 'vitest'; +import { renderHook } from '@testing-library/react'; +import { useMaskOps } from './useMaskOps'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { useDatasetStore } from '@/stores/datasetStore'; + +const SOURCE = 'local:x.tif'; + +beforeEach(() => { + useAnnotationStore.getState().reset(); + useDatasetStore.setState({ + meta: { nSlices: 5, height: 64, width: 64, dtype: 'uint8', isRgb: false, valueRange: [0, 255] }, + currentSlice: 0, + } as any); +}); + +describe('useMaskOps.applyCleanup', () => { + it('returns false with no sourceKey', () => { + const { result } = renderHook(() => useMaskOps(null, 1)); + expect(result.current.applyCleanup('fill')).toBe(false); + }); + + it('returns false with no active class', () => { + const { result } = renderHook(() => useMaskOps(SOURCE, null)); + expect(result.current.applyCleanup('fill')).toBe(false); + }); + + it('returns false with no dataset meta loaded', () => { + useDatasetStore.setState({ meta: null } as any); + const { result } = renderHook(() => useMaskOps(SOURCE, 1)); + expect(result.current.applyCleanup('fill')).toBe(false); + }); + + it('returns false when the active class has no shapes on the current slice', () => { + const { result } = renderHook(() => useMaskOps(SOURCE, 1)); + expect(result.current.applyCleanup('fill')).toBe(false); + }); + + it('re-rasterizes the active class shapes through a fill-holes pass', () => { + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 5, y: 5, w: 20, h: 20 }, + ]); + const { result } = renderHook(() => useMaskOps(SOURCE, 1)); + const ok = result.current.applyCleanup('fill'); + expect(ok).toBe(true); + const shapes = useAnnotationStore.getState().byImage[SOURCE]['0']; + expect(shapes.length).toBeGreaterThan(0); + expect(shapes.every((s) => s.classId === 1)).toBe(true); + }); + + it('only touches shapes of the active class, leaving other classes untouched', () => { + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 5, y: 5, w: 20, h: 20 }, + ]); + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 2, [ + { id: 's2', classId: 2, kind: 'rectangle', x: 40, y: 40, w: 10, h: 10 }, + ]); + const { result } = renderHook(() => useMaskOps(SOURCE, 1)); + result.current.applyCleanup('grow', 2); + const slice = useAnnotationStore.getState().byImage[SOURCE]['0']; + expect(slice.some((s) => s.classId === 2)).toBe(true); + }); + + it.each(['fill', 'islands', 'smooth', 'grow', 'shrink'] as const)('%s op runs without throwing', (op) => { + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 5, y: 5, w: 20, h: 20 }, + ]); + const { result } = renderHook(() => useMaskOps(SOURCE, 1)); + expect(() => result.current.applyCleanup(op)).not.toThrow(); + }); +}); + +describe('useMaskOps.copyToNext', () => { + it('returns false at the last slice', () => { + useDatasetStore.setState({ currentSlice: 4 } as any); + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 4, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + const { result } = renderHook(() => useMaskOps(SOURCE, 1)); + expect(result.current.copyToNext()).toBe(false); + }); + + it('copies the active class shapes to the next slice', () => { + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + const { result } = renderHook(() => useMaskOps(SOURCE, 1)); + expect(result.current.copyToNext()).toBe(true); + expect(useAnnotationStore.getState().byImage[SOURCE]['1']).toHaveLength(1); + }); +}); diff --git a/frontend/src/hooks/useOpenInAnnotate.test.tsx b/frontend/src/hooks/useOpenInAnnotate.test.tsx new file mode 100644 index 0000000..67c1936 --- /dev/null +++ b/frontend/src/hooks/useOpenInAnnotate.test.tsx @@ -0,0 +1,178 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook } from '@testing-library/react'; +import { useOpenInAnnotate } from './useOpenInAnnotate'; +import { useDatasetStore } from '@/stores/datasetStore'; +import { useClassStore } from '@/stores/classStore'; +import { useAnnotationStore } from '@/stores/annotationStore'; + +const navigateMock = vi.fn(); +vi.mock('react-router', () => ({ + useNavigate: () => navigateMock, +})); + +function metaResponse(overrides: Record = {}) { + return { + ok: true, + json: async () => ({ + n_slices: 5, + height: 100, + width: 200, + dtype: 'uint8', + is_rgb: false, + value_range: [0, 255], + keywords: [], + ...overrides, + }), + }; +} + +function draft404() { + return { status: 404, ok: false }; +} + +beforeEach(() => { + useDatasetStore.getState().reset(); + useClassStore.setState({ classes: [] }); + useAnnotationStore.getState().reset(); + navigateMock.mockClear(); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('useOpenInAnnotate', () => { + describe('openTiledArray', () => { + it('fetches meta, seeds the dataset store, and navigates to /annotate', async () => { + vi.stubGlobal('fetch', vi.fn() + .mockResolvedValueOnce(metaResponse()) + .mockResolvedValueOnce(draft404())); + const { result } = renderHook(() => useOpenInAnnotate()); + + await act(async () => { + await result.current.openTiledArray('browse/ds/img', 'http://srv'); + }); + + expect(useDatasetStore.getState().kind).toBe('tiled'); + expect(useDatasetStore.getState().source).toBe('browse/ds/img'); + expect(useDatasetStore.getState().serverUri).toBe('http://srv'); + expect(useDatasetStore.getState().meta?.nSlices).toBe(5); + expect(navigateMock).toHaveBeenCalledWith('/annotate'); + + const [url] = (fetch as any).mock.calls[0]; + expect(url).toContain('/api/image/meta?'); + expect(url).toContain('source=browse%2Fds%2Fimg'); + expect(url).toContain('kind=tiled'); + expect(url).toContain('server_uri='); + }); + + it('jumps to the requested initialSlice, clamped to the last slice', async () => { + vi.stubGlobal('fetch', vi.fn() + .mockResolvedValueOnce(metaResponse({ n_slices: 5 })) + .mockResolvedValueOnce(draft404())); + const { result } = renderHook(() => useOpenInAnnotate()); + + await act(async () => { + await result.current.openTiledArray('browse/ds/img', 'http://srv', 999); + }); + expect(useDatasetStore.getState().currentSlice).toBe(4); // clamped to n_slices-1 + }); + + it('does not override slice 0 when initialSlice is 0 (default)', async () => { + vi.stubGlobal('fetch', vi.fn() + .mockResolvedValueOnce(metaResponse()) + .mockResolvedValueOnce(draft404())); + const { result } = renderHook(() => useOpenInAnnotate()); + await act(async () => { + await result.current.openTiledArray('browse/ds/img', 'http://srv'); + }); + expect(useDatasetStore.getState().currentSlice).toBe(0); + }); + + it('seeds classes from meta.keywords when the draft has none', async () => { + vi.stubGlobal('fetch', vi.fn() + .mockResolvedValueOnce(metaResponse({ keywords: ['Cell', 'Pore', 'cell'] })) + .mockResolvedValueOnce(draft404())); + const { result } = renderHook(() => useOpenInAnnotate()); + await act(async () => { + await result.current.openTiledArray('browse/ds/img', 'http://srv'); + }); + const classes = useClassStore.getState().classes; + // 'cell' is deduped case-insensitively against 'Cell'. + expect(classes.map((c) => c.label)).toEqual(['Cell', 'Pore']); + }); + + it('adopts classes from a loaded draft instead of keyword-seeding', async () => { + vi.stubGlobal('fetch', vi.fn() + .mockResolvedValueOnce(metaResponse({ keywords: ['ShouldNotAppear'] })) + .mockResolvedValueOnce({ + status: 200, + ok: true, + json: async () => ({ + payload: { + classes: [{ classId: 1, label: 'FromDraft', color: '#000', isVisible: true }], + slices: { '0': [{ id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 1, h: 1 }] }, + split_by_slice: {}, + negative_slices: [], + }, + }), + })); + const { result } = renderHook(() => useOpenInAnnotate()); + await act(async () => { + await result.current.openTiledArray('browse/ds/img', 'http://srv'); + }); + expect(useClassStore.getState().classes).toEqual([ + { classId: 1, label: 'FromDraft', color: '#000', isVisible: true }, + ]); + const sourceKey = 'tiled:http://srv:browse/ds/img'; + expect(useAnnotationStore.getState().byImage[sourceKey]?.['0']).toHaveLength(1); + }); + + it('resets to empty classes when neither the draft nor keywords supply any', async () => { + useClassStore.setState({ classes: [{ classId: 9, label: 'Stale', color: '#fff', isVisible: true }] }); + vi.stubGlobal('fetch', vi.fn() + .mockResolvedValueOnce(metaResponse({ keywords: [] })) + .mockResolvedValueOnce(draft404())); + const { result } = renderHook(() => useOpenInAnnotate()); + await act(async () => { + await result.current.openTiledArray('browse/ds/img', 'http://srv'); + }); + expect(useClassStore.getState().classes).toEqual([]); + }); + + it('throws with the response body when the meta fetch fails', async () => { + vi.stubGlobal('fetch', vi.fn().mockResolvedValue({ ok: false, text: async () => 'not found' })); + const { result } = renderHook(() => useOpenInAnnotate()); + await expect(result.current.openTiledArray('browse/ds/img', 'http://srv')).rejects.toThrow('not found'); + expect(navigateMock).not.toHaveBeenCalled(); + }); + }); + + describe('openLocalFile', () => { + it('fetches meta with kind=local, seeds the dataset store, and navigates', async () => { + vi.stubGlobal('fetch', vi.fn() + .mockResolvedValueOnce(metaResponse()) + .mockResolvedValueOnce(draft404())); + const { result } = renderHook(() => useOpenInAnnotate()); + await act(async () => { + await result.current.openLocalFile('rel/path.tif'); + }); + expect(useDatasetStore.getState().kind).toBe('local'); + expect(useDatasetStore.getState().source).toBe('rel/path.tif'); + expect(useDatasetStore.getState().serverUri).toBeNull(); + expect(navigateMock).toHaveBeenCalledWith('/annotate'); + + const [url] = (fetch as any).mock.calls[0]; + expect(url).toContain('kind=local'); + expect(url).not.toContain('server_uri'); + }); + + it('throws with the response body when the meta fetch fails', async () => { + vi.stubGlobal('fetch', vi.fn().mockResolvedValue({ ok: false, text: async () => 'boom' })); + const { result } = renderHook(() => useOpenInAnnotate()); + await expect(result.current.openLocalFile('rel/path.tif')).rejects.toThrow('boom'); + expect(navigateMock).not.toHaveBeenCalled(); + }); + }); +}); diff --git a/frontend/src/hooks/useOpenInAnnotate.ts b/frontend/src/hooks/useOpenInAnnotate.ts index ff87fe5..8e9eeaf 100644 --- a/frontend/src/hooks/useOpenInAnnotate.ts +++ b/frontend/src/hooks/useOpenInAnnotate.ts @@ -86,6 +86,16 @@ export function useOpenInAnnotate() { dtype: meta.dtype, isRgb: meta.is_rgb, valueRange: meta.value_range, + globalValueRange: meta.global_value_range ?? null, + // Multiscale volumes: width/height/nSlices above are the FINEST level's + // (annotations live in full-res coordinates); these say what is drawn. + levelKey: meta.level_key ?? null, + levelIndex: meta.level_index ?? null, + levelCount: meta.level_count ?? null, + levelWidth: meta.level_width ?? null, + levelHeight: meta.level_height ?? null, + levelNSlices: meta.level_n_slices ?? null, + zDownsample: meta.z_downsample ?? null, }); // setDataset resets to slice 0; jump to the requested slice (clamped). if (initialSlice > 0) setSlice(Math.min(initialSlice, Math.max(0, meta.n_slices - 1))); @@ -120,6 +130,7 @@ export function useOpenInAnnotate() { dtype: meta.dtype, isRgb: meta.is_rgb, valueRange: meta.value_range, + globalValueRange: meta.global_value_range ?? null, }); const sourceKey = buildSourceKey('local', relPath); diff --git a/frontend/src/hooks/usePixelClassifier.slicePersist.test.ts b/frontend/src/hooks/usePixelClassifier.slicePersist.test.ts new file mode 100644 index 0000000..036b838 --- /dev/null +++ b/frontend/src/hooks/usePixelClassifier.slicePersist.test.ts @@ -0,0 +1,98 @@ +/** + * A trained model must survive a slice change (only sample/composition identity + * resets it), and predict() must target the CURRENT slice's feature bank rather + * than the bank the model happened to be trained on. + */ +import { describe, expect, it, vi, beforeEach } from 'vitest'; +import { act, renderHook, waitFor } from '@testing-library/react'; +import { usePixelClassifier } from './usePixelClassifier'; +import { useIpredStore } from '@/stores/ipredStore'; + +vi.mock('@/lib/pixelClf', () => ({ + thresholdProbaPngBlob: vi.fn(async () => new Blob()), +})); + +vi.mock('@/lib/ipredApi', async () => { + const actual = await vi.importActual('@/lib/ipredApi'); + return { + ...actual, + openIpredSession: vi.fn(async () => ({ session_id: 's1', project_id: 'p1' })), + ipredPreprocess: vi.fn(async () => { + throw new Error('ensureFeatureBank should not need to preprocess in this test'); + }), + ipredTrain: vi.fn(async () => ({ + model_id: 'model-1', + feature_id: 'bank-A', + trainer_id: 'catboost', + class_ids: [1, 2], + train_accuracy: 0.9, + n_train: 100, + n_cal: 20, + n_samples: 120, + params: { iterations: 50 }, + feature_importances: [], + })), + ipredInfer: vi.fn(async (payload: { feature_id?: string | null }) => ({ + run_id: 'run-1', + model_id: 'model-1', + feature_id: payload.feature_id ?? '', + alpha: 0.05, + class_ids: [1, 2], + counts: { singleton: 1, multi: 0, abstain: 0 }, + })), + }; +}); + +import { ipredInfer } from '@/lib/ipredApi'; + +describe('usePixelClassifier — model persists across slice changes', () => { + beforeEach(() => { + useIpredStore.getState().reset(); + vi.mocked(ipredInfer).mockClear(); + global.fetch = vi.fn(async () => ({ + ok: true, + blob: async () => new Blob(), + })) as unknown as typeof fetch; + global.URL.createObjectURL = vi.fn(() => 'blob:mock'); + global.URL.revokeObjectURL = vi.fn(); + }); + + it('keeps the model when only featureJobId changes, and predict() uses the current bank', async () => { + const { result, rerender } = renderHook( + (props: { featureJobId: string | null; sliceIndex: number }) => + usePixelClassifier({ + featureJobId: props.featureJobId, + resetKey: 'sample-1', // sourceKey only — does NOT change across slices + source: 'sample.tif', + kind: 'local', + serverUri: null, + sliceIndex: props.sliceIndex, + }), + { initialProps: { featureJobId: 'bank-A', sliceIndex: 5 } }, + ); + + await act(async () => { + await result.current.train([]); + }); + await waitFor(() => expect(result.current.model?.modelId).toBe('model-1')); + expect(result.current.model?.featureId).toBe('bank-A'); + + // Simulate switching to a different slice: featureJobId changes to that + // slice's own feature bank, resetKey (sourceKey) does not change. + rerender({ featureJobId: 'bank-B', sliceIndex: 8 }); + + // The model must still be there — only the run/preview state resets. + expect(result.current.model?.modelId).toBe('model-1'); + + await act(async () => { + await result.current.predict([]); + }); + await waitFor(() => expect(result.current.commitUrl).not.toBeNull()); + + // predict() must have targeted the CURRENT slice's bank (bank-B), not the + // bank the model was originally trained against (bank-A). + expect(ipredInfer).toHaveBeenCalledWith( + expect.objectContaining({ model_id: 'model-1', feature_id: 'bank-B' }), + ); + }); +}); diff --git a/frontend/src/hooks/usePixelClassifier.test.ts b/frontend/src/hooks/usePixelClassifier.test.ts new file mode 100644 index 0000000..1b4b779 --- /dev/null +++ b/frontend/src/hooks/usePixelClassifier.test.ts @@ -0,0 +1,549 @@ +/** + * usePixelClassifier — covers branches NOT already exercised by + * usePixelClassifier.slicePersist.test.ts (which only covers: model persists + * across a slice change, and predict() targets the current slice's bank). + * + * Here: train() error paths (no composition, unknown-feature-bank recovery, + * generic error), predict() guards + error path, proba class cycling + + * threshold clamping, saveThresholdedClass, dismiss, trainAcrossSlices / + * applyAcrossVolume (via the real useExportJob hook, driven through mocked + * fetch responses so no fake timers are needed — the first status poll + * always resolves 'done'/'error' immediately), the live volume-apply preview + * effect keyed off predictedRasterStore, and ensureSession reuse. + */ +import { describe, expect, it, vi, beforeEach } from 'vitest'; +import { act, renderHook, waitFor } from '@testing-library/react'; +import { usePixelClassifier } from './usePixelClassifier'; +import { useIpredStore } from '@/stores/ipredStore'; +import { usePredictedRasterStore } from '@/stores/predictedRasterStore'; + +vi.mock('@/lib/pixelClf', () => ({ + thresholdProbaPngBlob: vi.fn(async () => new Blob()), +})); + +vi.mock('@/lib/ipredApi', async () => { + const actual = await vi.importActual('@/lib/ipredApi'); + return { + ...actual, + openIpredSession: vi.fn(async () => ({ session_id: 's1', project_id: 'p1' })), + ipredPreprocess: vi.fn(async () => ({ + feature_id: 'bank-A', + project_id: 'p1', + setup_id: 'setup-1', + slice_index: 0, + n_channels: 3, + height: 10, + width: 10, + labels: ['a'], + cache_hit: false, + })), + ipredTrain: vi.fn(async () => ({ + model_id: 'model-1', + feature_id: 'bank-A', + trainer_id: 'catboost', + class_ids: [1, 2], + train_accuracy: 0.9, + n_train: 100, + n_cal: 20, + n_samples: 120, + params: { iterations: 50 }, + feature_importances: [{ label: 'intensity', importance: 3 }], + })), + ipredInfer: vi.fn(async () => ({ + run_id: 'run-1', + model_id: 'model-1', + feature_id: 'bank-A', + alpha: 0.05, + class_ids: [1, 2], + counts: { singleton: 1, multi: 0, abstain: 0 }, + })), + ipredThresholdClass: vi.fn(async () => ({ + run_id: 'run-1', + class_id: 1, + class_index: 0, + threshold: 0.5, + width: 10, + height: 10, + n_positive: 5, + label_map_b64: 'AAAA', + })), + }; +}); + +import { + openIpredSession, + ipredPreprocess, + ipredTrain, + ipredInfer, + ipredThresholdClass, +} from '@/lib/ipredApi'; + +const BASE_ARGS = { + featureJobId: 'bank-A' as string | null, + resetKey: 'sample-1' as string | null, + source: 'sample.tif' as string | null, + kind: 'local' as string | null, + serverUri: null as string | null, + sliceIndex: 0, +}; + +/** Routes fetch calls used by ipredRunCommitUrl/StatusUrl/ProbaUrl (blob PNGs) + * and by useExportJob's batch-train/apply start + status-poll routes. */ +function makeFetchRouter(opts?: { + batchTrainResult?: Record | null; + batchTrainError?: string; + batchApplyResult?: Record | null; + batchApplyError?: string; + pngOk?: boolean; +}) { + const { + batchTrainResult = null, + batchTrainError, + batchApplyResult = null, + batchApplyError, + pngOk = true, + } = opts ?? {}; + return vi.fn(async (input: RequestInfo | URL, init?: RequestInit) => { + const url = String(input); + if (url.includes('/api/ipred/batch/train') && init?.method === 'POST') { + return { ok: true, json: async () => ({ job_id: 'job-train-1' }) } as unknown as Response; + } + if (url.includes('/api/ipred/batch/apply') && init?.method === 'POST') { + return { ok: true, json: async () => ({ job_id: 'job-apply-1' }) } as unknown as Response; + } + if (url.includes('/api/export/status/job-train-1')) { + return { + ok: true, + json: async () => + batchTrainError + ? { state: 'error', error: batchTrainError } + : { state: 'done', result: batchTrainResult }, + } as unknown as Response; + } + if (url.includes('/api/export/status/job-apply-1')) { + return { + ok: true, + json: async () => + batchApplyError + ? { state: 'error', error: batchApplyError } + : { state: 'done', result: batchApplyResult }, + } as unknown as Response; + } + // commit.png / status.png / proba/N.png + return { ok: pngOk, blob: async () => new Blob() } as unknown as Response; + }); +} + +beforeEach(() => { + useIpredStore.getState().reset(); + usePredictedRasterStore.setState({ bySource: {} }); + vi.mocked(openIpredSession).mockClear(); + vi.mocked(ipredPreprocess).mockClear(); + vi.mocked(ipredTrain).mockClear(); + vi.mocked(ipredInfer).mockClear(); + vi.mocked(ipredThresholdClass).mockClear(); + global.fetch = makeFetchRouter(); + global.URL.createObjectURL = vi.fn(() => 'blob:mock'); + global.URL.revokeObjectURL = vi.fn(); +}); + +describe('usePixelClassifier — train()', () => { + it('sets an error and does not train without a preferred composition', async () => { + useIpredStore.setState({ preferredCompositionId: '' }); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + expect(result.current.error).toBe('Select a composition first.'); + expect(ipredTrain).not.toHaveBeenCalled(); + }); + + it('recovers from an "unknown feature" training error by clearing the model and flagging expiry', async () => { + vi.mocked(ipredTrain).mockRejectedValueOnce(new Error('unknown feature bank xyz')); + const onFeatureJobExpired = vi.fn(); + const { result } = renderHook(() => + usePixelClassifier({ ...BASE_ARGS, onFeatureJobExpired }), + ); + await act(async () => { + await result.current.train([]); + }); + expect(onFeatureJobExpired).toHaveBeenCalledTimes(1); + expect(result.current.model).toBeNull(); + expect(result.current.error).toMatch(/Feature bank missing/); + }); + + it('surfaces a generic training error message verbatim', async () => { + vi.mocked(ipredTrain).mockRejectedValueOnce(new Error('trainer exploded')); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + expect(result.current.error).toBe('trainer exploded'); + expect(result.current.model).toBeNull(); + }); + + it('trains successfully, mapping the ipred response into ClfTrainResult', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + expect(result.current.model).toMatchObject({ + modelId: 'model-1', + featureId: 'bank-A', + nSamples: 120, + nTrain: 100, + nCal: 20, + classIds: [1, 2], + trainAccuracy: 0.9, + nTrees: 50, + usesSam: false, + trainerId: 'catboost', + compositionId: 'comp-skimage-slimsam', + }); + expect(result.current.model?.featureImportances).toEqual([ + { label: 'intensity', importance: 3 }, + ]); + }); + + it('reuses an existing feature bank instead of calling ipredPreprocess', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + expect(ipredPreprocess).not.toHaveBeenCalled(); + expect(ipredTrain).toHaveBeenCalledWith( + expect.objectContaining({ feature_id: 'bank-A', session_id: 's1' }), + ); + }); + + it('auto-preprocesses and reports the new feature bank when no featureJobId is set', async () => { + const onFeatureReady = vi.fn(); + const { result } = renderHook(() => + usePixelClassifier({ ...BASE_ARGS, featureJobId: null, onFeatureReady }), + ); + await act(async () => { + await result.current.train([]); + }); + expect(ipredPreprocess).toHaveBeenCalledWith( + expect.objectContaining({ session_id: 's1', composition_id: 'comp-skimage-slimsam' }), + ); + expect(onFeatureReady).toHaveBeenCalledWith( + expect.objectContaining({ featureId: 'bank-A', width: 10, height: 10 }), + ); + expect(result.current.model?.modelId).toBe('model-1'); + }); + + it('reuses an already-open ipred session rather than opening a new one', async () => { + useIpredStore.getState().setIpredSession({ sessionId: 'existing-session', projectId: 'p9' }); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + expect(openIpredSession).not.toHaveBeenCalled(); + expect(ipredTrain).toHaveBeenCalledWith( + expect.objectContaining({ session_id: 'existing-session' }), + ); + }); +}); + +describe('usePixelClassifier — predict()', () => { + it('does nothing without a trained model', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.predict([]); + }); + expect(ipredInfer).not.toHaveBeenCalled(); + expect(result.current.commitUrl).toBeNull(); + }); + + it('predicts, publishes commit/status urls, counts, and the first proba channel', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.predict([]); + }); + expect(result.current.commitUrl).toBe('blob:mock'); + expect(result.current.statusUrl).toBe('blob:mock'); + expect(result.current.predictCounts).toEqual({ singleton: 1, multi: 0, abstain: 0 }); + expect(result.current.runId).toBe('run-1'); + expect(result.current.probaUrl).toBe('blob:mock'); + expect(result.current.activeProbaClassId).toBe(1); + }); + + it('revokes any partial preview and sets an error when the commit/status fetch fails', async () => { + global.fetch = makeFetchRouter({ pngOk: false }); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.predict([]); + }); + expect(result.current.error).toMatch(/Failed to fetch conformal prediction PNGs/); + expect(result.current.commitUrl).toBeNull(); + expect(result.current.runId).toBeNull(); + }); +}); + +describe('usePixelClassifier — proba class cycling + thresholds', () => { + async function trainAndPredict(result: { current: ReturnType }) { + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.predict([]); + }); + } + + it('selectProbaClass does nothing without a run/model', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.selectProbaClass(1); + }); + expect(result.current.probaClassIndex).toBe(0); + }); + + it('cycleProbaClass wraps around the class list in both directions', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await trainAndPredict(result); + expect(result.current.probaClassIndex).toBe(0); + + await act(async () => { + result.current.cycleProbaClass(-1); + }); + await waitFor(() => expect(result.current.probaClassIndex).toBe(1)); + expect(result.current.activeProbaClassId).toBe(2); + + await act(async () => { + result.current.cycleProbaClass(1); + }); + await waitFor(() => expect(result.current.probaClassIndex).toBe(0)); + expect(result.current.activeProbaClassId).toBe(1); + }); + + it('setProbaThreshold clamps to [0,1] and republishes the preview for the active class', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await trainAndPredict(result); + + await act(async () => { + result.current.setProbaThreshold(1.5); + }); + expect(result.current.activeProbaThreshold).toBe(1); + + await act(async () => { + result.current.setProbaThreshold(-0.5); + }); + expect(result.current.activeProbaThreshold).toBe(0); + }); + + it('setProbaThreshold is a no-op without a model', () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + act(() => { + result.current.setProbaThreshold(0.7); + }); + expect(result.current.activeProbaThreshold).toBe(0.5); + }); +}); + +describe('usePixelClassifier — saveThresholdedClass / dismiss', () => { + it('returns null and does not call the API without an active run', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + let out; + await act(async () => { + out = await result.current.saveThresholdedClass(); + }); + expect(out).toBeNull(); + expect(ipredThresholdClass).not.toHaveBeenCalled(); + }); + + it('calls ipredThresholdClass with the active class/threshold once predicted', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.predict([]); + }); + let out; + await act(async () => { + out = await result.current.saveThresholdedClass(); + }); + expect(ipredThresholdClass).toHaveBeenCalledWith('run-1', { class_id: 1, threshold: 0.5 }); + expect(out).toMatchObject({ run_id: 'run-1', class_id: 1 }); + expect(result.current.savingClass).toBe(false); + }); + + it('surfaces an error from a failed save without throwing', async () => { + vi.mocked(ipredThresholdClass).mockRejectedValueOnce(new Error('save failed')); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.predict([]); + }); + await act(async () => { + await result.current.saveThresholdedClass(); + }); + expect(result.current.error).toBe('save failed'); + }); + + it('dismiss() revokes the current prediction preview', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.predict([]); + }); + expect(result.current.commitUrl).not.toBeNull(); + act(() => { + result.current.dismiss(); + }); + expect(result.current.commitUrl).toBeNull(); + expect(result.current.runId).toBeNull(); + }); +}); + +describe('usePixelClassifier — identity resets', () => { + it('clears the trained model when resetKey (sample identity) changes', async () => { + const { result, rerender } = renderHook( + (props: { resetKey: string }) => usePixelClassifier({ ...BASE_ARGS, resetKey: props.resetKey }), + { initialProps: { resetKey: 'sample-1' } }, + ); + await act(async () => { + await result.current.train([]); + }); + expect(result.current.model?.modelId).toBe('model-1'); + + rerender({ resetKey: 'sample-2' }); + expect(result.current.model).toBeNull(); + }); + + it('clears the trained model when preferredCompositionId changes', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + expect(result.current.model?.modelId).toBe('model-1'); + + act(() => { + useIpredStore.getState().setPreferredCompositionId('comp-other'); + }); + expect(result.current.model).toBeNull(); + }); +}); + +describe('usePixelClassifier — trainAcrossSlices (multi-slice batch train)', () => { + it('errors when there are no annotated slices', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.trainAcrossSlices({}); + }); + expect(result.current.error).toBe('No annotated slices to train on.'); + }); + + it('errors without a preferred composition', async () => { + useIpredStore.setState({ preferredCompositionId: '' }); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.trainAcrossSlices({ 0: [] }); + }); + expect(result.current.error).toBe('Select a composition first.'); + }); + + it('starts the batch-train job and adopts the completed result as the model', async () => { + global.fetch = makeFetchRouter({ + batchTrainResult: { + model_id: 'model-multi', + feature_id: 'bank-multi', + trainer_id: 'catboost', + class_ids: [1, 2, 3], + train_accuracy: 0.8, + n_train: 300, + n_cal: 60, + n_samples: 360, + params: { iterations: 200 }, + feature_importances: [], + }, + }); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.trainAcrossSlices({ 0: [], 5: [] }); + }); + await waitFor(() => expect(result.current.model?.modelId).toBe('model-multi')); + expect(result.current.model?.classIds).toEqual([1, 2, 3]); + }); + + it('surfaces a batch-train job error', async () => { + global.fetch = makeFetchRouter({ batchTrainError: 'multi-train blew up' }); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.trainAcrossSlices({ 0: [], 5: [] }); + }); + await waitFor(() => expect(result.current.error).toBe('multi-train blew up')); + expect(result.current.model).toBeNull(); + }); +}); + +describe('usePixelClassifier — applyAcrossVolume (batch apply)', () => { + it('errors without a trained model', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.applyAcrossVolume([0, 1, 2]); + }); + expect(result.current.error).toBe('Train a model first.'); + }); + + it('is a no-op with an empty slice list even with a model', async () => { + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.applyAcrossVolume([]); + }); + expect(result.current.error).toBeNull(); + expect(result.current.volumeApplyJob.status).toBe('idle'); + }); + + it('surfaces a batch-apply job error', async () => { + global.fetch = makeFetchRouter({ batchApplyError: 'apply blew up' }); + const { result } = renderHook(() => usePixelClassifier(BASE_ARGS)); + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.applyAcrossVolume([0, 1]); + }); + await waitFor(() => expect(result.current.error).toBe('apply blew up')); + }); + + it('fetches and shows the live per-slice preview once the job result names a run for the current slice', async () => { + global.fetch = makeFetchRouter({ batchApplyResult: { runs: { '0': 'run-vol-1' } } }); + const { result } = renderHook(() => usePixelClassifier({ ...BASE_ARGS, sliceIndex: 0 })); + await act(async () => { + await result.current.train([]); + }); + await act(async () => { + await result.current.applyAcrossVolume([0, 1]); + }); + await waitFor(() => expect(result.current.commitUrl).toBe('blob:mock')); + expect(result.current.statusUrl).toBe('blob:mock'); + }); + + it('shows the committed pointer preview from predictedRasterStore once the job map is gone', async () => { + const { result } = renderHook(() => usePixelClassifier({ ...BASE_ARGS, resetKey: 'sample-9', sliceIndex: 3 })); + await act(async () => { + await result.current.train([]); + }); + act(() => { + usePredictedRasterStore.getState().setPointers('sample-9', { + '3': { runId: 'run-committed-1', classIds: [1, 2] }, + }); + }); + await waitFor(() => expect(result.current.commitUrl).toBe('blob:mock')); + }); +}); diff --git a/frontend/src/hooks/usePixelClassifier.ts b/frontend/src/hooks/usePixelClassifier.ts new file mode 100644 index 0000000..a8cb2d2 --- /dev/null +++ b/frontend/src/hooks/usePixelClassifier.ts @@ -0,0 +1,666 @@ +/** + * usePixelClassifier — train / conformal predict via ipred. + */ +import { useCallback, useEffect, useRef, useState } from 'react'; +import { useConnectionStore } from '@/stores/connectionStore'; +import { useIpredStore } from '@/stores/ipredStore'; +import { usePredictedRasterStore } from '@/stores/predictedRasterStore'; +import type { Shape } from '@/stores/annotationStore'; +import { useExportJob } from '@/hooks/useExportJob'; +import { + ipredInfer, + ipredPreprocess, + ipredRunCommitUrl, + ipredRunProbaUrl, + ipredRunStatusUrl, + ipredThresholdClass, + ipredTrain, + openIpredSession, + type IpredThresholdClassResult, + type IpredTrainResult, +} from '@/lib/ipredApi'; +import { thresholdProbaPngBlob } from '@/lib/pixelClf'; + +export interface ClfParams { + iterations: number; + depth: number; + learningRate: number; + /** Misfire level α for conformal sets (e.g. 0.05 = 5%). */ + alpha: number; +} + +export const DEFAULT_CLF_PARAMS: ClfParams = { + iterations: 200, + depth: 6, + learningRate: 0.1, + alpha: 0.05, +}; + +export interface ClfFeatureImportance { + label: string; + importance: number; +} + +export interface ClfPredictCounts { + singleton: number; + multi: number; + abstain: number; +} + +export interface ClfTrainResult { + modelId: string; + featureId: string; + nSamples: number; + nTrain: number; + nCal: number; + classIds: number[]; + trainAccuracy: number; + params: ClfParams; + nTrees: number; + usesSam: boolean; + featureImportances: ClfFeatureImportance[]; + trainerId: string; + compositionId: string | null; +} + +export interface UsePixelClassifierArgs { + /** Current ipred feature bank id (from Preprocess compute), if any. */ + featureJobId: string | null; + /** Sample/composition identity ONLY (e.g. sourceKey) — must NOT include the + * slice index, or a trained model gets wiped on every slice change. */ + resetKey: string | null; + source: string | null; + kind: string | null; + serverUri: string | null; + sliceIndex: number; + onFeatureJobExpired?: () => void; + /** Called when train auto-runs preprocess and gets a new feature bank. */ + onFeatureReady?: (info: { + featureId: string; + width: number; + height: number; + labels: string[]; + nChannels: number; + setupId: string; + cacheHit: boolean; + }) => void; +} + +const DEFAULT_PROBA_THRESHOLD = 0.5; + +/** Maps an ipred train response (single- or multi-slice — same shape, plus an + * optional `trained_slice_indices`) into the UI's ClfTrainResult. Shared so + * the multi-slice path doesn't duplicate `train()`'s mapping. */ +function toClfTrainResult( + data: IpredTrainResult, + params: ClfParams, + compositionId: string | null, +): ClfTrainResult { + return { + modelId: data.model_id, + featureId: data.feature_id, + nSamples: data.n_samples, + nTrain: data.n_train, + nCal: data.n_cal, + classIds: data.class_ids, + trainAccuracy: data.train_accuracy, + params: { ...params }, + nTrees: Number(data.params?.iterations ?? params.iterations), + usesSam: !!(data.params as { uses_sam?: boolean } | undefined)?.uses_sam, + featureImportances: (data.feature_importances ?? []).map((fi) => ({ + label: fi.label, + importance: fi.importance, + })), + trainerId: data.trainer_id, + compositionId, + }; +} + +export function usePixelClassifier({ + featureJobId, + resetKey, + source, + kind, + serverUri, + sliceIndex, + onFeatureJobExpired, + onFeatureReady, +}: UsePixelClassifierArgs) { + const preferredCompositionId = useIpredStore((s) => s.preferredCompositionId); + const preferredTrainerId = useIpredStore((s) => s.preferredTrainerId); + const preferredTrainerConfig = useIpredStore((s) => s.preferredTrainerConfig); + const ipredSessionId = useIpredStore((s) => s.ipredSessionId); + const setIpredSession = useIpredStore((s) => s.setIpredSession); + const localRoot = useConnectionStore((s) => s.localRoot); + + const [params, setParams] = useState(() => ({ + ...DEFAULT_CLF_PARAMS, + iterations: preferredTrainerConfig.iterations, + depth: preferredTrainerConfig.depth, + learningRate: preferredTrainerConfig.learning_rate, + })); + const [model, setModel] = useState(null); + const [commitUrl, setCommitUrl] = useState(null); + const [statusUrl, setStatusUrl] = useState(null); + const [predictCounts, setPredictCounts] = useState(null); + const [runId, setRunId] = useState(null); + const [probaClassIndex, setProbaClassIndex] = useState(0); + const [probaThresholds, setProbaThresholds] = useState>({}); + const [probaUrl, setProbaUrl] = useState(null); + const [training, setTraining] = useState(false); + const [predicting, setPredicting] = useState(false); + const [savingClass, setSavingClass] = useState(false); + const [error, setError] = useState(null); + const commitUrlRef = useRef(null); + const statusUrlRef = useRef(null); + const probaUrlRef = useRef(null); + /** run_id currently shown via the volume-apply live preview (see the effect + * below) — cleared in revokePredict() so a slice revisit after any revoke + * (slice change, sample change, job restart) always re-fetches rather than + * skipping because "we already showed this run_id once" while commitUrl + * itself has since gone back to null. */ + const volumeApplyPreviewRunIdRef = useRef(null); + /** Raw softmax PNG per class (before threshold preview). */ + const rawProbaBlobRef = useRef(null); + const probaThresholdsRef = useRef(probaThresholds); + probaThresholdsRef.current = probaThresholds; + const modelRef = useRef(model); + modelRef.current = model; + const onExpiredRef = useRef(onFeatureJobExpired); + onExpiredRef.current = onFeatureJobExpired; + const onFeatureReadyRef = useRef(onFeatureReady); + onFeatureReadyRef.current = onFeatureReady; + + // Keep Train knobs aligned with ipred trainer defaults when they change. + useEffect(() => { + setParams((p) => ({ + ...p, + iterations: preferredTrainerConfig.iterations, + depth: preferredTrainerConfig.depth, + learningRate: preferredTrainerConfig.learning_rate, + })); + }, [preferredTrainerConfig]); + + const revokeProba = useCallback(() => { + if (probaUrlRef.current) { + URL.revokeObjectURL(probaUrlRef.current); + probaUrlRef.current = null; + } + setProbaUrl(null); + rawProbaBlobRef.current = null; + }, []); + + const publishProbaPreview = useCallback(async (raw: Blob, threshold: number) => { + const preview = await thresholdProbaPngBlob(raw, threshold); + const url = URL.createObjectURL(preview); + if (probaUrlRef.current) URL.revokeObjectURL(probaUrlRef.current); + probaUrlRef.current = url; + setProbaUrl(url); + }, []); + + const revokePredict = useCallback(() => { + if (commitUrlRef.current) { + URL.revokeObjectURL(commitUrlRef.current); + commitUrlRef.current = null; + } + if (statusUrlRef.current) { + URL.revokeObjectURL(statusUrlRef.current); + statusUrlRef.current = null; + } + setCommitUrl(null); + setStatusUrl(null); + setPredictCounts(null); + setRunId(null); + setProbaClassIndex(0); + setProbaThresholds({}); + volumeApplyPreviewRunIdRef.current = null; + revokeProba(); + }, [revokeProba]); + + // Sample or composition identity changed — the trained model no longer applies. + // `resetKey` is the sample's sourceKey alone (no slice index baked in), so a + // plain slice change does NOT land here; see the effect below for that case. + useEffect(() => { + setModel(null); + revokePredict(); + setError(null); + }, [resetKey, preferredCompositionId, revokePredict]); + + // Feature bank changed (new slice, or a recompute) — any in-flight prediction + // preview is tied to the OLD bank and must go, but the trained model itself + // stays valid: it can be applied to whichever slice is on screen now (see + // `predict()`, which always resolves the CURRENT slice's bank via + // `ensureFeatureBank()` rather than the model's original training bank). + useEffect(() => { + revokePredict(); + setError(null); + }, [featureJobId, revokePredict]); + + useEffect( + () => () => { + if (commitUrlRef.current) URL.revokeObjectURL(commitUrlRef.current); + if (statusUrlRef.current) URL.revokeObjectURL(statusUrlRef.current); + if (probaUrlRef.current) URL.revokeObjectURL(probaUrlRef.current); + }, + [], + ); + + const ensureSession = useCallback(async (): Promise => { + if (ipredSessionId) return ipredSessionId; + if (!source || !kind) throw new Error('No sample open'); + const session = await openIpredSession({ + kind, + source, + server_uri: serverUri, + root: kind === 'local' ? localRoot : null, + }); + setIpredSession({ + sessionId: session.session_id, + projectId: session.project_id, + }); + return session.session_id; + }, [ipredSessionId, source, kind, serverUri, localRoot, setIpredSession]); + + const ensureFeatureBank = useCallback(async (): Promise => { + if (featureJobId) return featureJobId; + if (!preferredCompositionId) { + throw new Error('Select a composition first.'); + } + const sessionId = await ensureSession(); + const bank = await ipredPreprocess({ + session_id: sessionId, + composition_id: preferredCompositionId, + slice_index: sliceIndex, + }); + onFeatureReadyRef.current?.({ + featureId: bank.feature_id, + width: bank.width, + height: bank.height, + labels: bank.labels ?? [], + nChannels: bank.n_channels, + setupId: bank.setup_id, + cacheHit: bank.cache_hit, + }); + return bank.feature_id; + }, [featureJobId, preferredCompositionId, ensureSession, sliceIndex]); + + const loadProbaChannel = useCallback( + async (run: string, classIndex: number) => { + const res = await fetch(ipredRunProbaUrl(run, classIndex)); + if (!res.ok) throw new Error(`Failed to load class ${classIndex} probability map`); + const blob = await res.blob(); + rawProbaBlobRef.current = blob; + setProbaClassIndex(classIndex); + const m = modelRef.current; + const cid = m?.classIds[classIndex]; + const t = + cid !== undefined + ? (probaThresholdsRef.current[cid] ?? DEFAULT_PROBA_THRESHOLD) + : DEFAULT_PROBA_THRESHOLD; + await publishProbaPreview(blob, t); + }, + [publishProbaPreview], + ); + + const train = useCallback( + async (shapes: Shape[]) => { + if (training) return; + if (!preferredCompositionId) { + setError('Select a composition first.'); + return; + } + setTraining(true); + setError(null); + revokePredict(); + try { + const sessionId = await ensureSession(); + const featureId = await ensureFeatureBank(); + const data = await ipredTrain({ + session_id: sessionId, + shapes, + feature_id: featureId, + trainer_id: preferredTrainerId, + config: { + iterations: params.iterations, + depth: params.depth, + learning_rate: params.learningRate, + }, + }); + setModel(toClfTrainResult(data, params, preferredCompositionId)); + } catch (e) { + setModel(null); + const msg = e instanceof Error ? e.message : String(e); + if (/unknown feature|not found/i.test(msg)) { + onExpiredRef.current?.(); + setError('Feature bank missing. Compute or Train again (auto-preprocesses).'); + } else { + setError(msg); + } + } finally { + setTraining(false); + } + }, + [ + training, + preferredCompositionId, + preferredTrainerId, + params, + revokePredict, + ensureSession, + ensureFeatureBank, + ], + ); + + // ---- Batch operations: multi-slice train, whole-volume apply ---- + const multiTrainJobHook = useExportJob(); + const volumeApplyJobHook = useExportJob(); + // Guards against re-applying an already-handled job result on every render + // (the job's `state` object is recreated each poll tick even once done). + const multiTrainHandledRef = useRef(null); + const volumeApplyHandledRef = useRef(null); + + /** Train one model pooling labeled pixels across every slice in `perSliceShapes`. */ + const trainAcrossSlices = useCallback( + async (perSliceShapes: Record) => { + if (Object.keys(perSliceShapes).length === 0) { + setError('No annotated slices to train on.'); + return; + } + if (!preferredCompositionId) { + setError('Select a composition first.'); + return; + } + setError(null); + revokePredict(); + const sessionId = await ensureSession(); + multiTrainHandledRef.current = null; + await multiTrainJobHook.startIpredBatchTrain({ + session_id: sessionId, + slices: perSliceShapes, + composition_id: preferredCompositionId, + trainer_id: preferredTrainerId, + config: { + iterations: params.iterations, + depth: params.depth, + learning_rate: params.learningRate, + }, + }); + }, + [ + preferredCompositionId, + preferredTrainerId, + params, + revokePredict, + ensureSession, + multiTrainJobHook, + ], + ); + + // Adopt the completed multi-train job's result the same way `train()` does. + useEffect(() => { + const { status, result, error: jobError, jobId } = multiTrainJobHook.state; + if (!jobId || multiTrainHandledRef.current === jobId) return; + if (status === 'done' && result) { + multiTrainHandledRef.current = jobId; + setModel(toClfTrainResult(result as unknown as IpredTrainResult, params, preferredCompositionId)); + } else if (status === 'error') { + multiTrainHandledRef.current = jobId; + setModel(null); + setError(jobError ?? 'Multi-slice training failed.'); + } + }, [multiTrainJobHook.state, params, preferredCompositionId]); + + /** Run inference across many slices (e.g. the whole volume). Does not commit — + * turning the result's per-slice runs into shapes stays client-side in + * AnnotatePage, reusing the same PNG-vectorize path as the single-slice commit. */ + const applyAcrossVolume = useCallback( + async (sliceIndices: number[]) => { + if (!model) { + setError('Train a model first.'); + return; + } + if (sliceIndices.length === 0) return; + setError(null); + // Clear any prior job's frozen preview (single-slice OR a previous + // volume apply) before starting — otherwise switching slices right + // after kicking off a new run could briefly show a stale overlay left + // over from before this job's own results start landing. + revokePredict(); + const sessionId = await ensureSession(); + volumeApplyHandledRef.current = null; + await volumeApplyJobHook.startIpredBatchApply({ + session_id: sessionId, + model_id: model.modelId, + slice_indices: sliceIndices, + composition_id: preferredCompositionId, + alpha: params.alpha, + }); + }, + [model, preferredCompositionId, params.alpha, ensureSession, volumeApplyJobHook, revokePredict], + ); + + useEffect(() => { + const { status, error: jobError, jobId } = volumeApplyJobHook.state; + if (!jobId || volumeApplyHandledRef.current === jobId) return; + if (status === 'error') { + volumeApplyHandledRef.current = jobId; + setError(jobError ?? 'Volume apply failed.'); + } + }, [volumeApplyJobHook.state]); + + // Live per-slice preview during (or after) a volume-apply job: as soon as + // ipred_batch_jobs.py's result.runs has an entry for whichever slice is + // currently on screen, fetch and show that slice's commit/status overlay — + // reusing the exact same commitUrl/statusUrl the single-slice "Predict" + // button already drives, so AnnotationCanvas needs no new prop. Before this, + // switching slices during/after a volume apply showed nothing at all until + // the explicit "Commit" step vectorized everything into permanent shapes — + // there was no cheap way to just look at a slice's predicted result first. + // + // NOT a proba-channel preview: batch-apply runs are created with + // store_probabilities=false (ipred_batch_jobs.py's `_apply_one_slice`), so + // there is no proba.npy to load for these run ids — only commit/status. + // Reactive: must re-render this effect the instant "Commit predicted + // shapes" writes a pointer, not only on the next slice change — a plain + // `usePredictedRasterStore.getState()` read wouldn't re-fire the effect + // below when only the store (not sliceIndex/job result) changes. + const committedRunIdForSlice = usePredictedRasterStore( + (s) => (resetKey ? s.bySource[resetKey]?.[String(sliceIndex)]?.runId : undefined) ?? null, + ); + + useEffect(() => { + const result = volumeApplyJobHook.state.result as { runs?: Record } | null; + // Prefer the live job's own runs (covers "still running" and "just + // finished, not yet committed"); once "Commit predicted shapes" resets + // the job (see AnnotatePage's handleCommitVolumeApply), that map is gone + // and the durable predictedRasterStore pointer — set by Commit itself — + // becomes the only remaining source, so the SAME overlay keeps working + // after commit without ever having vectorized anything into Shape[]. + const targetRunId = result?.runs?.[String(sliceIndex)] ?? committedRunIdForSlice; + if (!targetRunId || targetRunId === volumeApplyPreviewRunIdRef.current) return; + let cancelled = false; + volumeApplyPreviewRunIdRef.current = targetRunId; + void (async () => { + try { + const [commitRes, statusRes] = await Promise.all([ + fetch(ipredRunCommitUrl(targetRunId)), + fetch(ipredRunStatusUrl(targetRunId)), + ]); + if (!commitRes.ok || !statusRes.ok || cancelled) return; + const [commitBlob, statusBlob] = await Promise.all([commitRes.blob(), statusRes.blob()]); + if (cancelled) return; + const cUrl = URL.createObjectURL(commitBlob); + const sUrl = URL.createObjectURL(statusBlob); + if (commitUrlRef.current) URL.revokeObjectURL(commitUrlRef.current); + if (statusUrlRef.current) URL.revokeObjectURL(statusUrlRef.current); + commitUrlRef.current = cUrl; + statusUrlRef.current = sUrl; + setCommitUrl(cUrl); + setStatusUrl(sUrl); + } catch { + // Best-effort live preview only — a fetch hiccup here shouldn't + // surface a hard error; the explicit Commit step remains the + // authoritative path regardless of whether this preview loaded. + } + })(); + return () => { + cancelled = true; + }; + }, [sliceIndex, volumeApplyJobHook.state.result, committedRunIdForSlice]); + + const predict = useCallback( + async (_shapes: Shape[]) => { + if (!model || predicting) return; + setPredicting(true); + setError(null); + try { + const sessionId = await ensureSession(); + // Always the CURRENT slice's bank, not model.featureId (the slice the model + // happened to be trained on) — the model persists across slices, so predict + // must target whichever slice is on screen, computing a bank if missing. + const featureId = await ensureFeatureBank(); + const run = await ipredInfer({ + session_id: sessionId, + model_id: model.modelId, + feature_id: featureId, + alpha: params.alpha, + }); + const [commitRes, statusRes] = await Promise.all([ + fetch(ipredRunCommitUrl(run.run_id)), + fetch(ipredRunStatusUrl(run.run_id)), + ]); + if (!commitRes.ok || !statusRes.ok) { + throw new Error('Failed to fetch conformal prediction PNGs'); + } + const [commitBlob, statusBlob] = await Promise.all([commitRes.blob(), statusRes.blob()]); + const cUrl = URL.createObjectURL(commitBlob); + const sUrl = URL.createObjectURL(statusBlob); + if (commitUrlRef.current) URL.revokeObjectURL(commitUrlRef.current); + if (statusUrlRef.current) URL.revokeObjectURL(statusUrlRef.current); + commitUrlRef.current = cUrl; + statusUrlRef.current = sUrl; + setCommitUrl(cUrl); + setStatusUrl(sUrl); + setPredictCounts(run.counts); + setRunId(run.run_id); + const classIds = run.class_ids?.length ? run.class_ids : model.classIds; + setModel((m) => (m ? { ...m, classIds } : m)); + const thresholds: Record = {}; + for (const cid of classIds) thresholds[cid] = DEFAULT_PROBA_THRESHOLD; + setProbaThresholds(thresholds); + await loadProbaChannel(run.run_id, 0); + } catch (e) { + revokePredict(); + setError(e instanceof Error ? e.message : String(e)); + } finally { + setPredicting(false); + } + }, + [model, predicting, params.alpha, ensureSession, ensureFeatureBank, revokePredict, loadProbaChannel], + ); + + const selectProbaClass = useCallback( + async (index: number) => { + if (!runId || !model) return; + const n = model.classIds.length; + if (n < 1) return; + const next = ((index % n) + n) % n; + setError(null); + try { + await loadProbaChannel(runId, next); + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + } + }, + [runId, model, loadProbaChannel], + ); + + const cycleProbaClass = useCallback( + (delta: number) => { + void selectProbaClass(probaClassIndex + delta); + }, + [selectProbaClass, probaClassIndex], + ); + + const setProbaThreshold = useCallback( + (threshold: number) => { + if (!model) return; + const classId = model.classIds[probaClassIndex]; + if (classId === undefined) return; + const t = Math.min(1, Math.max(0, threshold)); + setProbaThresholds((prev) => ({ ...prev, [classId]: t })); + const raw = rawProbaBlobRef.current; + if (raw) { + void publishProbaPreview(raw, t).catch((e) => { + setError(e instanceof Error ? e.message : String(e)); + }); + } + }, + [model, probaClassIndex, publishProbaPreview], + ); + + const activeProbaClassId = model?.classIds[probaClassIndex] ?? null; + const activeProbaThreshold = + activeProbaClassId !== null + ? (probaThresholds[activeProbaClassId] ?? DEFAULT_PROBA_THRESHOLD) + : DEFAULT_PROBA_THRESHOLD; + + const saveThresholdedClass = useCallback(async (): Promise => { + if (!runId || activeProbaClassId === null || savingClass) return null; + setSavingClass(true); + setError(null); + try { + return await ipredThresholdClass(runId, { + class_id: activeProbaClassId, + threshold: activeProbaThreshold, + }); + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + return null; + } finally { + setSavingClass(false); + } + }, [runId, activeProbaClassId, activeProbaThreshold, savingClass]); + + const dismiss = useCallback(() => { + revokePredict(); + }, [revokePredict]); + + return { + params, + setParams, + model, + commitUrl, + statusUrl, + predictUrl: commitUrl, + predictCounts, + runId, + probaUrl, + probaClassIndex, + activeProbaClassId, + activeProbaThreshold, + selectProbaClass, + cycleProbaClass, + setProbaThreshold, + saveThresholdedClass, + savingClass, + training, + predicting, + error, + train, + predict, + dismiss, + compositionId: preferredCompositionId, + trainerId: preferredTrainerId, + canTrainWithoutJob: !!preferredCompositionId && !!source && !!kind, + // ---- Batch operations ---- + trainAcrossSlices, + multiTrainJob: multiTrainJobHook.state, + multiTraining: multiTrainJobHook.state.status === 'running', + resetMultiTrainJob: multiTrainJobHook.reset, + applyAcrossVolume, + volumeApplyJob: volumeApplyJobHook.state, + volumeApplying: volumeApplyJobHook.state.status === 'running', + resetVolumeApplyJob: volumeApplyJobHook.reset, + }; +} diff --git a/frontend/src/hooks/useSam.test.ts b/frontend/src/hooks/useSam.test.ts new file mode 100644 index 0000000..db1ce6f --- /dev/null +++ b/frontend/src/hooks/useSam.test.ts @@ -0,0 +1,200 @@ +/** + * useSam wraps samClient, a singleton that owns a real Web Worker (spawned via + * `new Worker(new URL('./samWorker.ts', import.meta.url), ...)`). jsdom doesn't + * implement Worker, and vendoring a fake Worker global would still leave us + * exercising the worker's postMessage protocol rather than the hook's own logic. + * So — matching the pattern already used in Toolbar/index.test.tsx, which mocks + * `@/hooks/useSam` wholesale rather than deal with the worker boundary — we mock + * `@/lib/sam/samClient` here instead and test useSam's actual public surface + * (status subscription, ensureEncoded's caching/dedup, segment) against it. + */ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook, waitFor } from '@testing-library/react'; + +const subscribers = new Set<(s: string) => void>(); +let currentStatus = 'idle'; +const encode = vi.fn(); +const decode = vi.fn(); +const init = vi.fn(); +const getStatus = vi.fn(() => currentStatus); +const getBackend = vi.fn(() => 'wasm'); + +vi.mock('@/lib/sam/samClient', () => ({ + samClient: { + getStatus: () => getStatus(), + getBackend: () => getBackend(), + subscribe: (fn: (s: string) => void) => { + subscribers.add(fn); + fn(currentStatus); + return () => subscribers.delete(fn); + }, + init: (...args: unknown[]) => init(...args), + encode: (...args: unknown[]) => encode(...args), + decode: (...args: unknown[]) => decode(...args), + }, + webgpuAvailable: () => false, +})); + +import { useSam } from './useSam'; + +function setStatus(s: string) { + currentStatus = s; + subscribers.forEach((fn) => fn(s)); +} + +beforeEach(() => { + currentStatus = 'idle'; + subscribers.clear(); + encode.mockReset().mockResolvedValue(undefined); + decode.mockReset().mockResolvedValue({ mask: new Uint8Array([1]), width: 1, height: 1, score: 0.9 }); + init.mockReset().mockResolvedValue(undefined); + getStatus.mockClear(); + getBackend.mockClear(); +}); + +afterEach(() => { + cleanup(); +}); + +describe('useSam', () => { + it('reflects samClient.getStatus() at mount and updates on subscribe notifications', () => { + const { result } = renderHook(() => useSam(false)); + expect(result.current.status).toBe('idle'); + expect(result.current.supported).toBe(true); + + act(() => setStatus('ready')); + expect(result.current.status).toBe('ready'); + }); + + it('supported is false once status flips to unsupported', () => { + const { result } = renderHook(() => useSam(false)); + act(() => setStatus('unsupported')); + expect(result.current.supported).toBe(false); + }); + + it('calls samClient.init() when enabled and status is idle', () => { + renderHook(() => useSam(true)); + expect(init).toHaveBeenCalledTimes(1); + }); + + it('does not call init() when disabled', () => { + renderHook(() => useSam(false)); + expect(init).not.toHaveBeenCalled(); + }); + + it('does not call init() again when status is already past idle', () => { + currentStatus = 'ready'; + renderHook(() => useSam(true)); + expect(init).not.toHaveBeenCalled(); + }); + + it('swallows an init() rejection without throwing', async () => { + init.mockRejectedValue(new Error('no webgpu')); + expect(() => renderHook(() => useSam(true))).not.toThrow(); + await waitFor(() => expect(init).toHaveBeenCalled()); + }); + + it('ensureEncoded encodes once per key and returns true on success', async () => { + const { result } = renderHook(() => useSam(false)); + const makeSource = vi.fn(() => ({}) as CanvasImageSource); + + let ok = false; + await act(async () => { + ok = await result.current.ensureEncoded('slice-0', makeSource); + }); + expect(ok).toBe(true); + expect(encode).toHaveBeenCalledTimes(1); + expect(makeSource).toHaveBeenCalledTimes(1); + + // Same key again -> cached, no new encode/makeSource call. + await act(async () => { + ok = await result.current.ensureEncoded('slice-0', makeSource); + }); + expect(ok).toBe(true); + expect(encode).toHaveBeenCalledTimes(1); + expect(makeSource).toHaveBeenCalledTimes(1); + }); + + it('ensureEncoded re-encodes when the key changes (e.g. brightness/contrast changed)', async () => { + const { result } = renderHook(() => useSam(false)); + const makeSource = vi.fn(() => ({}) as CanvasImageSource); + + await act(async () => { + await result.current.ensureEncoded('slice-0', makeSource); + }); + await act(async () => { + await result.current.ensureEncoded('slice-0:bc=1', makeSource); + }); + expect(encode).toHaveBeenCalledTimes(2); + }); + + it('ensureEncoded returns false and sets error on a failed encode, and does not retry a known-bad key', async () => { + encode.mockRejectedValue(new Error('encode failed')); + const { result } = renderHook(() => useSam(false)); + const makeSource = vi.fn(() => ({}) as CanvasImageSource); + + let ok = true; + await act(async () => { + ok = await result.current.ensureEncoded('bad-key', makeSource); + }); + expect(ok).toBe(false); + expect(result.current.error).toBe('encode failed'); + expect(encode).toHaveBeenCalledTimes(1); + + // Known-bad key: no retry, no second makeSource/encode call. + await act(async () => { + ok = await result.current.ensureEncoded('bad-key', makeSource); + }); + expect(ok).toBe(false); + expect(encode).toHaveBeenCalledTimes(1); + expect(makeSource).toHaveBeenCalledTimes(1); + }); + + it('ensureEncoded awaits a concurrent in-flight encode for the same key rather than starting a second one', async () => { + let resolveEncode!: () => void; + encode.mockReturnValue(new Promise((resolve) => { resolveEncode = resolve; })); + const { result } = renderHook(() => useSam(false)); + const makeSource = vi.fn(() => ({}) as CanvasImageSource); + + let p1: Promise; + let p2: Promise; + act(() => { + p1 = result.current.ensureEncoded('k', makeSource); + p2 = result.current.ensureEncoded('k', makeSource); + }); + expect(encode).toHaveBeenCalledTimes(1); + + await act(async () => { + resolveEncode(); + await Promise.all([p1, p2]); + }); + expect(encode).toHaveBeenCalledTimes(1); + }); + + it('segment returns null immediately when there are no points and no box', async () => { + const { result } = renderHook(() => useSam(false)); + const mask = await result.current.segment([], null, 'auto', 0); + expect(mask).toBeNull(); + expect(decode).not.toHaveBeenCalled(); + }); + + it('segment decodes points/box via samClient and returns the mask', async () => { + const { result } = renderHook(() => useSam(false)); + const mask = await result.current.segment([{ x: 1, y: 2, label: 1 }], null, 'fine', 0.5); + expect(decode).toHaveBeenCalledWith([{ x: 1, y: 2, label: 1 }], null, 'fine', 0.5); + expect(mask).toEqual({ mask: new Uint8Array([1]), width: 1, height: 1, score: 0.9 }); + }); + + it('segment returns null (not throw) when decode rejects', async () => { + decode.mockRejectedValue(new Error('decode failed')); + const { result } = renderHook(() => useSam(false)); + const mask = await result.current.segment([{ x: 1, y: 2, label: 1 }], null, 'auto', 0); + expect(mask).toBeNull(); + }); + + it('exposes backend from samClient.getBackend()', () => { + getBackend.mockReturnValue('webgpu'); + const { result } = renderHook(() => useSam(false)); + expect(result.current.backend).toBe('webgpu'); + }); +}); diff --git a/frontend/src/hooks/useSave.test.tsx b/frontend/src/hooks/useSave.test.tsx new file mode 100644 index 0000000..cabec40 --- /dev/null +++ b/frontend/src/hooks/useSave.test.tsx @@ -0,0 +1,159 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, cleanup, renderHook, waitFor } from '@testing-library/react'; +import { useSave } from './useSave'; +import { useAnnotationStore } from '@/stores/annotationStore'; +import { useClassStore } from '@/stores/classStore'; + +const SOURCE = 'local:x.tif'; + +beforeEach(() => { + useAnnotationStore.getState().reset(); + useClassStore.setState({ classes: [] }); + vi.stubGlobal('fetch', vi.fn().mockResolvedValue({ ok: true, json: async () => [] })); +}); + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +describe('useSave', () => { + it('starts clean (not dirty)', async () => { + const { result } = renderHook(() => useSave(SOURCE)); + await waitFor(() => expect(fetch).toHaveBeenCalled()); + expect(result.current.isDirty).toBe(false); + }); + + it('becomes dirty after a store change', async () => { + const { result } = renderHook(() => useSave(SOURCE)); + await waitFor(() => expect(fetch).toHaveBeenCalled()); + act(() => { + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + ]); + }); + expect(result.current.isDirty).toBe(true); + }); + + it('resets dirty/versions/lastSavedAt when sourceKey changes', async () => { + const { result, rerender } = renderHook(({ src }) => useSave(src), { initialProps: { src: SOURCE } }); + await waitFor(() => expect(fetch).toHaveBeenCalled()); + act(() => { + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 1, []); + }); + expect(result.current.isDirty).toBe(true); + + act(() => { rerender({ src: 'local:other.tif' }); }); + expect(result.current.isDirty).toBe(false); + }); + + it('buildSavePayload returns null with no sourceKey', () => { + const { result } = renderHook(() => useSave(null)); + expect(result.current.buildSavePayload()).toBeNull(); + expect(result.current.saveSummary).toEqual({ shapeCount: 0, classCount: 0 }); + }); + + it('saveSummary reflects real shape/class counts for the source', async () => { + act(() => { + useClassStore.setState({ classes: [{ classId: 1, label: 'Cell', color: '#f00', isVisible: true }] }); + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 1, [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }, + { id: 's2', classId: 1, kind: 'rectangle', x: 5, y: 5, w: 2, h: 2 }, + ]); + }); + const { result } = renderHook(() => useSave(SOURCE)); + expect(result.current.saveSummary).toEqual({ shapeCount: 2, classCount: 1 }); + }); + + it('refreshVersions populates the versions list', async () => { + (fetch as any).mockResolvedValue({ ok: true, json: async () => [{ version: 1, saved_at: 't', shape_count: 0, class_count: 0 }] }); + const { result } = renderHook(() => useSave(SOURCE)); + await waitFor(() => expect(result.current.versions).toHaveLength(1)); + }); + + it('save() POSTs the payload, clears dirty, and sets lastSavedAt', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => [] }) // initial refreshVersions on mount + .mockResolvedValueOnce({ ok: true, json: async () => ({ saved_at: '2024-01-01T00:00:00Z' }) }) // save + .mockResolvedValueOnce({ ok: true, json: async () => [{ version: 1, saved_at: '2024-01-01T00:00:00Z', shape_count: 0, class_count: 0 }] }); // post-save refresh + + const { result } = renderHook(() => useSave(SOURCE)); + await waitFor(() => expect(fetch).toHaveBeenCalledTimes(1)); + + let ok = false; + await act(async () => { + ok = await result.current.save({ annotatedBy: 'Alice', notes: 'n' }); + }); + expect(ok).toBe(true); + expect(result.current.lastSavedAt).toBe('2024-01-01T00:00:00Z'); + expect(result.current.isDirty).toBe(false); + + const [url, init] = (fetch as any).mock.calls[1]; + expect(url).toContain('/api/annotations/save?source_key='); + const body = JSON.parse(init.body); + expect(body.annotated_by).toBe('Alice'); + expect(body.notes).toBe('n'); + }); + + it('save() returns false and does not throw on a failed request', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => [] }) + .mockResolvedValueOnce({ ok: false, status: 500 }); + const { result } = renderHook(() => useSave(SOURCE)); + await waitFor(() => expect(fetch).toHaveBeenCalledTimes(1)); + let ok = true; + await act(async () => { + ok = await result.current.save(); + }); + expect(ok).toBe(false); + }); + + it('fetchVersionPayload caches results by version number', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => [] }) + .mockResolvedValueOnce({ ok: true, json: async () => ({ payload: { classes: [], slices: {}, split_by_slice: {}, negative_slices: [] } }) }); + const { result } = renderHook(() => useSave(SOURCE)); + await waitFor(() => expect(fetch).toHaveBeenCalledTimes(1)); + + let payload1; + await act(async () => { payload1 = await result.current.fetchVersionPayload(1); }); + expect(payload1).toEqual({ classes: [], slices: {}, split_by_slice: {}, negative_slices: [] }); + + const callsBefore = (fetch as any).mock.calls.length; + await act(async () => { await result.current.fetchVersionPayload(1); }); + expect((fetch as any).mock.calls.length).toBe(callsBefore); // cached, no new fetch + }); + + it('restoreVersion loads a version into the stores and marks dirty', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => [] }) + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ + payload: { + classes: [{ classId: 1, label: 'Restored', color: '#000', isVisible: true }], + slices: { '0': [{ id: 'r1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 1, h: 1 }] }, + split_by_slice: {}, + negative_slices: [], + }, + }), + }); + const { result } = renderHook(() => useSave(SOURCE)); + await waitFor(() => expect(fetch).toHaveBeenCalledTimes(1)); + await act(async () => { await result.current.restoreVersion(1); }); + + expect(useClassStore.getState().classes[0].label).toBe('Restored'); + expect(useAnnotationStore.getState().byImage[SOURCE]?.['0']).toHaveLength(1); + expect(result.current.isDirty).toBe(true); + }); + + it('markClean clears the dirty flag', async () => { + const { result } = renderHook(() => useSave(SOURCE)); + await waitFor(() => expect(fetch).toHaveBeenCalled()); + act(() => { + useAnnotationStore.getState().replaceClassShapesOnSlice(SOURCE, 0, 1, []); + }); + act(() => result.current.markClean()); + expect(result.current.isDirty).toBe(false); + }); +}); diff --git a/frontend/src/hooks/useSave.ts b/frontend/src/hooks/useSave.ts index deb7d55..a292aac 100644 --- a/frontend/src/hooks/useSave.ts +++ b/frontend/src/hooks/useSave.ts @@ -84,6 +84,21 @@ export function useSave(sourceKey: string | null): UseSaveReturn { const cleanRef = useRef(true); const payloadCacheRef = useRef>(new Map()); + // Must run BEFORE the dirty-tracking effect below on the mount/sourceKey-change + // commit: it arms cleanRef so that same-commit effect can consume it and treat + // the freshly-loaded source's initial data as clean. Effects run in + // declaration order, so this one is declared first for that reason — + // reversing the order would re-arm cleanRef right after the dirty-tracker + // had already consumed it, silently swallowing the next real edit's dirty + // flag (a bug this ordering fixes). + useEffect(() => { + cleanRef.current = true; + setIsDirty(false); + setLastSavedAt(null); + setVersions([]); + payloadCacheRef.current.clear(); + }, [sourceKey]); + useEffect(() => { if (cleanRef.current) { cleanRef.current = false; @@ -93,14 +108,6 @@ export function useSave(sourceKey: string | null): UseSaveReturn { // eslint-disable-next-line react-hooks/exhaustive-deps }, [byImage, splitBySlice, negativeSlices, classes]); - useEffect(() => { - cleanRef.current = true; - setIsDirty(false); - setLastSavedAt(null); - setVersions([]); - payloadCacheRef.current.clear(); - }, [sourceKey]); - /** Assemble the current stores into a draft payload, or null if no sourceKey. */ const buildSavePayload = useCallback((): SaveDraftPayload | null => { if (!sourceKey) return null; diff --git a/frontend/src/hooks/useTrainCapability.test.tsx b/frontend/src/hooks/useTrainCapability.test.tsx new file mode 100644 index 0000000..5ffe5d4 --- /dev/null +++ b/frontend/src/hooks/useTrainCapability.test.tsx @@ -0,0 +1,65 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { renderHook, waitFor } from '@testing-library/react'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import type { ReactNode } from 'react'; +import { useTrainCapability } from './useTrainCapability'; + +function wrapper({ children }: { children: ReactNode }) { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return {children}; +} + +beforeEach(() => { + vi.stubGlobal('fetch', vi.fn()); +}); + +afterEach(() => { + vi.unstubAllGlobals(); +}); + +describe('useTrainCapability', () => { + it('returns the fallback shape while loading', () => { + (fetch as any).mockReturnValue(new Promise(() => {})); + const { result } = renderHook(() => useTrainCapability(), { wrapper }); + expect(result.current.isLoading).toBe(true); + expect(result.current.capability.torch_available).toBe(false); + expect(result.current.capability.dlsia).toEqual({ available: false }); + }); + + it('returns the parsed capability once loaded', async () => { + (fetch as any).mockResolvedValue({ + ok: true, + json: async () => ({ + torch_available: true, + torch_version: '2.1.0', + device: 'cpu', + dlsia: { available: true }, + denoise: { available: true, methods: [] }, + runs_dir: '/data/runs', + busy: false, + }), + }); + const { result } = renderHook(() => useTrainCapability(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + expect(result.current.capability.torch_available).toBe(true); + expect(result.current.capability.device).toBe('cpu'); + }); + + it('fills in a missing nested key (e.g. denoise) from the fallback instead of throwing', async () => { + (fetch as any).mockResolvedValue({ + ok: true, + json: async () => ({ torch_available: true, dlsia: { available: true } }), + }); + const { result } = renderHook(() => useTrainCapability(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + expect(result.current.capability.denoise).toEqual({ available: false, methods: [] }); + expect(result.current.capability.torch_available).toBe(true); + }); + + it('falls back to defaults entirely on a fetch failure', async () => { + (fetch as any).mockResolvedValue({ ok: false, status: 500 }); + const { result } = renderHook(() => useTrainCapability(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + expect(result.current.capability.torch_available).toBe(false); + }); +}); diff --git a/frontend/src/hooks/useTrainCapability.ts b/frontend/src/hooks/useTrainCapability.ts new file mode 100644 index 0000000..259b356 --- /dev/null +++ b/frontend/src/hooks/useTrainCapability.ts @@ -0,0 +1,86 @@ +/** + * useTrainCapability — Train-tab readiness: torch/dlsia availability and device. + * + * No DINOv3 checkpoint discovery here — that model family is deferred (see + * Phase 5.5 in the integration plan) and the backend's capability() probe + * doesn't report it at all. + */ +import { useQuery } from '@tanstack/react-query'; +import { API_BASE } from '@/config'; + +/** A classical denoise filter the server can actually run. `available` is + * per-method because some need an optional dependency (wavelet ⇢ PyWavelets), + * and `cost` drives whether the UI warns before a full-resolution preview. */ +export interface DenoiseMethodInfo { + method: string; + label: string; + cost: 'cheap' | 'moderate' | 'slow'; + description: string; + available: boolean; + z_radius: number; +} + +export interface TrainCapability { + torch_available: boolean; + torch_version: string | null; + device: string | null; + dlsia: { available: boolean }; + denoise: { available: boolean; methods: DenoiseMethodInfo[] }; + runs_dir: string; + busy: boolean; + error?: string; +} + +const FALLBACK: TrainCapability = { + torch_available: false, + torch_version: null, + device: null, + dlsia: { available: false }, + denoise: { available: false, methods: [] }, + runs_dir: '', + busy: false, +}; + +/** + * Merge a server response over FALLBACK so a key this build expects but the + * server doesn't send resolves to a safe default instead of `undefined`. + * + * `data ?? FALLBACK` alone only covers "no response at all". It does NOT cover + * a response that's missing a key — and consumers reach into nested keys + * (`capability.denoise.methods`), so a single absent key throws a TypeError + * mid-render, which unmounts the React tree and white-screens the whole app. + * That's reachable whenever the frontend is newer than the running backend: a + * dev server still running a process started before a capability key was + * added, or a stale deploy. A missing capability should degrade to "that + * feature is unavailable", never take down the page. + */ +function withDefaults(data: Partial | undefined): TrainCapability { + if (!data) return FALLBACK; + return { + ...FALLBACK, + ...data, + dlsia: { ...FALLBACK.dlsia, ...(data.dlsia ?? {}) }, + denoise: { ...FALLBACK.denoise, ...(data.denoise ?? {}) }, + }; +} + +/** Polls /api/train/capability every 10s — cheap enough to keep the + * CapabilityBanner and "busy" state current while the Train tab is open. */ +export function useTrainCapability() { + const query = useQuery({ + queryKey: ['trainCapability'], + queryFn: async ({ signal }) => { + const res = await fetch(`${API_BASE}/api/train/capability`, { signal }); + if (!res.ok) throw new Error(`Capability check failed: ${res.status}`); + return res.json(); + }, + staleTime: 5_000, + refetchInterval: 10_000, + }); + + return { + capability: withDefaults(query.data), + isLoading: query.isLoading, + refetch: query.refetch, + }; +} diff --git a/frontend/src/hooks/useTrainRuns.test.tsx b/frontend/src/hooks/useTrainRuns.test.tsx new file mode 100644 index 0000000..d9e16f0 --- /dev/null +++ b/frontend/src/hooks/useTrainRuns.test.tsx @@ -0,0 +1,111 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { act, renderHook, waitFor } from '@testing-library/react'; +import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; +import type { ReactNode } from 'react'; +import { useTrainRuns } from './useTrainRuns'; + +function wrapper({ children }: { children: ReactNode }) { + const client = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return {children}; +} + +beforeEach(() => { + vi.stubGlobal('fetch', vi.fn()); +}); + +afterEach(() => { + vi.unstubAllGlobals(); +}); + +const RUN = { + run_id: 'r1', model_family: 'dlsia_tunet', model_config: {}, classes: [], render: {}, + image_size: 256, hyperparams: {}, source_keys: [], created_at: '2024-01-01', metrics: {}, +}; + +describe('useTrainRuns', () => { + it('starts empty and loading', () => { + (fetch as any).mockReturnValue(new Promise(() => {})); + const { result } = renderHook(() => useTrainRuns(), { wrapper }); + expect(result.current.isLoading).toBe(true); + expect(result.current.runs).toEqual([]); + }); + + it('returns the loaded runs list', async () => { + (fetch as any).mockResolvedValue({ ok: true, json: async () => ({ runs: [RUN] }) }); + const { result } = renderHook(() => useTrainRuns(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + expect(result.current.runs).toEqual([RUN]); + }); + + it('returns an empty list (not an error) on a failed fetch', async () => { + (fetch as any).mockResolvedValue({ ok: false, status: 500 }); + const { result } = renderHook(() => useTrainRuns(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + expect(result.current.runs).toEqual([]); + }); + + it('defaults to an empty list when the response has no runs key', async () => { + (fetch as any).mockResolvedValue({ ok: true, json: async () => ({}) }); + const { result } = renderHook(() => useTrainRuns(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + expect(result.current.runs).toEqual([]); + }); + + it('deleteRun issues a DELETE and refreshes the list on success', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => ({ runs: [RUN] }) }) + .mockResolvedValueOnce({ ok: true }) + .mockResolvedValueOnce({ ok: true, json: async () => ({ runs: [] }) }); + const { result } = renderHook(() => useTrainRuns(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + + let ok: boolean = false; + await act(async () => { + ok = await result.current.deleteRun('r1'); + }); + expect(ok).toBe(true); + expect(fetch).toHaveBeenCalledWith(expect.stringContaining('/api/train/runs/r1'), { method: 'DELETE' }); + await waitFor(() => expect(result.current.runs).toEqual([])); + }); + + it('deleteRun returns false on a failed delete', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => ({ runs: [RUN] }) }) + .mockResolvedValueOnce({ ok: false }); + const { result } = renderHook(() => useTrainRuns(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + + let ok: boolean = true; + await act(async () => { + ok = await result.current.deleteRun('r1'); + }); + expect(ok).toBe(false); + }); + + it('deleteRun returns false when the request throws', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => ({ runs: [RUN] }) }) + .mockRejectedValueOnce(new Error('network down')); + const { result } = renderHook(() => useTrainRuns(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + + let ok: boolean = true; + await act(async () => { + ok = await result.current.deleteRun('r1'); + }); + expect(ok).toBe(false); + }); + + it('URL-encodes the run id in the delete request', async () => { + (fetch as any) + .mockResolvedValueOnce({ ok: true, json: async () => ({ runs: [] }) }) + .mockResolvedValueOnce({ ok: true }) + .mockResolvedValueOnce({ ok: true, json: async () => ({ runs: [] }) }); + const { result } = renderHook(() => useTrainRuns(), { wrapper }); + await waitFor(() => expect(result.current.isLoading).toBe(false)); + await act(async () => { + await result.current.deleteRun('run/with slash'); + }); + expect(fetch).toHaveBeenCalledWith(expect.stringContaining(encodeURIComponent('run/with slash')), expect.anything()); + }); +}); diff --git a/frontend/src/hooks/useTrainRuns.ts b/frontend/src/hooks/useTrainRuns.ts new file mode 100644 index 0000000..af34a29 --- /dev/null +++ b/frontend/src/hooks/useTrainRuns.ts @@ -0,0 +1,72 @@ +/** + * useTrainRuns — saved fine-tune runs (dlsia_tunet segmentation + the + * dlsia_denoiser family), for the inference run-picker. Invalidate + * `['trainRuns']` after a training job completes so a freshly-saved run shows + * up without a manual refresh. + */ +import { useQuery, useQueryClient } from '@tanstack/react-query'; +import { API_BASE } from '@/config'; +import type { AnnotationClass } from '@/stores/classStore'; + +export interface TrainRun { + run_id: string; + // Discriminator literals from backend/schemas.py's ModelConfig union. + model_family: 'dlsia_tunet' | 'dlsia_denoiser'; + model_config: Record; + classes: AnnotationClass[]; + render: Record; + /** + * Denoising baked into this run's INPUT pixels at training time (see + * schemas.DenoiseTrainOpts). Null for a run trained on raw pixels, and + * absent entirely on runs saved before the option existed — hence optional + * as well as nullable, so `run.denoise &&` is the only safe test. + * + * Read-only from the frontend's side: inference reapplies it off the run + * itself, never off the request, so nothing in the app should offer it as a + * choice at predict time. It's surfaced (RunsPanel) purely so two runs that + * expect different input can't look identical. + */ + denoise?: { method: string; strength: number } | null; + image_size: number; + hyperparams: Record; + source_keys: string[]; + created_at: string; + metrics: { + epochs_completed?: number; + final_train_loss?: number | null; + final_val_loss?: number | null; + val_miou?: number | null; + cancelled?: boolean; + }; +} + +export function useTrainRuns() { + const queryClient = useQueryClient(); + + const query = useQuery({ + queryKey: ['trainRuns'], + queryFn: async ({ signal }) => { + const res = await fetch(`${API_BASE}/api/train/runs`, { signal }); + if (!res.ok) return []; + const data = await res.json(); + return data.runs ?? []; + }, + staleTime: 10_000, + }); + + const invalidate = () => queryClient.invalidateQueries({ queryKey: ['trainRuns'] }); + + /** Permanently deletes a saved run (config, metrics, and weights) and refreshes the list. */ + const deleteRun = async (runId: string): Promise => { + try { + const res = await fetch(`${API_BASE}/api/train/runs/${encodeURIComponent(runId)}`, { method: 'DELETE' }); + if (!res.ok) return false; + invalidate(); + return true; + } catch { + return false; + } + }; + + return { runs: query.data ?? [], isLoading: query.isLoading, invalidate, deleteRun }; +} diff --git a/frontend/src/lib/apiError.test.ts b/frontend/src/lib/apiError.test.ts new file mode 100644 index 0000000..7323549 --- /dev/null +++ b/frontend/src/lib/apiError.test.ts @@ -0,0 +1,101 @@ +import { describe, expect, it } from 'vitest'; +import { formatApiError } from './apiError'; + +describe('formatApiError', () => { + it('formats the real 422 a bad patch size produces', () => { + const body = JSON.stringify({ + detail: [{ + type: 'greater_than_equal', + loc: ['body', 'model', 'dinov3_lora', 'hyperparams', 'image_size'], + msg: 'Input should be greater than or equal to 224', + input: 32, + ctx: { ge: 224 }, + }], + }); + + expect(formatApiError(body)).toBe('image size: Input should be greater than or equal to 224'); + }); + + it('passes through a plain string detail (HTTPException)', () => { + expect(formatApiError(JSON.stringify({ detail: 'Another job is already running' }))) + .toBe('Another job is already running'); + }); + + it('joins several validation errors', () => { + const body = JSON.stringify({ + detail: [ + { loc: ['body', 'epochs'], msg: 'Input should be less than or equal to 500' }, + { loc: ['body', 'lr'], msg: 'Input should be greater than 0' }, + ], + }); + + expect(formatApiError(body)).toBe( + 'epochs: Input should be less than or equal to 500; lr: Input should be greater than 0', + ); + }); + + it('omits the field when loc carries no usable name', () => { + expect(formatApiError(JSON.stringify({ detail: [{ msg: 'Something went wrong' }] }))) + .toBe('Something went wrong'); + }); + + it('falls back to raw text for a non-JSON body', () => { + expect(formatApiError('Internal Server Error')).toBe('Internal Server Error'); + }); + + it('falls back for JSON that is not a FastAPI error', () => { + expect(formatApiError(JSON.stringify({ oops: true }))).toBe('{"oops":true}'); + }); + + it('uses the fallback for an empty body', () => { + expect(formatApiError('', 'Request failed (500).')).toBe('Request failed (500).'); + expect(formatApiError(' ', 'Request failed (500).')).toBe('Request failed (500).'); + }); + + it('truncates a runaway body', () => { + expect(formatApiError('x'.repeat(5000)).length).toBe(500); + }); + + it('marks a truncated body with an ellipsis, not a silent cut', () => { + const result = formatApiError('x'.repeat(5000)); + expect(result.endsWith('…')).toBe(true); + }); + + it('does not add an ellipsis when the body is short enough as-is', () => { + expect(formatApiError('short message').endsWith('…')).toBe(false); + }); + + it('handles a detail that is a list of plain strings, not validation objects', () => { + expect(formatApiError(JSON.stringify({ detail: ['Field A is required', 'Field B is invalid'] }))) + .toBe('Field A is required; Field B is invalid'); + }); + + it('handles a single validation-error object, not wrapped in a list', () => { + const body = JSON.stringify({ + detail: { loc: ['body', 'image_size'], msg: 'Input should be greater than or equal to 224' }, + }); + expect(formatApiError(body)).toBe('image size: Input should be greater than or equal to 224'); + }); + + it('names a whole list element by its index when loc ends in a number', () => { + const body = JSON.stringify({ + detail: [{ loc: ['body', 'sources', 2], msg: 'field required' }], + }); + expect(formatApiError(body)).toBe('sources[2]: field required'); + }); + + it('falls back to a bare "item N" when a numeric loc has no preceding field name', () => { + const body = JSON.stringify({ detail: [{ loc: [3], msg: 'field required' }] }); + expect(formatApiError(body)).toBe('item 3: field required'); + }); + + it('still finds the leaf field name when a numeric index sits before it, not after', () => { + // The index (which source) isn't the leaf here — "color" is — so this must + // keep reporting the same leaf-only name it always has, unaffected by the + // numeric-leaf handling added for the case above. + const body = JSON.stringify({ + detail: [{ loc: ['body', 'sources', 2, 'color'], msg: 'invalid color' }], + }); + expect(formatApiError(body)).toBe('color: invalid color'); + }); +}); diff --git a/frontend/src/lib/apiError.ts b/frontend/src/lib/apiError.ts new file mode 100644 index 0000000..caaee9e --- /dev/null +++ b/frontend/src/lib/apiError.ts @@ -0,0 +1,87 @@ +/** + * apiError — turn a FastAPI error body into something a user can act on. + * + * FastAPI reports a rejected request as `{"detail": [{loc, msg, type, ...}]}`. + * Rendered raw that reads as a wall of JSON, which is what the job progress bar + * used to show. This extracts the field name and message instead. + */ + +interface ValidationItem { + loc?: unknown[]; + msg?: string; +} + +const TRUNCATE_AT = 500; + +/** Cut *text* to at most `limit` characters, marking that it was cut — a bare + * slice reads as the whole message, hiding that the rest was thrown away. */ +function truncate(text: string, limit = TRUNCATE_AT): string { + return text.length <= limit ? text : `${text.slice(0, limit - 1)}…`; +} + +/** + * Human-readable name for a pydantic `loc`, e.g. `image_size`, or `sources[2]` + * when the location is a whole list element rather than one of its fields + * (there's no field name to fall back on there, just the index). + */ +function fieldName(loc: unknown[] | undefined): string | null { + if (!Array.isArray(loc)) return null; + // Drop the leading "body" and any discriminated-union tag; a bare numeric + // index (a list position) is kept only long enough to attach it to the + // string segment right before it, then treated as the leaf itself. + const parts = loc.filter( + (p): p is string | number => p !== 'body' && (typeof p === 'string' || typeof p === 'number'), + ); + const leaf = parts[parts.length - 1]; + if (leaf === undefined) return null; + if (typeof leaf === 'number') { + const prev = parts[parts.length - 2]; + return (typeof prev === 'string' ? `${prev}[${leaf}]` : `item ${leaf}`).replace(/_/g, ' '); + } + return leaf.replace(/_/g, ' '); +} + +/** One `{msg, loc}`-shaped validation item (or a plain string) to a display line. */ +function formatValidationItem(item: unknown): string | null { + if (typeof item === 'string') return item.trim() || null; + const msg = typeof (item as ValidationItem)?.msg === 'string' ? (item as ValidationItem).msg! : ''; + if (!msg) return null; + const field = fieldName((item as ValidationItem)?.loc); + return field ? `${field}: ${msg}` : msg; +} + +/** + * Best-effort human-readable message for a non-OK API response body. + * + * Falls back to the raw text (trimmed) when the body isn't a shape we recognise, + * so nothing is ever swallowed — an unexpected error still reaches the user. + */ +export function formatApiError(body: string, fallback = 'Request failed.'): string { + const text = (body ?? '').trim(); + if (!text) return fallback; + + let parsed: unknown; + try { + parsed = JSON.parse(text); + } catch { + return truncate(text); + } + + const detail = (parsed as { detail?: unknown })?.detail; + if (typeof detail === 'string' && detail.trim()) return detail; + + // The common shape: a list of pydantic validation errors (or, less commonly, + // a list of plain message strings — some routes raise HTTPException with one). + if (Array.isArray(detail)) { + const messages = detail.map(formatValidationItem).filter((m): m is string => Boolean(m)); + if (messages.length) return messages.join('; '); + } + + // A single validation error/object, not wrapped in a list — same shape either way. + if (detail && typeof detail === 'object') { + const message = formatValidationItem(detail); + if (message) return message; + } + + return truncate(text); +} diff --git a/frontend/src/lib/bandTransferFunction.test.ts b/frontend/src/lib/bandTransferFunction.test.ts new file mode 100644 index 0000000..b8adee2 --- /dev/null +++ b/frontend/src/lib/bandTransferFunction.test.ts @@ -0,0 +1,100 @@ +import { describe, expect, it } from 'vitest'; +import { buildBandOpacityCurve, mapByteBandToViewerDomain } from './bandTransferFunction'; + +/** Every curve must have strictly increasing x across its points. */ +function expectMonotonic(points: readonly (readonly [number, number])[]) { + for (let i = 1; i < points.length; i++) { + expect(points[i][0]).toBeGreaterThan(points[i - 1][0]); + } +} + +/** Linear interpolation, mirroring how a GPU sampler would read the curve. */ +function opacityAt(points: readonly (readonly [number, number])[], x: number): number { + if (x <= points[0][0]) return points[0][1]; + for (let i = 1; i < points.length; i++) { + if (x <= points[i][0]) { + const [x0, o0] = points[i - 1]; + const [x1, o1] = points[i]; + const t = (x - x0) / (x1 - x0); + return o0 + (o1 - o0) * t; + } + } + return points[points.length - 1][1]; +} + +describe('buildBandOpacityCurve', () => { + it('is opaque in the middle of the band and transparent well outside it', () => { + const points = buildBandOpacityCurve(80 / 255, 180 / 255); + expectMonotonic(points); + expect(opacityAt(points, 130 / 255)).toBeCloseTo(1, 5); + expect(opacityAt(points, 0)).toBeCloseTo(0, 5); + expect(opacityAt(points, 1)).toBeCloseTo(0, 5); + }); + + it('stays within the 0-1 domain the viewer expects', () => { + const points = buildBandOpacityCurve(80 / 255, 180 / 255); + for (const [x] of points) { + expect(x).toBeGreaterThanOrEqual(0); + expect(x).toBeLessThanOrEqual(1); + } + }); + + it('swaps a reversed lo/hi so the band is always well-formed', () => { + const swapped = buildBandOpacityCurve(180 / 255, 80 / 255); + const normal = buildBandOpacityCurve(80 / 255, 180 / 255); + expect(swapped).toEqual(normal); + }); + + it('handles a band starting at 0 without producing non-monotonic points', () => { + const points = buildBandOpacityCurve(0, 100 / 255); + expectMonotonic(points); + expect(opacityAt(points, 50 / 255)).toBeCloseTo(1, 5); + }); + + it('handles a band ending at 1 without producing non-monotonic points', () => { + const points = buildBandOpacityCurve(150 / 255, 1); + expectMonotonic(points); + expect(opacityAt(points, 1)).toBeCloseTo(1, 5); + }); + + it('handles the full-range band (0-1) as fully opaque throughout', () => { + const points = buildBandOpacityCurve(0, 1); + expectMonotonic(points); + expect(opacityAt(points, 0.5)).toBeCloseTo(1, 5); + }); +}); + +describe('mapByteBandToViewerDomain', () => { + it('converts a byte band through the global range into the viewer range', () => { + // Global range matches the real petiole sample from live testing. + const globalRange: [number, number] = [-73.0, 71.3]; + const viewerRange: [number, number] = [-40, 40]; + const [lo, hi] = mapByteBandToViewerDomain(194, 215, globalRange, viewerRange); + // raw ≈ -73 + (194/255)*144.3 ≈ 36.8 ; -73 + (215/255)*144.3 ≈ 48.6 + // domain ≈ (36.8+40)/80 ≈ 0.960 ; (48.6+40)/80 clamped to 1 + expect(lo).toBeCloseTo(0.96, 1); + expect(hi).toBe(1); + }); + + it('falls back to treating the byte scale as already-physical when globalRange is missing', () => { + const viewerRange: [number, number] = [0, 255]; + const [lo, hi] = mapByteBandToViewerDomain(0, 255, null, viewerRange); + expect(lo).toBeCloseTo(0, 5); + expect(hi).toBeCloseTo(1, 5); + }); + + it('clamps to [0,1] when the converted value falls outside the viewer range', () => { + const globalRange: [number, number] = [0, 255]; + const viewerRange: [number, number] = [100, 200]; + const [lo, hi] = mapByteBandToViewerDomain(0, 255, globalRange, viewerRange); + expect(lo).toBe(0); + expect(hi).toBe(1); + }); + + it('ignores a degenerate (zero-span or inverted) global range and falls back to [0,255]', () => { + const viewerRange: [number, number] = [0, 255]; + const [lo, hi] = mapByteBandToViewerDomain(0, 255, [10, 10], viewerRange); + expect(lo).toBeCloseTo(0, 5); + expect(hi).toBeCloseTo(1, 5); + }); +}); diff --git a/frontend/src/lib/bandTransferFunction.ts b/frontend/src/lib/bandTransferFunction.ts new file mode 100644 index 0000000..62f75fa --- /dev/null +++ b/frontend/src/lib/bandTransferFunction.ts @@ -0,0 +1,98 @@ +/** + * bandTransferFunction — bridges the Sampler lasso's fitted intensity band + * (`thresholdFit.ts`'s `BandFit.lo`/`hi`) into the 3D viewer's opacity-curve + * transfer function, so a sparse, materially-distinct feature can be + * isolated in 3D by intensity alone — opaque inside the band, transparent + * outside — with zero upstream viewer changes (see the vendored + * `OpacityPoint = [intensity: 0-1, opacity: 0-1]` curve format). + * + * Deliberately narrow in scope: this only isolates by intensity VALUE across + * the whole volume, not by the traced SPATIAL region specifically. Good + * enough for a feature whose density is genuinely distinct from its + * surroundings (the same property the 2D threshold-lasso tool already + * exploits) — see the plan's own trade-off note. + * + * A real gotcha this module exists to fix: `BandFit.lo`/`hi` are 0-255 BYTE + * values from the 2D Annotate canvas's own PNG rendering — normalized via + * `GET /api/image/slice`'s norm="global" percentile stretch (see + * `ImageMeta.globalValueRange`), NOT the volume's raw physical intensity. + * The 3D viewer separately normalizes raw voxel values into its own [0,1] + * domain against a DIFFERENT per-dataset range (`WebGpuViewerInstance + * .getValueRange()`, a min/max-based estimate — see that method's own doc). + * These two ranges are typically close but never identical (different + * statistics, computed from different data), so a fitted band must be + * converted byte -> raw physical value -> 3D domain, not just divided by + * 255 — dividing by 255 assumes the two normalizations are the same, which + * silently produces a wildly wrong band (confirmed live: a materially-sparse + * 2D-highlighted feature came out as almost the entire volume in 3D). + */ + +export type OpacityPoint = readonly [intensity: number, opacity: number]; + +function clamp01(v: number): number { + return Math.min(1, Math.max(0, v)); +} + +/** + * Converts a fitted band's byte bounds (0-255, the 2D canvas's own + * percentile-stretched display space) into the 3D viewer's [0,1] domain. + * + * `globalRange` is `ImageMeta.globalValueRange` — the `[vmin, vmax]` the 2D + * backend actually stretched into 0-255; falls back to treating the byte + * scale as already-physical (`[0, 255]`) if unavailable (e.g. an RGB source, + * which the backend doesn't compute this for), matching this module's older, + * pre-conversion behavior. + * + * `viewerRange` is `WebGpuViewerInstance.getValueRange()` — the actual range + * the 3D viewer normalizes raw voxels against for this dataset. + */ +export function mapByteBandToViewerDomain( + loByte: number, + hiByte: number, + globalRange: readonly [number, number] | null | undefined, + viewerRange: readonly [number, number], +): [number, number] { + const [gMin, gMax] = + globalRange && Number.isFinite(globalRange[0]) && Number.isFinite(globalRange[1]) && globalRange[1] > globalRange[0] + ? globalRange + : [0, 255]; + const toRaw = (byte: number) => gMin + (byte / 255) * (gMax - gMin); + + const [vMin, vMax] = viewerRange; + const span = vMax - vMin || 1; + const toDomain = (raw: number) => clamp01((raw - vMin) / span); + + return [toDomain(toRaw(loByte)), toDomain(toRaw(hiByte))]; +} + +/** + * Builds an opacity curve that ramps to fully opaque across `[lo01, hi01]` + * (already in the viewer's own [0,1] domain — see `mapByteBandToViewerDomain`) + * and stays transparent everywhere else, with a small ramp at each edge + * (rather than a hard step) so the transfer function doesn't alias. + */ +export function buildBandOpacityCurve(lo01: number, hi01: number): OpacityPoint[] { + const EPS = 0.004; + const loN = clamp01(Math.min(lo01, hi01)); + const hiN = clamp01(Math.max(lo01, hi01)); + + const raw: OpacityPoint[] = [ + [0, 0], + [loN, loN <= EPS ? 1 : 0], + [Math.min(1, loN + EPS), 1], + [Math.max(0, hiN - EPS), 1], + [hiN, hiN >= 1 - EPS ? 1 : 0], + [1, hiN >= 1 - EPS ? 1 : 0], + ]; + + // Points must have strictly increasing x for the viewer's curve to be + // well-defined — clamping near 0/1 can collapse several of the above onto + // the same x. Sort, then keep only the first point at each x. + const sorted = [...raw].sort((a, b) => a[0] - b[0]); + const points: OpacityPoint[] = []; + for (const p of sorted) { + if (points.length > 0 && p[0] <= points[points.length - 1][0]) continue; + points.push(p); + } + return points; +} diff --git a/frontend/src/lib/bboxPrefilter.equivalence.test.ts b/frontend/src/lib/bboxPrefilter.equivalence.test.ts new file mode 100644 index 0000000..f8b4e1e --- /dev/null +++ b/frontend/src/lib/bboxPrefilter.equivalence.test.ts @@ -0,0 +1,326 @@ +/** + * Equivalence tests for the bounding-box pre-filters. + * + * The commit path skips shapes whose bounds cannot interact with the new + * geometry, which is only sound if it changes nothing. These tests pin that down + * by re-implementing the ORIGINAL unfiltered algorithms here as reference oracles + * and asserting the optimized versions agree with them, over randomized shape sets + * covering the cases the filter could plausibly get wrong: disjoint, exactly + * touching, nested, tiny-vs-huge, and brush-vs-vector mixes. + * + * If a future change to the filter makes it too aggressive, these fail. + */ +import { describe, it, expect } from 'vitest'; +import { clipShapesToOthers, clipShapesToOthersMask } from './clipToClasses'; +import { mergeNewWithSameClass, expandSameClassOverlap } from './mergeSameClass'; +import { gridFor, fullResGridFor, rasterizeShapes, rasterizeUnion } from './rasterize'; +import { maskToPolygonsWithHoles } from './magicwand'; +import { unionShapesToMultiPolygon, shapeToMultiPolygon, multiPolygonToShapes } from './polybool'; +import polygonClipping from 'polygon-clipping'; +import type { PolygonShape, Shape } from '@/stores/annotationStore'; + +const W = 200, H = 200; + +// ---- Deterministic PRNG so failures reproduce exactly ---------------------- +function mulberry32(seed: number) { + return () => { + seed |= 0; seed = (seed + 0x6d2b79f5) | 0; + let t = Math.imul(seed ^ (seed >>> 15), 1 | seed); + t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t; + return ((t ^ (t >>> 14)) >>> 0) / 4294967296; + }; +} + +/** A grab-bag of shape kinds at a given position/size. */ +function makeShape(rnd: () => number, id: string, classId: number, x: number, y: number, size: number): Shape { + const kind = Math.floor(rnd() * 4); + if (kind === 0) return { id, classId, kind: 'rectangle', x, y, w: size, h: size }; + if (kind === 1) return { id, classId, kind: 'ellipse', cx: x + size / 2, cy: y + size / 2, rx: size / 2, ry: size / 2 }; + if (kind === 2) { + return { + id, classId, kind: 'polygon', + points: [x, y, x + size, y, x + size, y + size, x, y + size], + }; + } + return { + id, classId, kind: 'brush', + strokes: [{ points: [x + 2, y + 2, x + size - 2, y + size - 2], radius: Math.max(2, size / 4), mode: 'paint' }], + }; +} + +/** Random slice populated with shapes across two classes, mixing overlaps and gaps. */ +function randomSlice(seed: number, n: number): Shape[] { + const rnd = mulberry32(seed); + const out: Shape[] = []; + for (let i = 0; i < n; i++) { + const size = 8 + Math.floor(rnd() * 40); + const x = Math.floor(rnd() * (W - size)); + const y = Math.floor(rnd() * (H - size)); + out.push(makeShape(rnd, `s${i}`, rnd() < 0.5 ? 1 : 2, x, y, size)); + } + return out; +} + +/** + * Compare clip results by the region they cover, not by their vertex lists. + * + * Vertex-list equality is the wrong oracle. When nothing is nearby, the optimized + * path skips `polygonClipping.difference` entirely; the naive path still runs it + * against distant polygons, and polygon-clipping re-emits the ring — from a + * different start vertex, and with coordinates perturbed in the last few decimal + * places. On a 200×200 ellipse that shows up as a *single* boundary pixel out of + * ~520 (0.19%). + * + * Note the direction of that error: the naive path is the one perturbing the + * geometry, by round-tripping a shape through a boolean op that cannot change it. + * The optimized path returns the original vertices untouched, so it is strictly + * the more faithful of the two — the pre-filter removes a source of drift rather + * than introducing one. + * + * So the assertion is high-IoU agreement, which catches any real region change + * (a wrongly-dropped neighbour moves IoU far more than a rounding wobble) while + * tolerating sub-pixel boundary noise neither implementation controls. + */ +function regionOf(shapes: PolygonShape[]): Uint8Array { + const { gw, gh, scale } = fullResGridFor(W, H); + return rasterizeShapes(shapes, gw, gh, scale); +} + +/** Intersection-over-union of two binary masks (1 when both are empty). */ +function iou(a: Uint8Array, b: Uint8Array): number { + let inter = 0, union = 0; + for (let i = 0; i < a.length; i++) { + const x = a[i], y = b[i]; + if (x || y) union++; + if (x && y) inter++; + } + return union === 0 ? 1 : inter / union; +} + +/** Assert two clip results describe the same regions. */ +function expectSameRegions(got: PolygonShape[], want: PolygonShape[]): void { + expect(got.length).toBe(want.length); + expect(got.map((s) => s.classId).sort()).toEqual(want.map((s) => s.classId).sort()); + expect(iou(regionOf(got), regionOf(want))).toBeGreaterThan(0.995); +} + +// ---- Reference oracles: the pre-filter-free originals ---------------------- + +function clipShapesToOthers_naive( + newShapes: Shape[], sliceShapes: Shape[], width: number, height: number, +): PolygonShape[] { + const out: PolygonShape[] = []; + const otherByClass = new Map>(); + const sameByClass = new Map>(); + const otherMP = (c: number) => { + let m = otherByClass.get(c); + if (!m) { m = unionShapesToMultiPolygon(sliceShapes.filter((s) => s.classId !== c), width, height); otherByClass.set(c, m); } + return m; + }; + const sameMP = (c: number) => { + let m = sameByClass.get(c); + if (!m) { m = unionShapesToMultiPolygon(sliceShapes.filter((s) => s.classId === c), width, height); sameByClass.set(c, m); } + return m; + }; + for (const shape of newShapes) { + try { + let mp = shapeToMultiPolygon(shape, width, height); + if (mp.length === 0) continue; + const other = otherMP(shape.classId); + const same = sameMP(shape.classId); + if (other.length) mp = polygonClipping.difference(mp, other); + if (mp.length === 0) continue; + if (same.length && polygonClipping.difference(mp, same).length === 0) continue; + const polys = multiPolygonToShapes(mp, shape.classId); + if (polys.length) polys[0].id = shape.id; + out.push(...polys); + } catch { + out.push(...clipShapesToOthersMask([shape], sliceShapes, width, height)); + } + } + return out; +} + +function mergeNewWithSameClass_naive( + newShapes: Shape[], sliceShapes: Shape[], width: number, height: number, +): { addCount: number; removeIds: string[] } { + const { gw, gh, scale } = fullResGridFor(width, height); + const removeIds: string[] = []; + let addCount = 0; + const byClass = new Map(); + for (const s of newShapes) { + const arr = byClass.get(s.classId); + if (arr) arr.push(s); else byClass.set(s.classId, [s]); + } + const masksIntersect = (a: Uint8Array, b: Uint8Array) => { + for (let i = 0; i < a.length; i++) if (a[i] && b[i]) return true; + return false; + }; + for (const [classId, group] of byClass) { + const existing = sliceShapes.filter((s) => s.classId === classId); + if (existing.length === 0) { addCount += group.length; continue; } + const newMask = rasterizeUnion(group, gw, gh, scale); + const overlapping = existing.filter((e) => masksIntersect(rasterizeShapes([e], gw, gh, scale), newMask)); + if (overlapping.length === 0) { addCount += group.length; continue; } + removeIds.push(...overlapping.map((e) => e.id)); + } + return { addCount, removeIds }; +} + +function expandSameClassOverlap_naive(seed: Shape[], all: Shape[], width: number, height: number): Shape[] { + if (seed.length === 0) return seed; + const { gw, gh, scale } = gridFor(width, height); + const maskCache = new Map(); + const maskOf = (s: Shape) => { + let m = maskCache.get(s.id); + if (!m) { m = rasterizeShapes([s], gw, gh, scale); maskCache.set(s.id, m); } + return m; + }; + const chosen = new Map(seed.map((s) => [s.id, s])); + const classes = new Set(seed.map((s) => s.classId)); + for (const classId of classes) { + const candidates = all.filter((s) => s.classId === classId && !chosen.has(s.id)); + const members = [...chosen.values()].filter((s) => s.classId === classId); + let changed = true; + while (changed && candidates.length > 0) { + changed = false; + const union = new Uint8Array(gw * gh); + for (const m of members) { + const mm = maskOf(m); + for (let i = 0; i < union.length; i++) if (mm[i]) union[i] = 1; + } + for (let i = candidates.length - 1; i >= 0; i--) { + let hit = false; + const cm = maskOf(candidates[i]); + for (let k = 0; k < union.length; k++) if (cm[k] && union[k]) { hit = true; break; } + if (hit) { + const c = candidates.splice(i, 1)[0]; + chosen.set(c.id, c); + members.push(c); + changed = true; + } + } + } + } + return [...chosen.values()]; +} + +// ---- The tests ------------------------------------------------------------- + +describe('clipShapesToOthers — bbox pre-filter changes nothing', () => { + for (const seed of [1, 7, 42, 99, 1234]) { + it(`matches the unfiltered result (seed ${seed})`, () => { + const slice = randomSlice(seed, 14); + const rnd = mulberry32(seed + 500); + const size = 20 + Math.floor(rnd() * 30); + const newShape = makeShape(rnd, 'new', 1, Math.floor(rnd() * (W - size)), Math.floor(rnd() * (H - size)), size); + expectSameRegions(clipShapesToOthers([newShape], slice, W, H), clipShapesToOthers_naive([newShape], slice, W, H)); + }); + } + + it('matches for a shape disjoint from everything', () => { + const slice: Shape[] = [{ id: 'a', classId: 2, kind: 'rectangle', x: 0, y: 0, w: 20, h: 20 }]; + const far: Shape = { id: 'new', classId: 1, kind: 'rectangle', x: 150, y: 150, w: 20, h: 20 }; + expectSameRegions(clipShapesToOthers([far], slice, W, H), clipShapesToOthers_naive([far], slice, W, H)); + }); + + it('matches for shapes whose edges exactly abut (the filter must not drop these)', () => { + const slice: Shape[] = [{ id: 'a', classId: 2, kind: 'rectangle', x: 50, y: 50, w: 30, h: 30 }]; + const touching: Shape = { id: 'new', classId: 1, kind: 'rectangle', x: 80, y: 50, w: 30, h: 30 }; + expectSameRegions(clipShapesToOthers([touching], slice, W, H), clipShapesToOthers_naive([touching], slice, W, H)); + }); + + it('matches for a large shape fully containing a small other-class one', () => { + const slice: Shape[] = [{ id: 'a', classId: 2, kind: 'rectangle', x: 90, y: 90, w: 10, h: 10 }]; + const big: Shape = { id: 'new', classId: 1, kind: 'rectangle', x: 20, y: 20, w: 160, h: 160 }; + expectSameRegions(clipShapesToOthers([big], slice, W, H), clipShapesToOthers_naive([big], slice, W, H)); + }); + + it('matches when several new shapes of one class commit together', () => { + const slice = randomSlice(21, 12); + const news: Shape[] = [ + { id: 'n1', classId: 1, kind: 'rectangle', x: 10, y: 10, w: 30, h: 30 }, + { id: 'n2', classId: 1, kind: 'rectangle', x: 140, y: 140, w: 30, h: 30 }, + ]; + expectSameRegions(clipShapesToOthers(news, slice, W, H), clipShapesToOthers_naive(news, slice, W, H)); + }); +}); + +describe('mergeNewWithSameClass — bbox pre-filter picks the same merge targets', () => { + for (const seed of [3, 11, 55, 808]) { + it(`selects the same overlapping shapes (seed ${seed})`, () => { + const slice = randomSlice(seed, 16); + const rnd = mulberry32(seed + 900); + const size = 20 + Math.floor(rnd() * 30); + const newShape = makeShape(rnd, 'new', 1, Math.floor(rnd() * (W - size)), Math.floor(rnd() * (H - size)), size); + const got = mergeNewWithSameClass([newShape], slice, W, H); + const want = mergeNewWithSameClass_naive([newShape], slice, W, H); + expect([...got.removeIds].sort()).toEqual([...want.removeIds].sort()); + }); + } + + it('finds a same-class neighbour that only just overlaps', () => { + const slice: Shape[] = [{ id: 'a', classId: 1, kind: 'rectangle', x: 50, y: 50, w: 30, h: 30 }]; + const overlapping: Shape = { id: 'new', classId: 1, kind: 'rectangle', x: 78, y: 50, w: 30, h: 30 }; + const got = mergeNewWithSameClass([overlapping], slice, W, H); + const want = mergeNewWithSameClass_naive([overlapping], slice, W, H); + expect([...got.removeIds].sort()).toEqual([...want.removeIds].sort()); + expect(got.removeIds).toContain('a'); + }); + + it('ignores a same-class shape that is merely nearby but not touching', () => { + const slice: Shape[] = [{ id: 'a', classId: 1, kind: 'rectangle', x: 10, y: 10, w: 20, h: 20 }]; + const apart: Shape = { id: 'new', classId: 1, kind: 'rectangle', x: 120, y: 120, w: 20, h: 20 }; + const got = mergeNewWithSameClass([apart], slice, W, H); + expect(got.removeIds).toEqual([]); + expect(got.removeIds).toEqual(mergeNewWithSameClass_naive([apart], slice, W, H).removeIds); + }); +}); + +describe('expandSameClassOverlap — bbox pre-filter finds the same cluster', () => { + for (const seed of [5, 17, 300]) { + it(`expands to the same set (seed ${seed})`, () => { + const slice = randomSlice(seed, 18).map((s) => ({ ...s, classId: 1 })); + const seedSel = [slice[0]]; + const got = expandSameClassOverlap(seedSel, slice, W, H).map((s) => s.id).sort(); + const want = expandSameClassOverlap_naive(seedSel, slice, W, H).map((s) => s.id).sort(); + expect(got).toEqual(want); + }); + } + + it('walks a transitive chain of overlaps', () => { + // A–B–C chained; D is isolated. Selecting A must pull in B and C, never D. + const chain: Shape[] = [ + { id: 'A', classId: 1, kind: 'rectangle', x: 10, y: 10, w: 30, h: 30 }, + { id: 'B', classId: 1, kind: 'rectangle', x: 35, y: 10, w: 30, h: 30 }, + { id: 'C', classId: 1, kind: 'rectangle', x: 60, y: 10, w: 30, h: 30 }, + { id: 'D', classId: 1, kind: 'rectangle', x: 150, y: 150, w: 30, h: 30 }, + ]; + const got = expandSameClassOverlap([chain[0]], chain, W, H).map((s) => s.id).sort(); + expect(got).toEqual(['A', 'B', 'C']); + expect(got).toEqual(expandSameClassOverlap_naive([chain[0]], chain, W, H).map((s) => s.id).sort()); + }); +}); + +describe('rasterizeShapes scratch-buffer reuse', () => { + it('produces the same mask as a fresh allocation', () => { + const { gw, gh, scale } = gridFor(W, H); + const shape: Shape = { id: 'x', classId: 1, kind: 'ellipse', cx: 100, cy: 100, rx: 40, ry: 25 }; + const fresh = rasterizeShapes([shape], gw, gh, scale); + const scratch = new Uint8Array(gw * gh).fill(1); // dirty on purpose + scratch.fill(0); + rasterizeShapes([shape], gw, gh, scale, scratch); + expect(Array.from(scratch)).toEqual(Array.from(fresh)); + }); + + it('accumulates when the buffer is intentionally not cleared', () => { + const { gw, gh, scale } = gridFor(W, H); + const a: Shape = { id: 'a', classId: 1, kind: 'rectangle', x: 10, y: 10, w: 20, h: 20 }; + const b: Shape = { id: 'b', classId: 1, kind: 'rectangle', x: 100, y: 100, w: 20, h: 20 }; + const buf = new Uint8Array(gw * gh); + rasterizeShapes([a], gw, gh, scale, buf); + rasterizeShapes([b], gw, gh, scale, buf); + const both = rasterizeShapes([a, b], gw, gh, scale); + expect(Array.from(buf)).toEqual(Array.from(both)); + }); +}); diff --git a/frontend/src/lib/blur.test.ts b/frontend/src/lib/blur.test.ts new file mode 100644 index 0000000..f49fd3a --- /dev/null +++ b/frontend/src/lib/blur.test.ts @@ -0,0 +1,68 @@ +import { describe, it, expect } from 'vitest'; +import { applyGaussianBlurRgba } from './blur'; + +function grayBuffer(w: number, h: number, fn: (x: number, y: number) => number): Uint8ClampedArray { + const d = new Uint8ClampedArray(w * h * 4); + for (let y = 0; y < h; y++) { + for (let x = 0; x < w; x++) { + const v = fn(x, y); + const i = (y * w + x) * 4; + d[i] = d[i + 1] = d[i + 2] = v; + d[i + 3] = 255; + } + } + return d; +} + +describe('applyGaussianBlurRgba', () => { + it('leaves a flat image unchanged (edges clamp, so no darkening at the border)', () => { + const d = grayBuffer(16, 16, () => 120); + applyGaussianBlurRgba(d, 16, 16, 2); + for (let i = 0; i < d.length; i += 4) expect(d[i]).toBeCloseTo(120, -0.5); + }); + + it('spreads an impulse into its neighbourhood', () => { + const d = grayBuffer(21, 21, (x, y) => (x === 10 && y === 10 ? 255 : 0)); + applyGaussianBlurRgba(d, 21, 21, 2); + const at = (x: number, y: number) => d[(y * 21 + x) * 4]; + expect(at(10, 10)).toBeLessThan(255); // peak flattened + expect(at(10, 10)).toBeGreaterThan(0); + expect(at(11, 10)).toBeGreaterThan(0); // energy moved outward + expect(at(10, 11)).toBeGreaterThan(0); + // Monotonically decreasing away from the impulse. + expect(at(10, 10)).toBeGreaterThan(at(12, 10)); + expect(at(12, 10)).toBeGreaterThan(at(15, 10)); + }); + + it('reduces the contrast of a step edge', () => { + const w = 32; + const d = grayBuffer(w, 4, (x) => (x < w / 2 ? 0 : 255)); + applyGaussianBlurRgba(d, w, 4, 3); + const at = (x: number) => d[(1 * w + x) * 4]; + // Straddling the edge, values are now intermediate rather than 0/255. + expect(at(w / 2 - 1)).toBeGreaterThan(0); + expect(at(w / 2)).toBeLessThan(255); + // Far from the edge the plateaus survive. + expect(at(0)).toBeLessThan(10); + expect(at(w - 1)).toBeGreaterThan(245); + }); + + it('is a no-op for sigma <= 0', () => { + const d = grayBuffer(8, 8, (x, y) => (x + y) * 8); + const before = new Uint8ClampedArray(d); + applyGaussianBlurRgba(d, 8, 8, 0); + expect(Array.from(d)).toEqual(Array.from(before)); + applyGaussianBlurRgba(d, 8, 8, -1); + expect(Array.from(d)).toEqual(Array.from(before)); + }); + + it('preserves alpha', () => { + const d = grayBuffer(12, 12, (x) => 20 * x); + applyGaussianBlurRgba(d, 12, 12, 1.5); + for (let i = 3; i < d.length; i += 4) expect(d[i]).toBe(255); + }); + + it('does not throw on empty input', () => { + expect(() => applyGaussianBlurRgba(new Uint8ClampedArray(0), 0, 0, 2)).not.toThrow(); + }); +}); diff --git a/frontend/src/lib/blur.ts b/frontend/src/lib/blur.ts new file mode 100644 index 0000000..a342a85 --- /dev/null +++ b/frontend/src/lib/blur.ts @@ -0,0 +1,158 @@ +/** + * Gaussian blur (display-only) — three successive box blurs, the standard + * approximation of a true Gaussian (by the central limit theorem, three passes + * are visually indistinguishable from a real Gaussian for σ ≳ 1 while staying + * O(n) per pass regardless of radius). + * + * Used as the first nonlinear display preprocessor so the intensity-driven tools + * (threshold brush, magic wand, fill) see a denoised image and similar features + * cohere into one selectable region. Mutates the RGBA Uint8ClampedArray in place + * (alpha untouched); edges clamp (out-of-bounds samples replicate the edge + * pixel). Pure/DOM-free for testing. + */ + +/** Box radii whose 3-pass cascade approximates a Gaussian of the given sigma. + * Standard Kovesi/Gwosdek construction: pick w_ideal from the box variance + * identity, split into `m` boxes of size wl and `3-m` of size wl+2. */ +function boxRadiiForGaussian(sigma: number, passes: number): number[] { + const wIdeal = Math.sqrt((12 * sigma * sigma) / passes + 1); + let wl = Math.floor(wIdeal); + if (wl % 2 === 0) wl--; + const wu = wl + 2; + const mIdeal = + (12 * sigma * sigma - passes * wl * wl - 4 * passes * wl - 3 * passes) / + (-4 * wl - 4); + const m = Math.round(mIdeal); + const radii: number[] = []; + for (let i = 0; i < passes; i++) { + const w = i < m ? wl : wu; + radii.push(Math.max(0, (w - 1) / 2)); + } + return radii; +} + +/** One horizontal box blur of radius `r` from `src` into `dst` (RGB only). */ +function boxBlurH( + src: Uint8ClampedArray, + dst: Uint8ClampedArray, + width: number, + height: number, + r: number, +): void { + const norm = 1 / (r + r + 1); + for (let y = 0; y < height; y++) { + const row = y * width; + for (let c = 0; c < 3; c++) { + // Seed the running sum for x=0 with edge-clamped left neighbours. + const first = src[row * 4 + c]; + let sum = (r + 1) * first; + for (let i = 0; i < r; i++) sum += src[(row + Math.min(i, width - 1)) * 4 + c]; + for (let x = 0; x < width; x++) { + const inIdx = Math.min(x + r, width - 1); + const outIdx = x - r - 1; + sum += src[(row + inIdx) * 4 + c]; + sum -= outIdx < 0 ? first : src[(row + outIdx) * 4 + c]; + dst[(row + x) * 4 + c] = sum * norm; + } + } + } +} + +/** One vertical box blur of radius `r` from `src` into `dst` (RGB only). */ +function boxBlurV( + src: Uint8ClampedArray, + dst: Uint8ClampedArray, + width: number, + height: number, + r: number, +): void { + const norm = 1 / (r + r + 1); + for (let x = 0; x < width; x++) { + for (let c = 0; c < 3; c++) { + const first = src[x * 4 + c]; + let sum = (r + 1) * first; + for (let i = 0; i < r; i++) sum += src[(Math.min(i, height - 1) * width + x) * 4 + c]; + for (let y = 0; y < height; y++) { + const inIdx = Math.min(y + r, height - 1); + const outIdx = y - r - 1; + sum += src[(inIdx * width + x) * 4 + c]; + sum -= outIdx < 0 ? first : src[(outIdx * width + x) * 4 + c]; + dst[(y * width + x) * 4 + c] = sum * norm; + } + } + } +} + +/** + * Blur a single-channel float field in place, same Gaussian approximation as the + * RGBA version above. + * + * The Threshold Brush's field is already grayscale, so the sampler's blur sweep + * can work on it directly instead of round-tripping through RGBA — which also + * means a candidate sigma is evaluated on exactly the values the brush gates on. + */ +export function applyGaussianBlurGray( + data: Float32Array, + width: number, + height: number, + sigma: number, +): void { + if (!(sigma > 0) || width === 0 || height === 0) return; + const scratch = new Float32Array(data.length); + + const boxH = (src: Float32Array, dst: Float32Array, r: number) => { + const norm = 1 / (r + r + 1); + for (let y = 0; y < height; y++) { + const row = y * width; + const first = src[row]; + let sum = (r + 1) * first; + for (let i = 0; i < r; i++) sum += src[row + Math.min(i, width - 1)]; + for (let x = 0; x < width; x++) { + sum += src[row + Math.min(x + r, width - 1)]; + sum -= x - r - 1 < 0 ? first : src[row + x - r - 1]; + dst[row + x] = sum * norm; + } + } + }; + const boxV = (src: Float32Array, dst: Float32Array, r: number) => { + const norm = 1 / (r + r + 1); + for (let x = 0; x < width; x++) { + const first = src[x]; + let sum = (r + 1) * first; + for (let i = 0; i < r; i++) sum += src[Math.min(i, height - 1) * width + x]; + for (let y = 0; y < height; y++) { + sum += src[Math.min(y + r, height - 1) * width + x]; + sum -= y - r - 1 < 0 ? first : src[(y - r - 1) * width + x]; + dst[y * width + x] = sum * norm; + } + } + }; + + for (const r of boxRadiiForGaussian(sigma, 3)) { + if (r <= 0) continue; + boxH(data, scratch, r); + boxV(scratch, data, r); + } +} + +/** + * Blur `data` (RGBA, row-major `width × height`) in place by a Gaussian of the + * given sigma in pixels. No-op for sigma ≤ 0 or an empty image. + */ +export function applyGaussianBlurRgba( + data: Uint8ClampedArray, + width: number, + height: number, + sigma: number, +): void { + if (!(sigma > 0) || width === 0 || height === 0) return; + const radii = boxRadiiForGaussian(sigma, 3); + // Scratch buffer for the horizontal half of each pass; the vertical half + // writes straight back into `data`, so the result always lands in place. + const scratch = new Uint8ClampedArray(data); + for (const r of radii) { + if (r <= 0) continue; + boxBlurH(data, scratch, width, height, r); + boxBlurV(scratch, data, width, height, r); + } +} diff --git a/frontend/src/lib/clipFallback.test.ts b/frontend/src/lib/clipFallback.test.ts new file mode 100644 index 0000000..87068a0 --- /dev/null +++ b/frontend/src/lib/clipFallback.test.ts @@ -0,0 +1,200 @@ +/** + * Clipping must never silently degrade to "not clipped". + * + * The boolean clip path can fail on awkward geometry — and complex Threshold + * Brush regions are exactly that. The old code returned `[]` from the union on + * failure, which callers read as "nothing to clip against", so the new annotation + * was committed overlapping its neighbour with no error anywhere. On a slice with + * threshold-painted classes that made clip-to-other-classes look simply broken for + * every subsequent annotation. + * + * These tests force the failure and assert the result is still clipped. + */ +import { describe, it, expect, vi, afterEach } from 'vitest'; +import { clipShapesToOthers } from './clipToClasses'; +import { unionShapesChecked } from './polybool'; +import { fullResGridFor, rasterizeShapes } from './rasterize'; +import { maskToPolygonsWithHoles } from './magicwand'; +import type { Shape } from '@/stores/annotationStore'; + +const W = 256, H = 256; + +afterEach(() => { vi.restoreAllMocks(); vi.resetModules(); }); + +/** Pixels where `a` and `b` both cover, ignoring shared-boundary slack. */ +function realOverlap(a: Shape[], b: Shape[]): number { + const { gw, gh, scale } = fullResGridFor(W, H); + const ma = rasterizeShapes(a, gw, gh, scale); + const mb = rasterizeShapes(b, gw, gh, scale); + // Erode the intersection by requiring a 4-neighbourhood hit, so a shared edge + // (which legitimately rasterizes into both) doesn't count as real overlap. + let n = 0; + for (let y = 1; y < gh - 1; y++) { + for (let x = 1; x < gw - 1; x++) { + const i = y * gw + x; + if (!(ma[i] && mb[i])) continue; + if (ma[i - 1] && mb[i - 1] && ma[i + 1] && mb[i + 1] && + ma[i - gw] && mb[i - gw] && ma[i + gw] && mb[i + gw]) n++; + } + } + return n; +} + +const otherClass: Shape = { id: 'o', classId: 2, kind: 'rectangle', x: 60, y: 60, w: 120, h: 120 }; +const incoming: Shape = { id: 'n', classId: 1, kind: 'rectangle', x: 100, y: 100, w: 120, h: 120 }; + +describe('clip never silently degrades to unclipped', () => { + it('clips normally when the boolean path works', () => { + const res = clipShapesToOthers([incoming], [otherClass], W, H); + expect(res.length).toBeGreaterThan(0); + expect(realOverlap(res, [otherClass])).toBe(0); + }); + + it('still clips when polygon-clipping throws on every boolean op', async () => { + // Simulate the failure mode: the library rejects this geometry outright. + vi.doMock('polygon-clipping', () => ({ + default: { + union: () => { throw new Error('boom'); }, + difference: () => { throw new Error('boom'); }, + }, + })); + vi.resetModules(); + const { clipShapesToOthers: clipFresh } = await import('./clipToClasses'); + + const res = clipFresh([incoming], [otherClass], W, H); + expect(res.length).toBeGreaterThan(0); // not dropped + expect(realOverlap(res, [otherClass])).toBe(0); // and genuinely clipped + }); + + it('still clips when only the union fails', async () => { + vi.doMock('polygon-clipping', () => ({ + default: { + union: () => { throw new Error('boom'); }, + difference: (a: unknown) => a, // difference "works" but is a no-op + }, + })); + vi.resetModules(); + const { clipShapesToOthers: clipFresh } = await import('./clipToClasses'); + + // Two other-class shapes so a union is actually required. + const others: Shape[] = [ + otherClass, + { id: 'o2', classId: 2, kind: 'rectangle', x: 150, y: 60, w: 60, h: 120 }, + ]; + const res = clipFresh([incoming], others, W, H); + expect(realOverlap(res, others)).toBe(0); + }); +}); + +describe('unionShapesChecked reports failure instead of hiding it', () => { + it('is ok for ordinary shapes', () => { + const { mp, ok } = unionShapesChecked([otherClass], W, H); + expect(ok).toBe(true); + expect(mp.length).toBeGreaterThan(0); + }); + + it('reports ok:false when a union throws, keeping what it can', async () => { + vi.doMock('polygon-clipping', () => ({ + default: { union: () => { throw new Error('boom'); }, difference: (a: unknown) => a }, + })); + vi.resetModules(); + const { unionShapesChecked: checked } = await import('./polybool'); + + const { mp, ok } = checked( + [otherClass, { id: 'o2', classId: 2, kind: 'rectangle', x: 10, y: 10, w: 20, h: 20 }], + W, H, + ); + expect(ok).toBe(false); // the caller can now react + expect(mp.length).toBeGreaterThan(0); // and we kept the first shape's geometry + }); + + it('is ok:true (empty) when there is genuinely nothing to union', () => { + const { mp, ok } = unionShapesChecked([], W, H); + expect(ok).toBe(true); + expect(mp).toEqual([]); + }); +}); + +describe('clipping against threshold-brush geometry', () => { + /** Many small speckled regions with holes, as a threshold stroke produces. */ + function thresholdShapes(classId: number): Shape[] { + const { gw, gh, scale } = fullResGridFor(W, H); + const mask = new Uint8Array(gw * gh); + let seed = 99; + const rnd = () => { + seed |= 0; seed = (seed + 0x6d2b79f5) | 0; + let t = Math.imul(seed ^ (seed >>> 15), 1 | seed); + t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t; + return ((t ^ (t >>> 14)) >>> 0) / 4294967296; + }; + for (let y = 40; y < 180; y++) for (let x = 40; x < 200; x++) if (rnd() > 0.45) mask[y * gw + x] = 1; + return maskToPolygonsWithHoles(mask, gw, gh, { minRegion: 4, scale }) + .filter((p) => p.points.length >= 6) + .map((p, i) => ({ + id: `t${i}`, classId, kind: 'polygon' as const, + points: p.points, ...(p.holes.length ? { holes: p.holes } : {}), + })); + } + + it('clips a new annotation against hundreds of speckled threshold regions', () => { + const others = thresholdShapes(2); + expect(others.length).toBeGreaterThan(50); // genuinely the hard case + const res = clipShapesToOthers([incoming], others, W, H); + expect(realOverlap(res, others)).toBe(0); + }); +}); + +describe('batched multi-shape clip matches per-shape clip', () => { + /** The per-shape path, forced by clipping each shape in its own call. */ + function perShape(news: Shape[], slice: Shape[]): Shape[] { + return news.flatMap((n) => clipShapesToOthers([n], slice, W, H)); + } + + const slice: Shape[] = [ + { id: 'o1', classId: 2, kind: 'rectangle', x: 40, y: 40, w: 90, h: 90 }, + { id: 'o2', classId: 2, kind: 'ellipse', cx: 180, cy: 170, rx: 45, ry: 30 }, + { id: 's1', classId: 1, kind: 'rectangle', x: 200, y: 30, w: 30, h: 30 }, + ]; + const news: Shape[] = [ + { id: 'n1', classId: 1, kind: 'rectangle', x: 80, y: 80, w: 80, h: 80 }, + { id: 'n2', classId: 1, kind: 'rectangle', x: 150, y: 140, w: 70, h: 70 }, + { id: 'n3', classId: 1, kind: 'rectangle', x: 30, y: 190, w: 50, h: 40 }, + ]; + + it('covers the same region', () => { + const batched = clipShapesToOthers(news, slice, W, H); + const single = perShape(news, slice); + const { gw, gh, scale } = fullResGridFor(W, H); + const a = rasterizeShapes(batched, gw, gh, scale); + const b = rasterizeShapes(single, gw, gh, scale); + let diff = 0, set = 0; + for (let i = 0; i < a.length; i++) { if (a[i] || b[i]) set++; if (a[i] !== b[i]) diff++; } + expect(diff / Math.max(1, set)).toBeLessThan(0.01); + }); + + it('still excludes the other classes', () => { + const batched = clipShapesToOthers(news, slice, W, H); + expect(realOverlap(batched, slice.filter((s) => s.classId === 2))).toBe(0); + }); + + it('clips a 100+ region threshold commit correctly', () => { + const { gw, gh, scale } = fullResGridFor(W, H); + const mask = new Uint8Array(gw * gh); + let seed = 3; + const rnd = () => { seed |= 0; seed = (seed + 0x6d2b79f5) | 0; + let t = Math.imul(seed ^ (seed >>> 15), 1 | seed); + t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t; + return ((t ^ (t >>> 14)) >>> 0) / 4294967296; }; + for (let y = 30; y < 220; y++) for (let x = 30; x < 220; x++) if (rnd() > 0.35) mask[y * gw + x] = 1; + const regions: Shape[] = maskToPolygonsWithHoles(mask, gw, gh, { minRegion: 4, scale }) + .filter((p) => p.points.length >= 6) + .map((p, i) => ({ id: `r${i}`, classId: 1, kind: 'polygon' as const, + points: p.points, ...(p.holes.length ? { holes: p.holes } : {}) })); + expect(regions.length).toBeGreaterThan(40); // genuinely a multi-region commit + + const blocker: Shape = { id: 'b', classId: 2, kind: 'rectangle', x: 90, y: 90, w: 80, h: 80 }; + const res = clipShapesToOthers(regions, [blocker], W, H); + expect(res.length).toBeGreaterThan(0); + expect(realOverlap(res, [blocker])).toBe(0); + }); +}); diff --git a/frontend/src/lib/clipToClasses.ts b/frontend/src/lib/clipToClasses.ts index 0691c58..2c9ad43 100644 --- a/frontend/src/lib/clipToClasses.ts +++ b/frontend/src/lib/clipToClasses.ts @@ -14,7 +14,7 @@ import polygonClipping, { type MultiPolygon } from 'polygon-clipping'; import type { PolygonShape, Shape } from '@/stores/annotationStore'; import { fullResGridFor, rasterizeShapes, rasterizeUnion } from '@/lib/rasterize'; import { maskToPolygonsWithHoles } from '@/lib/magicwand'; -import { shapeToMultiPolygon, multiPolygonToShapes, unionShapesToMultiPolygon } from '@/lib/polybool'; +import { shapeToMultiPolygon, multiPolygonToShapes, unionShapesChecked } from '@/lib/polybool'; /** Does the slice hold any shape of a different class than `classId`? */ export function hasOtherClass(sliceShapes: Shape[], classId: number): boolean { @@ -60,29 +60,118 @@ export function clipShapesToOthers( sliceShapes: Shape[], width: number, height: number, + upscale = 1, ): PolygonShape[] { const out: PolygonShape[] = []; - const otherByClass = new Map(); - const sameByClass = new Map(); + + // Unions are computed once per class and carry an `ok` flag. `ok: false` means + // some shape could not be folded in, so the union UNDER-covers the other classes + // — clipping against it would leave real overlap. That case must route to the + // mask path, not proceed: an under-covering union is indistinguishable from + // "nothing to clip against" if you only look at `mp.length`, which is precisely + // how complex geometry (e.g. speckled Threshold Brush regions) can silently + // disable clipping for every later annotation on the slice. + const otherByClass = new Map(); + const sameByClass = new Map(); const otherMP = (c: number) => { let m = otherByClass.get(c); - if (!m) { m = unionShapesToMultiPolygon(sliceShapes.filter((s) => s.classId !== c), width, height); otherByClass.set(c, m); } + if (!m) { m = unionShapesChecked(sliceShapes.filter((s) => s.classId !== c), width, height); otherByClass.set(c, m); } return m; }; const sameMP = (c: number) => { let m = sameByClass.get(c); - if (!m) { m = unionShapesToMultiPolygon(sliceShapes.filter((s) => s.classId === c), width, height); sameByClass.set(c, m); } + if (!m) { m = unionShapesChecked(sliceShapes.filter((s) => s.classId === c), width, height); sameByClass.set(c, m); } return m; }; - for (const shape of newShapes) { - const res = clipOneBoolean(shape, otherMP(shape.classId), sameMP(shape.classId), width, height); - if (res === null) out.push(...clipShapesToOthersMask([shape], sliceShapes, width, height)); - else out.push(...res); + // Shapes the boolean path can't handle are batched and clipped together at the + // end: `clipShapesToOthersMask` caches its class masks per CALL, so clipping + // them one at a time would re-rasterize every shape on the slice per shape. + const needsMask: Shape[] = []; + + // Group by class so a multi-shape commit can be clipped in ONE boolean pass. + // This matters enormously for the Threshold Brush, which commits every in-band + // region of a stroke at once — often 100+ polygons. Differencing them one at a + // time re-walks the whole other-class union per shape, so commit time grew as + // (new shapes × slice complexity) and reached seconds on a busy slice. Every + // other tool commits a single shape and never noticed. + const byClass = new Map(); + for (const s of newShapes) { + const arr = byClass.get(s.classId); + if (arr) arr.push(s); else byClass.set(s.classId, [s]); + } + + for (const [classId, group] of byClass) { + const other = otherMP(classId); + const same = sameMP(classId); + // Any doubt about the unions → mask clip, which rasterizes every shape + // independently and so cannot silently under-cover. + if (!other.ok || !same.ok) { needsMask.push(...group); continue; } + + // A single shape keeps the original per-shape path, which preserves the + // shape's id on its primary fragment. Batching regenerates ids, which is fine + // for freshly-committed regions but not worth changing for the common case. + if (group.length === 1) { + const res = clipOneBoolean(group[0], other.mp, same.mp, width, height); + if (res === null) needsMask.push(group[0]); else out.push(...res); + continue; + } + + const res = clipBatchBoolean(group, other.mp, same.mp, classId, width, height); + if (res === null) needsMask.push(...group); else out.push(...res); + } + + if (needsMask.length) { + out.push(...clipShapesToOthersMask(needsMask, sliceShapes, width, height, upscale)); } return out; } +/** + * Clip a whole group of same-class shapes in one boolean pass. + * + * The group is unioned first (disjoint regions — which is what a mask→polygon + * pass produces — stay separate, so the shape count is unchanged), then a single + * difference removes the other classes. Returns null on failure so the caller can + * fall back to the mask path, exactly like `clipOneBoolean`. + */ +function clipBatchBoolean( + shapes: Shape[], + otherMP: MultiPolygon, + sameMP: MultiPolygon, + classId: number, + width: number, + height: number, +): PolygonShape[] | null { + try { + const geoms: MultiPolygon[] = []; + for (const s of shapes) { + const g = shapeToMultiPolygon(s, width, height); + if (g.length) geoms.push(g); + } + if (geoms.length === 0) return []; + + let mp = geoms.length === 1 ? geoms[0] : polygonClipping.union(geoms[0], ...geoms.slice(1)); + if (otherMP.length) mp = polygonClipping.difference(mp, otherMP); + if (mp.length === 0) return []; + + const polys = multiPolygonToShapes(mp, classId); + if (!sameMP.length) return polys; + + // Drop fragments that add nothing beyond this class's existing labels. Done + // per fragment because a fragment is kept WHOLE if any part of it is new. + return polys.filter((p) => { + try { + return polygonClipping.difference(shapeToMultiPolygon(p, width, height), sameMP).length > 0; + } catch { + return true; // can't prove it's redundant — keep it + } + }); + } catch { + return null; + } +} + /** * Mask-based clip (fallback): rasterize each new shape, subtract the union of * other-class shapes, drop fully-redundant fragments, re-vectorize the remainder. @@ -93,10 +182,15 @@ export function clipShapesToOthersMask( sliceShapes: Shape[], width: number, height: number, + upscale = 1, ): PolygonShape[] { - const { gw, gh, scale } = fullResGridFor(width, height); + // `upscale` keeps sub-pixel geometry (e.g. a Threshold Brush region traced at 2x) + // from being re-snapped to the native pixel grid by this fallback round-trip. + const { gw, gh, scale } = fullResGridFor(width, height, upscale); const out: PolygonShape[] = []; + // See the note in `clipShapesToOthers`: the bounds pre-filter is reverted here + // too, so both clip paths behave exactly as they did before the perf pass. const otherMaskByClass = new Map(); const sameMaskByClass = new Map(); const otherMaskFor = (classId: number): Uint8Array => { diff --git a/frontend/src/lib/datasetStats.test.ts b/frontend/src/lib/datasetStats.test.ts new file mode 100644 index 0000000..44fa73d --- /dev/null +++ b/frontend/src/lib/datasetStats.test.ts @@ -0,0 +1,177 @@ +import { describe, expect, it } from 'vitest'; +import { computeSampleStats } from './datasetStats'; +import type { AnnotationClass } from '@/stores/classStore'; + +const CLASSES: AnnotationClass[] = [ + { classId: 1, label: 'Cell', color: '#f00', isVisible: true }, + { classId: 2, label: 'Wall', color: '#0f0', isVisible: true }, +]; + +describe('computeSampleStats', () => { + it('reports zero stats for a source with no slices at all', () => { + const stats = computeSampleStats({}, CLASSES, 3, [], 100, 100); + expect(stats.coverage.annotatedSlices).toBe(0); + expect(stats.classStats.every((c) => c.shapeCount === 0)).toBe(true); + }); + + it('counts shapes and pixel area per class', () => { + const stats = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 10, h: 10 }] }, + CLASSES, 3, [], 100, 100, + ); + const cell = stats.classStats.find((c) => c.classId === 1)!; + expect(cell.shapeCount).toBe(1); + expect(cell.pixelArea).toBeGreaterThan(0); + expect(stats.coverage.annotatedSlices).toBe(1); + }); + + it('flags a tiny (sliver) shape', () => { + const stats = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }] }, + CLASSES, 3, [], 100, 100, + ); + expect(stats.flags.some((f) => f.kind === 'sliver')).toBe(true); + }); + + it('does not flag a normal-sized shape as a sliver', () => { + const stats = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 50, h: 50 }] }, + CLASSES, 3, [], 100, 100, + ); + expect(stats.flags.some((f) => f.kind === 'sliver')).toBe(false); + }); + + it('flags a self-intersecting polygon', () => { + // A "bowtie" shape: crosses itself in the middle. + const bowtie = [0, 0, 100, 100, 100, 0, 0, 100]; + const stats = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'polygon', points: bowtie }] }, + CLASSES, 3, [], 200, 200, + ); + expect(stats.flags.some((f) => f.kind === 'self-intersection')).toBe(true); + }); + + it('does not flag a simple (non-intersecting) polygon', () => { + const square = [0, 0, 50, 0, 50, 50, 0, 50]; + const stats = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'polygon', points: square }] }, + CLASSES, 3, [], 200, 200, + ); + expect(stats.flags.some((f) => f.kind === 'self-intersection')).toBe(false); + }); + + it('flags an overlap between two different classes on the same slice', () => { + const stats = computeSampleStats( + { + '0': [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 50, h: 50 }, + { id: 's2', classId: 2, kind: 'rectangle', x: 25, y: 25, w: 50, h: 50 }, + ], + }, + CLASSES, 3, [], 200, 200, + ); + expect(stats.flags.some((f) => f.kind === 'overlap')).toBe(true); + }); + + it('does not flag disjoint shapes of different classes', () => { + const stats = computeSampleStats( + { + '0': [ + { id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 20, h: 20 }, + { id: 's2', classId: 2, kind: 'rectangle', x: 150, y: 150, w: 20, h: 20 }, + ], + }, + CLASSES, 3, [], 200, 200, + ); + expect(stats.flags.some((f) => f.kind === 'overlap')).toBe(false); + }); + + it('flags empty-unmarked slices (not annotated, not negative)', () => { + const stats = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 50, h: 50 }] }, + CLASSES, 3, [], 100, 100, + ); + expect(stats.coverage.emptyUnmarked).toEqual([1, 2]); + expect(stats.flags.filter((f) => f.kind === 'empty-unmarked')).toHaveLength(2); + }); + + it('does not flag a slice marked as negative', () => { + const stats = computeSampleStats({}, CLASSES, 2, ['0', '1'], 100, 100); + expect(stats.coverage.emptyUnmarked).toEqual([]); + expect(stats.coverage.negativeSlices).toBe(2); + }); + + it('suppresses empty-unmarked flags entirely when there are too many (>20)', () => { + const stats = computeSampleStats({}, CLASSES, 25, [], 100, 100); + expect(stats.coverage.emptyUnmarked).toHaveLength(25); + expect(stats.flags.filter((f) => f.kind === 'empty-unmarked')).toHaveLength(0); + }); + + it('computes area for an ellipse shape', () => { + const stats = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'ellipse', cx: 50, cy: 50, rx: 20, ry: 10 }] }, + CLASSES, 1, [], 100, 100, + ); + const cell = stats.classStats.find((c) => c.classId === 1)!; + expect(cell.pixelArea).toBeGreaterThan(0); + }); + + it('computes area for a brush stroke', () => { + const stats = computeSampleStats( + { + '0': [{ + id: 's1', classId: 1, kind: 'brush', + strokes: [{ mode: 'paint', points: [0, 0, 50, 50], radius: 5 }], + }], + }, + CLASSES, 1, [], 100, 100, + ); + const cell = stats.classStats.find((c) => c.classId === 1)!; + expect(cell.pixelArea).toBeGreaterThan(0); + }); + + it('ignores an erase-mode brush stroke for area purposes', () => { + const stats = computeSampleStats( + { + '0': [{ + id: 's1', classId: 1, kind: 'brush', + strokes: [{ mode: 'erase', points: [0, 0, 50, 50], radius: 5 }], + }], + }, + CLASSES, 1, [], 100, 100, + ); + const cell = stats.classStats.find((c) => c.classId === 1)!; + // Shape still counts, but analytic area contribution from an erase stroke is 0. + expect(cell.shapeCount).toBe(1); + }); + + it('a polygon hole reduces its net area', () => { + const outer = [0, 0, 100, 0, 100, 100, 0, 100]; + const hole = [25, 25, 75, 25, 75, 75, 25, 75]; + const statsWithHole = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'polygon', points: outer, holes: [hole] }] }, + CLASSES, 1, [], 200, 200, + ); + const statsNoHole = computeSampleStats( + { '0': [{ id: 's1', classId: 1, kind: 'polygon', points: outer }] }, + CLASSES, 1, [], 200, 200, + ); + const withHoleArea = statsWithHole.classStats.find((c) => c.classId === 1)!.pixelArea; + const noHoleArea = statsNoHole.classStats.find((c) => c.classId === 1)!.pixelArea; + expect(withHoleArea).toBeLessThan(noHoleArea); + }); + + it('ignores an empty-array slice entry (not annotated)', () => { + const stats = computeSampleStats({ '0': [] }, CLASSES, 1, [], 100, 100); + expect(stats.coverage.annotatedSlices).toBe(0); + }); + + it('labels a flag for a class not in the provided class list generically', () => { + const stats = computeSampleStats( + { '0': [{ id: 's1', classId: 99, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2 }] }, + CLASSES, 1, [], 100, 100, + ); + const sliver = stats.flags.find((f) => f.kind === 'sliver')!; + expect(sliver.message).toContain('class 99'); + }); +}); diff --git a/frontend/src/lib/denoiserTrainScope.test.ts b/frontend/src/lib/denoiserTrainScope.test.ts new file mode 100644 index 0000000..f59bcab --- /dev/null +++ b/frontend/src/lib/denoiserTrainScope.test.ts @@ -0,0 +1,70 @@ +import { describe, expect, it } from 'vitest'; +import { denoiserScopeBlockedReason, resolveDenoiserScopeIndices } from './denoiserTrainScope'; + +describe('resolveDenoiserScopeIndices', () => { + it('returns just the current slice for scope "current"', () => { + expect(resolveDenoiserScopeIndices('current', 5, 0, 0, 10)).toEqual([5]); + }); + + it('returns [] for "current" when the current slice is out of range', () => { + expect(resolveDenoiserScopeIndices('current', 20, 0, 0, 10)).toEqual([]); + expect(resolveDenoiserScopeIndices('current', -1, 0, 0, 10)).toEqual([]); + }); + + it('returns every index for scope "all"', () => { + expect(resolveDenoiserScopeIndices('all', 0, 0, 0, 4)).toEqual([0, 1, 2, 3]); + }); + + it('returns [] for an empty volume regardless of scope', () => { + expect(resolveDenoiserScopeIndices('all', 0, 0, 0, 0)).toEqual([]); + expect(resolveDenoiserScopeIndices('current', 0, 0, 0, 0)).toEqual([]); + expect(resolveDenoiserScopeIndices('range', 0, 0, 0, 0)).toEqual([]); + }); + + it('returns the inclusive range for scope "range"', () => { + expect(resolveDenoiserScopeIndices('range', 0, 2, 5, 10)).toEqual([2, 3, 4, 5]); + }); + + it('normalizes a reversed range (end typed before start)', () => { + expect(resolveDenoiserScopeIndices('range', 0, 5, 2, 10)).toEqual([2, 3, 4, 5]); + }); + + it('clamps a range to [0, nSlices)', () => { + expect(resolveDenoiserScopeIndices('range', 0, -3, 2, 5)).toEqual([0, 1, 2]); + expect(resolveDenoiserScopeIndices('range', 0, 3, 99, 5)).toEqual([3, 4]); + }); + + it('handles a single-slice range', () => { + expect(resolveDenoiserScopeIndices('range', 0, 3, 3, 10)).toEqual([3]); + }); +}); + +describe('denoiserScopeBlockedReason', () => { + it('blocks an empty scope regardless of scheme', () => { + // Enumerates EVERY scheme on purpose: this used to cover only two of them, + // so a newly-added scheme could silently skip the check. + expect(denoiserScopeBlockedReason([], 'n2v')).toMatch(/no slices/i); + expect(denoiserScopeBlockedReason([], 'n2n')).toMatch(/no slices/i); + expect(denoiserScopeBlockedReason([], 'ae')).toMatch(/no slices/i); + }); + + it('blocks Noise2Noise on a single slice', () => { + expect(denoiserScopeBlockedReason([4], 'n2n')).toMatch(/noise2noise/i); + }); + + it('allows Noise2Void on a single slice', () => { + expect(denoiserScopeBlockedReason([4], 'n2v')).toBeNull(); + }); + + it('allows the autoencoder on a single slice', () => { + // It reconstructs one slice through a bottleneck — unlike Noise2Noise it + // needs no second slice to pair with. + expect(denoiserScopeBlockedReason([4], 'ae')).toBeNull(); + expect(denoiserScopeBlockedReason([0, 1, 2], 'ae')).toBeNull(); + }); + + it('allows Noise2Noise once there are at least 2 slices', () => { + expect(denoiserScopeBlockedReason([4, 5], 'n2n')).toBeNull(); + expect(denoiserScopeBlockedReason([1, 2, 3], 'n2n')).toBeNull(); + }); +}); diff --git a/frontend/src/lib/denoiserTrainScope.ts b/frontend/src/lib/denoiserTrainScope.ts new file mode 100644 index 0000000..765feed --- /dev/null +++ b/frontend/src/lib/denoiserTrainScope.ts @@ -0,0 +1,65 @@ +/** + * denoiserTrainScope — resolves the Learned Denoiser panel's current/range/all + * scope selector into the concrete slice indices to train on, and validates + * the result before a training job is submitted. + * + * Mirrors ApplyModelPanel's fine-tune-scope selector (see its `trainScope`/ + * `trainSliceIndices`) as closely as possible, kept here as a standalone pure + * function (rather than inline component state, as ApplyModelPanel and + * InferencePanel each do) so it's independently testable and its validation + * message can't drift between call sites. + * + * Denoiser training needs no annotations at all — Noise2Noise and Noise2Void + * are both self-supervised on the raw slice pixels themselves — so the + * validation here is about the SCOPE being usable at all, not about + * annotation coverage. Noise2Noise specifically needs to pair up two + * independent-noise realizations of (approximately) the same structure, so a + * scope of a single slice can't train it; Noise2Void and the autoencoder train + * on single slices fine (the first masks pixels within one slice, the second + * reconstructs the slice through a bottleneck). + */ + +export type DenoiserTrainScope = 'current' | 'range' | 'all'; +export type DenoiserScheme = 'n2n' | 'n2v' | 'ae'; + +/** + * Concrete, ascending, deduplicated slice indices for `scope`, clamped to + * `[0, nSlices)`. Returns `[]` when the scope resolves to nothing usable + * (e.g. `currentSlice` outside `[0, nSlices)`, or an empty volume). + */ +export function resolveDenoiserScopeIndices( + scope: DenoiserTrainScope, + currentSlice: number, + rangeStart: number, + rangeEnd: number, + nSlices: number, +): number[] { + if (nSlices <= 0) return []; + if (scope === 'current') { + return currentSlice >= 0 && currentSlice < nSlices ? [currentSlice] : []; + } + if (scope === 'range') { + const lo = Math.max(0, Math.min(rangeStart, rangeEnd)); + const hi = Math.min(nSlices - 1, Math.max(rangeStart, rangeEnd)); + if (hi < lo) return []; + return Array.from({ length: hi - lo + 1 }, (_, i) => lo + i); + } + return Array.from({ length: nSlices }, (_, i) => i); +} + +/** + * Why `indices` can't be used to train `scheme`, phrased for display next to + * the Train button — or null when the scope is ready to submit. + */ +export function denoiserScopeBlockedReason( + indices: number[], + scheme: DenoiserScheme, +): string | null { + if (indices.length === 0) { + return 'No slices selected — choose a different scope.'; + } + if (scheme === 'n2n' && indices.length < 2) { + return 'Noise2Noise needs at least 2 slices to pair up — choose a range or all slices, or switch to Noise2Void.'; + } + return null; +} diff --git a/frontend/src/lib/displayPrefs.test.ts b/frontend/src/lib/displayPrefs.test.ts new file mode 100644 index 0000000..355f774 --- /dev/null +++ b/frontend/src/lib/displayPrefs.test.ts @@ -0,0 +1,39 @@ +import { afterEach, describe, expect, it } from 'vitest'; +import { loadDisplayPrefs, saveDisplayPrefs, type DisplayPrefs } from './displayPrefs'; + +const FULL: DisplayPrefs = { + brightness: 10, + contrast: -5, + levelsLo: 20, + levelsHi: 230, + gamma: 1.2, + colormap: 'viridis', + clahe: true, + sharpen: false, + blur: 1.5, +}; + +afterEach(() => { + localStorage.clear(); +}); + +describe('displayPrefs', () => { + it('returns an empty object when nothing has been saved', () => { + expect(loadDisplayPrefs()).toEqual({}); + }); + + it('round-trips a saved prefs object', () => { + saveDisplayPrefs(FULL); + expect(loadDisplayPrefs()).toEqual(FULL); + }); + + it('returns an empty object for corrupt stored JSON instead of throwing', () => { + localStorage.setItem('finch:displayPrefs', '{not json'); + expect(loadDisplayPrefs()).toEqual({}); + }); + + it('returns an empty object when the stored value is not an object', () => { + localStorage.setItem('finch:displayPrefs', '"just a string"'); + expect(loadDisplayPrefs()).toEqual({}); + }); +}); diff --git a/frontend/src/lib/displayPrefs.ts b/frontend/src/lib/displayPrefs.ts new file mode 100644 index 0000000..84eedff --- /dev/null +++ b/frontend/src/lib/displayPrefs.ts @@ -0,0 +1,43 @@ +/** + * Persists the purely cosmetic display sliders (brightness/contrast/levels/ + * gamma/colormap/clahe/sharpen/blur) across reloads, scoped globally as a + * viewer preference — not per-sample, not part of the draft. `denoise` is + * deliberately excluded: it's a real training input tuned to one volume's + * noise level, and carrying it over to a different volume would silently + * mis-filter it (see datasetStore.ts's own note on why it resets). + */ +import type { ColormapName } from '@/lib/colormaps'; + +export interface DisplayPrefs { + brightness: number; + contrast: number; + levelsLo: number; + levelsHi: number; + gamma: number; + colormap: ColormapName; + clahe: boolean; + sharpen: boolean; + blur: number; +} + +const STORAGE_KEY = 'finch:displayPrefs'; + +export function loadDisplayPrefs(): Partial { + try { + const raw = localStorage.getItem(STORAGE_KEY); + if (!raw) return {}; + const parsed = JSON.parse(raw); + return parsed && typeof parsed === 'object' ? parsed : {}; + } catch { + return {}; + } +} + +export function saveDisplayPrefs(prefs: DisplayPrefs): void { + try { + localStorage.setItem(STORAGE_KEY, JSON.stringify(prefs)); + } catch { + // Storage unavailable (private browsing, quota, etc.) — cosmetic prefs + // just won't persist this session; nothing to recover from. + } +} diff --git a/frontend/src/lib/displayTransform.test.ts b/frontend/src/lib/displayTransform.test.ts new file mode 100644 index 0000000..8b425af --- /dev/null +++ b/frontend/src/lib/displayTransform.test.ts @@ -0,0 +1,118 @@ +import { describe, it, expect } from 'vitest'; +import { + displayAffineFor, + baseToDisplay, + displayToBase, + displayBandToBase, + remapHistogramToDisplay, +} from './displayTransform'; + +const IDENTITY = displayAffineFor(0, 0, 0, 255); + +describe('displayAffineFor', () => { + it('is the identity with neutral settings', () => { + expect(IDENTITY.slope).toBeCloseTo(1, 6); + expect(IDENTITY.intercept).toBeCloseTo(0, 6); + for (const v of [0, 37, 128, 200, 255]) { + expect(baseToDisplay(v, IDENTITY)).toBeCloseTo(v, 4); + } + }); + + it('matches renderAdjusted for brightness + contrast + levels', () => { + // Mirror of the pipeline: brighten, contrast about mid-grey, then levels. + const brightness = 0.2, contrast = 40, lo = 20, hi = 200; + const affine = displayAffineFor(brightness, contrast, lo, hi); + const manual = (v: number) => { + const adjust = Math.pow((contrast + 100) / 100, 2); + let x = v + brightness * 255; + x = ((x / 255 - 0.5) * adjust + 0.5) * 255; + x = x <= lo ? 0 : x >= hi ? 255 : ((x - lo) * 255) / (hi - lo); + return x; + }; + for (const v of [0, 30, 64, 128, 190, 255]) { + expect(baseToDisplay(v, affine)).toBeCloseTo(Math.max(0, Math.min(255, manual(v))), 3); + } + }); +}); + +describe('displayToBase', () => { + it('round-trips baseToDisplay in the unclamped interior', () => { + const affine = displayAffineFor(0.1, 25, 10, 240); + for (const v of [60, 100, 128, 160]) { + const shown = baseToDisplay(v, affine); + expect(displayToBase(shown, affine)).toBeCloseTo(v, 3); + } + }); + + it('round-trips through gamma too', () => { + const affine = displayAffineFor(0, 0, 0, 255); + for (const gamma of [0.5, 1, 2.2]) { + for (const v of [40, 128, 210]) { + const shown = baseToDisplay(v, affine, gamma); + expect(displayToBase(shown, affine, gamma)).toBeCloseTo(v, 3); + } + } + }); + + it('returns null when the transform collapses (contrast -100)', () => { + const affine = displayAffineFor(0, -100, 0, 255); + expect(displayToBase(128, affine)).toBeNull(); + }); +}); + +describe('displayBandToBase', () => { + it('selects the same pixels as thresholding the adjusted image', () => { + const affine = displayAffineFor(0.15, 30, 15, 230); + const gamma = 1.4; + const lo = 90, hi = 180; + const band = displayBandToBase(lo, hi, affine, gamma); + // For every base value, gating in base space must agree with gating on the + // value the screen actually shows — that equivalence is the whole point. + for (let v = 0; v <= 255; v++) { + const shown = baseToDisplay(v, affine, gamma); + const viaDisplay = shown >= lo && shown <= hi; + const viaBase = v >= band.lo && v <= band.hi; + expect(viaBase).toBe(viaDisplay); + } + }); + + it('treats band edges at 0 / 255 as open', () => { + const affine = displayAffineFor(0.3, 0, 0, 255); + const band = displayBandToBase(0, 255, affine); + expect(band.lo).toBe(-Infinity); + expect(band.hi).toBe(Infinity); + }); + + it('is all-or-nothing when the transform collapses', () => { + const affine = displayAffineFor(0, -100, 0, 255); + // Everything renders mid-grey, so a band containing it selects everything. + const covering = displayBandToBase(100, 200, affine); + expect(covering.lo).toBe(-Infinity); + expect(covering.hi).toBe(Infinity); + const missing = displayBandToBase(0.5, 3, affine); + expect(missing.lo).toBeGreaterThan(missing.hi); // empty band + }); +}); + +describe('remapHistogramToDisplay', () => { + it('preserves total count', () => { + const bins = new Array(256).fill(0).map((_, i) => i); + const out = remapHistogramToDisplay(bins, displayAffineFor(0.1, 20, 0, 255)); + const sum = (a: number[]) => a.reduce((x, y) => x + y, 0); + expect(sum(out)).toBe(sum(bins)); + }); + + it('is a no-op for the identity transform', () => { + const bins = new Array(256).fill(0); + bins[10] = 5; bins[200] = 7; + expect(remapHistogramToDisplay(bins, IDENTITY)).toEqual(bins); + }); + + it('shifts mass brighter when brightness increases', () => { + const bins = new Array(256).fill(0); + bins[100] = 1000; + const out = remapHistogramToDisplay(bins, displayAffineFor(0.2, 0, 0, 255)); + expect(out[100]).toBe(0); + expect(out.findIndex((c) => c > 0)).toBeGreaterThan(100); + }); +}); diff --git a/frontend/src/lib/displayTransform.ts b/frontend/src/lib/displayTransform.ts new file mode 100644 index 0000000..eb7c0d9 --- /dev/null +++ b/frontend/src/lib/displayTransform.ts @@ -0,0 +1,118 @@ +/** + * displayTransform — the display chain the canvas applies on the GPU + * (brightness/contrast/levels as one affine, then gamma), expressed as plain math + * so tools can reason about it without re-rendering pixels. + * + * The canvas bakes nonlinear preprocessors (blur/CLAHE/sharpen) into an offscreen + * base and then applies THIS transform as an SVG filter. Anything that needs to + * know "which pixels look like X on screen" therefore has two options: re-render + * the whole image with the transform baked in (expensive, and it invalidates every + * cached field whenever a slider moves), or keep a field in un-adjusted base space + * and map the QUESTION through the inverse transform instead (O(1) per change). + * + * The Threshold Brush takes the second route: its intensity band is authored in + * displayed space, `displayToBase` maps the band's endpoints back to base space + * once per change, and the comparison happens against the cached base field. The + * selected pixel set is identical either way, because the transform is monotonic. + */ + +export interface DisplayAffine { + /** Per-channel slope on NORMALIZED [0,1] values (SVG feFuncX type="linear"). */ + slope: number; + /** Per-channel intercept on normalized values. */ + intercept: number; +} + +/** + * Fold brightness, contrast, and the levels window into one normalized affine — + * the exact transform `renderAdjusted` bakes and the canvas filter applies: + * Brighten adds `brightness*255`; Contrast scales around mid-grey by + * `((contrast+100)/100)^2`; Levels remaps `[lo,hi] → [0,255]`. + */ +export function displayAffineFor( + brightness: number, + contrast: number, + levelsLo: number, + levelsHi: number, +): DisplayAffine { + const b255 = brightness * 255; + const adjust = Math.pow((contrast + 100) / 100, 2); + const range = Math.max(1, levelsHi - levelsLo); + const slope = adjust * (255 / range); + const intercept = + ((255 / range) * (adjust * b255 + 127.5 * (1 - adjust)) - (255 * levelsLo) / range) / 255; + return { slope, intercept }; +} + +/** Base (un-adjusted) 0–255 value → what the screen shows, 0–255. */ +export function baseToDisplay(base: number, affine: DisplayAffine, gamma = 1): number { + const a = affine.slope * (base / 255) + affine.intercept; + const clamped = a < 0 ? 0 : a > 1 ? 1 : a; + const g = gamma === 1 ? clamped : Math.pow(clamped, gamma); + return g * 255; +} + +/** + * Displayed 0–255 value → the base value that produces it. The inverse of + * `baseToDisplay`, ignoring its output clamp: a displayed 0 or 255 corresponds to + * an unbounded range of base values, which callers handle by treating a band edge + * at the extreme as open (see `displayBandToBase`). + * + * Returns null when the transform collapses (contrast −100 flattens every input to + * a single grey, so no inverse exists). + */ +export function displayToBase(display: number, affine: DisplayAffine, gamma = 1): number | null { + if (Math.abs(affine.slope) < 1e-9) return null; + const g = Math.max(0, Math.min(1, display / 255)); + const a = gamma === 1 ? g : Math.pow(g, 1 / gamma); + return ((a - affine.intercept) / affine.slope) * 255; +} + +/** + * Map an intensity band authored in DISPLAYED space to the equivalent band in base + * space, so a threshold test against a cached base field selects exactly the pixels + * that look in-band on screen. + * + * Edges at 0 / 255 become open (∓Infinity): everything the display clamps to black + * or white is genuinely in-band. When the transform collapses, the band either + * covers everything or nothing depending on where the single output grey lands. + */ +export function displayBandToBase( + lo: number, + hi: number, + affine: DisplayAffine, + gamma = 1, +): { lo: number; hi: number } { + if (Math.abs(affine.slope) < 1e-9) { + // Every base value maps to the same displayed grey; the band is all-or-nothing. + const flat = baseToDisplay(128, affine, gamma); + return flat >= lo && flat <= hi + ? { lo: -Infinity, hi: Infinity } + : { lo: Infinity, hi: -Infinity }; + } + const bLo = lo <= 0 ? -Infinity : displayToBase(lo, affine, gamma) ?? -Infinity; + const bHi = hi >= 255 ? Infinity : displayToBase(hi, affine, gamma) ?? Infinity; + // A negative slope would invert the ordering; keep the band well-formed. + return affine.slope > 0 ? { lo: bLo, hi: bHi } : { lo: bHi, hi: bLo }; +} + +/** + * Re-bin a base-space luminance histogram into displayed space, so a picker drawn + * over it lines up with what the user sees (and with where the threshold band + * actually cuts). Pure 256-bin remap — no pixels touched. + */ +export function remapHistogramToDisplay( + bins: number[], + affine: DisplayAffine, + gamma = 1, +): number[] { + const out = new Array(256).fill(0); + for (let b = 0; b < bins.length; b++) { + const count = bins[b]; + if (!count) continue; + let d = Math.round(baseToDisplay(b, affine, gamma)); + d = d < 0 ? 0 : d > 255 ? 255 : d; + out[d] += count; + } + return out; +} diff --git a/frontend/src/lib/featureChannels.test.ts b/frontend/src/lib/featureChannels.test.ts new file mode 100644 index 0000000..178cc45 --- /dev/null +++ b/frontend/src/lib/featureChannels.test.ts @@ -0,0 +1,196 @@ +import { describe, it, expect } from 'vitest'; +import { + buildChannels, + fitProjection, + projectToScore, + CHANNEL_NAMES, + CHANNEL_COUNT, +} from './featureChannels'; +import { fitBand, sampleHistograms } from './thresholdFit'; + +const W = 64; +const H = 64; + +function mulberry32(seed: number) { + return () => { + seed |= 0; seed = (seed + 0x6d2b79f5) | 0; + let t = Math.imul(seed ^ (seed >>> 15), 1 | seed); + t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t; + return ((t ^ (t >>> 14)) >>> 0) / 4294967296; + }; +} + +/** Inside the central square is the target; outside is background. */ +function masks(): { pos: Uint8Array; neg: Uint8Array; inside: (i: number) => boolean } { + const pos = new Uint8Array(W * H); + const neg = new Uint8Array(W * H); + const inside = (i: number) => { + const x = i % W; + const y = (i / W) | 0; + return x >= 16 && x < 48 && y >= 16 && y < 48; + }; + for (let i = 0; i < W * H; i++) (inside(i) ? pos : neg)[i] = 1; + return { pos, neg, inside }; +} + +/** Best Dice achievable by thresholding a single channel directly. */ +function bandSkill(values: Float32Array, pos: Uint8Array, neg: Uint8Array): number { + const { posHist, negHist } = sampleHistograms(values, pos, neg); + return fitBand(posHist, negHist).skill; +} + +describe('buildChannels', () => { + it('returns one aligned array per named channel', () => { + const field = new Float32Array(W * H).fill(100); + const channels = buildChannels(field, W, H); + expect(channels).toHaveLength(CHANNEL_COUNT); + expect(CHANNEL_NAMES).toHaveLength(CHANNEL_COUNT); + for (const c of channels) expect(c).toHaveLength(W * H); + }); + + it('keeps intensity as the first channel, unmodified', () => { + const field = new Float32Array(W * H); + for (let i = 0; i < field.length; i++) field[i] = i % 200; + const [intensity] = buildChannels(field, W, H); + expect(Array.from(intensity)).toEqual(Array.from(field)); + }); + + it('gives a flat field no texture response', () => { + const field = new Float32Array(W * H).fill(77); + const [, dogFine, dogCoarse, localStd] = buildChannels(field, W, H); + for (let i = 0; i < field.length; i++) { + expect(Math.abs(dogFine[i])).toBeLessThan(1e-3); + expect(Math.abs(dogCoarse[i])).toBeLessThan(1e-3); + expect(localStd[i]).toBeLessThan(1e-2); + } + }); + + it('responds to fine texture where intensity is identical', () => { + // Both regions average 128; only the target is speckled. + const field = new Float32Array(W * H); + const { inside } = masks(); + const rnd = mulberry32(3); + for (let i = 0; i < field.length; i++) { + field[i] = inside(i) ? (rnd() > 0.5 ? 168 : 88) : 128; + } + const [, , , localStd] = buildChannels(field, W, H); + let insideStd = 0, insideN = 0, outsideStd = 0, outsideN = 0; + for (let i = 0; i < field.length; i++) { + if (inside(i)) { insideStd += localStd[i]; insideN++; } + else { outsideStd += localStd[i]; outsideN++; } + } + expect(insideStd / insideN).toBeGreaterThan((outsideStd / outsideN) * 5); + }); +}); + +describe('fitProjection + projectToScore', () => { + it('separates equal-brightness regions that differ only in texture', () => { + // The case Phase 1 cannot solve: identical means, different grain. + const field = new Float32Array(W * H); + const { pos, neg, inside } = masks(); + const rnd = mulberry32(11); + for (let i = 0; i < field.length; i++) { + field[i] = inside(i) ? (rnd() > 0.5 ? 170 : 86) : 128 + (rnd() - 0.5) * 2; + } + + const intensitySkill = bandSkill(field, pos, neg); + const channels = buildChannels(field, W, H); + const projection = fitProjection(channels, pos, neg); + expect(projection).not.toBeNull(); + const score = projectToScore(channels, projection!); + const projectedSkill = bandSkill(score, pos, neg); + + // A bimodal target against a mid-grey background IS partly reachable by an + // intensity band (it can grab one lobe), so the claim is not "intensity is + // useless" — it is that the projection is decisively better. + expect(intensitySkill).toBeLessThan(0.6); + expect(projectedSkill).toBeGreaterThan(0.9); + expect(projectedSkill).toBeGreaterThan(intensitySkill + 0.4); + }); + + it('leans on the texture channel when that is what distinguishes the regions', () => { + // Same fixture as above: equal means, different grain. The fitted weights + // should concentrate on localStd — evidence the projection is separating for + // the right reason rather than getting lucky. + const field = new Float32Array(W * H); + const { pos, neg, inside } = masks(); + const rnd = mulberry32(11); + for (let i = 0; i < field.length; i++) { + field[i] = inside(i) ? (rnd() > 0.5 ? 170 : 86) : 128 + (rnd() - 0.5) * 2; + } + const projection = fitProjection(buildChannels(field, W, H), pos, neg)!; + const stdIndex = CHANNEL_NAMES.indexOf('localStd'); + const dominant = projection.weights + .map((w, i) => ({ w: Math.abs(w), i })) + .sort((a, b) => b.w - a.w)[0].i; + expect(dominant).toBe(stdIndex); + }); + + it('separates a material whose brightness drifts across the image', () => { + // Target is always +25 above its surroundings, but absolute level ramps + // across x — so no single intensity band covers it. + const field = new Float32Array(W * H); + const { pos, neg, inside } = masks(); + for (let i = 0; i < field.length; i++) { + const x = i % W; + const ramp = 40 + (x / W) * 150; + field[i] = inside(i) ? ramp + 25 : ramp; + } + + const intensitySkill = bandSkill(field, pos, neg); + const channels = buildChannels(field, W, H); + const projection = fitProjection(channels, pos, neg)!; + const projectedSkill = bandSkill(projectToScore(channels, projection), pos, neg); + + expect(projectedSkill).toBeGreaterThan(intensitySkill); + expect(projectedSkill).toBeGreaterThan(0.5); + }); + + it('does not do worse than intensity on a cleanly separable target', () => { + const field = new Float32Array(W * H); + const { pos, neg, inside } = masks(); + for (let i = 0; i < field.length; i++) field[i] = inside(i) ? 200 : 60; + + const channels = buildChannels(field, W, H); + const projection = fitProjection(channels, pos, neg)!; + const projectedSkill = bandSkill(projectToScore(channels, projection), pos, neg); + expect(projectedSkill).toBeGreaterThan(0.9); + }); + + it('returns normalised weights', () => { + const field = new Float32Array(W * H); + const { pos, neg, inside } = masks(); + for (let i = 0; i < field.length; i++) field[i] = inside(i) ? 180 : 70; + const projection = fitProjection(buildChannels(field, W, H), pos, neg)!; + expect(Math.hypot(...projection.weights)).toBeCloseTo(1, 6); + }); + + it('returns null when a sample is empty', () => { + const field = new Float32Array(W * H).fill(120); + const channels = buildChannels(field, W, H); + const empty = new Uint8Array(W * H); + const full = new Uint8Array(W * H).fill(1); + expect(fitProjection(channels, empty, full)).toBeNull(); + expect(fitProjection(channels, full, empty)).toBeNull(); + }); + + it('returns null when nothing distinguishes the samples', () => { + // Identical constant field: every channel is constant, so no axis separates. + const field = new Float32Array(W * H).fill(120); + const { pos, neg } = masks(); + expect(fitProjection(buildChannels(field, W, H), pos, neg)).toBeNull(); + }); + + it('produces scores inside the 0–255 range the band fitter expects', () => { + const field = new Float32Array(W * H); + const { pos, neg, inside } = masks(); + const rnd = mulberry32(5); + for (let i = 0; i < field.length; i++) field[i] = (inside(i) ? 150 : 90) + (rnd() - 0.5) * 40; + const channels = buildChannels(field, W, H); + const score = projectToScore(channels, fitProjection(channels, pos, neg)!); + for (let i = 0; i < score.length; i++) { + expect(score[i]).toBeGreaterThanOrEqual(0); + expect(score[i]).toBeLessThanOrEqual(255); + } + }); +}); diff --git a/frontend/src/lib/featureChannels.ts b/frontend/src/lib/featureChannels.ts new file mode 100644 index 0000000..b247685 --- /dev/null +++ b/frontend/src/lib/featureChannels.ts @@ -0,0 +1,186 @@ +/** + * featureChannels — extra per-pixel descriptors for the Threshold Brush's fit. + * + * Phase 1 gates on intensity alone, which fails in two specific ways that show up + * constantly on tomography: + * + * 1. *Overlapping intensity.* Two materials share a grey range but differ in + * texture or grain scale. No band on intensity can separate them. + * 2. *Drifting intensity.* The same material is darker on one side of the slice + * (beam hardening, illumination), so one global band fits the middle and + * misses both ends. + * + * Each channel below answers one of those. They are all built from Gaussian + * blurs, which are separable and O(n) — the same job an FFT band-pass would do + * for this purpose, without the transform. + * + * Everything is pure and canvas-free so the maths can be tested directly. + */ +import { applyGaussianBlurGray } from '@/lib/blur'; + +/** Channels evaluated for every sampled pixel, in a fixed order. */ +export const CHANNEL_NAMES = ['intensity', 'dogFine', 'dogCoarse', 'localStd', 'meanRatio'] as const; +export type ChannelName = (typeof CHANNEL_NAMES)[number]; +export const CHANNEL_COUNT = CHANNEL_NAMES.length; + +/** Scales (sigma, px) the texture channels are built at. */ +const SIGMA_FINE = 1; +const SIGMA_MID = 2.5; +const SIGMA_COARSE = 6; + +function blurred(src: Float32Array, w: number, h: number, sigma: number): Float32Array { + const out = Float32Array.from(src); + applyGaussianBlurGray(out, w, h, sigma); + return out; +} + +/** + * Build every channel for a field crop. + * + * Returns one Float32Array per channel, all `w*h` long and index-aligned with + * `field`, so a pixel's feature vector is `channels.map(c => c[i])`. + */ +export function buildChannels(field: Float32Array, w: number, h: number): Float32Array[] { + const fine = blurred(field, w, h, SIGMA_FINE); + const mid = blurred(field, w, h, SIGMA_MID); + const coarse = blurred(field, w, h, SIGMA_COARSE); + + const n = field.length; + const dogFine = new Float32Array(n); + const dogCoarse = new Float32Array(n); + const localStd = new Float32Array(n); + const meanRatio = new Float32Array(n); + + // Local variance via E[x²] − E[x]², both from the same Gaussian window. + const sq = new Float32Array(n); + for (let i = 0; i < n; i++) sq[i] = field[i] * field[i]; + const sqMean = blurred(sq, w, h, SIGMA_MID); + + for (let i = 0; i < n; i++) { + // Difference-of-Gaussians: a band-pass. Responds to structure at a scale, + // which is how two materials of equal brightness but different grain are + // told apart (failure mode 1). + dogFine[i] = fine[i] - mid[i]; + dogCoarse[i] = mid[i] - coarse[i]; + + const variance = Math.max(0, sqMean[i] - mid[i] * mid[i]); + localStd[i] = Math.sqrt(variance); + + // Brightness relative to the neighbourhood rather than absolute. Constant for + // a material even where the illumination drifts (failure mode 2). + const denom = Math.abs(coarse[i]) + 1e-3; + meanRatio[i] = (field[i] - coarse[i]) / denom; + } + + return [field, dogFine, dogCoarse, localStd, meanRatio]; +} + +/** Mean and standard deviation of `values` at the indices where `mask` is set. */ +function moments(values: Float32Array, mask: Uint8Array): { mean: number; count: number } { + let sum = 0; + let count = 0; + for (let i = 0; i < mask.length; i++) { + if (!mask[i]) continue; + sum += values[i]; + count++; + } + return { mean: count ? sum / count : 0, count }; +} + +/** + * Weights projecting the channels onto the single axis that best separates the + * two samples — a diagonal-covariance LDA (a.k.a. naive Fisher discriminant). + * + * Full LDA would invert the channel covariance matrix; with five channels and + * possibly few sampled pixels that inverse is easily ill-conditioned, and a + * silently unstable projection is worse than a slightly suboptimal one. Using + * per-channel variance only is stable for any sample size and, once each channel + * is standardised, loses little. + * + * Returns weights aligned with `channels`, plus the per-channel standardisation + * needed to apply them to new pixels. + */ +export function fitProjection( + channels: Float32Array[], + pos: Uint8Array, + neg: Uint8Array, +): { weights: number[]; centers: number[]; scales: number[] } | null { + const weights: number[] = []; + const centers: number[] = []; + const scales: number[] = []; + + for (const channel of channels) { + const p = moments(channel, pos); + const n = moments(channel, neg); + if (p.count === 0 || n.count === 0) return null; + + // Standardise on the pooled spread so channels with different units (grey + // levels vs a ratio) contribute comparably. + let ss = 0; + let total = 0; + for (const [mask, m] of [[pos, p.mean], [neg, n.mean]] as const) { + for (let i = 0; i < mask.length; i++) { + if (!mask[i]) continue; + const d = channel[i] - m; + ss += d * d; + total++; + } + } + const variance = total > 1 ? ss / (total - 1) : 0; + const sd = Math.sqrt(variance); + const center = (p.mean + n.mean) / 2; + + if (!(sd > 1e-6)) { + // Constant channel — carries no information; drop it rather than dividing + // by ~0 and letting numerical noise dominate the projection. + weights.push(0); + centers.push(center); + scales.push(1); + continue; + } + weights.push((p.mean - n.mean) / sd); + centers.push(center); + scales.push(sd); + } + + const magnitude = Math.hypot(...weights); + if (!(magnitude > 1e-9)) return null; // no channel separates anything + return { weights: weights.map((w) => w / magnitude), centers, scales }; +} + +/** + * Project every pixel onto the fitted axis and rescale to the 0–255 range the + * band fitter and the overlay already speak — so Phase 2 changes what the number + * *means* without changing any of the machinery that consumes it. + */ +export function projectToScore( + channels: Float32Array[], + projection: { weights: number[]; centers: number[]; scales: number[] }, +): Float32Array { + const n = channels[0].length; + const raw = new Float32Array(n); + const { weights, centers, scales } = projection; + + for (let c = 0; c < channels.length; c++) { + const w = weights[c]; + if (w === 0) continue; + const channel = channels[c]; + const center = centers[c]; + const scale = scales[c] || 1; + for (let i = 0; i < n; i++) raw[i] += w * ((channel[i] - center) / scale); + } + + // Robust rescale: 1st–99th percentile onto 0–255, so a few extreme pixels can't + // squash everything else into one bin. + const sorted = Float32Array.from(raw).sort(); + const lo = sorted[Math.floor(0.01 * (sorted.length - 1))]; + const hi = sorted[Math.floor(0.99 * (sorted.length - 1))]; + const span = hi - lo; + const out = new Float32Array(n); + if (!(span > 1e-9)) return out; + for (let i = 0; i < n; i++) { + const v = ((raw[i] - lo) / span) * 255; + out[i] = v < 0 ? 0 : v > 255 ? 255 : v; + } + return out; +} diff --git a/frontend/src/lib/featureManifold.test.ts b/frontend/src/lib/featureManifold.test.ts new file mode 100644 index 0000000..e354150 --- /dev/null +++ b/frontend/src/lib/featureManifold.test.ts @@ -0,0 +1,36 @@ +import { describe, expect, it } from 'vitest'; +import { colorizeManifoldHeatmap, manifoldMarkerRect } from './featureManifold'; + +describe('colorizeManifoldHeatmap', () => { + it('returns a canvas matching the grayscale map size', () => { + const w = 4; + const h = 3; + const gray = new Uint8Array(w * h); + gray[0] = 255; + gray[1] = 0; + const canvas = colorizeManifoldHeatmap(gray, w, h, 0.5); + expect(canvas.width).toBe(w); + expect(canvas.height).toBe(h); + }); +}); + +describe('manifoldMarkerRect', () => { + it('uses server box coordinates when present (no edge shift)', () => { + const r = manifoldMarkerRect( + { x: 10, y: 10, box_size: 64, box: { x0: 0, y0: 0, x1: 32, y1: 32 } }, + { side: 64, width: 400, height: 300 }, + ); + expect(r).toEqual({ x: 0, y: 0, width: 32, height: 32 }); + }); + + it('clips reconstructed squares instead of shifting them', () => { + const r = manifoldMarkerRect( + { x: 10, y: 50 }, + { side: 64, width: 200, height: 200 }, + ); + expect(r.x).toBe(0); + expect(r.width).toBe(42); // 10 + 32, clipped — not shifted to 64 wide + expect(r.y).toBe(18); + expect(r.height).toBe(64); + }); +}); diff --git a/frontend/src/lib/featureManifold.ts b/frontend/src/lib/featureManifold.ts new file mode 100644 index 0000000..fdf2f0e --- /dev/null +++ b/frontend/src/lib/featureManifold.ts @@ -0,0 +1,92 @@ +/** + * Feature-manifold heatmap helpers for Suggest Labels overlay. + */ +import { loadLabelPng } from '@/lib/pixelClf'; + +export interface ManifoldPoint { + x: number; + y: number; + cluster: number; + score?: number; + /** Exclusion / packing radius — spacing only, not box half-size. */ + radius?: number; + dist?: number; + /** Full square side in image pixels (preferred for drawing). */ + box_size?: number; + /** Axis-aligned box; x1/y1 are exclusive max coords (width = x1 − x0). */ + box?: { x0: number; y0: number; x1: number; y1: number }; +} + +/** + * Axis-aligned marker rect in image pixels for a suggest-label center. + * + * Prefer the server `box` (already clipped, non-overlapping). Never use + * exclusion `radius` for geometry. When reconstructing, **clip** — do not + * shift to preserve side length (shifting causes overlays to overlap). + */ +export function manifoldMarkerRect( + point: Pick, + opts: { side?: number; width: number; height: number }, +): { x: number; y: number; width: number; height: number } { + if (point.box) { + const { x0, y0, x1, y1 } = point.box; + return { + x: x0, + y: y0, + width: Math.max(1, x1 - x0), + height: Math.max(1, y1 - y0), + }; + } + const side = Math.max(8, opts.side ?? point.box_size ?? 64); + const half = side / 2; + const x0 = Math.max(0, point.x - half); + const y0 = Math.max(0, point.y - half); + const x1 = Math.min(opts.width, point.x + half); + const y1 = Math.min(opts.height, point.y + half); + return { + x: x0, + y: y0, + width: Math.max(1, x1 - x0), + height: Math.max(1, y1 - y0), + }; +} + +/** Cyan→magenta coverage colormap; alpha scales with score. */ +export function colorizeManifoldHeatmap( + gray: Uint8Array, + width: number, + height: number, + opacity = 0.45, +): HTMLCanvasElement { + const canvas = document.createElement('canvas'); + canvas.width = width; + canvas.height = height; + const ctx = canvas.getContext('2d'); + if (!ctx) return canvas; + const img = ctx.createImageData(width, height); + const out = img.data; + const aScale = Math.round(Math.min(1, Math.max(0, opacity)) * 255); + for (let i = 0; i < gray.length; i++) { + const t = gray[i] / 255; + // low residual (dark) → cool; high residual interestingness → warm + const r = Math.round(40 + 200 * t); + const g = Math.round(180 * (1 - t) + 40 * t); + const b = Math.round(220 * (1 - 0.3 * t)); + const p = i * 4; + out[p] = r; + out[p + 1] = g; + out[p + 2] = b; + out[p + 3] = Math.round(aScale * (0.15 + 0.85 * t)); + } + ctx.putImageData(img, 0, 0); + return canvas; +} + +/** Load grayscale coverage PNG and colorize. */ +export async function loadManifoldHeatmapCanvas( + url: string, + opacity = 0.45, +): Promise { + const { data, width, height } = await loadLabelPng(url); + return colorizeManifoldHeatmap(data, width, height, opacity); +} diff --git a/frontend/src/lib/gatherTrainingSources.ts b/frontend/src/lib/gatherTrainingSources.ts new file mode 100644 index 0000000..1003971 --- /dev/null +++ b/frontend/src/lib/gatherTrainingSources.ts @@ -0,0 +1,72 @@ +/** + * gatherTrainingSources — turn a set of selected, annotated sourceKeys into + * the /api/train/start `sources` list. + * + * Mirrors DownloadModal.tsx's multi-source export scope-building: training + * data is drawn from `annotationStore.byImage` (every sample touched this + * session), not fetched from the backend for samples annotated in an earlier + * session — same scoping DownloadModal itself uses for "All annotated samples". + * + * Unlike the fork this was ported from, this app has no per-source class + * taxonomy to validate: `classStore` is a single global, flat class list + * shared by every sample, so there is nothing to reconcile across sources — + * the caller passes that one list straight to `/api/train/start` alongside + * whatever this returns. + */ +import { parseSourceKey } from './sourceKey'; +import type { Shape } from '@/stores/annotationStore'; + +export interface TrainingCandidate { + sourceKey: string; + shapeCount: number; +} + +/** Every sourceKey in `byImage` with at least one shape, sorted for stable display. */ +export function listAnnotatedSourceKeys( + byImage: Record>, +): TrainingCandidate[] { + return Object.keys(byImage) + .map((sourceKey) => ({ + sourceKey, + shapeCount: Object.values(byImage[sourceKey]).reduce((n, shapes) => n + shapes.length, 0), + })) + .filter((c) => c.shapeCount > 0) + .sort((a, b) => a.sourceKey.localeCompare(b.sourceKey)); +} + +export interface TrainingSourceItem { + kind: 'tiled' | 'local'; + source: string; + server_uri: string | null; + slices: Record; + split_by_slice: Record; + negative_slices: string[]; +} + +/** + * Build the per-sample training payload for `selectedKeys`. + * + * @throws Error if `selectedKeys` is empty. + */ +export function gatherTrainingSources( + selectedKeys: string[], + byImage: Record>, + splitBySlice: Record>, + negativeSlices: Record, +): TrainingSourceItem[] { + if (selectedKeys.length === 0) { + throw new Error('Select at least one annotated sample to train on.'); + } + + return selectedKeys.map((sk) => { + const { kind, source, serverUri } = parseSourceKey(sk); + return { + kind, + source, + server_uri: serverUri, + slices: byImage[sk] ?? {}, + split_by_slice: splitBySlice[sk] ?? {}, + negative_slices: negativeSlices[sk] ?? [], + }; + }); +} diff --git a/frontend/src/lib/geometry.bbox.test.ts b/frontend/src/lib/geometry.bbox.test.ts new file mode 100644 index 0000000..e8353e1 --- /dev/null +++ b/frontend/src/lib/geometry.bbox.test.ts @@ -0,0 +1,121 @@ +import { describe, it, expect } from 'vitest'; +import { shapeBBox, unionBBox, bboxIntersects, bboxNear, type BBox } from './geometry'; +import { rasterizeShapes, gridFor } from './rasterize'; +import type { Shape } from '@/stores/annotationStore'; + +/** Every set pixel of `shape`, rasterized, must fall inside `box`. This is the + * property the pre-filters depend on: the bbox is a true superset. */ +function bboxCoversRaster(shape: Shape, box: BBox, width = 200, height = 200): boolean { + const { gw, gh, scale } = gridFor(width, height); + const mask = rasterizeShapes([shape], gw, gh, scale); + for (let gy = 0; gy < gh; gy++) { + for (let gx = 0; gx < gw; gx++) { + if (!mask[gy * gw + gx]) continue; + // Grid cell (gx,gy) covers image pixels [gx*scale, (gx+1)*scale). + const x0 = gx * scale, y0 = gy * scale; + const x1 = x0 + scale, y1 = y0 + scale; + // Allow one cell of slack — rasterization rounds outward. + if (x1 < box.x - scale || x0 > box.x + box.w + scale) return false; + if (y1 < box.y - scale || y0 > box.y + box.h + scale) return false; + } + } + return true; +} + +describe('shapeBBox', () => { + it('bounds a rectangle, normalizing negative extents', () => { + const s: Shape = { id: 'r', classId: 0, kind: 'rectangle', x: 30, y: 40, w: -10, h: 20 }; + expect(shapeBBox(s)).toEqual({ x: 20, y: 40, w: 10, h: 20 }); + }); + + it('bounds an ellipse by its radii', () => { + const s: Shape = { id: 'e', classId: 0, kind: 'ellipse', cx: 50, cy: 60, rx: 10, ry: 5 }; + expect(shapeBBox(s)).toEqual({ x: 40, y: 55, w: 20, h: 10 }); + }); + + it('bounds a polygon by its vertices', () => { + const s: Shape = { id: 'p', classId: 0, kind: 'polygon', points: [10, 10, 40, 12, 25, 35] }; + expect(shapeBBox(s)).toEqual({ x: 10, y: 10, w: 30, h: 25 }); + }); + + it('includes polygon holes (superset is always safe)', () => { + const s: Shape = { + id: 'p', classId: 0, kind: 'polygon', + points: [0, 0, 100, 0, 100, 100, 0, 100], + holes: [[20, 20, 40, 20, 40, 40, 20, 40]], + }; + const b = shapeBBox(s); + expect(b).toEqual({ x: 0, y: 0, w: 100, h: 100 }); + }); + + it('expands brush strokes by their radius', () => { + const s: Shape = { + id: 'b', classId: 0, kind: 'brush', + strokes: [{ points: [50, 50, 60, 50], radius: 5, mode: 'paint' }], + }; + expect(shapeBBox(s)).toEqual({ x: 45, y: 45, w: 20, h: 10 }); + }); + + it('covers every rasterized pixel, for each shape kind', () => { + const shapes: Shape[] = [ + { id: 'r', classId: 0, kind: 'rectangle', x: 20, y: 30, w: 60, h: 40 }, + { id: 'e', classId: 0, kind: 'ellipse', cx: 100, cy: 90, rx: 30, ry: 18 }, + { id: 'p', classId: 0, kind: 'polygon', points: [10, 10, 90, 20, 70, 80, 15, 60] }, + { id: 'b', classId: 0, kind: 'brush', strokes: [{ points: [30, 30, 120, 140, 60, 150], radius: 9, mode: 'paint' }] }, + ]; + for (const s of shapes) { + expect(bboxCoversRaster(s, shapeBBox(s))).toBe(true); + } + }); + + it('returns an empty box for a degenerate shape', () => { + const s: Shape = { id: 'p', classId: 0, kind: 'polygon', points: [] }; + const b = shapeBBox(s); + expect(b.w).toBeLessThan(0); + expect(bboxNear(b, { x: 0, y: 0, w: 10, h: 10 })).toBe(false); + }); +}); + +describe('unionBBox', () => { + it('covers all inputs', () => { + const shapes: Shape[] = [ + { id: 'a', classId: 0, kind: 'rectangle', x: 0, y: 0, w: 10, h: 10 }, + { id: 'b', classId: 0, kind: 'rectangle', x: 90, y: 80, w: 10, h: 20 }, + ]; + expect(unionBBox(shapes)).toEqual({ x: 0, y: 0, w: 100, h: 100 }); + }); + + it('is empty for no shapes, and skips degenerate members', () => { + expect(unionBBox([]).w).toBeLessThan(0); + const shapes: Shape[] = [ + { id: 'empty', classId: 0, kind: 'polygon', points: [] }, + { id: 'a', classId: 0, kind: 'rectangle', x: 5, y: 5, w: 10, h: 10 }, + ]; + expect(unionBBox(shapes)).toEqual({ x: 5, y: 5, w: 10, h: 10 }); + }); +}); + +describe('bboxIntersects / bboxNear', () => { + const a: BBox = { x: 0, y: 0, w: 10, h: 10 }; + + it('detects overlap and separation', () => { + expect(bboxIntersects(a, { x: 5, y: 5, w: 10, h: 10 })).toBe(true); + expect(bboxIntersects(a, { x: 20, y: 0, w: 5, h: 5 })).toBe(false); + }); + + it('bboxNear accepts exactly-touching boxes that the strict test rejects', () => { + const touching: BBox = { x: 10, y: 0, w: 5, h: 5 }; + expect(bboxIntersects(a, touching)).toBe(false); + expect(bboxNear(a, touching, 1)).toBe(true); + }); + + it('bboxNear honours the pad', () => { + const gap: BBox = { x: 13, y: 0, w: 5, h: 5 }; + expect(bboxNear(a, gap, 1)).toBe(false); + expect(bboxNear(a, gap, 4)).toBe(true); + }); + + it('bboxNear rejects degenerate boxes', () => { + expect(bboxNear(a, { x: 0, y: 0, w: -1, h: -1 })).toBe(false); + }); +}); diff --git a/frontend/src/lib/geometry.ts b/frontend/src/lib/geometry.ts index f126cfa..c750475 100644 --- a/frontend/src/lib/geometry.ts +++ b/frontend/src/lib/geometry.ts @@ -3,12 +3,111 @@ * All shape coordinates are stored in IMAGE pixels. * The Stage transform (scaleX/scaleY/x/y) is display-only. */ +import type { Shape } from '@/stores/annotationStore'; export interface Point { x: number; y: number; } +/** Axis-aligned bounding box in image pixels. */ +export interface BBox { + x: number; + y: number; + w: number; + h: number; +} + +/** An empty box that intersects nothing (used for degenerate/empty shapes). */ +const EMPTY_BBOX: BBox = { x: 0, y: 0, w: -1, h: -1 }; + +/** Extend `acc` (as [minX,minY,maxX,maxY]) to cover a flat [x,y,…] ring. + * Tolerates a missing/short array: shapes restored from an old draft or version + * payload are not guaranteed to match the current type exactly, and a throw here + * would abort an entire commit. */ +function growByFlat(acc: number[], pts: number[] | undefined, pad = 0): void { + if (!pts) return; + for (let i = 0; i + 1 < pts.length; i += 2) { + const x = pts[i], y = pts[i + 1]; + if (!Number.isFinite(x) || !Number.isFinite(y)) continue; + if (x - pad < acc[0]) acc[0] = x - pad; + if (y - pad < acc[1]) acc[1] = y - pad; + if (x + pad > acc[2]) acc[2] = x + pad; + if (y + pad > acc[3]) acc[3] = y + pad; + } +} + +/** + * Axis-aligned bounds of a shape, in image pixels. + * + * Used to skip expensive geometry work (boolean ops, full-resolution + * rasterization) for shapes that cannot possibly interact — see the callers in + * `clipToClasses` / `mergeSameClass`. Correctness rests on this being a true + * SUPERSET of the shape's covered pixels, so it deliberately over-covers: + * + * - Brush strokes expand by their radius; `erase` strokes are included even + * though they only subtract, since a superset is always safe. + * - Polygon holes are included for the same reason — they can only clear pixels, + * but counting them costs nothing and removes a class of edge case. + * - Vector `erased` carve-outs likewise cannot extend the shape and are ignored. + */ +export function shapeBBox(shape: Shape): BBox { + if (shape.kind === 'rectangle') { + return normalizeRect(shape.x, shape.y, shape.w, shape.h); + } + if (shape.kind === 'ellipse') { + const rx = Math.abs(shape.rx), ry = Math.abs(shape.ry); + return { x: shape.cx - rx, y: shape.cy - ry, w: rx * 2, h: ry * 2 }; + } + + const acc = [Infinity, Infinity, -Infinity, -Infinity]; + if (shape.kind === 'polygon') { + growByFlat(acc, shape.points); + for (const hole of shape.holes ?? []) growByFlat(acc, hole); + } else { + // Brush: every stroke's polyline, fattened by its own radius. + for (const st of shape.strokes ?? []) growByFlat(acc, st?.points, st?.radius ?? 0); + } + if (!Number.isFinite(acc[0])) return EMPTY_BBOX; + return { x: acc[0], y: acc[1], w: acc[2] - acc[0], h: acc[3] - acc[1] }; +} + +/** Union of several shapes' bounds (empty when the list is empty). */ +export function unionBBox(shapes: Shape[]): BBox { + const acc = [Infinity, Infinity, -Infinity, -Infinity]; + for (const s of shapes) { + const b = shapeBBox(s); + if (b.w < 0 || b.h < 0) continue; + if (b.x < acc[0]) acc[0] = b.x; + if (b.y < acc[1]) acc[1] = b.y; + if (b.x + b.w > acc[2]) acc[2] = b.x + b.w; + if (b.y + b.h > acc[3]) acc[3] = b.y + b.h; + } + if (!Number.isFinite(acc[0])) return EMPTY_BBOX; + return { x: acc[0], y: acc[1], w: acc[2] - acc[0], h: acc[3] - acc[1] }; +} + +/** True if two AABBs overlap (strict — touching edges do not count). */ +export function bboxIntersects(a: BBox, b: BBox): boolean { + return a.x < b.x + b.w && a.x + a.w > b.x && a.y < b.y + b.h && a.y + a.h > b.y; +} + +/** + * Overlap test with a tolerance, for deciding whether two shapes can interact. + * + * Prefer this over `bboxIntersects` when the answer gates real geometry work: a + * pair whose bounds merely touch can still produce adjacent set pixels once + * rasterized, and a shape excluded here is never examined again. `pad` should be + * at least one grid cell of whatever raster the caller compares on. + */ +export function bboxNear(a: BBox, b: BBox, pad = 1): boolean { + if (a.w < 0 || a.h < 0 || b.w < 0 || b.h < 0) return false; + return ( + a.x - pad < b.x + b.w && a.x + a.w + pad > b.x && + a.y - pad < b.y + b.h && a.y + a.h + pad > b.y + ); +} + export interface StageTransform { scaleX: number; scaleY: number; diff --git a/frontend/src/lib/importPredictions.test.ts b/frontend/src/lib/importPredictions.test.ts new file mode 100644 index 0000000..80d866d --- /dev/null +++ b/frontend/src/lib/importPredictions.test.ts @@ -0,0 +1,73 @@ +import { describe, expect, it } from 'vitest'; +import { remapPredictedShapes, type RunClass } from './importPredictions'; +import type { AnnotationClass } from '@/stores/classStore'; +import type { Shape } from '@/stores/annotationStore'; + +const shape = (id: string, classId: number): Shape => ({ + id, classId, kind: 'polygon', points: [0, 0, 10, 0, 5, 10], +}); + +describe('remapPredictedShapes', () => { + it('maps a run class to an existing current class by case-insensitive label', () => { + const runClasses: RunClass[] = [{ classId: 1, label: 'Pore', color: '#ff0000' }]; + const current: AnnotationClass[] = [{ classId: 7, label: 'pore', color: '#000000', isVisible: true }]; + const slices = { '0': [shape('s1', 1)] }; + + const result = remapPredictedShapes(runClasses, current, slices); + + expect(result.classes).toEqual(current); // no new class appended + expect(result.slices['0'][0].classId).toBe(7); // remapped to the existing class's id + }); + + it('appends a run class with no label match, keeping its color', () => { + const runClasses: RunClass[] = [{ classId: 1, label: 'Void', color: '#00ff00' }]; + const current: AnnotationClass[] = [{ classId: 3, label: 'Pore', color: '#000000', isVisible: true }]; + const slices = { '0': [shape('s1', 1)] }; + + const result = remapPredictedShapes(runClasses, current, slices); + + expect(result.classes).toHaveLength(2); + const appended = result.classes[1]; + expect(appended.label).toBe('Void'); + expect(appended.color).toBe('#00ff00'); + expect(appended.classId).toBe(4); // max existing (3) + 1 + expect(result.slices['0'][0].classId).toBe(4); + }); + + it('assigns classId 1 when the current class list is empty', () => { + const runClasses: RunClass[] = [{ classId: 1, label: 'Pore', color: '#ff0000' }]; + const result = remapPredictedShapes(runClasses, [], { '0': [shape('s1', 1)] }); + expect(result.classes[0].classId).toBe(1); + }); + + it('handles multiple run classes, some matched and some appended', () => { + const runClasses: RunClass[] = [ + { classId: 1, label: 'Pore', color: '#111111' }, + { classId: 2, label: 'Crack', color: '#222222' }, + ]; + const current: AnnotationClass[] = [{ classId: 5, label: 'pore', color: '#000000', isVisible: true }]; + const slices = { '0': [shape('s1', 1), shape('s2', 2)] }; + + const result = remapPredictedShapes(runClasses, current, slices); + + expect(result.classes).toHaveLength(2); + expect(result.slices['0'].find((s) => s.id === 's1')?.classId).toBe(5); // matched existing + expect(result.slices['0'].find((s) => s.id === 's2')?.classId).toBe(6); // appended + }); + + it('preserves shape geometry, only remapping classId', () => { + const runClasses: RunClass[] = [{ classId: 1, label: 'Pore', color: '#ff0000' }]; + const current: AnnotationClass[] = []; + const result = remapPredictedShapes(runClasses, current, { '0': [shape('s1', 1)] }); + expect(result.slices['0'][0]).toMatchObject({ id: 's1', kind: 'polygon', points: [0, 0, 10, 0, 5, 10] }); + }); + + it('stamps every returned shape origin: "predicted"', () => { + const runClasses: RunClass[] = [{ classId: 1, label: 'Pore', color: '#ff0000' }]; + const current: AnnotationClass[] = []; + const result = remapPredictedShapes(runClasses, current, { + '0': [shape('s1', 1), shape('s2', 1)], + }); + expect(result.slices['0'].every((s) => s.origin === 'predicted')).toBe(true); + }); +}); diff --git a/frontend/src/lib/importPredictions.ts b/frontend/src/lib/importPredictions.ts new file mode 100644 index 0000000..ad3c5b6 --- /dev/null +++ b/frontend/src/lib/importPredictions.ts @@ -0,0 +1,68 @@ +/** + * importPredictions — remap a saved run's predicted shapes onto the + * currently-open sample's class list before importing them as annotations. + * + * A run's classes are a *snapshot* from whenever it was trained — the + * currently open sample's class list may have since been renamed, reordered, + * or extended. Policy: match by case-insensitive label; a run class with no + * matching label is appended to the current class list (keeping the run's + * color); every predicted shape's classId is remapped to the resolved id. + * Deterministic and non-destructive — existing classes/shapes are untouched. + * + * Every returned shape is stamped `origin: 'predicted'` (see `ShapeOrigin` in + * annotationStore) so it renders dashed and can be filtered separately from + * hand-drawn shapes, same as iPred's committed predictions. + */ +import type { AnnotationClass } from '@/stores/classStore'; +import type { Shape } from '@/stores/annotationStore'; + +export interface RunClass { + classId: number; + label: string; + color: string; + isVisible?: boolean; +} + +export interface RemappedPredictions { + /** The (possibly extended) class list to apply via classStore.setClasses. */ + classes: AnnotationClass[]; + /** Predicted shapes with classId remapped to match `classes`. */ + slices: Record; +} + +/** Remap a run's predicted shapes onto `currentClasses`, appending any of the + * run's classes that have no case-insensitive label match. */ +export function remapPredictedShapes( + runClasses: RunClass[], + currentClasses: AnnotationClass[], + predictedSlices: Record, +): RemappedPredictions { + const byLabel = new Map(currentClasses.map((c) => [c.label.trim().toLowerCase(), c.classId])); + const nextClasses = [...currentClasses]; + let nextId = nextClasses.length ? Math.max(...nextClasses.map((c) => c.classId)) + 1 : 1; + + const runIdToResolvedId = new Map(); + for (const runClass of runClasses) { + const key = runClass.label.trim().toLowerCase(); + const existingId = byLabel.get(key); + if (existingId !== undefined) { + runIdToResolvedId.set(runClass.classId, existingId); + continue; + } + const newId = nextId++; + nextClasses.push({ classId: newId, label: runClass.label, color: runClass.color, isVisible: true }); + byLabel.set(key, newId); + runIdToResolvedId.set(runClass.classId, newId); + } + + const slices: Record = {}; + for (const [sliceKey, shapes] of Object.entries(predictedSlices)) { + slices[sliceKey] = shapes.map((shape) => ({ + ...shape, + classId: runIdToResolvedId.get(shape.classId) ?? shape.classId, + origin: 'predicted', + })); + } + + return { classes: nextClasses, slices }; +} diff --git a/frontend/src/lib/ipredApi.test.ts b/frontend/src/lib/ipredApi.test.ts new file mode 100644 index 0000000..8942d7f --- /dev/null +++ b/frontend/src/lib/ipredApi.test.ts @@ -0,0 +1,198 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { + getIpredComposition, getIpredSetup, ipredChannelUrl, ipredHealth, ipredInfer, + ipredPreprocess, ipredRunCommitUrl, ipredRunProbaUrl, ipredRunStatusUrl, + ipredThresholdClass, ipredTrain, listIpredCompositions, listIpredModules, + listIpredSetups, listIpredTrainers, openIpredSession, previewIpredComposition, + upsertIpredComposition, upsertIpredSetup, +} from './ipredApi'; + +beforeEach(() => { + vi.stubGlobal('fetch', vi.fn()); +}); + +afterEach(() => { + vi.unstubAllGlobals(); +}); + +function ok(body: unknown) { + return { ok: true, json: async () => body }; +} + +function errRes(status: number, body: string, statusText = 'Error') { + return { ok: false, status, statusText, text: async () => body }; +} + +describe('ipredHealth', () => { + it('returns the parsed health body', async () => { + (fetch as any).mockResolvedValue(ok({ status: 'ok' })); + expect(await ipredHealth()).toEqual({ status: 'ok' }); + expect(fetch).toHaveBeenCalledWith(expect.stringContaining('/api/ipred/health')); + }); + + it('throws a parsed error with the detail string on failure', async () => { + (fetch as any).mockResolvedValue(errRes(503, JSON.stringify({ detail: 'ipred unreachable' }))); + await expect(ipredHealth()).rejects.toThrow('ipred unreachable'); + }); + + it('falls back to raw text when the error body is not JSON', async () => { + (fetch as any).mockResolvedValue(errRes(500, 'plain text failure')); + await expect(ipredHealth()).rejects.toThrow('plain text failure'); + }); + + it('falls back to statusText when the error body is empty', async () => { + (fetch as any).mockResolvedValue(errRes(500, '', 'Internal Server Error')); + await expect(ipredHealth()).rejects.toThrow('Internal Server Error'); + }); +}); + +describe('openIpredSession', () => { + it('POSTs the payload and returns the session', async () => { + (fetch as any).mockResolvedValue(ok({ session_id: 's1', project_id: 'p1' })); + const result = await openIpredSession({ kind: 'local', source: 'x.tif' }); + expect(result).toEqual({ session_id: 's1', project_id: 'p1' }); + const [url, init] = (fetch as any).mock.calls[0]; + expect(url).toContain('/api/ipred/sessions'); + expect(init.method).toBe('POST'); + expect(JSON.parse(init.body)).toEqual({ kind: 'local', source: 'x.tif' }); + }); +}); + +describe('listIpredSetups', () => { + it('returns the setups array', async () => { + (fetch as any).mockResolvedValue(ok({ setups: [{ id: 's1' }] })); + expect(await listIpredSetups()).toEqual([{ id: 's1' }]); + }); + + it('defaults to an empty array when the setups key is missing', async () => { + (fetch as any).mockResolvedValue(ok({})); + expect(await listIpredSetups()).toEqual([]); + }); +}); + +describe('getIpredSetup', () => { + it('encodes the setup id in the URL', async () => { + (fetch as any).mockResolvedValue(ok({ id: 'a/b' })); + await getIpredSetup('a/b'); + expect(fetch).toHaveBeenCalledWith(expect.stringContaining(encodeURIComponent('a/b'))); + }); +}); + +describe('upsertIpredSetup', () => { + it('POSTs the setup payload', async () => { + (fetch as any).mockResolvedValue(ok({ id: 's1' })); + await upsertIpredSetup({ name: 'setup', kind: 'procedure' }); + const [, init] = (fetch as any).mock.calls[0]; + expect(init.method).toBe('POST'); + }); +}); + +describe('listIpredTrainers', () => { + it('returns the trainers array, defaulting to empty', async () => { + (fetch as any).mockResolvedValue(ok({})); + expect(await listIpredTrainers()).toEqual([]); + (fetch as any).mockResolvedValue(ok({ trainers: ['catboost'] })); + expect(await listIpredTrainers()).toEqual(['catboost']); + }); +}); + +describe('listIpredModules', () => { + it('returns the modules array, defaulting to empty', async () => { + (fetch as any).mockResolvedValue(ok({})); + expect(await listIpredModules()).toEqual([]); + }); +}); + +describe('listIpredCompositions / getIpredComposition / upsertIpredComposition', () => { + it('lists compositions, defaulting to empty', async () => { + (fetch as any).mockResolvedValue(ok({})); + expect(await listIpredCompositions()).toEqual([]); + }); + + it('gets one composition by id', async () => { + (fetch as any).mockResolvedValue(ok({ id: 'c1', name: 'comp', nodes: [], outputs: [] })); + const result = await getIpredComposition('c1'); + expect(result.id).toBe('c1'); + }); + + it('upserts a composition via POST', async () => { + (fetch as any).mockResolvedValue(ok({ id: 'c2', name: 'x', nodes: [], outputs: [] })); + await upsertIpredComposition({ name: 'x', nodes: [], outputs: [] }); + const [url, init] = (fetch as any).mock.calls[0]; + expect(url).toContain('/api/ipred/compositions'); + expect(init.method).toBe('POST'); + }); +}); + +describe('previewIpredComposition', () => { + it('defaults the name to "preview" when omitted', async () => { + (fetch as any).mockResolvedValue(ok({ preview_labels: [] })); + await previewIpredComposition({ nodes: [], outputs: [] }); + const [, init] = (fetch as any).mock.calls[0]; + expect(JSON.parse(init.body).name).toBe('preview'); + }); + + it('keeps an explicit name', async () => { + (fetch as any).mockResolvedValue(ok({ preview_labels: [] })); + await previewIpredComposition({ name: 'draft', nodes: [], outputs: [] }); + const [, init] = (fetch as any).mock.calls[0]; + expect(JSON.parse(init.body).name).toBe('draft'); + }); +}); + +describe('ipredPreprocess', () => { + it('POSTs and returns the preprocess result', async () => { + (fetch as any).mockResolvedValue(ok({ + feature_id: 'f1', project_id: 'p1', setup_id: 's1', slice_index: 0, + n_channels: 3, height: 64, width: 64, labels: [], cache_hit: false, + })); + const result = await ipredPreprocess({ session_id: 's1' }); + expect(result.feature_id).toBe('f1'); + }); +}); + +describe('URL builders', () => { + it('ipredChannelUrl encodes the feature id', () => { + expect(ipredChannelUrl('feat/1', 2)).toContain(encodeURIComponent('feat/1')); + expect(ipredChannelUrl('feat1', 2)).toContain('/channels/2'); + }); + + it('ipredRunCommitUrl/StatusUrl/ProbaUrl build the expected paths', () => { + expect(ipredRunCommitUrl('run/1')).toContain(`${encodeURIComponent('run/1')}/commit.png`); + expect(ipredRunStatusUrl('run1')).toContain('run1/status.png'); + expect(ipredRunProbaUrl('run1', 3)).toContain('run1/proba/3.png'); + }); +}); + +describe('ipredTrain / ipredInfer', () => { + it('ipredTrain POSTs the payload and returns the result', async () => { + (fetch as any).mockResolvedValue(ok({ + model_id: 'm1', feature_id: 'f1', trainer_id: 'catboost', class_ids: [1], + train_accuracy: 0.9, n_train: 10, n_cal: 5, n_samples: 15, params: {}, + })); + const result = await ipredTrain({ session_id: 's1', shapes: [] }); + expect(result.model_id).toBe('m1'); + }); + + it('ipredInfer POSTs the payload and returns the result', async () => { + (fetch as any).mockResolvedValue(ok({ + run_id: 'r1', model_id: 'm1', feature_id: 'f1', alpha: 0.05, class_ids: [1], + counts: { singleton: 1, multi: 0, abstain: 0 }, + })); + const result = await ipredInfer({ session_id: 's1' }); + expect(result.run_id).toBe('r1'); + }); +}); + +describe('ipredThresholdClass', () => { + it('POSTs to the run-scoped threshold-class endpoint', async () => { + (fetch as any).mockResolvedValue(ok({ + run_id: 'r1', class_id: 1, class_index: 0, threshold: 0.5, + width: 10, height: 10, n_positive: 5, label_map_b64: 'abc', + })); + const result = await ipredThresholdClass('r1', { class_id: 1, threshold: 0.5 }); + expect(result.n_positive).toBe(5); + const [url] = (fetch as any).mock.calls[0]; + expect(url).toContain('/runs/r1/threshold-class'); + }); +}); diff --git a/frontend/src/lib/ipredApi.ts b/frontend/src/lib/ipredApi.ts new file mode 100644 index 0000000..240b99d --- /dev/null +++ b/frontend/src/lib/ipredApi.ts @@ -0,0 +1,313 @@ +/** + * Thin client for Annotate → ipred proxy (`/api/ipred/*`). + * Frontend never talks to the ipred port directly. + */ +import { API_BASE } from '@/config'; + +export interface IpredSession { + session_id: string; + project_id: string; + current_feature_id?: string | null; + current_model_id?: string | null; + current_run_id?: string | null; +} + +export interface FeatureSetup { + id: string; + name: string; + kind: 'procedure' | 'weights' | string; + builtin?: boolean; + procedure_id?: string; + params?: Record; + encoder_setup_id?: string; + weights_path?: string; + weights_format?: string; + inference?: Record; + content_hash?: string; +} + +export interface IpredPreprocessResult { + feature_id: string; + project_id: string; + setup_id: string; + slice_index: number; + n_channels: number; + height: number; + width: number; + labels: string[]; + cache_hit: boolean; + blob_dir?: string; +} + +async function parseError(res: Response): Promise { + const text = await res.text(); + try { + const j = JSON.parse(text) as { detail?: unknown }; + if (typeof j.detail === 'string') return new Error(j.detail); + return new Error(text || res.statusText); + } catch { + return new Error(text || res.statusText); + } +} + +export async function ipredHealth(): Promise<{ status: string; service?: string }> { + const res = await fetch(`${API_BASE}/api/ipred/health`); + if (!res.ok) throw await parseError(res); + return res.json() as Promise<{ status: string; service?: string }>; +} + +export async function openIpredSession(payload: { + kind: string; + source: string; + server_uri?: string | null; + root?: string | null; +}): Promise { + const res = await fetch(`${API_BASE}/api/ipred/sessions`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(payload), + }); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} + +export async function listIpredSetups(): Promise { + const res = await fetch(`${API_BASE}/api/ipred/setups`); + if (!res.ok) throw await parseError(res); + const body = (await res.json()) as { setups: FeatureSetup[] }; + return body.setups ?? []; +} + +export async function getIpredSetup(setupId: string): Promise { + const res = await fetch(`${API_BASE}/api/ipred/setups/${encodeURIComponent(setupId)}`); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} + +export async function upsertIpredSetup(payload: { + name: string; + kind: string; + procedure_id?: string; + params?: Record; + encoder_setup_id?: string | null; + weights_path?: string; + weights_format?: string; + inference?: Record; + setup_id?: string; +}): Promise { + const res = await fetch(`${API_BASE}/api/ipred/setups`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(payload), + }); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} + +export async function listIpredTrainers(): Promise { + const res = await fetch(`${API_BASE}/api/ipred/trainers`); + if (!res.ok) throw await parseError(res); + const body = (await res.json()) as { trainers: string[] }; + return body.trainers ?? []; +} + +export interface FeatureModuleInfo { + id: string; + name: string; + description: string; + runtime: string; + ready: boolean; + accepts_input_from: boolean; + produces_channels: boolean; + produces_embedding: boolean; + params_schema: Record; +} + +export interface CompositionNode { + id: string; + module: string; + params?: Record; + input_from?: string; +} + +export interface CompositionDoc { + id: string; + name: string; + kind?: string; + builtin?: boolean; + nodes: CompositionNode[]; + outputs: string[]; + content_hash?: string; + preview_labels?: string[]; +} + +export async function listIpredModules(): Promise { + const res = await fetch(`${API_BASE}/api/ipred/modules`); + if (!res.ok) throw await parseError(res); + const body = (await res.json()) as { modules: FeatureModuleInfo[] }; + return body.modules ?? []; +} + +export async function listIpredCompositions(): Promise { + const res = await fetch(`${API_BASE}/api/ipred/compositions`); + if (!res.ok) throw await parseError(res); + const body = (await res.json()) as { compositions: CompositionDoc[] }; + return body.compositions ?? []; +} + +export async function getIpredComposition(id: string): Promise { + const res = await fetch( + `${API_BASE}/api/ipred/compositions/${encodeURIComponent(id)}`, + ); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} + +export async function upsertIpredComposition(payload: { + name: string; + nodes: CompositionNode[]; + outputs: string[]; + composition_id?: string; + builtin?: boolean; +}): Promise { + const res = await fetch(`${API_BASE}/api/ipred/compositions`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(payload), + }); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} + +export async function previewIpredComposition(payload: { + name?: string; + nodes: CompositionNode[]; + outputs: string[]; +}): Promise<{ preview_labels: string[] }> { + const res = await fetch(`${API_BASE}/api/ipred/compositions/preview`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ name: payload.name ?? 'preview', ...payload }), + }); + if (!res.ok) throw await parseError(res); + return res.json() as Promise<{ preview_labels: string[] }>; +} + +export async function ipredPreprocess(payload: { + session_id: string; + feature_setup_id?: string; + composition_id?: string; + slice_index?: number; + array_ref?: string; +}): Promise { + const res = await fetch(`${API_BASE}/api/ipred/preprocess`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(payload), + }); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} + +export function ipredChannelUrl(featureId: string, index: number): string { + return `${API_BASE}/api/ipred/features/${encodeURIComponent(featureId)}/channels/${index}`; +} + +export interface IpredFeatureImportance { + label: string; + importance: number; +} + +export interface IpredTrainResult { + model_id: string; + feature_id: string; + trainer_id: string; + class_ids: number[]; + train_accuracy: number; + n_train: number; + n_cal: number; + n_samples: number; + params: Record; + /** Trainer-specific (CatBoost); omit or empty when unavailable. */ + feature_importances?: IpredFeatureImportance[]; +} + +export interface IpredInferResult { + run_id: string; + model_id: string; + feature_id: string; + alpha: number; + class_ids: number[]; + counts: { singleton: number; multi: number; abstain: number }; + q_by_class?: Record; +} + +export async function ipredTrain(payload: { + session_id: string; + shapes: unknown[]; + feature_id?: string | null; + trainer_id?: string; + config?: Record; +}): Promise { + const res = await fetch(`${API_BASE}/api/ipred/train`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(payload), + }); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} + +export async function ipredInfer(payload: { + session_id: string; + model_id?: string | null; + feature_id?: string | null; + alpha?: number; +}): Promise { + const res = await fetch(`${API_BASE}/api/ipred/infer`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(payload), + }); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} + +export function ipredRunCommitUrl(runId: string): string { + return `${API_BASE}/api/ipred/runs/${encodeURIComponent(runId)}/commit.png`; +} + +export function ipredRunStatusUrl(runId: string): string { + return `${API_BASE}/api/ipred/runs/${encodeURIComponent(runId)}/status.png`; +} + +export function ipredRunProbaUrl(runId: string, classIndex: number): string { + return `${API_BASE}/api/ipred/runs/${encodeURIComponent(runId)}/proba/${classIndex}.png`; +} + +export interface IpredThresholdClassResult { + run_id: string; + class_id: number; + class_index: number; + threshold: number; + width: number; + height: number; + n_positive: number; + label_map_b64: string; +} + +export async function ipredThresholdClass( + runId: string, + payload: { class_id: number; threshold: number }, +): Promise { + const res = await fetch( + `${API_BASE}/api/ipred/runs/${encodeURIComponent(runId)}/threshold-class`, + { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify(payload), + }, + ); + if (!res.ok) throw await parseError(res); + return res.json() as Promise; +} diff --git a/frontend/src/lib/livewire.test.ts b/frontend/src/lib/livewire.test.ts index 58cc6b2..7bf662c 100644 --- a/frontend/src/lib/livewire.test.ts +++ b/frontend/src/lib/livewire.test.ts @@ -1,5 +1,5 @@ -import { describe, it, expect } from 'vitest'; -import { dijkstra, tracePath, imageToGrid, simplifyPath, type CostMap } from './livewire'; +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { buildCostMap, dijkstra, tracePath, imageToGrid, simplifyPath, type CostMap } from './livewire'; /** 5x5 grid, scale 1, with a zero-cost "edge" corridor along row y=2. */ function corridorMap(): CostMap { @@ -32,6 +32,22 @@ describe('livewire', () => { expect(path[path.length - 2]).toBeCloseTo(4.5, 5); }); + it('maps and traces in native coords at a fractional (upscaled) scale', () => { + // 6x6 grid at scale 0.5 = a 3x3 NATIVE image sampled at 2x. + const gw = 6, gh = 6; + const cost = new Float32Array(gw * gh).fill(1); + for (let x = 0; x < gw; x++) cost[2 * gw + x] = 0; + const cm: CostMap = { gw, gh, scale: 0.5, cost }; + // Native x=1 → grid cell 2 (1 / 0.5). + expect(imageToGrid(cm, 1, 1)).toBe(2 * gw + 2); + const prev = dijkstra(cm, imageToGrid(cm, 0, 1)); + const path = tracePath(cm, prev, imageToGrid(cm, 2.5, 1)); + // Vertices are cell centers in NATIVE coords: gy=2 → 2*0.5 + 0.25 = 1.25. + for (let i = 1; i < path.length; i += 2) expect(path[i]).toBeCloseTo(1.25, 5); + // And they stay inside the 3-px-wide native image. + for (let i = 0; i < path.length; i += 2) expect(path[i]).toBeLessThanOrEqual(3); + }); + it('simplifyPath keeps endpoints and thins the middle', () => { const dense = [0, 0, 1, 0, 2, 0, 3, 0, 4, 0]; // 5 points const simplified = simplifyPath(dense, 2); @@ -39,4 +55,90 @@ describe('livewire', () => { expect(simplified.slice(-2)).toEqual([4, 0]); expect(simplified.length).toBeLessThan(dense.length); }); + + describe('buildCostMap', () => { + const stores = new WeakMap(); + + function seedSource(width: number, height: number, fill: (i: number) => number): HTMLCanvasElement { + const canvas = document.createElement('canvas'); + canvas.width = width; + canvas.height = height; + const data = new Uint8ClampedArray(width * height * 4); + for (let i = 0; i < data.length; i++) data[i] = fill(i); + stores.set(canvas, data); + return canvas; + } + + beforeEach(() => { + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockImplementation(function ( + this: HTMLCanvasElement, + ): any { + const canvas = this; + return { + drawImage: (source: HTMLCanvasElement, _sx: number, _sy: number, w: number, h: number) => { + const src = stores.get(source); + const dst = new Uint8ClampedArray(w * h * 4); + if (src) dst.set(src.subarray(0, dst.length)); + stores.set(canvas, dst); + }, + getImageData: (_x: number, _y: number, w: number, h: number) => ({ + data: stores.get(canvas) ?? new Uint8ClampedArray(w * h * 4), + width: w, + height: h, + }), + }; + }); + }); + + afterEach(() => vi.restoreAllMocks()); + + it('downsamples so the long side stays within maxDim', () => { + const source = seedSource(1000, 500, () => 128); + const cm = buildCostMap(source, 1000, 500, 100); + expect(cm).not.toBeNull(); + expect(Math.max(cm!.gw, cm!.gh)).toBeLessThanOrEqual(100); + }); + + it('produces near-zero cost on a strong edge and near-one cost on a flat region', () => { + // A 10x10 image split vertically: black left half, white right half — + // a strong vertical edge at x=5. + const source = seedSource(10, 10, (i) => { + if (i % 4 === 3) return 255; // alpha + const px = Math.floor(i / 4); + const x = px % 10; + return x < 5 ? 0 : 255; + }); + const cm = buildCostMap(source, 10, 10, 512)!; + expect(cm).not.toBeNull(); + // A flat region far from the edge should have cost close to 1 (no gradient). + const flatIdx = imageToGrid(cm, 1, 5); + // The edge column should have a lower cost than the flat region. + const edgeIdx = imageToGrid(cm, 5, 5); + expect(cm.cost[edgeIdx]).toBeLessThan(cm.cost[flatIdx]); + }); + + it('scales the grid by the upscale factor', () => { + const source = seedSource(4, 4, () => 100); + const cm1 = buildCostMap(source, 4, 4, 512, 1)!; + const cm2 = buildCostMap(source, 4, 4, 512, 2)!; + expect(cm2.gw).toBeGreaterThan(cm1.gw); + }); + + it('returns null when the canvas cannot produce a 2D context', () => { + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockReturnValue(null as any); + const source = seedSource(4, 4, () => 0); + expect(buildCostMap(source, 4, 4)).toBeNull(); + }); + + it('returns null when getImageData throws (tainted canvas)', () => { + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockImplementation(function (): any { + return { + drawImage: () => {}, + getImageData: () => { throw new Error('tainted'); }, + }; + }); + const source = seedSource(4, 4, () => 0); + expect(buildCostMap(source, 4, 4)).toBeNull(); + }); + }); }); diff --git a/frontend/src/lib/livewire.ts b/frontend/src/lib/livewire.ts index 0da38ca..dacfbc1 100644 --- a/frontend/src/lib/livewire.ts +++ b/frontend/src/lib/livewire.ts @@ -10,21 +10,31 @@ export interface CostMap { gw: number; gh: number; - /** Image pixels per grid cell (downsample factor). */ + /** Image pixels per grid cell. >1 when downsampled; FRACTIONAL (e.g. 0.5) when + * the caller asked for an upscaled working resolution. */ scale: number; /** Per-node traversal cost in [~0, 1]; ~0 on strong edges. */ cost: Float32Array; } /** Build an edge-cost map from an image (or preprocessed canvas), downsampled so - * the long side <= maxDim. */ + * the long side <= maxDim. + * + * `imgW`/`imgH` are always NATIVE image pixels. `upscale` (1, 2, 4) raises the + * working resolution — both the grid and the maxDim cap scale with it — so a + * traced path can land on sub-pixel coordinates. `image` should already be + * rendered at that working resolution; it is resampled into the grid either way. */ export function buildCostMap( image: CanvasImageSource, imgW: number, imgH: number, maxDim = 512, + upscale = 1, ): CostMap | null { - const scale = Math.max(1, Math.ceil(Math.max(imgW, imgH) / maxDim)); + const u = Math.max(1, upscale); + // Native cell size from the maxDim cap, then `u` sub-cells per native cell. + const cell = Math.max(1, Math.ceil(Math.max(imgW, imgH) / maxDim)); + const scale = cell / u; const gw = Math.max(1, Math.floor(imgW / scale)); const gh = Math.max(1, Math.floor(imgH / scale)); diff --git a/frontend/src/lib/magicwand.test.ts b/frontend/src/lib/magicwand.test.ts index 25c34b9..257e917 100644 --- a/frontend/src/lib/magicwand.test.ts +++ b/frontend/src/lib/magicwand.test.ts @@ -1,5 +1,5 @@ -import { describe, it, expect } from 'vitest'; -import { magicSelect, maskToPolygons, maskToPolygonsWithHoles, gradientField, type GrayField } from './magicwand'; +import { afterEach, beforeEach, describe, it, expect, vi } from 'vitest'; +import { buildField, magicSelect, maskToPolygons, maskToPolygonsWithHoles, gradientField, otsuThreshold, type GrayField } from './magicwand'; /** 40x40 grid (scale 1): background 0 with two value-200 blocks. */ function twoBlocks(): GrayField { @@ -59,6 +59,36 @@ describe('magicwand', () => { expect(Math.max(...walled[0].filter((_, i) => i % 2 === 0))).toBeLessThan(21); }); + it('blocked keeps the flood from crossing a cell already claimed by another class', () => { + // Uniform intensity field — nothing but `blocked` should stop the flood. + const gw = 40, gh = 40; + const gray = new Float32Array(gw * gh).fill(50); + const field: GrayField = { gw, gh, scale: 1, gray }; + + // Without `blocked` the whole uniform field floods (crosses x=20). + const open = magicSelect(field, 5, 5, { toleranceFrac: 0.5, mode: 'contiguous', smooth: 0, minRegion: 8 }); + expect(open.length).toBe(1); + expect(Math.max(...open[0].filter((_, i) => i % 2 === 0))).toBeGreaterThan(25); + + // A vertical "wall" of another class's pixels at x=20 stops the flood at it. + const blocked = new Uint8Array(gw * gh); + for (let y = 0; y < gh; y++) blocked[y * gw + 20] = 1; + const walled = magicSelect(field, 5, 5, { toleranceFrac: 0.5, mode: 'contiguous', smooth: 0, minRegion: 8, blocked }); + expect(walled.length).toBe(1); + expect(Math.max(...walled[0].filter((_, i) => i % 2 === 0))).toBeLessThan(21); + }); + + it('blocked does not prevent flooding when the seed itself sits on a blocked cell', () => { + const gw = 20, gh = 20; + const gray = new Float32Array(gw * gh).fill(50); + const field: GrayField = { gw, gh, scale: 1, gray }; + const blocked = new Uint8Array(gw * gh); + blocked[9 * gw + 9] = 1; // the seed cell itself + + const polys = magicSelect(field, 9, 9, { toleranceFrac: 0.5, mode: 'contiguous', smooth: 0, minRegion: 8, blocked }); + expect(polys.length).toBe(1); // seed is exempt, matching edgeStop's own wall exemption + }); + it('fills a uniform region with edgeStop on (noise must not wall the flood)', () => { // Left half uniform (100, with faint noise), right half 200, sharp edge at x=30. const gw = 60, gh = 40; @@ -169,3 +199,188 @@ describe('maskToPolygonsWithHoles', () => { expect(out[0].holes.length).toBe(0); }); }); + +describe('fractional scale (upscaled working resolution)', () => { + /** 40x40 grid at scale 0.5 — i.e. a 20x20 NATIVE image sampled at 2x. */ + function upscaledBlock(): GrayField { + const gw = 40, gh = 40; + const gray = new Float32Array(gw * gh); + // Grid cells 8..24 → native image coords 4..12. + for (let y = 8; y < 24; y++) for (let x = 8; x < 24; x++) gray[y * gw + x] = 200; + return { gw, gh, scale: 0.5, gray }; + } + + it('magicSelect seeds and returns polygons in NATIVE image coords', () => { + const field = upscaledBlock(); + // Seed at native (8,8) → grid (16,16), inside the block. + const polys = magicSelect(field, 8, 8, { toleranceFrac: 0.2, mode: 'contiguous', smooth: 0, minRegion: 8 }); + expect(polys.length).toBe(1); + const xs = polys[0].filter((_, i) => i % 2 === 0); + const ys = polys[0].filter((_, i) => i % 2 === 1); + // Native extent of the block is 4..12, not the 8..24 grid extent. + expect(Math.min(...xs)).toBeGreaterThanOrEqual(3); + expect(Math.max(...xs)).toBeLessThanOrEqual(13); + expect(Math.min(...ys)).toBeGreaterThanOrEqual(3); + expect(Math.max(...ys)).toBeLessThanOrEqual(13); + }); + + it('yields sub-pixel (half-integer) vertices a native-resolution grid could not', () => { + const field = upscaledBlock(); + const polys = magicSelect(field, 8, 8, { toleranceFrac: 0.2, mode: 'contiguous', smooth: 0, minRegion: 8 }); + const coords = polys[0]; + expect(coords.some((v) => !Number.isInteger(v))).toBe(true); + }); +}); + +/** + * jsdom has no real 2D canvas, so `buildField` (which draws the source image + * into a small canvas and reads it back) needs `getContext('2d')` stubbed. + * drawImage is a no-op; getImageData returns a fixed RGBA pattern (128 gray + * everywhere) — buildField only cares that SOME data comes back and is + * converted to a `gray` field of the right size, not the actual pixel values. + */ +function stubCanvas2d() { + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockImplementation(function ( + this: HTMLCanvasElement, + ): any { + return { + drawImage: () => {}, + getImageData: (_x: number, _y: number, w: number, h: number) => ({ + data: new Uint8ClampedArray(w * h * 4).fill(128), + width: w, + height: h, + }), + } as unknown as CanvasRenderingContext2D; + }); +} + +describe('buildField', () => { + afterEach(() => vi.restoreAllMocks()); + + it('builds a downsampled gray + gradient field sized to fit maxDim', () => { + stubCanvas2d(); + const field = buildField({} as CanvasImageSource, 3200, 1600, 1600); + expect(field).not.toBeNull(); + // scale = ceil(3200/1600) = 2 -> gw = 3200/2 = 1600, gh = 1600/2 = 800 + expect(field!.gw).toBe(1600); + expect(field!.gh).toBe(800); + expect(field!.scale).toBe(2); + expect(field!.gray.length).toBe(1600 * 800); + expect(field!.grad).toBeInstanceOf(Float32Array); + }); + + it('applies upscale to raise grid resolution and divide the effective scale', () => { + stubCanvas2d(); + const field = buildField({} as CanvasImageSource, 3200, 1600, 1600, 2); + // scale = ceil(3200/1600)/2 = 1 + expect(field!.scale).toBe(1); + expect(field!.gw).toBe(3200); + expect(field!.gh).toBe(1600); + }); + + it('omits the gradient field when needGradient is false', () => { + stubCanvas2d(); + const field = buildField({} as CanvasImageSource, 800, 800, 1600, 1, false); + expect(field!.grad).toBeUndefined(); + }); + + it('returns null when a 2D context is unavailable', () => { + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockReturnValue(null); + const field = buildField({} as CanvasImageSource, 800, 800); + expect(field).toBeNull(); + }); + + it('returns null on a tainted canvas (getImageData throws)', () => { + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockImplementation(function (): any { + return { + drawImage: () => {}, + getImageData: () => { + throw new Error('SecurityError'); + }, + }; + }); + const field = buildField({} as CanvasImageSource, 800, 800); + expect(field).toBeNull(); + }); +}); + +describe('magicSelect smoothing (boxBlur pre-filter + Chaikin rounding)', () => { + it('produces a valid, smaller/rounder contour when smooth > 0', () => { + const field = twoBlocks(); + const smoothed = magicSelect(field, 9, 9, { toleranceFrac: 0.2, mode: 'contiguous', smooth: 3, minRegion: 8 }); + expect(smoothed.length).toBe(1); + expect(smoothed[0].length).toBeGreaterThanOrEqual(6); + }); + + it('caps Chaikin iterations at 4 for very high smooth values without throwing', () => { + const field = twoBlocks(); + expect(() => + magicSelect(field, 9, 9, { toleranceFrac: 0.2, mode: 'contiguous', smooth: 10, minRegion: 8 }), + ).not.toThrow(); + }); +}); + +describe('maskToPolygonsWithHoles — multiple outer regions', () => { + it('assigns a hole to its owning outer via the centroid test when there are 2+ regions', () => { + const W = 100, H = 60; + const m = new Uint8Array(W * H); + // Region A (donut) at x[10,40), region B (solid) at x[60,90). + for (let y = 10; y < 50; y++) for (let x = 10; x < 40; x++) m[y * W + x] = 1; + for (let y = 22; y < 38; y++) for (let x = 18; x < 32; x++) m[y * W + x] = 0; // hole in A + for (let y = 10; y < 50; y++) for (let x = 60; x < 90; x++) m[y * W + x] = 1; // solid B + + const out = maskToPolygonsWithHoles(m, W, H, { minRegion: 8 }); + expect(out.length).toBe(2); + const withHole = out.find((r) => r.holes.length > 0); + const withoutHole = out.find((r) => r.holes.length === 0); + expect(withHole).toBeDefined(); + expect(withoutHole).toBeDefined(); + }); + + it('falls back to the bbox-containment test when the hole centroid lands outside every outer (concave, multi-region)', () => { + const W = 100, H = 60; + const m = new Uint8Array(W * H); + // Region A: concave "C" block with a notch, plus an enclosed hole whose + // centroid falls in the notch (outside A) — same shape as the single-region + // concave test, but with an unrelated solid region B elsewhere so `result.length` + // is 2 and the direct single-region shortcut cannot apply. + for (let y = 8; y < 52; y++) for (let x = 8; x < 52; x++) m[y * W + x] = 1; // block A + for (let y = 8; y < 30; y++) for (let x = 40; x < 52; x++) m[y * W + x] = 0; // notch (concavity) + for (let y = 34; y < 46; y++) for (let x = 16; x < 28; x++) m[y * W + x] = 0; // enclosed hole in A + for (let y = 8; y < 52; y++) for (let x = 70; x < 95; x++) m[y * W + x] = 1; // solid region B + + const out = maskToPolygonsWithHoles(m, W, H, { minRegion: 8 }); + expect(out.length).toBe(2); + const withHole = out.find((r) => r.holes.length > 0); + expect(withHole).toBeDefined(); + }); +}); + +describe('otsuThreshold', () => { + it('splits a clean bimodal histogram between the modes', () => { + const bins = new Array(256).fill(0); + bins[40] = 1000; // dark mode + bins[200] = 1000; // bright mode + const t = otsuThreshold(bins); + expect(t).toBeGreaterThan(40); + expect(t).toBeLessThan(200); + }); + + it('handles broad overlapping modes', () => { + const bins = new Array(256).fill(0); + for (let i = 30; i < 70; i++) bins[i] = 100; + for (let i = 150; i < 220; i++) bins[i] = 100; + const t = otsuThreshold(bins); + expect(t).toBeGreaterThanOrEqual(69); + expect(t).toBeLessThanOrEqual(150); + }); + + it('returns a safe default for degenerate histograms', () => { + expect(otsuThreshold([])).toBe(128); + expect(otsuThreshold(new Array(256).fill(0))).toBe(128); + // Single-valued: no valid two-class split, so the default stands. + const one = new Array(256).fill(0); + one[77] = 500; + expect(otsuThreshold(one)).toBe(128); + }); +}); diff --git a/frontend/src/lib/magicwand.ts b/frontend/src/lib/magicwand.ts index f193f52..aaeb385 100644 --- a/frontend/src/lib/magicwand.ts +++ b/frontend/src/lib/magicwand.ts @@ -13,7 +13,8 @@ import { fillHoles } from '@/lib/morphology'; export interface GrayField { gw: number; gh: number; - /** Image pixels per grid cell (downsample factor). */ + /** Image pixels per grid cell. >1 when downsampled; FRACTIONAL (e.g. 0.5) when + * the caller asked for an upscaled working resolution. */ scale: number; gray: Float32Array; /** Sobel gradient magnitude normalised to ~[0,1] (edge barrier for flood). */ @@ -71,14 +72,24 @@ export function gradientField(gray: Float32Array, gw: number, gh: number): Float /** Build a grayscale + gradient field from an image (long side ≤ maxDim). * * 1600 keeps a 2560px slice at half-resolution (scale 2) — plenty of detail for - * the edge-aware flood while keeping the per-click work ~4x cheaper than full res. */ + * the edge-aware flood while keeping the per-click work ~4x cheaper than full res. + * + * `imgW`/`imgH` are always NATIVE image pixels. `upscale` (1, 2, 4) multiplies the + * grid resolution and divides `scale` to match, so the returned polygons carry + * sub-pixel coordinates while still being expressed in native image space. Pass an + * `image` already rendered at that working resolution to get the extra detail; + * it is resampled into the grid either way. */ export function buildField( image: CanvasImageSource, imgW: number, imgH: number, maxDim = 1600, + upscale = 1, + needGradient = true, ): GrayField | null { - const scale = Math.max(1, Math.ceil(Math.max(imgW, imgH) / maxDim)); + const u = Math.max(1, upscale); + // Native cell size from the maxDim cap, then `u` sub-cells per native cell. + const scale = Math.max(1, Math.ceil(Math.max(imgW, imgH) / maxDim)) / u; const gw = Math.max(1, Math.floor(imgW / scale)); const gh = Math.max(1, Math.floor(imgH / scale)); const canvas = document.createElement('canvas'); @@ -97,7 +108,57 @@ export function buildField( for (let i = 0; i < gw * gh; i++) { gray[i] = 0.299 * data[i * 4] + 0.587 * data[i * 4 + 1] + 0.114 * data[i * 4 + 2]; } - return { gw, gh, scale, gray, grad: gradientField(gray, gw, gh) }; + // The gradient is only needed by the edge-aware flood (`edgeStop`); a pure + // intensity consumer like the Threshold Brush skips a full Sobel pass and half + // the memory by opting out — which matters at a full-resolution 2x/4x grid. + return { gw, gh, scale, gray, ...(needGradient ? { grad: gradientField(gray, gw, gh) } : {}) }; +} + +/** + * Otsu's method: the 0–255 level that best splits a luminance histogram into two + * classes (maximises between-class variance). Drives the Threshold Brush's "Auto" + * button, matching ImageJ's default auto-threshold. Returns 128 for a degenerate + * (empty or single-valued) histogram. + */ +export function otsuThreshold(bins: number[]): number { + const n = bins.length; + if (n === 0) return 128; + let total = 0; + let sumAll = 0; + for (let i = 0; i < n; i++) { + total += bins[i]; + sumAll += i * bins[i]; + } + if (total === 0) return 128; + + let sumB = 0; + let wB = 0; + let bestVar = -1; + // Between-class variance is flat across the empty gap separating two modes, so + // every level in that gap is equally optimal. Average the tied plateau (as + // ImageJ does) to land in the middle of the gap rather than hard against the + // dark mode, which is what a user dragging "Auto" expects. + let tieSum = 0; + let tieCount = 0; + for (let t = 0; t < n; t++) { + wB += bins[t]; + if (wB === 0) continue; + const wF = total - wB; + if (wF === 0) break; + sumB += t * bins[t]; + const mB = sumB / wB; + const mF = (sumAll - sumB) / wF; + const between = wB * wF * (mB - mF) * (mB - mF); + if (between > bestVar * (1 + 1e-12)) { + bestVar = between; + tieSum = t; + tieCount = 1; + } else if (Math.abs(between - bestVar) <= bestVar * 1e-12) { + tieSum += t; + tieCount++; + } + } + return tieCount > 0 ? Math.round(tieSum / tieCount) : 128; } /** Robust intensity spread (2nd–98th percentile) for tolerance scaling. */ @@ -239,6 +300,11 @@ interface SelectOpts { /** 0–1 edge barrier (contiguous only): higher = flood stops at weaker edges. */ edgeStop?: number; minRegion?: number; // min component size in grid pixels + /** Same `gw*gh` grid as `field` (contiguous mode only): cells already claimed + * by a different annotated class the flood must not cross, e.g. so filling + * one region doesn't spill across an already-labeled wall into a neighbor. + * `1` = blocked. Like `edgeStop`'s wall, the seed cell itself is exempt. */ + blocked?: Uint8Array; } interface MaskPolyOpts { @@ -394,7 +460,7 @@ export function magicSelect( field: GrayField, seedXimg: number, seedYimg: number, - { toleranceFrac, mode, smooth = 0, edgeStop = 0, minRegion = 12 }: SelectOpts, + { toleranceFrac, mode, smooth = 0, edgeStop = 0, minRegion = 12, blocked }: SelectOpts, ): number[][] { const { gw, gh, scale, grad } = field; // Smoothing drives a pre-blur (denoise so the boundary is less ragged) here; @@ -410,7 +476,8 @@ export function magicSelect( // flood won't cross — keeps a void's selection bounded by its rim instead of // leaking across a soft/ringy edge. Disabled when edgeStop is 0 or no grad. const wallLimit = edgeStop > 0 ? 1 - edgeStop : Infinity; - const isWall = (i: number) => grad !== undefined && grad[i] >= wallLimit; + const isWall = (i: number) => + (grad !== undefined && grad[i] >= wallLimit) || (blocked !== undefined && blocked[i] === 1); const mask = new Uint8Array(gw * gh); if (mode === 'global') { diff --git a/frontend/src/lib/measure.test.ts b/frontend/src/lib/measure.test.ts new file mode 100644 index 0000000..b323d28 --- /dev/null +++ b/frontend/src/lib/measure.test.ts @@ -0,0 +1,95 @@ +import { describe, expect, it } from 'vitest'; +import { measureRegion } from './measure'; +import type { Shape } from '@/stores/annotationStore'; + +describe('measureRegion', () => { + it('returns an all-zero/null measurement for no shapes', () => { + const m = measureRegion([], 100, 100); + expect(m).toEqual({ count: 0, areaPx: 0, perimeterPx: null, centroid: null, bbox: null }); + }); + + it('measures a single rectangle: area, perimeter, centroid, bbox', () => { + const rect: Shape = { id: 'r1', classId: 1, kind: 'rectangle', x: 10, y: 10, w: 20, h: 10 }; + const m = measureRegion([rect], 100, 100); + expect(m.count).toBe(1); + // area ~ 20*10 = 200 (rasterized, so allow a little slack) + expect(m.areaPx).toBeGreaterThan(150); + expect(m.areaPx).toBeLessThanOrEqual(231); + // analytic perimeter = 2*(20+10) = 60 + expect(m.perimeterPx).toBe(60); + expect(m.centroid).not.toBeNull(); + expect(m.centroid!.x).toBeCloseTo(20, 0); // center x ~ 10+20/2 + expect(m.centroid!.y).toBeCloseTo(15, 0); + expect(m.bbox).not.toBeNull(); + // bbox is (maxX-minX+1)*scale, so it can run 1px wider than the nominal size. + expect(m.bbox!.w).toBeGreaterThanOrEqual(20); + expect(m.bbox!.w).toBeLessThanOrEqual(21); + expect(m.bbox!.h).toBeGreaterThanOrEqual(10); + expect(m.bbox!.h).toBeLessThanOrEqual(11); + }); + + it('sums perimeter and area across multiple shapes (count reflects shapes, not area)', () => { + const r1: Shape = { id: 'r1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 10, h: 10 }; + const r2: Shape = { id: 'r2', classId: 1, kind: 'rectangle', x: 50, y: 50, w: 10, h: 10 }; + const m = measureRegion([r1, r2], 100, 100); + expect(m.count).toBe(2); + expect(m.perimeterPx).toBe(80); // 40 + 40 + }); + + it('computes ellipse perimeter via the Ramanujan approximation', () => { + const ellipse: Shape = { id: 'e1', classId: 1, kind: 'ellipse', cx: 50, cy: 50, rx: 10, ry: 10 }; + const m = measureRegion([ellipse], 100, 100); + // A circle of radius 10: circumference = 2*pi*10 ≈ 62.83 + expect(m.perimeterPx).toBeCloseTo(2 * Math.PI * 10, 0); + }); + + it('computes polygon perimeter as the closed-ring length, including holes', () => { + // A 10x10 square polygon (ring perimeter 40). + const square: Shape = { + id: 'p1', classId: 1, kind: 'polygon', + points: [0, 0, 10, 0, 10, 10, 0, 10], + }; + const m1 = measureRegion([square], 100, 100); + expect(m1.perimeterPx).toBeCloseTo(40, 5); + + // Same square with a 4x4 hole adds the hole's ring perimeter (16). + const withHole: Shape = { + id: 'p2', classId: 1, kind: 'polygon', + points: [0, 0, 10, 0, 10, 10, 0, 10], + holes: [[3, 3, 7, 3, 7, 7, 3, 7]], + }; + const m2 = measureRegion([withHole], 100, 100); + expect(m2.perimeterPx).toBeCloseTo(40 + 16, 5); + }); + + it('returns null perimeter when any shape lacks a defined perimeter (brush)', () => { + const rect: Shape = { id: 'r1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 10, h: 10 }; + const brushShape: Shape = { + id: 'b1', classId: 1, kind: 'brush', + strokes: [{ points: [20, 20, 30, 30], radius: 3, mode: 'paint' }], + }; + const m = measureRegion([rect, brushShape], 100, 100); + expect(m.perimeterPx).toBeNull(); + // Geometry (area/centroid/bbox) is still computed from the union of both shapes. + expect(m.count).toBe(2); + expect(m.areaPx).toBeGreaterThan(0); + expect(m.bbox).not.toBeNull(); + }); + + it('treats overlapping shapes as a union for area/centroid/bbox (not double-counted)', () => { + const a: Shape = { id: 'a', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 10, h: 10 }; + const b: Shape = { id: 'b', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 10, h: 10 }; // identical, fully overlapping + const single = measureRegion([a], 100, 100); + const overlapping = measureRegion([a, b], 100, 100); + expect(overlapping.areaPx).toBeCloseTo(single.areaPx, 5); + // count still reflects the number of shapes passed in, not the union region count. + expect(overlapping.count).toBe(2); + }); + + it('degenerate zero-size rectangle yields a definite (non-crashing) result with zero perimeter contribution area', () => { + const zero: Shape = { id: 'z', classId: 1, kind: 'rectangle', x: 5, y: 5, w: 0, h: 0 }; + const m = measureRegion([zero], 100, 100); + expect(m.count).toBe(1); + expect(m.perimeterPx).toBe(0); + }); +}); diff --git a/frontend/src/lib/mergeSameClass.ts b/frontend/src/lib/mergeSameClass.ts index b55ab8d..67d2468 100644 --- a/frontend/src/lib/mergeSameClass.ts +++ b/frontend/src/lib/mergeSameClass.ts @@ -9,6 +9,7 @@ */ import type { Shape } from '@/stores/annotationStore'; import { gridFor, fullResGridFor, rasterizeShapes, rasterizeUnion } from '@/lib/rasterize'; +import { shapeBBox, unionBBox, bboxNear, type BBox } from '@/lib/geometry'; import { maskToPolygonsWithHoles } from '@/lib/magicwand'; import { unionShapesToPolygons } from '@/lib/polybool'; import { v4 as uuidv4 } from 'uuid'; @@ -27,6 +28,48 @@ function masksIntersect(a: Uint8Array, b: Uint8Array): boolean { return false; } +/** + * True if two 0/1 masks share a set pixel, scanning only the rows/columns where + * both shapes' bounds overlap. Equivalent to `masksIntersect` — outside the + * overlap rect at least one mask is empty by construction — but it avoids walking + * a multi-megapixel grid to answer a question about a small region. + */ +function masksIntersectIn( + a: Uint8Array, b: Uint8Array, gw: number, gh: number, rect: BBox, scale: number, +): boolean { + const x0 = Math.max(0, Math.floor(rect.x / scale)); + const y0 = Math.max(0, Math.floor(rect.y / scale)); + const x1 = Math.min(gw - 1, Math.ceil((rect.x + rect.w) / scale)); + const y1 = Math.min(gh - 1, Math.ceil((rect.y + rect.h) / scale)); + for (let y = y0; y <= y1; y++) { + const row = y * gw; + for (let x = x0; x <= x1; x++) { + const i = row + x; + if (a[i] && b[i]) return true; + } + } + return false; +} + +/** Grow a box by `pad` on every side. */ +function inflate(b: BBox, pad: number): BBox { + return { x: b.x - pad, y: b.y - pad, w: b.w + pad * 2, h: b.h + pad * 2 }; +} + +/** + * Intersection of two boxes, each first grown by `pad` (empty w/h < 0 when they + * remain disjoint). The padding absorbs rasterization rounding: two shapes whose + * boxes merely abut can still land set pixels in the same grid cell. + */ +function bboxOverlap(a: BBox, b: BBox, pad = 0): BBox { + const A = inflate(a, pad), B = inflate(b, pad); + const x = Math.max(A.x, B.x); + const y = Math.max(A.y, B.y); + const w = Math.min(A.x + A.w, B.x + B.w) - x; + const h = Math.min(A.y + A.h, B.y + B.h) - y; + return { x, y, w, h }; +} + /** * Expand a seed selection to the transitive closure of same-class shapes that * overlap it — so selecting one region in an overlapping same-class cluster @@ -46,6 +89,12 @@ export function expandSameClassOverlap( if (!m) { m = rasterizeShapes([s], gw, gh, scale); maskCache.set(s.id, m); } return m; }; + const boxCache = new Map(); + const boxOf = (s: Shape): BBox => { + let b = boxCache.get(s.id); + if (!b) { b = shapeBBox(s); boxCache.set(s.id, b); } + return b; + }; const chosen = new Map(seed.map((s) => [s.id, s])); const classes = new Set(seed.map((s) => s.classId)); @@ -56,14 +105,19 @@ export function expandSameClassOverlap( let changed = true; while (changed && candidates.length > 0) { changed = false; - // Union mask of the current cluster members for this class. + // Union mask + bounds of the current cluster members for this class. The + // bounds let a candidate be rejected without rasterizing it at all, which + // matters because this loop is quadratic in cluster size. const union = new Uint8Array(gw * gh); for (const m of members) { const mm = maskOf(m); for (let i = 0; i < union.length; i++) if (mm[i]) union[i] = 1; } + const unionBounds = unionBBox(members); for (let i = candidates.length - 1; i >= 0; i--) { - if (masksIntersect(maskOf(candidates[i]), union)) { + const cand = candidates[i]; + if (!bboxNear(boxOf(cand), unionBounds, scale + 1)) continue; + if (masksIntersect(maskOf(cand), union)) { const c = candidates.splice(i, 1)[0]; chosen.set(c.id, c); members.push(c); @@ -87,9 +141,11 @@ export function mergeNewWithSameClass( sliceShapes: Shape[], width: number, height: number, + upscale = 1, ): MergeResult { - // Full resolution so the merged region preserves existing geometry (no erosion). - const { gw, gh, scale } = fullResGridFor(width, height); + // Full resolution so the merged region preserves existing geometry (no erosion); + // `upscale` additionally preserves sub-pixel detail from an upscaled working grid. + const { gw, gh, scale } = fullResGridFor(width, height, upscale); const add: Shape[] = []; const removeIds: string[] = []; @@ -101,12 +157,29 @@ export function mergeNewWithSameClass( } for (const [classId, group] of byClass) { - const existing = sliceShapes.filter((s) => s.classId === classId); + const groupBounds = unionBBox(group); + // Bounds pre-filter before ANY rasterization: a shape whose box is disjoint + // from the new geometry cannot share a set pixel with it, so the mask test + // below would always say no. Without this, committing one stroke rasterizes + // every same-class shape on the slice into its own full-resolution buffer + // (millions of cells each) purely to be told they don't touch. + const existing = sliceShapes.filter( + (s) => s.classId === classId && bboxNear(shapeBBox(s), groupBounds, scale + 1), + ); if (existing.length === 0) { add.push(...group); continue; } - // Cheap mask test to decide WHICH existing shapes to merge (fast, tolerant). + // Cheap mask test to decide WHICH of the remaining candidates to merge. + // One reused scratch buffer instead of an allocation per candidate, and the + // comparison only walks the rows/columns where the two boxes actually overlap. const newMask = rasterizeUnion(group, gw, gh, scale); - const overlapping = existing.filter((e) => masksIntersect(rasterizeShapes([e], gw, gh, scale), newMask)); + const scratch = new Uint8Array(gw * gh); + const overlapping = existing.filter((e) => { + const rect = bboxOverlap(shapeBBox(e), groupBounds, scale + 1); + if (rect.w < 0 || rect.h < 0) return false; + scratch.fill(0); + rasterizeShapes([e], gw, gh, scale, scratch); + return masksIntersectIn(scratch, newMask, gw, gh, rect, scale); + }); if (overlapping.length === 0) { add.push(...group); continue; } // Combine via a true polygon boolean union so the existing shapes keep their diff --git a/frontend/src/lib/perf.test.ts b/frontend/src/lib/perf.test.ts new file mode 100644 index 0000000..fd7ee0e --- /dev/null +++ b/frontend/src/lib/perf.test.ts @@ -0,0 +1,90 @@ +import { describe, it, expect, beforeEach, afterEach, vi } from 'vitest'; +import { mark, time, snapshot, resetPerf, perfEnabled, initPerf } from './perf'; + +/** initPerf reads the flag from URL/localStorage; force it on for these tests. */ +function enable() { + window.localStorage.setItem('perf', '1'); + initPerf(); +} + +beforeEach(() => { + window.localStorage.clear(); + resetPerf(); +}); + +afterEach(() => { + window.localStorage.clear(); + initPerf(); + resetPerf(); +}); + +describe('perf gating', () => { + it('is disabled by default and records nothing', () => { + initPerf(); + expect(perfEnabled()).toBe(false); + mark('commit', 123); + expect(snapshot()).toEqual([]); + }); + + it('still runs the timed function when disabled', () => { + initPerf(); + const fn = vi.fn(() => 42); + expect(time('commit', fn)).toBe(42); + expect(fn).toHaveBeenCalledOnce(); + expect(snapshot()).toEqual([]); + }); + + it('enables via localStorage', () => { + enable(); + expect(perfEnabled()).toBe(true); + mark('commit', 5); + expect(snapshot()).toHaveLength(1); + }); +}); + +describe('perf statistics', () => { + beforeEach(enable); + + it('reports p50/p95 over recorded samples', () => { + for (let i = 1; i <= 100; i++) mark('commit', i); + const [stat] = snapshot(); + expect(stat.label).toBe('commit'); + expect(stat.count).toBe(100); + expect(stat.p50).toBeGreaterThanOrEqual(50); + expect(stat.p50).toBeLessThanOrEqual(52); + expect(stat.p95).toBeGreaterThanOrEqual(95); + expect(stat.last).toBe(100); + }); + + it('sorts slowest p95 first', () => { + mark('commit', 1); + mark('clip', 100); + mark('merge', 50); + expect(snapshot().map((s) => s.label)).toEqual(['clip', 'merge', 'commit']); + }); + + it('keeps a bounded window of samples', () => { + for (let i = 0; i < 500; i++) mark('commit', i); + expect(snapshot()[0].count).toBeLessThanOrEqual(120); + }); + + it('records the duration of a timed function and returns its value', () => { + const out = time('merge', () => 'result'); + expect(out).toBe('result'); + const [stat] = snapshot(); + expect(stat.label).toBe('merge'); + expect(stat.count).toBe(1); + expect(stat.last).toBeGreaterThanOrEqual(0); + }); + + it('records a sample even when the timed function throws', () => { + expect(() => time('clip', () => { throw new Error('boom'); })).toThrow('boom'); + expect(snapshot().find((s) => s.label === 'clip')?.count).toBe(1); + }); + + it('resetPerf clears everything', () => { + mark('commit', 10); + resetPerf(); + expect(snapshot()).toEqual([]); + }); +}); diff --git a/frontend/src/lib/perf.ts b/frontend/src/lib/perf.ts new file mode 100644 index 0000000..457f1f2 --- /dev/null +++ b/frontend/src/lib/perf.ts @@ -0,0 +1,131 @@ +/** + * perf — opt-in timing instrumentation for the annotation workspace. + * + * Enabled by `?perf=1` in the URL (sticky for the session) or by setting + * `localStorage.perf = '1'`; disabled by `?perf=0`. When off, `time()` is a + * function call plus a boolean test and `mark()` returns immediately, so + * instrumented paths can stay in the hot code without measurable cost. + * + * Exists because the interesting slices are the user's, not any we can synthesize: + * the only way to confirm an optimization on a 4000² volume with a thousand shapes + * is to measure it there. Rolling p50/p95 rather than a single number, because the + * costs that hurt are the occasional long frames, not the average. + */ + +const SAMPLE_LIMIT = 120; + +export type PerfLabel = + | 'commit' // full stroke/shape commit (clip + merge + store write) + | 'clip' // clipShapesToOthers + | 'merge' // mergeNewWithSameClass + | 'layer-cache' // Konva shapes-layer cache() rebuild + | 'threshold-field' // full-resolution threshold field build + | 'overlay-field' // downsampled overlay field build + | 'overlay-paint' // threshold overlay repaint + | 'magic-field' // magic-wand gray field build + | 'cost-map'; // livewire cost map build + +interface Series { + samples: number[]; + count: number; + /** Ring-buffer write position. */ + pos: number; +} + +const series = new Map(); +let enabled = false; +/** Bumped on every record so subscribers can re-read cheaply. */ +let version = 0; +const listeners = new Set<() => void>(); + +/** Read the enable flag from the URL / localStorage. Safe to call repeatedly. */ +export function initPerf(): boolean { + if (typeof window === 'undefined') return false; + try { + const param = new URLSearchParams(window.location.search).get('perf'); + if (param === '1') window.localStorage.setItem('perf', '1'); + if (param === '0') window.localStorage.removeItem('perf'); + enabled = window.localStorage.getItem('perf') === '1'; + } catch { + enabled = false; // private mode / storage denied + } + return enabled; +} + +export function perfEnabled(): boolean { + return enabled; +} + +// Notifications are throttled: samples can arrive many times per frame during a +// stroke, and re-rendering the HUD on each one would perturb the very timings it +// reports. Recording stays exact; only the display lags by up to NOTIFY_MS. +const NOTIFY_MS = 250; +let notifyTimer: ReturnType | null = null; + +function scheduleNotify(): void { + if (notifyTimer != null || listeners.size === 0) return; + notifyTimer = setTimeout(() => { + notifyTimer = null; + for (const fn of listeners) fn(); + }, NOTIFY_MS); +} + +/** Record one timing sample (ms). No-op when disabled. */ +export function mark(label: PerfLabel, ms: number): void { + if (!enabled) return; + let s = series.get(label); + if (!s) { s = { samples: new Array(SAMPLE_LIMIT).fill(0), count: 0, pos: 0 }; series.set(label, s); } + s.samples[s.pos] = ms; + s.pos = (s.pos + 1) % SAMPLE_LIMIT; + if (s.count < SAMPLE_LIMIT) s.count++; + version++; + scheduleNotify(); +} + +/** Time a synchronous function, recording under `label`. Returns its result. */ +export function time(label: PerfLabel, fn: () => T): T { + if (!enabled) return fn(); + const t0 = performance.now(); + try { + return fn(); + } finally { + mark(label, performance.now() - t0); + } +} + +export interface PerfStat { + label: PerfLabel; + count: number; + p50: number; + p95: number; + last: number; +} + +/** Snapshot of every recorded series, sorted slowest-p95 first. */ +export function snapshot(): PerfStat[] { + const out: PerfStat[] = []; + for (const [label, s] of series) { + if (s.count === 0) continue; + const vals = s.samples.slice(0, s.count).sort((a, b) => a - b); + const at = (q: number) => vals[Math.min(vals.length - 1, Math.floor(q * vals.length))]; + const lastIdx = (s.pos - 1 + SAMPLE_LIMIT) % SAMPLE_LIMIT; + out.push({ label, count: s.count, p50: at(0.5), p95: at(0.95), last: s.samples[lastIdx] }); + } + return out.sort((a, b) => b.p95 - a.p95); +} + +/** Subscribe to new samples (for useSyncExternalStore). */ +export function subscribe(fn: () => void): () => void { + listeners.add(fn); + return () => { listeners.delete(fn); }; +} + +export function getVersion(): number { + return version; +} + +export function resetPerf(): void { + series.clear(); + version++; + for (const fn of listeners) fn(); +} diff --git a/frontend/src/lib/pixelClf.gaps.test.ts b/frontend/src/lib/pixelClf.gaps.test.ts new file mode 100644 index 0000000..161e602 --- /dev/null +++ b/frontend/src/lib/pixelClf.gaps.test.ts @@ -0,0 +1,298 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { + applyProbaThresholdRgba, + classProbaMaskSetName, + colorizeConformalOverlay, + colorizeLabelMap, + labelMapToPolygonShapes, + loadLabelPng, + parseHexColor, + thresholdProbaPngBlob, +} from './pixelClf'; + +/** + * jsdom has no real 2D canvas; stub `getContext('2d')` with an in-memory + * implementation of just the methods pixelClf.ts uses (createImageData / + * putImageData / getImageData), so `colorizeLabelMap` / `colorizeConformalOverlay` + * actually run their pixel-coloring loops instead of hitting the `!ctx` early + * return that made the sibling `pixelClf.overlay.test.ts` a no-op smoke test. + */ +function stubCanvasContext() { + const stores = new WeakMap(); + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockImplementation(function ( + this: HTMLCanvasElement, + ): any { + const canvas = this; + return { + createImageData: (w: number, h: number) => { + const data = new Uint8ClampedArray(w * h * 4); + stores.set(canvas, data); + return { data, width: w, height: h }; + }, + putImageData: () => {}, + getImageData: (_x: number, _y: number, w: number, h: number) => ({ + data: stores.get(canvas) ?? new Uint8ClampedArray(w * h * 4), + width: w, + height: h, + }), + drawImage: () => {}, + } as unknown as CanvasRenderingContext2D; + }); +} + +beforeEach(() => stubCanvasContext()); +afterEach(() => vi.restoreAllMocks()); + +describe('parseHexColor', () => { + it('parses 3-digit shorthand hex by doubling each digit', () => { + expect(parseHexColor('#0f8')).toEqual([0, 255, 136]); + }); +}); + +describe('classProbaMaskSetName', () => { + it('clamps thresholds outside [0,1]', () => { + expect(classProbaMaskSetName(1, 1.5, 'x')).toBe('x p≥100%'); + expect(classProbaMaskSetName(1, -0.5, 'x')).toBe('x p≥0%'); + }); + + it('trims whitespace-only labels and falls back to the class id', () => { + expect(classProbaMaskSetName(4, 0.2, ' ')).toBe('class 4 p≥20%'); + }); +}); + +describe('applyProbaThresholdRgba edge thresholds', () => { + it('threshold 0 colors every pixel (nothing is below the cut)', () => { + const rgba = new Uint8ClampedArray([0, 0, 0, 255, 128, 128, 128, 255]); + applyProbaThresholdRgba(rgba, 0); + expect(rgba[3]).toBe(255); + expect(rgba[7]).toBe(255); + }); + + it('threshold 1 makes everything but a full-white pixel transparent (tiny span floor)', () => { + const rgba = new Uint8ClampedArray([200, 200, 200, 255, 255, 255, 255, 255]); + applyProbaThresholdRgba(rgba, 1); + expect(rgba[3]).toBe(0); // p=200/255 < 1 -> transparent + expect(rgba[7]).toBe(255); // p=1 -> exactly at/above cut + }); +}); + +describe('colorizeLabelMap', () => { + it('colors known classes and leaves 0/unknown classes transparent', () => { + const width = 2, height = 1; + const labels = new Uint8Array([1, 2]); // 2 = unknown (no color registered) + const colors = new Map([[1, '#ff0000']]); + const canvas = colorizeLabelMap(labels, width, height, colors, 90); + const ctx = canvas.getContext('2d')!; + const { data } = ctx.getImageData(0, 0, width, height); + // Pixel 0 -> class 1 -> red at alpha 90. + expect([data[0], data[1], data[2], data[3]]).toEqual([255, 0, 0, 90]); + // Pixel 1 -> class 2 has no registered color -> left transparent (all zero). + expect([data[4], data[5], data[6], data[7]]).toEqual([0, 0, 0, 0]); + }); + + it('leaves class-0 (background) pixels transparent', () => { + const labels = new Uint8Array([0]); + const canvas = colorizeLabelMap(labels, 1, 1, new Map([[1, '#00ff00']])); + const ctx = canvas.getContext('2d')!; + const { data } = ctx.getImageData(0, 0, 1, 1); + expect([...data]).toEqual([0, 0, 0, 0]); + }); +}); + +describe('colorizeConformalOverlay', () => { + it('colors a singleton pixel with its class color at fixed alpha', () => { + const commit = new Uint8Array([1]); + const status = new Uint8Array([1]); // singleton + const canvas = colorizeConformalOverlay(commit, status, 1, 1, new Map([[1, '#00ff00']])); + const { data } = canvas.getContext('2d')!.getImageData(0, 0, 1, 1); + expect([...data]).toEqual([0, 255, 0, 120]); + }); + + it('hides a singleton pixel whose class is filtered out by classVisible', () => { + const commit = new Uint8Array([1]); + const status = new Uint8Array([1]); + const canvas = colorizeConformalOverlay(commit, status, 1, 1, new Map([[1, '#00ff00']]), { + classVisible: () => false, + }); + const { data } = canvas.getContext('2d')!.getImageData(0, 0, 1, 1); + expect([...data]).toEqual([0, 0, 0, 0]); + }); + + it('leaves a singleton pixel transparent when its class has no registered color', () => { + const commit = new Uint8Array([9]); + const status = new Uint8Array([1]); + const canvas = colorizeConformalOverlay(commit, status, 1, 1, new Map([[1, '#00ff00']])); + const { data } = canvas.getContext('2d')!.getImageData(0, 0, 1, 1); + expect([...data]).toEqual([0, 0, 0, 0]); + }); + + it('renders a multi-class hatch pattern that alternates by (x+y) parity, and can be hidden', () => { + // 4x1: status=2 (multi) at every pixel; (x+y)&3 < 2 is true for x=0,1 ("on") + // and false for x=2,3 ("off"), giving both hatch colors in one row. + const commit = new Uint8Array([0, 0, 0, 0]); + const status = new Uint8Array([2, 2, 2, 2]); + const canvas = colorizeConformalOverlay(commit, status, 4, 1, new Map()); + const { data } = canvas.getContext('2d')!.getImageData(0, 0, 4, 1); + expect([data[0], data[1], data[2], data[3]]).toEqual([245, 158, 11, 140]); // x=0: on + expect([data[8], data[9], data[10], data[11]]).toEqual([180, 120, 40, 70]); // x=2: off + + const hidden = colorizeConformalOverlay(commit, status, 4, 1, new Map(), { showMulti: false }); + const hiddenData = hidden.getContext('2d')!.getImageData(0, 0, 4, 1).data; + expect([...hiddenData]).toEqual(new Array(16).fill(0)); + }); + + it('renders the abstain color for status 0 by default, and can hide it', () => { + const commit = new Uint8Array([0]); + const status = new Uint8Array([0]); + const canvas = colorizeConformalOverlay(commit, status, 1, 1, new Map()); + const { data } = canvas.getContext('2d')!.getImageData(0, 0, 1, 1); + expect([...data]).toEqual([30, 30, 40, 90]); + + const hidden = colorizeConformalOverlay(commit, status, 1, 1, new Map(), { showAbstain: false }); + const hiddenData = hidden.getContext('2d')!.getImageData(0, 0, 1, 1).data; + expect([...hiddenData]).toEqual([0, 0, 0, 0]); + }); +}); + +describe('labelMapToPolygonShapes default options', () => { + it('defaults origin to "predicted" when not specified', () => { + const w = 16, h = 16; + const labels = new Uint8Array(w * h); + for (let y = 4; y < 12; y++) for (let x = 4; x < 12; x++) labels[y * w + x] = 1; + const shapes = labelMapToPolygonShapes(labels, w, h, [1], { minRegion: 4, smooth: 0 }); + expect(shapes[0].origin).toBe('predicted'); + }); + + it('honors an explicit "human" origin override', () => { + const w = 16, h = 16; + const labels = new Uint8Array(w * h); + for (let y = 4; y < 12; y++) for (let x = 4; x < 12; x++) labels[y * w + x] = 1; + const shapes = labelMapToPolygonShapes(labels, w, h, [1], { minRegion: 4, smooth: 0, origin: 'human' }); + expect(shapes[0].origin).toBe('human'); + }); + + it('skips classes with no matching pixels entirely', () => { + const w = 8, h = 8; + const labels = new Uint8Array(w * h); // all zero + const shapes = labelMapToPolygonShapes(labels, w, h, [1, 2], { minRegion: 1 }); + expect(shapes).toHaveLength(0); + }); + + it('drops tiny polygons below minRegion vertex count (<6 flat coords) even if any() was true', () => { + // A 1-pixel "region" with default minRegion (64) is dropped by maskToPolygonsWithHoles + // before points.length is even checked, so this also covers the minRegion path. + const w = 32, h = 32; + const labels = new Uint8Array(w * h); + labels[10 * w + 10] = 1; // single pixel + const shapes = labelMapToPolygonShapes(labels, w, h, [1]); + expect(shapes).toHaveLength(0); + }); +}); + +/** + * thresholdProbaPngBlob / loadLabelPng both decode a PNG via `new Image()` and + * read it back through a 2D canvas — jsdom has neither a real `Image` decoder + * nor `HTMLCanvasElement.toBlob`, so both are stubbed here. `stubCanvasContext` + * above already covers getImageData/putImageData/drawImage; `installImageStub` + * drives `Image`'s onload/onerror synchronously (as a microtask) so `await new + * Promise(...)` in the source resolves without needing real image bytes. + */ +function installImageStub(opts: { fail?: boolean; width?: number; height?: number } = {}) { + const { fail = false, width = 2, height = 1 } = opts; + class StubImage { + onload: (() => void) | null = null; + onerror: (() => void) | null = null; + naturalWidth = 0; + naturalHeight = 0; + private _src = ''; + set src(v: string) { + this._src = v; + queueMicrotask(() => { + if (fail) { + this.onerror?.(); + return; + } + this.naturalWidth = width; + this.naturalHeight = height; + this.onload?.(); + }); + } + get src() { + return this._src; + } + } + vi.stubGlobal('Image', StubImage); +} + +describe('thresholdProbaPngBlob', () => { + beforeEach(() => { + global.URL.createObjectURL = vi.fn(() => 'blob:mock-url'); + global.URL.revokeObjectURL = vi.fn(); + }); + afterEach(() => { + vi.unstubAllGlobals(); + }); + + it('decodes, thresholds, and re-encodes as a PNG blob, then revokes the object URL', async () => { + installImageStub({ width: 2, height: 1 }); + HTMLCanvasElement.prototype.toBlob = function (cb: BlobCallback) { + cb(new Blob(['ok'], { type: 'image/png' })); + }; + const out = await thresholdProbaPngBlob(new Blob(), 0.5); + expect(out).toBeInstanceOf(Blob); + expect(URL.revokeObjectURL).toHaveBeenCalledWith('blob:mock-url'); + }); + + it('rejects when the image fails to decode', async () => { + installImageStub({ fail: true }); + await expect(thresholdProbaPngBlob(new Blob(), 0.5)).rejects.toThrow( + 'Failed to decode probability PNG', + ); + // The object URL is still revoked even on failure (finally block). + expect(URL.revokeObjectURL).toHaveBeenCalledWith('blob:mock-url'); + }); + + it('rejects when the 2D context is unavailable', async () => { + installImageStub(); + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockReturnValue(null); + await expect(thresholdProbaPngBlob(new Blob(), 0.5)).rejects.toThrow('2D context unavailable'); + }); + + it('rejects when canvas.toBlob yields no blob', async () => { + installImageStub(); + HTMLCanvasElement.prototype.toBlob = function (cb: BlobCallback) { + cb(null); + }; + await expect(thresholdProbaPngBlob(new Blob(), 0.5)).rejects.toThrow('toBlob failed'); + }); +}); + +describe('loadLabelPng', () => { + afterEach(() => { + vi.unstubAllGlobals(); + }); + + it('decodes a grayscale classId PNG into a Uint8Array + dimensions', async () => { + installImageStub({ width: 2, height: 1 }); + const { data, width, height } = await loadLabelPng('http://example.test/label.png'); + expect(width).toBe(2); + expect(height).toBe(1); + expect(data).toBeInstanceOf(Uint8Array); + expect(data.length).toBe(2); + }); + + it('rejects when the image fails to load', async () => { + installImageStub({ fail: true }); + await expect(loadLabelPng('http://example.test/bad.png')).rejects.toThrow( + 'Failed to load prediction PNG', + ); + }); + + it('rejects when the 2D context is unavailable', async () => { + installImageStub(); + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockReturnValue(null); + await expect(loadLabelPng('http://example.test/label.png')).rejects.toThrow( + '2D context unavailable', + ); + }); +}); diff --git a/frontend/src/lib/pixelClf.overlay.test.ts b/frontend/src/lib/pixelClf.overlay.test.ts new file mode 100644 index 0000000..cb55a76 --- /dev/null +++ b/frontend/src/lib/pixelClf.overlay.test.ts @@ -0,0 +1,23 @@ +import { describe, expect, it } from 'vitest'; +import { colorizeConformalOverlay } from './pixelClf'; + +describe('colorizeConformalOverlay filters', () => { + it('builds a canvas with class / multi / abstain filters without throwing', () => { + const w = 2; + const h = 2; + const commit = new Uint8Array([1, 2, 0, 0]); + const status = new Uint8Array([1, 1, 2, 0]); + const colors = new Map([ + [1, '#ff0000'], + [2, '#00ff00'], + ]); + + const canvas = colorizeConformalOverlay(commit, status, w, h, colors, { + classVisible: (cid) => cid === 1, + showMulti: false, + showAbstain: false, + }); + expect(canvas.width).toBe(2); + expect(canvas.height).toBe(2); + }); +}); diff --git a/frontend/src/lib/pixelClf.test.ts b/frontend/src/lib/pixelClf.test.ts new file mode 100644 index 0000000..8275bd3 --- /dev/null +++ b/frontend/src/lib/pixelClf.test.ts @@ -0,0 +1,100 @@ +import { describe, expect, it } from 'vitest'; +import { + applyProbaThresholdRgba, + classProbaMaskSetName, + labelMapToPolygonShapes, + parseHexColor, +} from './pixelClf'; + +describe('parseHexColor', () => { + it('parses 6-digit hex', () => { + expect(parseHexColor('#ff0000')).toEqual([255, 0, 0]); + }); +}); + +describe('classProbaMaskSetName', () => { + it('formats label and percent threshold', () => { + expect(classProbaMaskSetName(2, 0.65, 'membrane')).toBe('membrane p≥65%'); + }); + + it('falls back to class id when label missing', () => { + expect(classProbaMaskSetName(3, 0.5, null)).toBe('class 3 p≥50%'); + }); +}); + +describe('applyProbaThresholdRgba', () => { + it('makes pixels below the cut transparent and colors the rest viridis', () => { + const rgba = new Uint8ClampedArray([ + 100, 100, 100, 255, + 255, 255, 255, 255, + ]); + applyProbaThresholdRgba(rgba, 0.5); // cut = 0.5 + expect(rgba[3]).toBe(0); // below: transparent + expect(rgba[7]).toBe(255); // above: opaque + // p=1 remaps to viridis end (yellowish) + expect(rgba[4]).toBeGreaterThan(200); + expect(rgba[5]).toBeGreaterThan(200); + }); +}); + +describe('labelMapToPolygonShapes', () => { + it('emits a polygon for a solid class blob', () => { + const w = 16; + const h = 16; + const labels = new Uint8Array(w * h); + for (let y = 4; y < 12; y++) { + for (let x = 4; x < 12; x++) labels[y * w + x] = 1; + } + const shapes = labelMapToPolygonShapes(labels, w, h, [1], { minRegion: 4, smooth: 0 }); + expect(shapes.length).toBeGreaterThanOrEqual(1); + expect(shapes[0].classId).toBe(1); + expect(shapes[0].kind).toBe('polygon'); + expect(shapes[0].points.length).toBeGreaterThanOrEqual(6); + }); + + it('skips pixels covered by preserveShapes', () => { + const w = 16; + const h = 16; + const labels = new Uint8Array(w * h).fill(1); + const preserve = [ + { + id: 'a', + classId: 1, + kind: 'rectangle' as const, + x: 0, + y: 0, + w: 16, + h: 16, + }, + ]; + const shapes = labelMapToPolygonShapes(labels, w, h, [1], { + minRegion: 4, + preserveShapes: preserve, + smooth: 0, + }); + expect(shapes).toHaveLength(0); + }); + + it('carves holes so nested class regions do not fill through each other', () => { + const w = 32; + const h = 32; + const labels = new Uint8Array(w * h); + // Outer class 1 border; inner 10×10 is class 2 + for (let y = 2; y < 30; y++) { + for (let x = 2; x < 30; x++) labels[y * w + x] = 1; + } + for (let y = 11; y < 21; y++) { + for (let x = 11; x < 21; x++) labels[y * w + x] = 2; + } + const shapes = labelMapToPolygonShapes(labels, w, h, [1, 2], { + minRegion: 4, + smooth: 0, + }); + const c1 = shapes.filter((s) => s.classId === 1); + const c2 = shapes.filter((s) => s.classId === 2); + expect(c1.length).toBeGreaterThanOrEqual(1); + expect(c2.length).toBeGreaterThanOrEqual(1); + // Class 1 must have a hole for the class 2 island (pixel-accurate commit). + expect(c1.some((s) => (s.holes?.length ?? 0) > 0)).toBe(true); + }); +}); diff --git a/frontend/src/lib/pixelClf.ts b/frontend/src/lib/pixelClf.ts new file mode 100644 index 0000000..c7b6aec --- /dev/null +++ b/frontend/src/lib/pixelClf.ts @@ -0,0 +1,289 @@ +/** + * Pixel-classifier helpers: label PNG → colorized overlay / polygon shapes. + */ +import { v4 as uuidv4 } from 'uuid'; +import { maskToPolygonsWithHoles } from '@/lib/magicwand'; +import { rasterizeUnion } from '@/lib/rasterize'; +import type { PolygonShape, Shape } from '@/stores/annotationStore'; +import { viridisRgb } from '@/lib/viridis'; + +/** Parse #rgb / #rrggbb into [r,g,b]. */ +export function parseHexColor(hex: string): [number, number, number] { + const h = hex.replace('#', ''); + if (h.length === 3) { + return [ + parseInt(h[0] + h[0], 16), + parseInt(h[1] + h[1], 16), + parseInt(h[2] + h[2], 16), + ]; + } + return [ + parseInt(h.slice(0, 2), 16) || 0, + parseInt(h.slice(2, 4), 16) || 0, + parseInt(h.slice(4, 6), 16) || 0, + ]; +} + +/** Mask-set name for a thresholded softmax class cache. */ +export function classProbaMaskSetName( + classId: number, + threshold: number, + classLabel?: string | null, +): string { + const t = Math.round(Math.min(1, Math.max(0, threshold)) * 100); + const label = (classLabel ?? '').trim() || `class ${classId}`; + return `${label} p≥${t}%`; +} + +/** + * Colorize grayscale probability RGBA with viridis; pixels below ``threshold`` + * are transparent. Values at/above the cut are remapped to [0,1] so the full + * colormap spans the visible range. + */ +export function applyProbaThresholdRgba( + rgba: Uint8ClampedArray, + threshold: number, +): void { + const t = Math.min(1, Math.max(0, threshold)); + const span = Math.max(1e-6, 1 - t); + for (let i = 0; i < rgba.length; i += 4) { + const p = rgba[i] / 255; + if (p < t) { + rgba[i] = 0; + rgba[i + 1] = 0; + rgba[i + 2] = 0; + rgba[i + 3] = 0; + continue; + } + const [r, g, b] = viridisRgb((p - t) / span); + rgba[i] = r; + rgba[i + 1] = g; + rgba[i + 2] = b; + rgba[i + 3] = 255; + } +} + +/** Render a softmax grayscale PNG as viridis, clipped below ``threshold``. */ +export async function thresholdProbaPngBlob( + sourceBlob: Blob, + threshold: number, +): Promise { + const url = URL.createObjectURL(sourceBlob); + try { + const img = await new Promise((resolve, reject) => { + const el = new Image(); + el.onload = () => resolve(el); + el.onerror = () => reject(new Error('Failed to decode probability PNG')); + el.src = url; + }); + const width = img.naturalWidth; + const height = img.naturalHeight; + const canvas = document.createElement('canvas'); + canvas.width = width; + canvas.height = height; + const ctx = canvas.getContext('2d'); + if (!ctx) throw new Error('2D context unavailable'); + ctx.drawImage(img, 0, 0); + const imageData = ctx.getImageData(0, 0, width, height); + applyProbaThresholdRgba(imageData.data, threshold); + ctx.putImageData(imageData, 0, 0); + const out = await new Promise((resolve, reject) => { + canvas.toBlob( + (b) => (b ? resolve(b) : reject(new Error('toBlob failed'))), + 'image/png', + ); + }); + return out; + } finally { + URL.revokeObjectURL(url); + } +} + +/** Load a grayscale classId PNG into Uint8Array + dimensions. */ +export async function loadLabelPng( + url: string, +): Promise<{ data: Uint8Array; width: number; height: number }> { + const img = await new Promise((resolve, reject) => { + const el = new Image(); + el.onload = () => resolve(el); + el.onerror = () => reject(new Error('Failed to load prediction PNG')); + el.src = url; + }); + const width = img.naturalWidth; + const height = img.naturalHeight; + const canvas = document.createElement('canvas'); + canvas.width = width; + canvas.height = height; + const ctx = canvas.getContext('2d'); + if (!ctx) throw new Error('2D context unavailable'); + ctx.drawImage(img, 0, 0); + const rgba = ctx.getImageData(0, 0, width, height).data; + const data = new Uint8Array(width * height); + for (let i = 0, p = 0; i < data.length; i++, p += 4) data[i] = rgba[p]; + return { data, width, height }; +} + +/** + * Build a translucent RGBA canvas (one pixel per classId with known color). + * Unknown / 0 stays transparent. + */ +export function colorizeLabelMap( + labels: Uint8Array, + width: number, + height: number, + colorByClass: Map, + alpha = 110, +): HTMLCanvasElement { + const canvas = document.createElement('canvas'); + canvas.width = width; + canvas.height = height; + const ctx = canvas.getContext('2d'); + if (!ctx) return canvas; + const img = ctx.createImageData(width, height); + const out = img.data; + const rgbCache = new Map(); + for (const [cid, hex] of colorByClass) rgbCache.set(cid, parseHexColor(hex)); + + for (let i = 0; i < labels.length; i++) { + const cid = labels[i]; + if (!cid) continue; + const rgb = rgbCache.get(cid); + if (!rgb) continue; + const p = i * 4; + out[p] = rgb[0]; + out[p + 1] = rgb[1]; + out[p + 2] = rgb[2]; + out[p + 3] = alpha; + } + ctx.putImageData(img, 0, 0); + return canvas; +} + +/** + * Build conformal overlay: singleton = class color, multi = amber hatch, + * abstain = dark translucent. + * + * Optional filters hide prediction classes / multi / abstain for the Layers panel. + */ +export function colorizeConformalOverlay( + commit: Uint8Array, + status: Uint8Array, + width: number, + height: number, + colorByClass: Map, + opts?: { + classVisible?: (classId: number) => boolean; + showMulti?: boolean; + showAbstain?: boolean; + }, +): HTMLCanvasElement { + const canvas = document.createElement('canvas'); + canvas.width = width; + canvas.height = height; + const ctx = canvas.getContext('2d'); + if (!ctx) return canvas; + const img = ctx.createImageData(width, height); + const out = img.data; + const rgbCache = new Map(); + for (const [cid, hex] of colorByClass) rgbCache.set(cid, parseHexColor(hex)); + const classVisible = opts?.classVisible ?? (() => true); + const showMulti = opts?.showMulti !== false; + const showAbstain = opts?.showAbstain !== false; + + for (let i = 0; i < commit.length; i++) { + const st = status[i]; + const p = i * 4; + const x = i % width; + const y = (i / width) | 0; + if (st === 1) { + // singleton + const cid = commit[i]; + if (!classVisible(cid)) continue; + const rgb = rgbCache.get(cid); + if (!rgb) continue; + out[p] = rgb[0]; + out[p + 1] = rgb[1]; + out[p + 2] = rgb[2]; + out[p + 3] = 120; + } else if (st === 2) { + if (!showMulti) continue; + // multi — diagonal hatch in amber + const on = ((x + y) & 3) < 2; + out[p] = on ? 245 : 180; + out[p + 1] = on ? 158 : 120; + out[p + 2] = on ? 11 : 40; + out[p + 3] = on ? 140 : 70; + } else if (showAbstain) { + // abstain + out[p] = 30; + out[p + 1] = 30; + out[p + 2] = 40; + out[p + 3] = 90; + } + } + ctx.putImageData(img, 0, 0); + return canvas; +} + +/** + * Convert a classId label map into polygon shapes. + * Each pixel contributes to at most one class (labels are already argmax). + * Uses outer+hole rings so one class does not fill through another (or through + * preserved scribble gaps) — matching the pixel overlay. + * Pixels covered by *preserveShapes* are skipped so Commit keeps annotations. + */ +export function labelMapToPolygonShapes( + labels: Uint8Array, + width: number, + height: number, + classIds: number[], + { + minRegion = 64, + preserveShapes, + smooth = 1, + origin = 'predicted', + }: { + minRegion?: number; + preserveShapes?: Shape[]; + smooth?: number; + /** Stamped onto every returned shape — see `ShapeOrigin` in annotationStore. + * Defaults to 'predicted' since every current caller commits iPred output. */ + origin?: 'human' | 'predicted'; + } = {}, +): PolygonShape[] { + const preserve = + preserveShapes && preserveShapes.length > 0 + ? rasterizeUnion(preserveShapes, width, height, 1) + : null; + + const shapes: PolygonShape[] = []; + for (const classId of classIds) { + const mask = new Uint8Array(width * height); + let any = false; + for (let i = 0; i < labels.length; i++) { + if (preserve && preserve[i]) continue; + if (labels[i] === classId) { + mask[i] = 1; + any = true; + } + } + if (!any) continue; + const polys = maskToPolygonsWithHoles(mask, width, height, { + minRegion, + scale: 1, + smooth, + }); + for (const { points, holes } of polys) { + if (points.length < 6) continue; + shapes.push({ + id: uuidv4(), + classId, + kind: 'polygon', + points, + origin, + ...(holes.length ? { holes } : {}), + }); + } + } + return shapes; +} diff --git a/frontend/src/lib/polybool.ts b/frontend/src/lib/polybool.ts index 9346d55..ae21bb3 100644 --- a/frontend/src/lib/polybool.ts +++ b/frontend/src/lib/polybool.ts @@ -116,14 +116,75 @@ export function eraseStampToMultiPolygon(points: number[], radius: number, width .map((p) => [flatToRing(p.points), ...p.holes.map(flatToRing)]); } +/** + * Build a MultiPolygon from already-vectorized regions (flat outer ring + flat + * hole rings), e.g. the Threshold Brush's mask→polygon output, so it can be + * boolean-subtracted from shapes. + */ +export function regionsToMultiPolygon( + regions: Array<{ points: number[]; holes: number[][] }>, +): MultiPolygon { + return regions + .filter((p) => p.points.length >= 6) + .map((p) => [flatToRing(p.points), ...p.holes.map(flatToRing)]); +} + /** Boolean-union a set of shapes into one MultiPolygon (empty if none / on failure). */ export function unionShapesToMultiPolygon(shapes: Shape[], width: number, height: number): MultiPolygon { - const geoms = shapes.map((s) => shapeToMultiPolygon(s, width, height)).filter((g) => g.length > 0); - if (geoms.length === 0) return []; + return unionShapesChecked(shapes, width, height).mp; +} + +/** + * Union with an explicit success flag. + * + * The plain version returns `[]` both when there is nothing to union AND when the + * boolean op failed — and callers gate on `mp.length`, so a failure reads exactly + * like "nothing to clip against" and clipping is silently skipped. That is how a + * single awkward polygon (a speckled Threshold Brush region, say) can disable + * clip-to-other-classes for a whole slice, with no error anywhere. + * + * This version unions incrementally so one unusable shape only costs that shape, + * and reports `ok: false` when anything was dropped, letting callers fall back to + * the mask-based clip instead of quietly leaving the new annotation unclipped. + */ +export function unionShapesChecked( + shapes: Shape[], + width: number, + height: number, +): { mp: MultiPolygon; ok: boolean } { + const geoms: MultiPolygon[] = []; + let ok = true; + for (const s of shapes) { + try { + const g = shapeToMultiPolygon(s, width, height); + if (g.length > 0) geoms.push(g); + } catch { + ok = false; // this shape can't be expressed as geometry at all + } + } + if (geoms.length === 0) return { mp: [], ok }; + if (geoms.length === 1) return { mp: geoms[0], ok }; + + // Fast path: one variadic union. polygon-clipping sweeps all inputs together, + // which is dramatically cheaper than folding them in pairwise — the pairwise + // version re-sweeps a growing accumulator once per shape, so a slice with a few + // hundred speckled regions turns a single sweep into hundreds of them. try { - return geoms.length === 1 ? geoms[0] : polygonClipping.union(geoms[0], ...geoms.slice(1)); + return { mp: polygonClipping.union(geoms[0], ...geoms.slice(1)), ok }; } catch { - return []; + // Something in the set is unusable. Fold in one at a time so we lose only the + // offending shape(s) rather than the whole union, and flag it so callers can + // choose the mask path instead of clipping against an under-covering union. + let acc: MultiPolygon = []; + for (const g of geoms) { + if (acc.length === 0) { acc = g; continue; } + try { + acc = polygonClipping.union(acc, g); + } catch { + ok = false; // keep what we have; this one is dropped + } + } + return { mp: acc, ok }; } } diff --git a/frontend/src/lib/rasterize.test.ts b/frontend/src/lib/rasterize.test.ts index 381e014..a878e21 100644 --- a/frontend/src/lib/rasterize.test.ts +++ b/frontend/src/lib/rasterize.test.ts @@ -1,5 +1,5 @@ import { describe, it, expect } from 'vitest'; -import { rasterizeShapes, gridFor } from './rasterize'; +import { rasterizeShapes, gridFor, fullResGridFor, stampStroke } from './rasterize'; import { maskToPolygons } from './magicwand'; import type { Shape } from '@/stores/annotationStore'; @@ -69,3 +69,52 @@ describe('rasterizeShapes', () => { expect(Math.max(...ys)).toBeGreaterThan(74); }); }); + +describe('fullResGridFor upscale', () => { + it('is unchanged at 1x', () => { + expect(fullResGridFor(200, 100)).toEqual({ gw: 200, gh: 100, scale: 1 }); + }); + + it('doubles the grid and halves the scale at 2x', () => { + expect(fullResGridFor(200, 100, 2)).toEqual({ gw: 400, gh: 200, scale: 0.5 }); + }); + + it('keeps the downsample guard for very large images', () => { + // >4096 falls back to the 1600 cap (scale 4 natively), then 2x halves it. + const g = fullResGridFor(6400, 6400, 2); + expect(g.scale).toBe(2); + expect(g.gw).toBe(3200); + }); + + it('round-trips a rect through an upscaled grid', () => { + const rect: Shape = { id: 'r', classId: 0, kind: 'rectangle', x: 20, y: 20, w: 60, h: 60 }; + const { gw, gh, scale } = fullResGridFor(200, 200, 2); + const mask = rasterizeShapes([rect], gw, gh, scale); + // 60x60 image px at 2x = 120x120 grid cells. + expect(area(mask)).toBeGreaterThan(14000); + expect(area(mask)).toBeLessThan(14700); + const polys = maskToPolygons(mask, gw, gh, { minRegion: 4, scale }); + const xs = polys[0].filter((_, i) => i % 2 === 0); + expect(Math.min(...xs)).toBeGreaterThan(18); // still native coords, not doubled + expect(Math.max(...xs)).toBeLessThan(82); + }); +}); + +describe('stampStroke gate', () => { + it('restricts the stamp to gated cells (Threshold Brush band)', () => { + const gw = 20, gh = 20; + // Gate lets only the left half through. + const gate = new Uint8Array(gw * gh); + for (let y = 0; y < gh; y++) for (let x = 0; x < 10; x++) gate[y * gw + x] = 1; + + const ungated = new Uint8Array(gw * gh); + stampStroke(ungated, gw, gh, [5, 10, 15, 10], 3, 1, 1); + const gated = new Uint8Array(gw * gh); + stampStroke(gated, gw, gh, [5, 10, 15, 10], 3, 1, 1, gate); + + expect(area(gated)).toBeGreaterThan(0); + expect(area(gated)).toBeLessThan(area(ungated)); + // Nothing landed outside the gate. + for (let i = 0; i < gated.length; i++) if (gated[i]) expect(gate[i]).toBe(1); + }); +}); diff --git a/frontend/src/lib/rasterize.ts b/frontend/src/lib/rasterize.ts index 5f7c6ba..99ea545 100644 --- a/frontend/src/lib/rasterize.ts +++ b/frontend/src/lib/rasterize.ts @@ -27,19 +27,39 @@ export function gridFor(width: number, height: number, maxDim = 1600): MaskGrid * full native resolution (scale 1) for reasonably sized images, so re-vectorizing * an unchanged region is ~idempotent and existing nodes don't erode or shift a * little each time a stroke is added. Only very large images (>4096 px) downsample. + * + * `upscale` (1, 2, 4) raises the working resolution so sub-pixel geometry — e.g. a + * Threshold Brush region traced at 2× — survives a clip/merge round-trip instead of + * being re-snapped to the native pixel grid. `scale` becomes fractional (1/upscale); + * `rasterizeShapes` divides by it and `maskToPolygons*` multiplies by it, so nothing + * downstream needs to change. */ -export function fullResGridFor(width: number, height: number): MaskGrid { +export function fullResGridFor(width: number, height: number, upscale = 1): MaskGrid { + const u = Math.max(1, upscale); const emax = Math.max(width, height); - return gridFor(width, height, emax <= 4096 ? emax : 1600); + const base = gridFor(width, height, emax <= 4096 ? emax : 1600); + if (u === 1) return base; + const scale = base.scale / u; + return { + gw: Math.max(1, Math.floor(width / scale)), + gh: Math.max(1, Math.floor(height / scale)), + scale, + }; } /** * Rasterize *shapes* into a `gw × gh` binary mask (Uint8Array of 0/1). Paint * strokes and shape bodies set 1; brush erase strokes and vector `erased` * carve-outs set 0, applied per-shape so a later shape can repaint. + * + * Pass `out` to render into an existing buffer instead of allocating one. Callers + * that rasterize many shapes in a loop (overlap tests) reuse a single scratch + * array this way — at full resolution each allocation is multiple megabytes, and + * the garbage adds up fast. `out` must be `gw*gh` long and is NOT cleared: clear + * it yourself when reusing, or leave it to accumulate a union deliberately. */ -export function rasterizeShapes(shapes: Shape[], gw: number, gh: number, scale = 1): Uint8Array { - const mask = new Uint8Array(gw * gh); +export function rasterizeShapes(shapes: Shape[], gw: number, gh: number, scale = 1, out?: Uint8Array): Uint8Array { + const mask = out ?? new Uint8Array(gw * gh); const s = scale || 1; for (const shape of shapes) { if (shape.kind === 'polygon') { @@ -141,8 +161,12 @@ function fillEllipse(mask: Uint8Array, gw: number, gh: number, cx: number, cy: n } } -/** Stamp a round-capped thick polyline (image-coord points) with value *v*. */ -function stampStroke(mask: Uint8Array, gw: number, gh: number, imgPts: number[], r: number, scale: number, v: number): void { +/** Stamp a round-capped thick polyline (image-coord points) with value *v*. + * `r` is the radius in GRID cells (i.e. image radius / scale). + * + * `gate`, when given, restricts the stamp to cells where `gate[i]` is non-zero — + * this is what makes the Threshold Brush paint only inside its intensity band. */ +export function stampStroke(mask: Uint8Array, gw: number, gh: number, imgPts: number[], r: number, scale: number, v: number, gate?: Uint8Array): void { const rad = Math.max(0.5, r); const r2 = rad * rad; const pts: number[] = []; @@ -160,7 +184,10 @@ function stampStroke(mask: Uint8Array, gw: number, gh: number, imgPts: number[], t = t < 0 ? 0 : t > 1 ? 1 : t; const cxp = ax + t * dx, cyp = ay + t * dy; const ddx = x - cxp, ddy = y - cyp; - if (ddx * ddx + ddy * ddy <= r2) mask[y * gw + x] = v; + if (ddx * ddx + ddy * ddy > r2) continue; + const i = y * gw + x; + if (gate && !gate[i]) continue; + mask[i] = v; } } }; diff --git a/frontend/src/lib/runCompatibility.test.ts b/frontend/src/lib/runCompatibility.test.ts new file mode 100644 index 0000000..33eb572 --- /dev/null +++ b/frontend/src/lib/runCompatibility.test.ts @@ -0,0 +1,79 @@ +import { describe, expect, it } from 'vitest'; +import { canContinueFineTuning, fineTuningBlockedReason, isDenoiserRun, isSegmentationRun } from './runCompatibility'; +import type { AnnotationClass } from '@/stores/classStore'; + +const cls = (...labels: string[]): AnnotationClass[] => + labels.map((label, i) => ({ classId: i + 1, label, color: '#ff0000', isVisible: true })); + +const runCls = (...labels: string[]) => labels.map((label) => ({ label })); + +describe('canContinueFineTuning', () => { + it('accepts identical class lists', () => { + expect(canContinueFineTuning(runCls('background', 'leaf'), cls('background', 'leaf'))).toBe(true); + }); + + it('ignores case and surrounding whitespace, matching the backend', () => { + expect(canContinueFineTuning(runCls('Background', ' Leaf '), cls('background', 'leaf'))).toBe(true); + }); + + it('rejects a different class count', () => { + expect(canContinueFineTuning(runCls('blob'), cls('background', 'leaf'))).toBe(false); + }); + + it('rejects reordered classes', () => { + // The saved head's channel c means the run's class c, so order is load-bearing. + expect(canContinueFineTuning(runCls('background', 'leaf'), cls('leaf', 'background'))).toBe(false); + }); + + it('rejects a rename at the same count', () => { + // The dangerous case: nothing downstream would notice on its own. + expect(canContinueFineTuning(runCls('blob'), cls('leaf'))).toBe(false); + }); + + it('treats two empty lists as compatible', () => { + expect(canContinueFineTuning([], [])).toBe(true); + }); +}); + +describe('fineTuningBlockedReason', () => { + it('returns null when the run can be continued', () => { + expect(fineTuningBlockedReason(runCls('leaf'), cls('leaf'))).toBeNull(); + }); + + it('names both class lists on a count mismatch, so the user can act on it', () => { + const reason = fineTuningBlockedReason(runCls('blob'), cls('background', 'leaf')); + expect(reason).toContain('blob'); + expect(reason).toContain('background, leaf'); + }); + + it('names both class lists on a same-count mismatch', () => { + const reason = fineTuningBlockedReason(runCls('blob'), cls('leaf')); + expect(reason).toContain('blob'); + expect(reason).toContain('leaf'); + }); + + it('points the user at applying instead, which does work across taxonomies', () => { + const reason = fineTuningBlockedReason(runCls('blob'), cls('leaf')); + expect(reason).toMatch(/apply it instead/i); + }); +}); + +describe('isDenoiserRun / isSegmentationRun', () => { + it('classifies a dlsia_denoiser run as a denoiser run, not a segmentation run', () => { + expect(isDenoiserRun({ model_family: 'dlsia_denoiser' })).toBe(true); + expect(isSegmentationRun({ model_family: 'dlsia_denoiser' })).toBe(false); + }); + + it('classifies both segmentation families as segmentation runs, not denoiser runs', () => { + for (const family of ['dinov3_lora', 'dlsia_tunet']) { + expect(isSegmentationRun({ model_family: family })).toBe(true); + expect(isDenoiserRun({ model_family: family })).toBe(false); + } + }); + + it('are exact complements of each other', () => { + for (const family of ['dinov3_lora', 'dlsia_tunet', 'dlsia_denoiser', 'something_unknown']) { + expect(isSegmentationRun({ model_family: family })).toBe(!isDenoiserRun({ model_family: family })); + } + }); +}); diff --git a/frontend/src/lib/runCompatibility.ts b/frontend/src/lib/runCompatibility.ts new file mode 100644 index 0000000..8de1298 --- /dev/null +++ b/frontend/src/lib/runCompatibility.ts @@ -0,0 +1,77 @@ +/** + * runCompatibility — can a saved run be *continued* (fine-tuned further) on the + * currently-open sample's class list? + * + * Applying a run for inference works across taxonomies: predictions are remapped + * by label and unmatched run classes are appended (see importPredictions.ts). + * Continuing to TRAIN one does not. The saved segmentation head has exactly one + * output channel per class, in the run's own class order (the backend maps + * channel `c` to `classes[c].classId`), so the class list has to line up + * positionally or those weights mean something different. + * + * Mirrors `check_resume_compatible` in backend/train_jobs.py — that remains the + * authority; this exists so the UI can say so before the user commits to a save + * and a request, rather than surfacing it as an error afterwards. + */ +import type { AnnotationClass } from '@/stores/classStore'; + +export interface RunClassLike { + label: string; +} + +const normalize = (label: string) => label.trim().toLowerCase(); + +/** + * True when *runClasses* and *currentClasses* are the same labels in the same + * order (case- and whitespace-insensitive, matching the backend). + */ +export function canContinueFineTuning( + runClasses: RunClassLike[], + currentClasses: AnnotationClass[], +): boolean { + if (runClasses.length !== currentClasses.length) return false; + return runClasses.every((rc, i) => normalize(rc.label) === normalize(currentClasses[i].label)); +} + +/** + * Why a run can't be continued, phrased for display next to the button, or null + * when it can. Deliberately concrete about both lists — "incompatible" alone + * leaves the user with nothing to act on. + */ +export function fineTuningBlockedReason( + runClasses: RunClassLike[], + currentClasses: AnnotationClass[], +): string | null { + if (canContinueFineTuning(runClasses, currentClasses)) return null; + const runLabels = runClasses.map((c) => c.label).join(', ') || '—'; + const currentLabels = currentClasses.map((c) => c.label).join(', ') || '—'; + if (runClasses.length !== currentClasses.length) { + return `This run was trained on ${runClasses.length} class(es) (${runLabels}) but this image has ${currentClasses.length} (${currentLabels}). Continuing would not fit its saved weights — apply it instead, or train a new run.`; + } + return `This run's classes (${runLabels}) don't line up with this image's (${currentLabels}). Continuing would train against the wrong classes — apply it instead, or train a new run.`; +} + +/** + * Segmentation vs. denoiser runs, for pickers that only make sense for one or + * the other. `/api/train/runs` lists every run this server has ever produced + * regardless of task — segmentation (DINOv3+LoRA, dlsia TUNet) and, once the + * backend's denoiser-training work lands, self-supervised denoiser runs + * (Noise2Noise/Noise2Void) all come back in the same flat list. Fine-tune, + * apply-to-image, and inference pickers only understand segmentation runs + * (they key off a class list a denoiser run doesn't have); the Learned + * Denoiser panel's run picker is the mirror image. + * + * Matches `schemas.DlsiaDenoiserConfig.model_family`, the discriminator literal + * on the `ModelConfig` union. + */ +export interface RunFamilyLike { + model_family: string; +} + +export function isDenoiserRun(run: RunFamilyLike): boolean { + return run.model_family === 'dlsia_denoiser'; +} + +export function isSegmentationRun(run: RunFamilyLike): boolean { + return !isDenoiserRun(run); +} diff --git a/frontend/src/lib/sam/adjust.test.ts b/frontend/src/lib/sam/adjust.test.ts new file mode 100644 index 0000000..db09ba3 --- /dev/null +++ b/frontend/src/lib/sam/adjust.test.ts @@ -0,0 +1,133 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { renderAdjusted, renderPreprocessOnly } from './adjust'; + +/** + * jsdom has no real 2D canvas. Stub getContext('2d') with an in-memory pixel + * buffer per canvas, and make drawImage copy the source canvas's own stored + * buffer into the destination — so renderAdjusted's internal + * drawImage-then-getImageData round trip actually sees real pixel data + * instead of zeros, matching the pattern established in pixelClf.gaps.test.ts. + */ +const stores = new WeakMap(); + +function seedCanvas(width: number, height: number, fill: (i: number) => number): HTMLCanvasElement { + const canvas = document.createElement('canvas'); + canvas.width = width; + canvas.height = height; + const data = new Uint8ClampedArray(width * height * 4); + for (let i = 0; i < data.length; i++) data[i] = fill(i); + stores.set(canvas, data); + return canvas; +} + +function readPixels(canvas: HTMLCanvasElement): Uint8ClampedArray { + return stores.get(canvas)!; +} + +beforeEach(() => { + vi.spyOn(HTMLCanvasElement.prototype, 'getContext').mockImplementation(function ( + this: HTMLCanvasElement, + ): any { + const canvas = this; + return { + imageSmoothingEnabled: true, + imageSmoothingQuality: 'high', + drawImage: (source: HTMLCanvasElement, _sx: number, _sy: number, w: number, h: number) => { + const src = stores.get(source); + const dst = new Uint8ClampedArray(w * h * 4); + if (src) dst.set(src.subarray(0, dst.length)); + stores.set(canvas, dst); + }, + getImageData: (_x: number, _y: number, w: number, h: number) => ({ + data: stores.get(canvas) ?? new Uint8ClampedArray(w * h * 4), + width: w, + height: h, + }), + putImageData: (imageData: { data: Uint8ClampedArray }) => { + stores.set(canvas, imageData.data); + }, + }; + }); +}); + +afterEach(() => vi.restoreAllMocks()); + +describe('renderAdjusted', () => { + it('returns the raw canvas unchanged when everything is a no-op', () => { + const source = seedCanvas(2, 2, () => 100); + const out = renderAdjusted(source, 2, 2, 0, 0, 0, 255); + expect([...readPixels(out)]).toEqual([...readPixels(source)]); + }); + + it('brightens pixels when brightness > 0', () => { + const source = seedCanvas(1, 1, (i) => (i % 4 === 3 ? 255 : 100)); + const out = renderAdjusted(source, 1, 1, 0.2, 0, 0, 255); + const px = readPixels(out); + expect(px[0]).toBeGreaterThan(100); + expect(px[3]).toBe(255); // alpha untouched + }); + + it('darkens pixels when brightness < 0', () => { + const source = seedCanvas(1, 1, (i) => (i % 4 === 3 ? 255 : 100)); + const out = renderAdjusted(source, 1, 1, -0.2, 0, 0, 255); + expect(readPixels(out)[0]).toBeLessThan(100); + }); + + it('applies a levels remap clamping outside [lo,hi]', () => { + const source = seedCanvas(3, 1, (i) => { + const px = Math.floor(i / 4); + return i % 4 === 3 ? 255 : [10, 128, 250][px]; + }); + const out = renderAdjusted(source, 3, 1, 0, 0, 50, 200); + const px = readPixels(out); + expect(px[0]).toBe(0); // below lo -> 0 + expect(px[8]).toBe(255); // above hi -> 255 + expect(px[4]).toBeGreaterThan(0); // in range, remapped + expect(px[4]).toBeLessThan(255); + }); + + it('bakes a gaussian blur when preprocess.blur > 0', () => { + // A single bright pixel among dark neighbors should spread out after blur. + const source = seedCanvas(5, 5, (i) => { + const px = Math.floor(i / 4); + const isCenter = px === 12; // middle of 5x5 + return i % 4 === 3 ? 255 : isCenter ? 255 : 0; + }); + const out = renderAdjusted(source, 5, 5, 0, 0, 0, 255, { blur: 1 }); + const px = readPixels(out); + // A neighboring pixel (index 11, just left of center) should pick up some brightness. + expect(px[11 * 4]).toBeGreaterThan(0); + }); + + it('scales blur sigma by the working-resolution upscale factor', () => { + const source = seedCanvas(3, 3, (i) => (i % 4 === 3 ? 255 : (Math.floor(i / 4) === 4 ? 255 : 0))); + // Just confirm it runs without throwing at a non-1 upscale and returns a bigger canvas. + const out = renderAdjusted(source, 3, 3, 0, 0, 0, 255, { blur: 1 }, 2); + expect(out.width).toBe(6); + expect(out.height).toBe(6); + }); + + it('leaves the canvas untouched with all-zero preprocess flags and no blur', () => { + const source = seedCanvas(2, 2, () => 77); + const out = renderAdjusted(source, 2, 2, 0, 0, 0, 255, { blur: 0, clahe: false, sharpen: false }); + expect([...readPixels(out)]).toEqual([...readPixels(source)]); + }); +}); + +describe('renderPreprocessOnly', () => { + it('applies only the nonlinear preprocessors, ignoring brightness/contrast/levels', () => { + const source = seedCanvas(5, 5, (i) => { + const px = Math.floor(i / 4); + return i % 4 === 3 ? 255 : px === 12 ? 255 : 0; + }); + const out = renderPreprocessOnly(source, 5, 5, { blur: 1 }); + const px = readPixels(out); + expect(px[11 * 4]).toBeGreaterThan(0); // blur still applied + }); + + it('returns an unmodified canvas when no preprocessors are requested', () => { + const source = seedCanvas(2, 2, () => 50); + const out = renderPreprocessOnly(source, 2, 2, {}); + expect([...readPixels(out)]).toEqual([...readPixels(source)]); + }); +}); diff --git a/frontend/src/lib/sam/adjust.ts b/frontend/src/lib/sam/adjust.ts index 4067e2f..5d163c2 100644 --- a/frontend/src/lib/sam/adjust.ts +++ b/frontend/src/lib/sam/adjust.ts @@ -1,8 +1,12 @@ import { claheRgba } from '@/lib/clahe'; import { applySharpen } from '@/lib/sharpen'; +import { applyGaussianBlurRgba } from '@/lib/blur'; /** Nonlinear display preprocessors baked before the linear brightness/contrast/levels. */ export interface PreprocessOpts { + /** Gaussian blur sigma in NATIVE image pixels (0 = off). Rescaled internally when + * rendering at an upscaled working resolution, so σ means the same thing at any. */ + blur?: number; /** Adaptive (local) contrast — CLAHE. Replaces the old global stretch. */ clahe?: boolean; sharpen?: boolean; @@ -14,10 +18,17 @@ export interface PreprocessOpts { * the user actually sees. Windowing a low-contrast tomography slice before * encoding is one of the biggest levers on mask quality. * - * Order: nonlinear preprocessors first (CLAHE → Sharpen), then the linear chain - * mirroring the canvas filter exactly (Brighten → Contrast → Levels): Brighten - * adds `brightness*255`; Contrast scales around mid-grey by `((contrast+100)/100)^2`; - * Levels remaps `[lo,hi] → [0,255]`. Returns a canvas at native resolution. + * Order: nonlinear preprocessors first (Blur → CLAHE → Sharpen), then the linear + * chain mirroring the canvas filter exactly (Brighten → Contrast → Levels): + * Brighten adds `brightness*255`; Contrast scales around mid-grey by + * `((contrast+100)/100)^2`; Levels remaps `[lo,hi] → [0,255]`. Blur runs first so + * it denoises the raw slice; a following Sharpen can then deliberately re-crisp + * the edges the blur softened. + * + * `upscale` (1, 2, 4) resamples to `width*upscale × height*upscale` with smooth + * interpolation — the "working resolution" for the drawing tools. `width`/`height` + * stay in NATIVE image pixels and the blur sigma is rescaled to match, so callers + * and σ mean the same thing at any working resolution. */ export function renderAdjusted( image: CanvasImageSource, @@ -28,25 +39,33 @@ export function renderAdjusted( levelsLo = 0, levelsHi = 255, preprocess?: PreprocessOpts, + upscale = 1, ): HTMLCanvasElement { + const u = Math.max(1, upscale); + const w = Math.round(width * u); + const h = Math.round(height * u); const canvas = document.createElement('canvas'); - canvas.width = width; - canvas.height = height; + canvas.width = w; + canvas.height = h; const ctx = canvas.getContext('2d', { willReadFrequently: true })!; - ctx.drawImage(image, 0, 0, width, height); + ctx.imageSmoothingEnabled = true; + ctx.imageSmoothingQuality = 'high'; + ctx.drawImage(image, 0, 0, w, h); - const doPre = !!(preprocess && (preprocess.clahe || preprocess.sharpen)); + const blurSigma = (preprocess?.blur ?? 0) * u; // σ is given in native pixels + const doPre = !!(preprocess && (blurSigma > 0 || preprocess.clahe || preprocess.sharpen)); const bcNoop = brightness === 0 && contrast === 0; const levelsNoop = levelsLo <= 0 && levelsHi >= 255; if (!doPre && bcNoop && levelsNoop) return canvas; - const imageData = ctx.getImageData(0, 0, width, height); + const imageData = ctx.getImageData(0, 0, w, h); const d = imageData.data; // Uint8ClampedArray → assignments auto-clamp to [0,255] // Nonlinear preprocessors first (baked once), in a fixed order. if (doPre) { - if (preprocess!.clahe) claheRgba(d, width, height); - if (preprocess!.sharpen) applySharpen(d, width, height); + if (blurSigma > 0) applyGaussianBlurRgba(d, w, h, blurSigma); + if (preprocess!.clahe) claheRgba(d, w, h); + if (preprocess!.sharpen) applySharpen(d, w, h); } if (!(bcNoop && levelsNoop)) { @@ -73,16 +92,17 @@ export function renderAdjusted( } /** - * Bake ONLY the nonlinear preprocessors (CLAHE / Sharpen) into a canvas — used as - * the Konva display base, since brightness/contrast/levels/gamma/colormap stay on - * the GPU SVG filter applied over it. Returns the source untouched when no - * preprocessor is active. + * Bake ONLY the nonlinear preprocessors (Blur / CLAHE / Sharpen) into a canvas — + * used as the Konva display base, since brightness/contrast/levels/gamma/colormap + * stay on the GPU SVG filter applied over it. `upscale` also resamples to the + * working resolution (see `renderAdjusted`). */ export function renderPreprocessOnly( image: CanvasImageSource, width: number, height: number, preprocess: PreprocessOpts, + upscale = 1, ): HTMLCanvasElement { - return renderAdjusted(image, width, height, 0, 0, 0, 255, preprocess); + return renderAdjusted(image, width, height, 0, 0, 0, 255, preprocess, upscale); } diff --git a/frontend/src/lib/sam/samClient.test.ts b/frontend/src/lib/sam/samClient.test.ts new file mode 100644 index 0000000..a680cd9 --- /dev/null +++ b/frontend/src/lib/sam/samClient.test.ts @@ -0,0 +1,177 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; + +/** + * samClient exports only a module-level singleton (no class, no reset), so + * each test re-imports the module fresh via vi.resetModules() to get an + * un-contaminated instance — otherwise `init()`'s cached promise/worker would + * leak state across tests. + */ +class FakeWorker { + static instances: FakeWorker[] = []; + onmessage: ((e: MessageEvent) => void) | null = null; + onerror: ((e: any) => void) | null = null; + terminated = false; + posted: any[] = []; + + constructor(_url: unknown, _opts?: unknown) { + FakeWorker.instances.push(this); + } + + postMessage(msg: any) { + this.posted.push(msg); + } + + terminate() { + this.terminated = true; + } + + // Test helper: simulate the worker sending a message back. + emit(data: any) { + this.onmessage?.({ data } as MessageEvent); + } +} + +beforeEach(() => { + FakeWorker.instances = []; + vi.stubGlobal('Worker', FakeWorker as any); + vi.stubGlobal('createImageBitmap', vi.fn(async () => ({ close: () => {} }) as any)); + vi.resetModules(); +}); + +afterEach(() => { + vi.unstubAllGlobals(); +}); + +async function freshClient() { + const mod = await import('./samClient'); + return mod; +} + +describe('samClient.init', () => { + it('starts idle before init', async () => { + const { samClient } = await freshClient(); + expect(samClient.getStatus()).toBe('idle'); + }); + + it('spawns exactly one worker across repeated init() calls (idempotent)', async () => { + const { samClient } = await freshClient(); + const p1 = samClient.init(); + const p2 = samClient.init(); + expect(p1).toBe(p2); + expect(FakeWorker.instances).toHaveLength(1); + }); + + it('resolves once the worker reports status "ready", and notifies subscribers', async () => { + const { samClient } = await freshClient(); + const seen: string[] = []; + samClient.subscribe((s) => seen.push(s)); + const p = samClient.init(); + FakeWorker.instances[0].emit({ type: 'status', status: 'loading-model' }); + FakeWorker.instances[0].emit({ type: 'status', status: 'ready', backend: 'wasm' }); + await expect(p).resolves.toBeUndefined(); + expect(samClient.getStatus()).toBe('ready'); + expect(samClient.getBackend()).toBe('wasm'); + expect(seen).toContain('loading-model'); + expect(seen).toContain('ready'); + }); + + it('rejects when the worker reports "unsupported", and allows retrying init after', async () => { + const { samClient } = await freshClient(); + const p = samClient.init(); + FakeWorker.instances[0].emit({ type: 'status', status: 'unsupported', message: 'no webgpu/wasm' }); + await expect(p).rejects.toThrow('no webgpu/wasm'); + expect(samClient.getStatus()).toBe('unsupported'); + + // A later init() attempt spawns a fresh worker rather than reusing the dead one. + await Promise.resolve(); // let the internal .catch() teardown run + const p2 = samClient.init(); + expect(FakeWorker.instances.length).toBeGreaterThanOrEqual(2); + FakeWorker.instances[FakeWorker.instances.length - 1].emit({ type: 'status', status: 'ready' }); + await expect(p2).resolves.toBeUndefined(); + }); + + it('rejects and sets "unsupported" when the Worker constructor throws', async () => { + vi.stubGlobal('Worker', class { + constructor() { throw new Error('workers disabled'); } + } as any); + vi.resetModules(); + const { samClient } = await freshClient(); + await expect(samClient.init()).rejects.toThrow('workers disabled'); + expect(samClient.getStatus()).toBe('unsupported'); + }); + + it('rejects on a worker onerror event', async () => { + const { samClient } = await freshClient(); + const p = samClient.init(); + FakeWorker.instances[0].onerror?.({ message: 'boom' }); + await expect(p).rejects.toThrow('boom'); + expect(samClient.getStatus()).toBe('unsupported'); + }); +}); + +describe('samClient.encode/decode', () => { + it('encode() initializes, sets status to encoding then back to ready', async () => { + const { samClient } = await freshClient(); + const statuses: string[] = []; + samClient.subscribe((s) => statuses.push(s)); + + const encodePromise = samClient.encode({} as any); + FakeWorker.instances[0].emit({ type: 'status', status: 'ready' }); + // Let init() resolve and encode() proceed to post the encode message. + await Promise.resolve(); + await Promise.resolve(); + const encodeMsg = FakeWorker.instances[0].posted.find((m) => m.type === 'encode'); + expect(encodeMsg).toBeTruthy(); + FakeWorker.instances[0].emit({ type: 'encoded', id: encodeMsg.id }); + await encodePromise; + + expect(statuses).toContain('encoding'); + expect(samClient.getStatus()).toBe('ready'); + }); + + it('decode() resolves a SamMask from the worker reply', async () => { + const { samClient } = await freshClient(); + const p = samClient.init(); + FakeWorker.instances[0].emit({ type: 'status', status: 'ready' }); + await p; + + const decodePromise = samClient.decode( + [{ x: 1, y: 2, label: 1 }], null, 'auto', 0.5, + ); + const decodeMsg = FakeWorker.instances[0].posted.find((m) => m.type === 'decode'); + expect(decodeMsg).toBeTruthy(); + const buf = new Uint8Array([1, 0, 1, 0]).buffer; + FakeWorker.instances[0].emit({ + type: 'decoded', id: decodeMsg.id, mask: buf, width: 2, height: 2, score: 0.9, + }); + const result = await decodePromise; + expect(result.width).toBe(2); + expect(result.height).toBe(2); + expect(result.score).toBe(0.9); + expect([...result.mask]).toEqual([1, 0, 1, 0]); + }); + + it('rejects a pending request when the worker reports an error for its id', async () => { + const { samClient } = await freshClient(); + const p = samClient.init(); + FakeWorker.instances[0].emit({ type: 'status', status: 'ready' }); + await p; + + const decodePromise = samClient.decode([], null, 'auto', 0.5); + const decodeMsg = FakeWorker.instances[0].posted.find((m) => m.type === 'decode'); + FakeWorker.instances[0].emit({ type: 'error', id: decodeMsg.id, message: 'decode failed' }); + await expect(decodePromise).rejects.toThrow('decode failed'); + }); +}); + +describe('webgpuAvailable', () => { + it('reflects navigator.gpu presence', async () => { + const { webgpuAvailable } = await freshClient(); + const original = (navigator as any).gpu; + (navigator as any).gpu = {}; + expect(webgpuAvailable()).toBe(true); + delete (navigator as any).gpu; + expect(webgpuAvailable()).toBe(false); + if (original !== undefined) (navigator as any).gpu = original; + }); +}); diff --git a/frontend/src/lib/sam/samWorker.ts b/frontend/src/lib/sam/samWorker.ts index e422b46..37afb69 100644 --- a/frontend/src/lib/sam/samWorker.ts +++ b/frontend/src/lib/sam/samWorker.ts @@ -44,7 +44,7 @@ const MODEL_SOURCES: ModelSource[] = [ ]; env.allowLocalModels = true; -env.localModelPath = '/models/'; +env.localModelPath = `${import.meta.env.BASE_URL}models/`; if (env.backends?.onnx?.wasm) env.backends.onnx.wasm.numThreads = 1; interface PromptPoint { x: number; y: number; label: 0 | 1 } diff --git a/frontend/src/lib/sourceKey.test.ts b/frontend/src/lib/sourceKey.test.ts index 3ddf469..db7105c 100644 --- a/frontend/src/lib/sourceKey.test.ts +++ b/frontend/src/lib/sourceKey.test.ts @@ -1,9 +1,31 @@ import { describe, it, expect } from 'vitest'; -import { buildSourceKey, isAnnotatedPath } from './sourceKey'; +import { buildSourceKey, isAnnotatedPath, parseSourceKey } from './sourceKey'; const URI = 'http://127.0.0.1:8010'; const key = (path: string) => buildSourceKey('tiled', path, URI); +describe('parseSourceKey', () => { + it('round-trips a tiled key whose server URI has a port (colons in serverUri)', () => { + const sk = buildSourceKey('tiled', 'browse/rec20260221_135217_petiole22', URI); + expect(parseSourceKey(sk)).toEqual({ kind: 'tiled', source: 'browse/rec20260221_135217_petiole22', serverUri: URI }); + }); + + it('round-trips a tiled key with a nested path', () => { + const sk = buildSourceKey('tiled', 'browse/ds/img_00003', URI); + expect(parseSourceKey(sk)).toEqual({ kind: 'tiled', source: 'browse/ds/img_00003', serverUri: URI }); + }); + + it('round-trips a local key', () => { + const sk = buildSourceKey('local', 'some/rel/path.tif'); + expect(parseSourceKey(sk)).toEqual({ kind: 'local', source: 'some/rel/path.tif', serverUri: null }); + }); + + it('handles a tiled key with no server URI', () => { + const sk = buildSourceKey('tiled', 'browse/ds', null); + expect(parseSourceKey(sk)).toEqual({ kind: 'tiled', source: 'browse/ds', serverUri: null }); + }); +}); + describe('isAnnotatedPath', () => { it('matches the sample own key', () => { const keys = new Set([key('browse/ds')]); diff --git a/frontend/src/lib/sourceKey.ts b/frontend/src/lib/sourceKey.ts index 5b0072e..00d7ffd 100644 --- a/frontend/src/lib/sourceKey.ts +++ b/frontend/src/lib/sourceKey.ts @@ -18,6 +18,18 @@ export function buildSourceKey( return kind === 'tiled' ? `tiled:${serverUri ?? ''}:${path}` : `local:${path}`; } +/** Inverse of {@link buildSourceKey}: recovers {kind, source, serverUri} from a canonical sourceKey. */ +export function parseSourceKey(sk: string): { kind: 'tiled' | 'local'; source: string; serverUri: string | null } { + if (sk.startsWith('tiled:')) { + const rest = sk.slice('tiled:'.length); + // The serverUri itself contains colons (`http://host:port`), so the + // separator before the tiled path is the LAST colon, not the first. + const sep = rest.lastIndexOf(':'); + return { kind: 'tiled', serverUri: rest.slice(0, sep) || null, source: rest.slice(sep + 1) }; + } + return { kind: 'local', serverUri: null, source: sk.slice('local:'.length) }; +} + /** * True if *path* itself, or anything below it, has annotations. * diff --git a/frontend/src/lib/thresholdFit.test.ts b/frontend/src/lib/thresholdFit.test.ts new file mode 100644 index 0000000..1de14e6 --- /dev/null +++ b/frontend/src/lib/thresholdFit.test.ts @@ -0,0 +1,386 @@ +import { describe, it, expect } from 'vitest'; +import { + fitBand, + describeFit, + sampleHistograms, + backgroundRing, + fitWithBlurSweep, + combineSigma, + ringWidthFor, + MAX_RING_WIDTH, + type BandFit, +} from './thresholdFit'; +import { dilate } from './morphology'; +import { applyGaussianBlurGray } from './blur'; +import { displayAffineFor, baseToDisplay, displayBandToBase, displayToBase } from './displayTransform'; + +/** Histogram with `count` samples placed at each listed bin. */ +function hist(entries: Array<[number, number]>): Float64Array { + const h = new Float64Array(256); + for (const [bin, count] of entries) h[bin] += count; + return h; +} + +/** Spread `count` samples uniformly across [lo,hi]. */ +function band(lo: number, hi: number, count: number): Float64Array { + const h = new Float64Array(256); + const per = count / (hi - lo + 1); + for (let i = lo; i <= hi; i++) h[i] += per; + return h; +} + +/** Reference implementation: the same score, computed the slow obvious way. */ +function bruteForce(pos: ArrayLike, neg: ArrayLike): { lo: number; hi: number; dice: number } { + let best = { lo: 0, hi: 0, dice: -1 }; + let totalPos = 0; + for (let i = 0; i < 256; i++) totalPos += pos[i]; + for (let lo = 0; lo < 256; lo++) { + for (let hi = lo; hi < 256; hi++) { + let tp = 0; + let fp = 0; + for (let i = lo; i <= hi; i++) { tp += pos[i]; fp += neg[i]; } + if (tp === 0) continue; + const dice = (2 * tp) / (2 * tp + fp + (totalPos - tp)); + if (dice > best.dice) best = { lo, hi, dice }; + } + } + return best; +} + +describe('fitBand — cleanly separated material', () => { + it('recovers the planted band almost exactly', () => { + // Feature at 100–140, surroundings well away at 10–40. + const fit = fitBand(band(100, 140, 1000), band(10, 40, 4000)); + expect(fit.ok).toBe(true); + expect(fit.lo).toBeLessThanOrEqual(100); + expect(fit.hi).toBeGreaterThanOrEqual(140); + expect(fit.dice).toBeGreaterThan(0.99); + expect(fit.coverage).toBeGreaterThan(0.99); + expect(fit.leakage).toBeLessThan(0.01); + }); + + it('excludes a nearby but distinct background band', () => { + const fit = fitBand(band(120, 160, 1000), band(60, 110, 1000)); + // The gap is at 111–119; the fitted low edge must sit inside it. + expect(fit.lo).toBeGreaterThan(110); + expect(fit.dice).toBeGreaterThan(0.95); + }); +}); + +describe('fitBand — overlapping material', () => { + it('reports weak separation instead of a falsely confident band', () => { + // Same range for both: no band can separate them. + const fit = fitBand(band(80, 160, 1000), band(80, 160, 1000)); + expect(fit.ok).toBe(true); + // Dice looks deceptively respectable here (0.67 is the FLOOR against an + // equal-sized ring, not a decent result) — skill is what exposes it. + expect(fit.dice).toBeCloseTo(fit.baselineDice, 6); + expect(fit.skill).toBeCloseTo(0, 6); + expect(describeFit(fit).quality).toBe('weak'); + }); + + it('still returns the best achievable band on partial overlap', () => { + // Positives 100–160, negatives 130–200: the useful part is 100–129. + const fit = fitBand(band(100, 160, 1000), band(130, 200, 1000)); + expect(fit.dice).toBeGreaterThan(0.5); + expect(fit.dice).toBeLessThan(0.95); + expect(fit.lo).toBeLessThanOrEqual(100); + // It should stop before swallowing the whole negative range. + expect(fit.hi).toBeLessThan(200); + }); + + it('grades fits by skill, not by raw Dice', () => { + const mk = (skill: number): BandFit => ({ + lo: 0, hi: 1, dice: 0.9, baselineDice: 0.67, skill, coverage: 1, leakage: 0, ok: true, + }); + expect(describeFit(mk(0.95)).quality).toBe('good'); + expect(describeFit(mk(0.5)).quality).toBe('fair'); + // High Dice but no skill => still weak, which is the whole point. + expect(describeFit(mk(0.05)).quality).toBe('weak'); + }); + + it('reports skill relative to the ring size, so grading is size-independent', () => { + // Same inseparable overlap, but a 4x larger ring: Dice collapses to 0.33 + // while skill stays at 0 — the grade must not move. + const small = fitBand(band(80, 160, 1000), band(80, 160, 1000)); + const large = fitBand(band(80, 160, 1000), band(80, 160, 4000)); + expect(small.dice).toBeGreaterThan(large.dice + 0.2); + expect(small.skill).toBeCloseTo(0, 6); + expect(large.skill).toBeCloseTo(0, 6); + expect(describeFit(small).quality).toBe(describeFit(large).quality); + }); +}); + +describe('fitBand — agrees with a brute-force reference', () => { + function mulberry32(seed: number) { + return () => { + seed |= 0; seed = (seed + 0x6d2b79f5) | 0; + let t = Math.imul(seed ^ (seed >>> 15), 1 | seed); + t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t; + return ((t ^ (t >>> 14)) >>> 0) / 4294967296; + }; + } + + for (const seed of [1, 17, 99, 2024]) { + it(`matches on random histograms (seed ${seed})`, () => { + const rnd = mulberry32(seed); + const pos = new Float64Array(256); + const neg = new Float64Array(256); + for (let i = 0; i < 256; i++) { + pos[i] = Math.floor(rnd() * 30); + neg[i] = Math.floor(rnd() * 30); + } + const got = fitBand(pos, neg); + const want = bruteForce(pos, neg); + // The prefix-sum search must find the same optimum as the naive one. + expect(got.dice).toBeCloseTo(want.dice, 10); + }); + } +}); + +describe('fitBand — degenerate inputs', () => { + it('reports not-ok when nothing was sampled', () => { + const fit = fitBand(new Float64Array(256), band(10, 20, 100)); + expect(fit.ok).toBe(false); + expect(fit.dice).toBe(0); + }); + + it('handles an empty negative ring (selects the whole sample)', () => { + const fit = fitBand(band(50, 90, 500), new Float64Array(256)); + expect(fit.ok).toBe(true); + expect(fit.dice).toBeCloseTo(1, 6); + expect(fit.leakage).toBe(0); + }); + + it('handles a single-valued sample', () => { + const fit = fitBand(hist([[123, 400]]), hist([[7, 400]])); + expect(fit.lo).toBeLessThanOrEqual(123); + expect(fit.hi).toBeGreaterThanOrEqual(123); + expect(fit.dice).toBeCloseTo(1, 6); + }); + + it('never returns NaN when both inputs are empty', () => { + const fit = fitBand(new Float64Array(256), new Float64Array(256)); + expect(fit.ok).toBe(false); + expect(Number.isNaN(fit.dice)).toBe(false); + }); + + it('handles identical positive and negative sets without dividing by zero', () => { + const same = band(100, 100, 50); + const fit = fitBand(same, same); + expect(Number.isFinite(fit.dice)).toBe(true); + expect(fit.dice).toBeCloseTo(2 / 3, 6); // 2TP/(2TP+FP+FN) with FP=TP, FN=0 + }); +}); + +describe('sampleHistograms', () => { + it('bins each mask separately and ignores unmasked pixels', () => { + const field = [10, 10, 200, 200, 50, 50]; + const pos = [1, 1, 0, 0, 0, 0]; + const neg = [0, 0, 1, 1, 0, 0]; + const { posHist, negHist, posCount, negCount } = sampleHistograms(field, pos, neg); + expect(posCount).toBe(2); + expect(negCount).toBe(2); + expect(posHist[10]).toBe(2); + expect(negHist[200]).toBe(2); + expect(posHist[50]).toBe(0); // outside both masks + }); + + it('clamps out-of-range and skips non-finite values', () => { + const field = [-40, 999, NaN, 128]; + const pos = [1, 1, 1, 1]; + const neg = [0, 0, 0, 0]; + const { posHist, posCount } = sampleHistograms(field, pos, neg); + expect(posHist[0]).toBe(1); // -40 clamped up + expect(posHist[255]).toBe(1); // 999 clamped down + expect(posHist[128]).toBe(1); + // A non-finite pixel is not a usable sample, so it is skipped entirely + // rather than being counted toward the sample size. + expect(posCount).toBe(3); + }); + + it('treats a pixel in both masks as positive', () => { + const { posHist, negHist } = sampleHistograms([77], [1], [1]); + expect(posHist[77]).toBe(1); + expect(negHist[77]).toBe(0); + }); +}); + +describe('backgroundRing', () => { + it('surrounds the mask without overlapping it', () => { + const gw = 40, gh = 40; + const mask = new Uint8Array(gw * gh); + for (let y = 15; y < 25; y++) for (let x = 15; x < 25; x++) mask[y * gw + x] = 1; + + const ring = backgroundRing(mask, gw, gh, dilate, 3); + for (let i = 0; i < mask.length; i++) { + if (mask[i]) expect(ring[i]).toBe(0); // disjoint from the sample + } + let ringCount = 0; + for (let i = 0; i < ring.length; i++) if (ring[i]) ringCount++; + expect(ringCount).toBeGreaterThan(0); + }); + + it('scales its width with the region when not given one', () => { + const gw = 80, gh = 80; + const small = new Uint8Array(gw * gh); + for (let y = 40; y < 44; y++) for (let x = 40; x < 44; x++) small[y * gw + x] = 1; + const big = new Uint8Array(gw * gh); + for (let y = 20; y < 60; y++) for (let x = 20; x < 60; x++) big[y * gw + x] = 1; + + const area = (m: Uint8Array) => m.reduce((n, v) => n + (v ? 1 : 0), 0); + const ringSmall = area(backgroundRing(small, gw, gh, dilate)); + const ringBig = area(backgroundRing(big, gw, gh, dilate)); + expect(ringBig).toBeGreaterThan(ringSmall); + }); + + it('returns an empty ring for an empty mask', () => { + const empty = new Uint8Array(100); + const ring = backgroundRing(empty, 10, 10, dilate); + expect(ring.some((v) => v === 1)).toBe(false); + }); +}); + +describe('fitWithBlurSweep', () => { + /** A speckled target: alternating pixels stray into the background's range. */ + function speckled(w: number, h: number, targetVal: number, bgVal: number, noise: number) { + const field = new Float32Array(w * h); + const pos = new Uint8Array(w * h); + const neg = new Uint8Array(w * h); + let seed = 7; + const rnd = () => { + seed |= 0; seed = (seed + 0x6d2b79f5) | 0; + let t = Math.imul(seed ^ (seed >>> 15), 1 | seed); + t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t; + return ((t ^ (t >>> 14)) >>> 0) / 4294967296; + }; + for (let y = 0; y < h; y++) { + for (let x = 0; x < w; x++) { + const i = y * w + x; + const inside = x >= w / 4 && x < (3 * w) / 4 && y >= h / 4 && y < (3 * h) / 4; + field[i] = (inside ? targetVal : bgVal) + (rnd() - 0.5) * 2 * noise; + if (inside) pos[i] = 1; else neg[i] = 1; + } + } + return { field, pos, neg }; + } + + it('prefers blur when the target is speckled into the background range', () => { + // Heavy noise: per-pixel the two classes overlap badly, but their MEANS differ, + // so smoothing should separate them. + const { field, pos, neg } = speckled(64, 64, 150, 110, 55); + const swept = fitWithBlurSweep(field, 64, 64, pos, neg, applyGaussianBlurGray); + const unblurred = fitBand(...(() => { + const { posHist, negHist } = sampleHistograms(field, pos, neg); + return [posHist, negHist] as const; + })()); + expect(swept.extraSigma).toBeGreaterThan(0); + expect(swept.skill).toBeGreaterThan(unblurred.skill); + }); + + it('leaves a cleanly separated target unblurred', () => { + const { field, pos, neg } = speckled(64, 64, 200, 60, 2); + const swept = fitWithBlurSweep(field, 64, 64, pos, neg, applyGaussianBlurGray); + expect(swept.extraSigma).toBe(0); + expect(swept.skill).toBeGreaterThan(0.9); + }); + + it('does not mutate the caller’s field', () => { + const { field, pos, neg } = speckled(32, 32, 150, 110, 40); + const before = Float32Array.from(field); + fitWithBlurSweep(field, 32, 32, pos, neg, applyGaussianBlurGray); + expect(Array.from(field)).toEqual(Array.from(before)); + }); + + it('returns not-ok when there is nothing to sample', () => { + const field = new Float32Array(64); + const empty = new Uint8Array(64); + const swept = fitWithBlurSweep(field, 8, 8, empty, empty, applyGaussianBlurGray); + expect(swept.ok).toBe(false); + }); +}); + +describe('combineSigma', () => { + it('adds blurs in quadrature', () => { + expect(combineSigma(0, 2)).toBeCloseTo(2, 6); + expect(combineSigma(3, 4)).toBeCloseTo(5, 6); + }); +}); + +describe('base <-> display round trip (the sampler’s load-bearing assumption)', () => { + // The fit runs on the field (BASE space) but the band is stored in DISPLAYED + // space. If that conversion drifted, the applied band would quietly select + // different pixels than the ones that were fitted. + it('recovers the fitted base band through any display transform', () => { + for (const [b, c, lo, hi, gamma] of [ + [0, 0, 0, 255, 1], + [0.2, 30, 10, 240, 1], + [-0.1, -20, 0, 255, 1.8], + [0.05, 15, 40, 200, 0.6], + ] as const) { + const affine = displayAffineFor(b, c, lo, hi); + for (const [baseLo, baseHi] of [[60, 90], [100, 180], [20, 240]] as const) { + const dLo = baseToDisplay(baseLo, affine, gamma); + const dHi = baseToDisplay(baseHi, affine, gamma); + // Skip settings that clamp this band — the sampler detects and reports + // that case rather than applying an unrepresentable band. + if (Math.round(dHi) - Math.round(dLo) < 1) continue; + // Skip saturating endpoints too: displayBandToBase deliberately reopens + // an edge at 0/255 to ∓Infinity, which the sampler detects and refuses + // rather than applying a silently wider band. + if (Math.round(dLo) <= 0 || Math.round(dHi) >= 255) continue; + const back = displayBandToBase(Math.round(dLo), Math.round(dHi), affine, gamma); + // Integer storage costs up to half a displayed level. Its width in base + // units varies along the range once gamma is involved, so measure it at + // the endpoint rather than assuming the affine slope. + const width = (d: number) => + Math.abs((displayToBase(d + 0.5, affine, gamma) ?? 0) - (displayToBase(d - 0.5, affine, gamma) ?? 0)); + const tol = Math.max(0.51, width(Math.round(dLo)), width(Math.round(dHi))) + 0.01; + expect(Math.abs(back.lo - baseLo)).toBeLessThanOrEqual(tol); + expect(Math.abs(back.hi - baseHi)).toBeLessThanOrEqual(tol); + } + } + }); + + it('collapses a narrow band under heavy negative contrast', () => { + // Low contrast compresses the whole range, so neighbouring base values land + // on the same displayed value and no band can distinguish them. + const affine = displayAffineFor(0, -60, 0, 255); + const dLo = baseToDisplay(128, affine, 1); + const dHi = baseToDisplay(130, affine, 1); + expect(Math.round(dHi) - Math.round(dLo)).toBeLessThan(1); + }); + + it('reopens a saturated endpoint, which is why the sampler must check', () => { + // A bright fitted edge pushed past white comes back as +Infinity — applying + // that band would select everything above the fitted range. + const affine = displayAffineFor(0.5, 40, 0, 255); + const dHi = Math.round(baseToDisplay(200, affine, 1)); + expect(dHi).toBeGreaterThanOrEqual(255); + const back = displayBandToBase(100, dHi, affine, 1); + expect(Number.isFinite(back.hi)).toBe(false); + }); +}); + +describe('ringWidthFor', () => { + it('caps the ring so a large sample cannot explode the dilation cost', () => { + // sqrt(area)/2 for a 800x800 region is 400 — 400 full passes over the grid, + // which is what made a big lasso hang. + const big = new Uint8Array(1000 * 1000).fill(1); + expect(ringWidthFor(big)).toBe(MAX_RING_WIDTH); + }); + + it('still scales down for small samples', () => { + const small = new Uint8Array(100 * 100); + for (let y = 40; y < 60; y++) for (let x = 40; x < 60; x++) small[y * 100 + x] = 1; + const w = ringWidthFor(small); + expect(w).toBeGreaterThanOrEqual(4); + expect(w).toBeLessThan(MAX_RING_WIDTH); + }); + + it('has a floor so a tiny sample still gets a usable ring', () => { + const tiny = new Uint8Array(100); + tiny[55] = 1; + expect(ringWidthFor(tiny)).toBe(4); + }); +}); diff --git a/frontend/src/lib/thresholdFit.ts b/frontend/src/lib/thresholdFit.ts new file mode 100644 index 0000000..627e90e --- /dev/null +++ b/frontend/src/lib/thresholdFit.ts @@ -0,0 +1,276 @@ +/** + * thresholdFit — derive a Threshold Brush intensity band from a lassoed example. + * + * The user lassos one representative feature; everything inside is a positive + * example and a ring just outside it supplies negatives. The question "which band + * best fills that region and not its neighbours?" is then a supervised threshold + * fit with an *exact* answer, not something needing a search heuristic: + * + * - histogram both sets into 256 bins, + * - prefix-sum them, so any band's counts are two subtractions, + * - evaluate every (lo, hi) pair and keep the best. + * + * That is 256² ≈ 33k ordered pairs of integer work — well under a millisecond — + * so the returned band is the true optimum over all bands, with no tuning + * constants and no iteration to converge. + * + * Deliberately pure and canvas-free: the correctness of this file is what makes + * the feature trustworthy, and it can be tested directly. + */ + +/** How well a band matched the example, and what it selected. */ +export interface BandFit { + /** Inclusive bin bounds of the best band, in the field's own 0–255 space. */ + lo: number; + hi: number; + /** Dice overlap with the positive sample, 0–1. */ + dice: number; + /** + * Dice of the trivial "select everything" band — the score achievable with no + * separation at all. It depends on how big the ring is relative to the sample + * (equal sizes give 0.67; a 4x larger ring gives 0.33), which is exactly why + * raw Dice cannot be graded on its own. + */ + baselineDice: number; + /** + * How far the fit closed the gap between that baseline and a perfect match: + * `(dice - baseline) / (1 - baseline)`, clamped to 0–1. This is the number to + * judge quality by — 0 means the band did no better than selecting everything, + * however flattering its Dice looks. + */ + skill: number; + /** Fraction of the positive sample the band captures (recall). */ + coverage: number; + /** Fraction of selected pixels that came from the negative ring. */ + leakage: number; + /** False when there was nothing to fit (empty sample). */ + ok: boolean; +} + +const BINS = 256; + +/** An empty/degenerate result — a band that selects nothing. */ +const NO_FIT: BandFit = { + lo: 0, hi: 0, dice: 0, baselineDice: 0, skill: 0, coverage: 0, leakage: 0, ok: false, +}; + +/** Running totals of `hist`, where `cum[i]` counts bins 0..i inclusive. */ +function prefixSum(hist: ArrayLike): Float64Array { + const cum = new Float64Array(BINS); + let running = 0; + for (let i = 0; i < BINS; i++) { + running += hist[i] ?? 0; + cum[i] = running; + } + return cum; +} + +/** Count within `[lo,hi]` inclusive, from a prefix sum. */ +function rangeCount(cum: Float64Array, lo: number, hi: number): number { + return cum[hi] - (lo > 0 ? cum[lo - 1] : 0); +} + +/** + * Best intensity band separating `posHist` (the lassoed feature) from `negHist` + * (its surroundings), maximising Dice overlap with the positives. + * + * Dice — 2TP / (2TP + FP + FN) — is used rather than raw accuracy because the + * negative ring usually outnumbers the positives, and accuracy would then be + * maximised by selecting almost nothing. + * + * @param posHist 256-bin histogram of values inside the lasso. + * @param negHist 256-bin histogram of values in the surrounding ring. + */ +export function fitBand(posHist: ArrayLike, negHist: ArrayLike): BandFit { + const cumPos = prefixSum(posHist); + const cumNeg = prefixSum(negHist); + const totalPos = cumPos[BINS - 1]; + const totalNeg = cumNeg[BINS - 1]; + if (totalPos <= 0) return { ...NO_FIT }; + + // The score a band gets for free by selecting the entire range. Quality has to + // be measured against this, not against zero. + const baselineDice = (2 * totalPos) / (2 * totalPos + totalNeg); + + let best: BandFit = { ...NO_FIT, ok: true }; + let bestDice = -1; + + for (let lo = 0; lo < BINS; lo++) { + // Widening `hi` only ever adds counts, so both terms grow monotonically — + // but Dice itself is not monotonic, so every hi must still be evaluated. + for (let hi = lo; hi < BINS; hi++) { + const tp = rangeCount(cumPos, lo, hi); + if (tp === 0) continue; // selects none of the sample — cannot be best + const fp = rangeCount(cumNeg, lo, hi); + const fn = totalPos - tp; + const dice = (2 * tp) / (2 * tp + fp + fn); + if (dice > bestDice) { + bestDice = dice; + best = { + lo, + hi, + dice, + baselineDice, + skill: baselineDice < 1 ? Math.max(0, (dice - baselineDice) / (1 - baselineDice)) : 0, + coverage: tp / totalPos, + leakage: tp + fp > 0 ? fp / (tp + fp) : 0, + ok: true, + }; + } + } + } + + return bestDice < 0 ? { ...NO_FIT } : best; +} + +/** + * Plain-language reading of a fit, so a weak result is visible as weak rather + * than being silently applied and discovered later while painting. + */ +export function describeFit(fit: BandFit): { label: string; quality: 'good' | 'fair' | 'weak' } { + if (!fit.ok) return { label: 'Nothing sampled', quality: 'weak' }; + // Graded on `skill`, not `dice`: a feature that cannot be separated at all still + // scores 0.67 Dice against an equal-sized ring, and calling that "fair" would be + // exactly the false confidence this readout exists to prevent. + if (fit.skill >= 0.7) { + return { label: 'Strong — this band matches your sample closely', quality: 'good' }; + } + if (fit.skill >= 0.4) { + return { label: 'Fair — expect some over- or under-fill', quality: 'fair' }; + } + return { + label: 'Weak — intensity alone barely separates this from its surroundings', + quality: 'weak', + }; +} + +/** + * Histogram a field over two masks at once. + * + * `field` values are the 0–255 greys the Threshold Brush gates on; `pos` and + * `neg` are same-length 0/1 masks over the same grid. + */ +export function sampleHistograms( + field: ArrayLike, + pos: ArrayLike, + neg: ArrayLike, +): { posHist: Float64Array; negHist: Float64Array; posCount: number; negCount: number } { + const posHist = new Float64Array(BINS); + const negHist = new Float64Array(BINS); + let posCount = 0; + let negCount = 0; + const n = Math.min(field.length, pos.length, neg.length); + for (let i = 0; i < n; i++) { + const inPos = pos[i]; + const inNeg = neg[i]; + if (!inPos && !inNeg) continue; + let v = Math.round(field[i]); + if (!Number.isFinite(v)) continue; + v = v < 0 ? 0 : v > 255 ? 255 : v; + if (inPos) { posHist[v]++; posCount++; } + else { negHist[v]++; negCount++; } + } + return { posHist, negHist, posCount, negCount }; +} + +/** Sigmas the sampler tries, in addition to whatever blur is already applied. */ +export const BLUR_CANDIDATES: readonly number[] = [0, 0.5, 1, 1.5, 2, 3]; + +/** Result of a sweep: the winning band and the extra blur it needed. */ +export interface SweepResult extends BandFit { + /** Additional sigma applied on top of the current setting (0 = none). */ + extraSigma: number; +} + +/** + * Fit a band across several candidate blur levels and keep the best. + * + * Blur is swept because — unlike the display sliders — it genuinely changes what + * is separable: it is a spatial operation, so it can pull a speckled feature's + * values together and away from its surroundings. Sigmas are *additional* blur on + * top of whatever the user has already set, since the field handed in is the one + * the brush currently uses. + * + * @param field Crop of the threshold field (grays), row-major `w × h`. + * @param pos Positive mask over the same crop. + * @param neg Negative (ring) mask over the same crop. + * @param blurFn Injected gray blur (`blur.applyGaussianBlurGray`) so this module + * stays dependency-free and directly testable. + */ +export function fitWithBlurSweep( + field: Float32Array, + w: number, + h: number, + pos: Uint8Array, + neg: Uint8Array, + blurFn: (data: Float32Array, w: number, h: number, sigma: number) => void, + sigmas: readonly number[] = BLUR_CANDIDATES, +): SweepResult { + let best: SweepResult = { ...fitBand(new Float64Array(BINS), new Float64Array(BINS)), extraSigma: 0 }; + let bestSkill = -1; + + for (const sigma of sigmas) { + const candidate = sigma > 0 ? Float32Array.from(field) : field; + if (sigma > 0) blurFn(candidate, w, h, sigma); + const { posHist, negHist } = sampleHistograms(candidate, pos, neg); + const fit = fitBand(posHist, negHist); + if (!fit.ok) continue; + // Compare on skill, not Dice: Dice shifts with class balance, and blurring + // does not change the balance — but skill is the number the UI grades on, so + // optimising anything else could pick a sigma the readout then calls worse. + if (fit.skill > bestSkill) { + bestSkill = fit.skill; + best = { ...fit, extraSigma: sigma }; + } + } + return best; +} + +/** Combine two Gaussian blurs: sigmas add in quadrature. */ +export function combineSigma(current: number, extra: number): number { + return Math.sqrt(current * current + extra * extra); +} + +/** Widest background ring, in grid cells. See `ringWidthFor`. */ +export const MAX_RING_WIDTH = 24; + +/** + * How wide a background ring to grow around a sample. + * + * Scales with the region so a small lasso gets a proportionate neighbourhood, + * but is **capped**: the ring is meant to be the immediate surroundings, and + * dilation costs one pass per cell of width. Uncapped, lassoing an 800x800 + * region asked for a 400-cell ring — 400 passes over the grid, which is why a + * big sample used to hang. A capped ring is both faster and more correct: what + * a feature touches is a local question. + */ +export function ringWidthFor(mask: ArrayLike): number { + let area = 0; + for (let i = 0; i < mask.length; i++) if (mask[i]) area++; + return Math.min(MAX_RING_WIDTH, Math.max(4, Math.round(Math.sqrt(area) / 2))); +} + +/** + * Ring of background around `mask`: `dilate(mask, iters) \ mask`. + * + * Fitting against the immediate neighbourhood rather than the whole slice is + * what makes the result mean "fill this and not what it touches" — using the + * entire image as negatives would penalise the very lookalikes elsewhere that + * the user wants highlighted. + * + * @param dilateFn Injected (`morphology.dilate`) to keep this module canvas- and + * dependency-free for testing. + */ +export function backgroundRing( + mask: Uint8Array, + gw: number, + gh: number, + dilateFn: (m: Uint8Array, gw: number, gh: number, iters: number) => Uint8Array, + iters?: number, +): Uint8Array { + const width = iters ?? ringWidthFor(mask); + const grown = dilateFn(mask, gw, gh, width); + const ring = new Uint8Array(mask.length); + for (let i = 0; i < ring.length; i++) ring[i] = grown[i] && !mask[i] ? 1 : 0; + return ring; +} diff --git a/frontend/src/lib/thresholdRegularize.test.ts b/frontend/src/lib/thresholdRegularize.test.ts new file mode 100644 index 0000000..4e70b84 --- /dev/null +++ b/frontend/src/lib/thresholdRegularize.test.ts @@ -0,0 +1,81 @@ +/** + * A per-pixel intensity gate speckles, and every speck becomes its own polygon — + * which is what made a single Threshold Brush stroke commit hundreds of shapes + * and dominate clip cost. These pin the opening + small-component pass that the + * commit path applies before vectorizing. + */ +import { describe, it, expect } from 'vitest'; +import { dilate, erode, removeSmallComponents } from './morphology'; +import { maskToPolygonsWithHoles } from './magicwand'; + +const W = 256; +const H = 256; + +function mulberry32(seed: number) { + return () => { + seed |= 0; seed = (seed + 0x6d2b79f5) | 0; + let t = Math.imul(seed ^ (seed >>> 15), 1 | seed); + t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t; + return ((t ^ (t >>> 14)) >>> 0) / 4294967296; + }; +} + +/** The commit path's regularization, in one place. */ +function regularize(mask: Uint8Array, minRegion = 12): Uint8Array { + return removeSmallComponents(dilate(erode(mask, W, H, 1), W, H, 1), W, H, minRegion); +} + +const area = (m: Uint8Array) => m.reduce((n, v) => n + (v ? 1 : 0), 0); +const polys = (m: Uint8Array) => + maskToPolygonsWithHoles(m, W, H, { minRegion: 4, scale: 1 }).filter((p) => p.points.length >= 6); + +/** A solid square feature plus scattered thresholding noise. */ +function speckledFeature(noise: number) { + const mask = new Uint8Array(W * H); + const rnd = mulberry32(42); + for (let y = 80; y < 176; y++) for (let x = 80; x < 176; x++) mask[y * W + x] = 1; + for (let i = 0; i < mask.length; i++) if (rnd() > 1 - noise) mask[i] = 1; + return mask; +} + +describe('threshold stroke regularization', () => { + it('collapses speckle into the real feature', () => { + const raw = speckledFeature(0.25); + const cleaned = regularize(raw); + // 96x96 = 9216 true pixels; the raw gate is far larger because of speckle. + expect(area(raw)).toBeGreaterThan(9216 * 1.5); + expect(area(cleaned)).toBeGreaterThan(9216 * 0.95); + expect(area(cleaned)).toBeLessThan(9216 * 1.05); + }); + + it('turns hundreds of speck polygons into one', () => { + const raw = speckledFeature(0.25); + expect(polys(raw).length).toBeGreaterThan(50); + expect(polys(regularize(raw)).length).toBe(1); + }); + + it('leaves a clean region essentially unchanged', () => { + const clean = new Uint8Array(W * H); + for (let y = 60; y < 200; y++) for (let x = 60; x < 200; x++) clean[y * W + x] = 1; + const out = regularize(clean); + expect(area(out)).toBeGreaterThan(area(clean) * 0.97); + expect(polys(out).length).toBe(1); + }); + + it('erases a hairline stroke entirely — which is why the commit keeps a fallback', () => { + // A 1px line cannot survive an erosion. The commit path detects this and + // falls back to the raw mask rather than discarding the user's stroke. + const hairline = new Uint8Array(W * H); + for (let x = 40; x < 200; x++) hairline[128 * W + x] = 1; + expect(area(regularize(hairline))).toBe(0); + expect(area(hairline)).toBeGreaterThan(0); + }); + + it('preserves a genuine hole rather than filling it', () => { + const donut = new Uint8Array(W * H); + for (let y = 60; y < 200; y++) for (let x = 60; x < 200; x++) donut[y * W + x] = 1; + for (let y = 110; y < 150; y++) for (let x = 110; x < 150; x++) donut[y * W + x] = 0; + const out = regularize(donut); + expect(polys(out)[0].holes.length).toBe(1); + }); +}); diff --git a/frontend/src/lib/trainConstraints.ts b/frontend/src/lib/trainConstraints.ts new file mode 100644 index 0000000..223cb6d --- /dev/null +++ b/frontend/src/lib/trainConstraints.ts @@ -0,0 +1,50 @@ +/** + * trainConstraints — the bounds the backend enforces on training hyperparameters, + * mirrored client-side so the Train tab can show the valid range and reject a bad + * value before it becomes an opaque 422. + * + * Keep in sync with `TunetHyperParams` in backend/schemas.py — that remains the + * authority; this only exists to fail earlier and more legibly. + * + * dlsia TUNet only — DINOv3 LoRA is deferred (see Phase 5.5 in the integration + * plan), so there is only one family to constrain here. + */ + +export type TrainFamily = 'dlsia_tunet'; + +export interface SizeConstraint { + min: number; + max: number; + /** The side must be divisible by this. */ + multipleOf: number; + /** Why the divisor exists, for the error message. */ + reason: string; +} + +export const IMAGE_SIZE_CONSTRAINTS: Record = { + // TUNet halves the size once per depth level; 64 covers the deepest allowed net. + dlsia_tunet: { min: 64, max: 2048, multipleOf: 64, reason: "TUNet's downsampling steps" }, +}; + +/** + * Validate the image/patch size for `family`, returning a message to show the + * user, or null when the value is acceptable. + * + * `label` names the field as the UI currently labels it ("Patch size" when + * tiling, "Image size" otherwise) so the message matches what's on screen. + */ +export function validateImageSize(family: TrainFamily, value: number, label: string): string | null { + const c = IMAGE_SIZE_CONSTRAINTS[family]; + if (!Number.isFinite(value) || !Number.isInteger(value)) { + return `${label} must be a whole number of pixels.`; + } + if (value < c.min || value > c.max) { + return `${label} must be between ${c.min} and ${c.max} px for this model (you entered ${value}).`; + } + if (value % c.multipleOf !== 0) { + const lower = Math.floor(value / c.multipleOf) * c.multipleOf; + const nearest = Math.max(c.min, lower < c.min ? c.min : lower); + return `${label} must be a multiple of ${c.multipleOf} (${c.reason}) — try ${nearest}.`; + } + return null; +} diff --git a/frontend/src/lib/trainDenoiseOption.test.ts b/frontend/src/lib/trainDenoiseOption.test.ts new file mode 100644 index 0000000..da0eea2 --- /dev/null +++ b/frontend/src/lib/trainDenoiseOption.test.ts @@ -0,0 +1,105 @@ +import { describe, expect, it } from 'vitest'; +import { + denoiseMethodLabel, + trainDenoiseBlockedReason, + trainDenoisePayload, + trainDenoiseSummary, +} from './trainDenoiseOption'; + +/** Shape of `capability.denoise.methods` entries this module cares about. */ +const METHODS = [ + { method: 'none', label: 'None' }, + { method: 'gaussian', label: 'Gaussian' }, + { method: 'tv', label: 'Total variation' }, + { method: 'nlm', label: 'Non-local means' }, +]; + +describe('trainDenoiseBlockedReason', () => { + it('blocks "none" — there is no filter to apply', () => { + expect(trainDenoiseBlockedReason('none')).toMatch(/pick a denoise filter/i); + }); + + it('blocks "model" — a learned denoiser is not a preprocessor', () => { + expect(trainDenoiseBlockedReason('model')).toMatch(/learned denoiser/i); + }); + + it('allows any classical filter', () => { + expect(trainDenoiseBlockedReason('tv')).toBeNull(); + expect(trainDenoiseBlockedReason('gaussian')).toBeNull(); + expect(trainDenoiseBlockedReason('median3d')).toBeNull(); + }); +}); + +describe('trainDenoisePayload', () => { + it('adds no key at all when the option is off', () => { + // Not `{denoise: null}`: an un-denoised request must stay byte-identical + // to what the app sent before this option existed. + expect(trainDenoisePayload(false, { method: 'tv', strength: 0.6 })).toEqual({}); + expect(Object.keys(trainDenoisePayload(false, { method: 'tv', strength: 0.6 }))).toEqual([]); + }); + + it('adds no key when off even with denoising set to none', () => { + expect(trainDenoisePayload(false, { method: 'none', strength: 0.5 })).toEqual({}); + }); + + it('sends {method, strength} when on with a classical filter', () => { + expect(trainDenoisePayload(true, { method: 'tv', strength: 0.6 })) + .toEqual({ denoise: { method: 'tv', strength: 0.6 } }); + }); + + it('sends nothing when on but the method is "none"', () => { + expect(trainDenoisePayload(true, { method: 'none', strength: 0.5 })).toEqual({}); + }); + + it('sends nothing when on but the method is "model"', () => { + // The UI disables the checkbox in this case; the payload builder refuses + // independently so a stale checked box can't reach the server. + expect(trainDenoisePayload(true, { method: 'model', strength: 0.5 })).toEqual({}); + }); + + it('strips preview-only fields the strict backend schema would reject', () => { + const payload = trainDenoisePayload(true, { + method: 'nlm', strength: 0.4, crop: 768, runId: 'run-123', + } as { method: string; strength: number }); + expect(payload).toEqual({ denoise: { method: 'nlm', strength: 0.4 } }); + expect(Object.keys((payload as { denoise: object }).denoise).sort()).toEqual(['method', 'strength']); + }); + + it('passes the strength through unrounded at the extremes', () => { + expect(trainDenoisePayload(true, { method: 'tv', strength: 0 })) + .toEqual({ denoise: { method: 'tv', strength: 0 } }); + expect(trainDenoisePayload(true, { method: 'tv', strength: 1 })) + .toEqual({ denoise: { method: 'tv', strength: 1 } }); + }); +}); + +describe('trainDenoiseSummary', () => { + it('names the filter and strength as a percentage', () => { + expect(trainDenoiseSummary({ method: 'tv', strength: 0.6 }, METHODS)).toBe('Total variation, 60%'); + }); + + it('rounds the percentage to a whole number', () => { + expect(trainDenoiseSummary({ method: 'gaussian', strength: 0.333 }, METHODS)).toBe('Gaussian, 33%'); + }); + + it('returns null for a method that cannot be trained on', () => { + expect(trainDenoiseSummary({ method: 'none', strength: 0.5 }, METHODS)).toBeNull(); + expect(trainDenoiseSummary({ method: 'model', strength: 0.5 }, METHODS)).toBeNull(); + }); + + it('falls back to the raw method id when the server did not describe it', () => { + expect(trainDenoiseSummary({ method: 'wavelet', strength: 0.5 }, METHODS)).toBe('wavelet, 50%'); + expect(trainDenoiseSummary({ method: 'tv', strength: 0.5 }, [])).toBe('tv, 50%'); + }); +}); + +describe('denoiseMethodLabel', () => { + it('maps a known method id to its label', () => { + expect(denoiseMethodLabel('nlm', METHODS)).toBe('Non-local means'); + }); + + it('falls back to the id itself rather than rendering blank', () => { + expect(denoiseMethodLabel('median3d', METHODS)).toBe('median3d'); + expect(denoiseMethodLabel('tv', [])).toBe('tv'); + }); +}); diff --git a/frontend/src/lib/trainDenoiseOption.ts b/frontend/src/lib/trainDenoiseOption.ts new file mode 100644 index 0000000..35b6f25 --- /dev/null +++ b/frontend/src/lib/trainDenoiseOption.ts @@ -0,0 +1,100 @@ +/** + * trainDenoiseOption — decides what (if anything) the "Train on denoised + * input" checkbox contributes to a `/api/train/start` request. + * + * Both places that start a segmentation job (TrainPage's "Start training" and + * ApplyModelPanel's "Fine-tune & apply") offer this checkbox and must agree + * exactly on the answer, so the decision lives here as a pure function rather + * than inline in either component — same reasoning as `denoiserTrainScope`. + * + * The setting itself comes from the Annotate tab's shared denoise store, so the + * user trains on the very filter they were looking at. Two properties matter: + * + * - OFF must be bit-identical to the behaviour that predates this option: + * no `denoise` key in the request body at all. `TrainRequest.denoise` + * defaults to None, so an absent key and an explicit null mean the same + * thing server-side, but omitting it keeps un-denoised requests literally + * unchanged. + * - ON must send ONLY `{method, strength}`. `DenoiseOpts` also carries `crop` + * (a preview-only centre crop) and `runId`; `DenoiseTrainOpts` is a strict + * model, so forwarding the whole object would be rejected outright — and + * `crop` would be nonsense for training even if it weren't. + * + * `'none'` and `'model'` are both rejected: there is nothing to apply for the + * former, and the backend explicitly does not support a learned denoiser as a + * preprocessor for another model (it would need its own run and a GPU pass per + * slice). Guarding here as well as in the UI means a stale checked box can + * never smuggle an invalid method into a request. + */ + +/** The `denoise` field of a `/api/train/start` body (schemas.DenoiseTrainOpts). */ +export interface TrainDenoiseOpts { + method: string; + strength: number; +} + +/** Just the parts of `hooks/useImageSlice`'s `DenoiseOpts` that matter here. */ +interface DenoiseSettings { + method: string; + strength: number; +} + +/** Method + human label, as served by `capability.denoise.methods`. */ +interface MethodLabel { + method: string; + label: string; +} + +/** + * Why the current denoise setting can't be baked into a training run, phrased + * for display next to the checkbox — or null when it can. + */ +export function trainDenoiseBlockedReason(method: string): string | null { + if (method === 'none') { + return 'Pick a denoise filter in the Annotate tab first — there is nothing to apply yet.'; + } + if (method === 'model') { + return 'A learned denoiser cannot be used as a preprocessor for another model. Pick a classical filter in the Annotate tab.'; + } + return null; +} + +/** + * The `denoise` fragment to spread into a `/api/train/start` body: `{}` when + * the option is off or the setting is unusable, `{ denoise: {method, strength} }` + * when it applies. + * + * Returning a spreadable fragment rather than a nullable value keeps the two + * call sites from each re-deriving "and what do I do when it's off?" — the + * whole point being that OFF adds no key whatsoever. + */ +export function trainDenoisePayload( + enabled: boolean, + denoise: DenoiseSettings, +): { denoise: TrainDenoiseOpts } | Record { + if (!enabled) return {}; + if (trainDenoiseBlockedReason(denoise.method) !== null) return {}; + return { denoise: { method: denoise.method, strength: denoise.strength } }; +} + +/** + * Which filter the checkbox would actually bake in, e.g. `"Total variation, + * 60%"` — so "Train on denoised input" is never ambiguous about *what*. + * Null when the setting is unusable (the blocked reason is shown instead). + * + * `methods` is `capability.denoise.methods`; an unknown method falls back to + * its raw name rather than rendering blank, since a frontend newer than the + * server (or vice versa) shouldn't produce a nameless setting. + */ +export function trainDenoiseSummary( + denoise: DenoiseSettings, + methods: MethodLabel[], +): string | null { + if (trainDenoiseBlockedReason(denoise.method) !== null) return null; + return `${denoiseMethodLabel(denoise.method, methods)}, ${Math.round(denoise.strength * 100)}%`; +} + +/** Human label for a denoise method id, falling back to the id itself. */ +export function denoiseMethodLabel(method: string, methods: MethodLabel[]): string { + return methods.find((m) => m.method === method)?.label ?? method; +} diff --git a/frontend/src/lib/trainModelConfig.ts b/frontend/src/lib/trainModelConfig.ts new file mode 100644 index 0000000..dc67572 --- /dev/null +++ b/frontend/src/lib/trainModelConfig.ts @@ -0,0 +1,71 @@ +/** + * trainModelConfig — builds the `model` payload TrainPage sends to + * `/api/train/start` and `/api/train/estimate-batch`, and validates it before + * either request goes out. Both call sites used to assemble this same shape + * (and re-check the same precondition) by hand; this is the one place it + * happens now, so the two can't drift out of sync with each other. + * + * dlsia TUNet only — DINOv3 LoRA (and its checkpoint picker) is deferred, so + * there is only one model family to build a config for. + */ +import type { HyperparamsState } from '@/components/train/HyperparamsPanel'; +import { validateImageSize } from '@/lib/trainConstraints'; + +export interface TunetModelConfig { + model_family: 'dlsia_tunet'; + hyperparams: { + epochs: number; + lr: number; + depth: number; + base_channels: number; + growth_rate: number; + batch_size: number; + image_size: number; + flip_augment: boolean; + tiling: boolean; + }; +} + +export type TrainModelConfig = TunetModelConfig; + +/** Build the `model` field of a train/estimate-batch request. */ +export function buildModelConfig(hp: HyperparamsState): TrainModelConfig { + return { + model_family: 'dlsia_tunet', + hyperparams: { + epochs: hp.epochs, + lr: hp.lr, + depth: hp.depth, + base_channels: hp.base_channels, + growth_rate: hp.growth_rate, + batch_size: hp.batch_size, + image_size: hp.image_size, + flip_augment: hp.flip_augment, + tiling: hp.tiling, + }, + }; +} + +/** + * Validate a training config before either "Start training" or "Estimate + * batch size" fires. Returns a message to show the user, or null when the + * config is ready to submit. + */ +export function validateTrainConfig(hp: HyperparamsState): string | null { + return validateImageSize('dlsia_tunet', hp.image_size, hp.tiling ? 'Patch size' : 'Image size'); +} + +/** + * A string that changes iff any input to the batch-size probe's memory + * footprint changes: the image/patch size (tiling toggles whether patches are + * cut at native resolution, which changes the tensor shapes the probe + * actually measures). + * + * TrainPage stashes this when a probe starts and compares it against the + * current config when the probe finishes — a mismatch means the user changed + * something mid-measurement, so the result no longer describes what's about + * to be submitted and must not be silently adopted into `batch_size`. + */ +export function trainConfigSignature(hp: Pick): string { + return JSON.stringify([hp.image_size, hp.tiling]); +} diff --git a/frontend/src/lib/viridis.ts b/frontend/src/lib/viridis.ts new file mode 100644 index 0000000..14ce9bd --- /dev/null +++ b/frontend/src/lib/viridis.ts @@ -0,0 +1,269 @@ +/** + * Viridis (Matplotlib) 256-stop RGB LUT for probability overlays. + */ +export const VIRIDIS_RGB: ReadonlyArray = [ + [68, 1, 84], + [68, 2, 86], + [69, 4, 87], + [69, 5, 89], + [70, 7, 90], + [70, 8, 92], + [70, 10, 93], + [70, 11, 94], + [71, 13, 96], + [71, 14, 97], + [71, 16, 99], + [71, 17, 100], + [71, 19, 101], + [72, 20, 103], + [72, 22, 104], + [72, 23, 105], + [72, 24, 106], + [72, 26, 108], + [72, 27, 109], + [72, 28, 110], + [72, 29, 111], + [72, 31, 112], + [72, 32, 113], + [72, 33, 115], + [72, 35, 116], + [72, 36, 117], + [72, 37, 118], + [72, 38, 119], + [72, 40, 120], + [72, 41, 121], + [71, 42, 122], + [71, 44, 122], + [71, 45, 123], + [71, 46, 124], + [71, 47, 125], + [70, 48, 126], + [70, 50, 126], + [70, 51, 127], + [70, 52, 128], + [69, 53, 129], + [69, 55, 129], + [69, 56, 130], + [68, 57, 131], + [68, 58, 131], + [68, 59, 132], + [67, 61, 132], + [67, 62, 133], + [66, 63, 133], + [66, 64, 134], + [66, 65, 134], + [65, 66, 135], + [65, 68, 135], + [64, 69, 136], + [64, 70, 136], + [63, 71, 136], + [63, 72, 137], + [62, 73, 137], + [62, 74, 137], + [62, 76, 138], + [61, 77, 138], + [61, 78, 138], + [60, 79, 138], + [60, 80, 139], + [59, 81, 139], + [59, 82, 139], + [58, 83, 139], + [58, 84, 140], + [57, 85, 140], + [57, 86, 140], + [56, 88, 140], + [56, 89, 140], + [55, 90, 140], + [55, 91, 141], + [54, 92, 141], + [54, 93, 141], + [53, 94, 141], + [53, 95, 141], + [52, 96, 141], + [52, 97, 141], + [51, 98, 141], + [51, 99, 141], + [50, 100, 142], + [50, 101, 142], + [49, 102, 142], + [49, 103, 142], + [49, 104, 142], + [48, 105, 142], + [48, 106, 142], + [47, 107, 142], + [47, 108, 142], + [46, 109, 142], + [46, 110, 142], + [46, 111, 142], + [45, 112, 142], + [45, 113, 142], + [44, 113, 142], + [44, 114, 142], + [44, 115, 142], + [43, 116, 142], + [43, 117, 142], + [42, 118, 142], + [42, 119, 142], + [42, 120, 142], + [41, 121, 142], + [41, 122, 142], + [41, 123, 142], + [40, 124, 142], + [40, 125, 142], + [39, 126, 142], + [39, 127, 142], + [39, 128, 142], + [38, 129, 142], + [38, 130, 142], + [38, 130, 142], + [37, 131, 142], + [37, 132, 142], + [37, 133, 142], + [36, 134, 142], + [36, 135, 142], + [35, 136, 142], + [35, 137, 142], + [35, 138, 141], + [34, 139, 141], + [34, 140, 141], + [34, 141, 141], + [33, 142, 141], + [33, 143, 141], + [33, 144, 141], + [33, 145, 140], + [32, 146, 140], + [32, 146, 140], + [32, 147, 140], + [31, 148, 140], + [31, 149, 139], + [31, 150, 139], + [31, 151, 139], + [31, 152, 139], + [31, 153, 138], + [31, 154, 138], + [30, 155, 138], + [30, 156, 137], + [30, 157, 137], + [31, 158, 137], + [31, 159, 136], + [31, 160, 136], + [31, 161, 136], + [31, 161, 135], + [31, 162, 135], + [32, 163, 134], + [32, 164, 134], + [33, 165, 133], + [33, 166, 133], + [34, 167, 133], + [34, 168, 132], + [35, 169, 131], + [36, 170, 131], + [37, 171, 130], + [37, 172, 130], + [38, 173, 129], + [39, 173, 129], + [40, 174, 128], + [41, 175, 127], + [42, 176, 127], + [44, 177, 126], + [45, 178, 125], + [46, 179, 124], + [47, 180, 124], + [49, 181, 123], + [50, 182, 122], + [52, 182, 121], + [53, 183, 121], + [55, 184, 120], + [56, 185, 119], + [58, 186, 118], + [59, 187, 117], + [61, 188, 116], + [63, 188, 115], + [64, 189, 114], + [66, 190, 113], + [68, 191, 112], + [70, 192, 111], + [72, 193, 110], + [74, 193, 109], + [76, 194, 108], + [78, 195, 107], + [80, 196, 106], + [82, 197, 105], + [84, 197, 104], + [86, 198, 103], + [88, 199, 101], + [90, 200, 100], + [92, 200, 99], + [94, 201, 98], + [96, 202, 96], + [99, 203, 95], + [101, 203, 94], + [103, 204, 92], + [105, 205, 91], + [108, 205, 90], + [110, 206, 88], + [112, 207, 87], + [115, 208, 86], + [117, 208, 84], + [119, 209, 83], + [122, 209, 81], + [124, 210, 80], + [127, 211, 78], + [129, 211, 77], + [132, 212, 75], + [134, 213, 73], + [137, 213, 72], + [139, 214, 70], + [142, 214, 69], + [144, 215, 67], + [147, 215, 65], + [149, 216, 64], + [152, 216, 62], + [155, 217, 60], + [157, 217, 59], + [160, 218, 57], + [162, 218, 55], + [165, 219, 54], + [168, 219, 52], + [170, 220, 50], + [173, 220, 48], + [176, 221, 47], + [178, 221, 45], + [181, 222, 43], + [184, 222, 41], + [186, 222, 40], + [189, 223, 38], + [192, 223, 37], + [194, 223, 35], + [197, 224, 33], + [200, 224, 32], + [202, 225, 31], + [205, 225, 29], + [208, 225, 28], + [210, 226, 27], + [213, 226, 26], + [216, 226, 25], + [218, 227, 25], + [221, 227, 24], + [223, 227, 24], + [226, 228, 24], + [229, 228, 25], + [231, 228, 25], + [234, 229, 26], + [236, 229, 27], + [239, 229, 28], + [241, 229, 29], + [244, 230, 30], + [246, 230, 32], + [248, 230, 33], + [251, 231, 35], + [253, 231, 37], +]; + +/** Sample viridis; ``t`` in [0, 1] → RGB. */ +export function viridisRgb(t: number): [number, number, number] { + const x = Math.min(1, Math.max(0, t)); + const i = Math.min(255, Math.max(0, Math.round(x * 255))); + const rgb = VIRIDIS_RGB[i]!; + return [rgb[0], rgb[1], rgb[2]]; +} diff --git a/frontend/src/lib/volumeMaskPreview.test.ts b/frontend/src/lib/volumeMaskPreview.test.ts new file mode 100644 index 0000000..9c3de28 --- /dev/null +++ b/frontend/src/lib/volumeMaskPreview.test.ts @@ -0,0 +1,57 @@ +import { describe, it, expect } from 'vitest'; +import { buildLiveMaskVolume } from './volumeMaskPreview'; +import type { Shape } from '@/stores/annotationStore'; + +const rect = (classId: number, x: number, y: number, w: number, h: number): Shape => ({ + id: `${classId}-${x}-${y}`, + kind: 'rectangle', + classId, + x, y, w, h, +}); + +describe('buildLiveMaskVolume', () => { + it('returns null when there are no shapes anywhere', () => { + expect(buildLiveMaskVolume({}, 64, 64, 4)).toBeNull(); + expect(buildLiveMaskVolume({ '0': [] }, 64, 64, 4)).toBeNull(); + }); + + it('returns null when nSlices is not positive', () => { + expect(buildLiveMaskVolume({ '0': [rect(1, 0, 0, 4, 4)] }, 64, 64, 0)).toBeNull(); + }); + + it('produces dims in [width, height, depth] order matching the real dataset shape', () => { + const volume = buildLiveMaskVolume({ '0': [rect(1, 0, 0, 8, 8)] }, 64, 32, 10, 1000); + expect(volume).not.toBeNull(); + expect(volume!.dims).toEqual([64, 32, 10]); + expect(volume!.data.length).toBe(64 * 32 * 10); + }); + + it('leaves unannotated slices as all-background (0)', () => { + const volume = buildLiveMaskVolume({ '2': [rect(1, 0, 0, 8, 8)] }, 16, 16, 4, 1000); + const [w, h] = volume!.dims; + const sliceBytes = w * h; + const slice0 = volume!.data.subarray(0, sliceBytes); + expect(slice0.every((v) => v === 0)).toBe(true); + }); + + it('paints the annotated slice with the shape\'s class id', () => { + const volume = buildLiveMaskVolume({ '1': [rect(3, 0, 0, 8, 8)] }, 16, 16, 4, 1000); + const [w, h] = volume!.dims; + const sliceBytes = w * h; + const slice1 = volume!.data.subarray(sliceBytes, sliceBytes * 2); + expect(slice1[0]).toBe(3); + expect(Math.max(...slice1)).toBe(3); + }); + + it('resolves overlapping different-class shapes by ascending class id (higher wins)', () => { + const shapes = [rect(1, 0, 0, 16, 16), rect(5, 0, 0, 16, 16)]; + const volume = buildLiveMaskVolume({ '0': shapes }, 16, 16, 1, 1000); + expect(volume!.data[0]).toBe(5); + }); + + it('downsamples to at most maxDim on the longest in-plane edge', () => { + const volume = buildLiveMaskVolume({ '0': [rect(1, 0, 0, 100, 100)] }, 2048, 1024, 2, 256); + const [w, h] = volume!.dims; + expect(Math.max(w, h)).toBeLessThanOrEqual(256); + }); +}); diff --git a/frontend/src/lib/volumeMaskPreview.ts b/frontend/src/lib/volumeMaskPreview.ts new file mode 100644 index 0000000..a2bc094 --- /dev/null +++ b/frontend/src/lib/volumeMaskPreview.ts @@ -0,0 +1,88 @@ +/** + * Client-side, no-network rasterization of a sample's CURRENT annotation + * shapes (every slice, live in the annotation store) into a class-id volume + * — feeds the 3D view's "Fast (iPred)" mask layer via + * `WebGpuViewerInstance.loadMaskFromArray(slot, data, dims)` without + * requiring a "Sync masks to Tiled" round trip first. + * + * Unlike a Tiled-backed mask, this does NOT need to match the primary + * volume's own chosen texture resolution: `sampleMask()` in the viewer's WGSL + * maps a ray position to the mask's OWN voxel grid independently of the + * primary's (`frame.maskCtl.yzw` carries the mask's dims) — the two are only + * required to describe the same normalized [0,1]^3 box, not the same voxel + * count. That's what makes a coarse, fast, purely-client-side rasterization + * a legitimate live preview rather than something that needs to negotiate + * resolution with the renderer. + * + * Deliberately approximate, not a port of the backend's exact rasterizer + * (`coco_export.shape_to_mask`'s per-SHAPE "last one wins" order): shapes are + * grouped by class and each class's union is painted, so overlapping shapes + * of DIFFERENT classes resolve by ascending class id rather than draw order. + * Good enough for "does this look roughly right in 3D," not a source of + * truth — "Push to Tiled" is what produces the precise result. + */ +import { rasterizeUnion, gridFor } from './rasterize'; +import type { Shape } from '@/stores/annotationStore'; + +/** Longest in-plane edge the live preview rasterizes at. Coarser than the + * default 2-D mask grid (1600) on purpose — this runs entirely on the main + * thread for every slice up front, and a 3-D GPU mask texture reads it at + * whatever resolution it's given (see this module's own doc), so there is no + * fidelity reason to go higher for a quick preview. */ +const LIVE_PREVIEW_MAX_DIM = 256; + +export interface LiveMaskVolume { + data: Uint8Array; + /** `[width, height, depth]` — the order `loadMaskFromArray` expects. */ + dims: readonly [number, number, number]; +} + +/** + * Rasterize every slice of `byImageForSource` (already scoped to one sample) + * into one `width*height*depth` class-id volume, `depth` = `nSlices` (every + * slice in the sample gets a slot, annotated or not, so the volume's aspect + * ratio matches the real dataset rather than just the annotated subset). + * + * Returns `null` if there is nothing to rasterize (no shapes anywhere) — + * callers should treat that as "nothing to load," not an error. + */ +export function buildLiveMaskVolume( + byImageForSource: Record, + imageWidth: number, + imageHeight: number, + nSlices: number, + maxDim: number = LIVE_PREVIEW_MAX_DIM, +): LiveMaskVolume | null { + const hasAnyShape = Object.values(byImageForSource).some((shapes) => shapes.length > 0); + if (!hasAnyShape || nSlices <= 0) return null; + + const { gw, gh, scale } = gridFor(imageWidth, imageHeight, maxDim); + const depth = Math.max(1, Math.floor(nSlices)); + const data = new Uint8Array(gw * gh * depth); + + for (let z = 0; z < depth; z++) { + const shapes = byImageForSource[String(z)]; + if (!shapes || shapes.length === 0) continue; + + const byClass = new Map(); + for (const shape of shapes) { + const list = byClass.get(shape.classId); + if (list) list.push(shape); + else byClass.set(shape.classId, [shape]); + } + + const slice = data.subarray(z * gw * gh, (z + 1) * gw * gh); + // Ascending class id so a higher id visually "wins" ties — arbitrary but + // deterministic, matching mask_pyramid.majority_downsample's own + // tie-break convention on the backend. + for (const classId of [...byClass.keys()].sort((a, b) => a - b)) { + if (classId <= 0 || classId > 255) continue; // 0 is background by convention + const mask = rasterizeUnion(byClass.get(classId)!, gw, gh, scale); + for (let i = 0; i < mask.length; i++) { + if (mask[i]) slice[i] = classId; + } + } + } + + return { data, dims: [gw, gh, depth] }; +} diff --git a/frontend/src/lib/zarrUrl.test.ts b/frontend/src/lib/zarrUrl.test.ts new file mode 100644 index 0000000..36f1c7e --- /dev/null +++ b/frontend/src/lib/zarrUrl.test.ts @@ -0,0 +1,101 @@ +import { describe, it, expect } from 'vitest'; +import { buildZarrUrl, buildMaskZarrUrl, zarrRootFor, describeUnavailable } from './zarrUrl'; + +describe('zarrRootFor', () => { + it('appends the zarr v2 root to a bare origin', () => { + expect(zarrRootFor('http://127.0.0.1:8010')).toBe('http://127.0.0.1:8010/zarr/v2'); + }); + + it('tolerates a trailing slash', () => { + expect(zarrRootFor('http://127.0.0.1:8010/')).toBe('http://127.0.0.1:8010/zarr/v2'); + }); + + it('replaces a REST prefix rather than nesting under it', () => { + // .../api/v1/zarr/v2/... would 404 in a way that reads as missing data. + expect(zarrRootFor('https://tiled.example.org/api/v1')).toBe('https://tiled.example.org/zarr/v2'); + }); + + it('does not hardcode a port', () => { + // start_all.sh reassigns Tiled's port when 8010 is busy; the URI is the + // source of truth, so whatever port it carries must survive. + expect(zarrRootFor('http://127.0.0.1:8777')).toContain(':8777'); + }); +}); + +describe('buildZarrUrl', () => { + const server = 'http://127.0.0.1:8010'; + + it('addresses a tiled node under the zarr root', () => { + expect(buildZarrUrl('tiled', 'scans/petiole22', server)).toEqual({ + url: 'http://127.0.0.1:8010/zarr/v2/scans/petiole22', + reason: null, + }); + }); + + it('preserves path separators while escaping each segment', () => { + const { url } = buildZarrUrl('tiled', 'my scans/sample #1', server); + expect(url).toBe('http://127.0.0.1:8010/zarr/v2/my%20scans/sample%20%231'); + }); + + it('ignores empty path segments', () => { + const { url } = buildZarrUrl('tiled', '/scans//petiole22/', server); + expect(url).toBe('http://127.0.0.1:8010/zarr/v2/scans/petiole22'); + }); + + it.each([ + ['no source open', null, null, server, 'no-source'], + ['no path', 'tiled', null, server, 'no-source'], + ['local file', 'local', 'foo.tif', null, 'local-source'], + ['server unresolved', 'tiled', 'scans/x', null, 'no-server'], + ])('reports %s rather than guessing a URL', (_label, kind, source, uri, reason) => { + const result = buildZarrUrl(kind as string | null, source as string | null, uri as string | null); + expect(result.url).toBeNull(); + expect(result.reason).toBe(reason); + }); + + it('never emits credentials in the URL', () => { + // Anonymous read access is the contract; a key in the query string would + // put a write-capable credential in browser history and referrers. + const { url } = buildZarrUrl('tiled', 'scans/x', server); + expect(url).not.toMatch(/api_key|token|Authorization/i); + }); +}); + +describe('buildMaskZarrUrl', () => { + const server = 'http://127.0.0.1:8010'; + + it('appends the __masks/semantic sibling to the dataset path', () => { + expect(buildMaskZarrUrl('tiled', 'browse/260_R1_Oct_slice_2', server)).toEqual({ + url: 'http://127.0.0.1:8010/zarr/v2/browse/260_R1_Oct_slice_2__masks/semantic', + reason: null, + }); + }); + + it('keeps the __masks suffix on the leaf segment only, not the whole path', () => { + const { url } = buildMaskZarrUrl('tiled', 'a/b/c', server); + expect(url).toBe('http://127.0.0.1:8010/zarr/v2/a/b/c__masks/semantic'); + }); + + it('points at a distinct container for the deep-model suffix', () => { + const { url } = buildMaskZarrUrl('tiled', 'browse/x', server, '_deep'); + expect(url).toBe('http://127.0.0.1:8010/zarr/v2/browse/x__masks_deep/semantic'); + }); + + it.each([ + ['no source open', null, null, server, 'no-source'], + ['local file', 'local', 'foo.tif', null, 'local-source'], + ['server unresolved', 'tiled', 'scans/x', null, 'no-server'], + ])('reports %s rather than guessing a URL', (_label, kind, source, uri, reason) => { + const result = buildMaskZarrUrl(kind as string | null, source as string | null, uri as string | null); + expect(result.url).toBeNull(); + expect(result.reason).toBe(reason); + }); +}); + +describe('describeUnavailable', () => { + it('explains every reason', () => { + for (const reason of ['no-source', 'local-source', 'no-server'] as const) { + expect(describeUnavailable(reason).length).toBeGreaterThan(0); + } + }); +}); diff --git a/frontend/src/lib/zarrUrl.ts b/frontend/src/lib/zarrUrl.ts new file mode 100644 index 0000000..ef1f83a --- /dev/null +++ b/frontend/src/lib/zarrUrl.ts @@ -0,0 +1,143 @@ +/** + * zarrUrl — address a Tiled node as a Zarr store for the WebGPU volume viewer. + * + * Tiled 0.2.12 mounts a Zarr v2 router at `/zarr/v2`, so every array already in + * the catalog is *already* readable as a Zarr store — `.zattrs`/`.zarray` are + * synthesized from the Tiled structure and `/{i.j.k}` serves a chunk. Nothing + * has to be exported or copied: the 3D view reads the same catalog Browse and + * Annotate read. + * + * Two properties this module exists to preserve: + * + * 1. **The port is never hardcoded.** `start_all.sh` picks a free port for + * Tiled (default 8010, reassigned when busy) and the resolved URI reaches + * the frontend through `GET /api/config/servers` / `datasetStore.serverUri`. + * A literal `8010` anywhere here breaks the moment that fallback fires. + * 2. **No API key goes to the browser.** `tiled/config.yml` sets + * `allow_anonymous_access: true` and anonymous access is read-only; the + * generated key authorizes *writes* and stays server-side by design. Zarr + * chunk reads are reads, so the renderer's plain `fetch()` needs no + * credentials — and must not be given any. + * + * Pure and dependency-free so the URL algebra is unit-testable; the React side + * lives in `useZarrUrl`. + */ + +/** Why a source cannot be opened as a Zarr volume, when it cannot. */ +export type ZarrUnavailable = + /** Local files are read off disk by the backend; they are not in the catalog. */ + | 'local-source' + /** Nothing is open yet. */ + | 'no-source' + /** Kind is 'tiled' but no server URI has been resolved yet. */ + | 'no-server' + /** Still asking the backend which node holds the volume. */ + | 'resolving' + /** The dataset exists but has no multiscale volume to render. */ + | 'no-volume'; + +export interface ZarrUrlResult { + /** Zarr store root, or `null` when unavailable. */ + url: string | null; + reason: ZarrUnavailable | null; +} + +/** Human-readable explanation for a `ZarrUnavailable`, for the empty state. */ +export function describeUnavailable(reason: ZarrUnavailable): string { + switch (reason) { + case 'no-source': + return 'Open a dataset from Browse to view it in 3D.'; + case 'local-source': + return 'The 3D view reads from the Tiled catalog. Ingest this file first, then reopen it.'; + case 'no-server': + return 'Still resolving the Tiled server address — one moment.'; + case 'resolving': + return 'Looking for this dataset’s 3D volume…'; + case 'no-volume': + return 'No 3D volume has been built for this dataset yet.'; + } +} + +/** + * The `/zarr/v2` root for a Tiled server URI. + * + * Accepts a URI with or without a trailing `/api/v1`: our backend reports the + * bare origin (`http://127.0.0.1:8010`), but a Tiled URI copied from elsewhere + * often carries the REST prefix, and silently producing + * `.../api/v1/zarr/v2/...` would 404 in a way that looks like missing data + * rather than a malformed URL. + */ +export function zarrRootFor(serverUri: string): string { + const trimmed = serverUri.trim().replace(/\/+$/, ''); + const base = trimmed.replace(/\/api\/v\d+$/, ''); + return `${base}/zarr/v2`; +} + +/** + * Zarr store URL for a dataset's volume. + * + * @param kind `datasetStore.kind` — 'tiled' or 'local'. + * @param source Tiled path of the node that **holds the volume** — not + * necessarily the open dataset. A per-slice TIFF stack keeps its volume in a + * `__volume` sidecar, and a pyramid level's volume is its parent, so this must + * be the path `GET /api/volume/resolve` returned. Pointing the viewer at + * whatever happens to be open is what produces + * `openOmeZarr: missing multiscales in root .zattrs`. + * @param serverUri Resolved Tiled server URI (see the port note above). + */ +export function buildZarrUrl( + kind: string | null, + source: string | null, + serverUri: string | null, +): ZarrUrlResult { + if (!kind || !source) return { url: null, reason: 'no-source' }; + if (kind !== 'tiled') return { url: null, reason: 'local-source' }; + if (!serverUri) return { url: null, reason: 'no-server' }; + + // Encode each segment separately: the path is a `/`-joined chain of Tiled + // keys, and encodeURIComponent on the whole thing would escape the separators. + const path = source + .split('/') + .filter(Boolean) + .map(encodeURIComponent) + .join('/'); + + return { url: `${zarrRootFor(serverUri)}/${path}`, reason: null }; +} + +/** + * Zarr store URL for a sample's mask/annotation volume — the + * `__masks` sibling container `tiled_mask_sync.write_masks_to_tiled` + * writes, registered as a real OME-NGFF multiscale node by + * `mask_pyramid.register_mask_pyramid` so this is directly loadable via the + * volume viewer's mask layer (`loadMask(slot, url)`). + * + * @param source Tiled path of the annotated dataset — same `source` passed to + * `buildZarrUrl` for the primary volume, NOT the `__masks` container itself; + * the `__masks`/`semantic` suffix is appended here. + * @param suffix Distinguishes independent mask producers for the same source + * that must not merge into one container — `''` for the manual "sync masks + * to Tiled" action (iPred's fast results), `'_deep'` for a dlsia run's + * "Write masks to Tiled" (`infer_jobs.py`'s `container_suffix="_deep"`). + * Must match the backend suffix exactly or this points at an empty/missing + * container. + */ +export function buildMaskZarrUrl( + kind: string | null, + source: string | null, + serverUri: string | null, + suffix: '' | '_deep' = '', +): ZarrUrlResult { + if (!kind || !source) return { url: null, reason: 'no-source' }; + if (kind !== 'tiled') return { url: null, reason: 'local-source' }; + if (!serverUri) return { url: null, reason: 'no-server' }; + + const parts = source.split('/').filter(Boolean); + const stem = parts.pop(); + if (!stem) return { url: null, reason: 'no-source' }; + const path = [...parts, `${stem}__masks${suffix}`, 'semantic'] + .map(encodeURIComponent) + .join('/'); + + return { url: `${zarrRootFor(serverUri)}/${path}`, reason: null }; +} diff --git a/frontend/src/main.tsx b/frontend/src/main.tsx index 7b498a9..71a0290 100644 --- a/frontend/src/main.tsx +++ b/frontend/src/main.tsx @@ -3,13 +3,18 @@ import { createRoot } from 'react-dom/client'; import { BrowserRouter } from 'react-router'; import { QueryClient, QueryClientProvider } from '@tanstack/react-query'; import App from './app/App'; +import { installImageSliceGc } from './hooks/useImageSlice'; import './app/index.css'; const queryClient = new QueryClient(); +// Image slices are cached as blob object URLs, which the browser keeps alive until +// explicitly revoked. Release them as their queries leave the cache. +installImageSliceGc(queryClient); + createRoot(document.getElementById('root')!).render( - + diff --git a/frontend/src/stores/annotationStore.test.ts b/frontend/src/stores/annotationStore.test.ts new file mode 100644 index 0000000..3b9d954 --- /dev/null +++ b/frontend/src/stores/annotationStore.test.ts @@ -0,0 +1,389 @@ +import { describe, it, expect, beforeEach } from 'vitest'; +import { useAnnotationStore } from './annotationStore'; +import type { Shape, BrushShape } from './annotationStore'; + +const SK = 'tiled::sample'; + +const rect = (id: string, classId: number, overrides: Partial = {}): Shape => ({ + id, classId, kind: 'rectangle', x: 0, y: 0, w: 2, h: 2, ...overrides, +} as Shape); + +const poly = (id: string, classId: number): Shape => ({ + id, classId, kind: 'polygon', points: [0, 0, 10, 0, 10, 10], +}); + +const brush = (id: string, classId: number): BrushShape => ({ + id, classId, kind: 'brush', strokes: [], +}); + +beforeEach(() => { + useAnnotationStore.getState().reset(); + useAnnotationStore.temporal.getState().clear(); +}); + +describe('addShape / addShapes / addShapesAcrossSlices', () => { + it('addShape appends a single shape to the (source, slice)', () => { + useAnnotationStore.getState().addShape(SK, 0, rect('a', 1)); + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(1); + expect(useAnnotationStore.getState().byImage[SK]['0'][0].id).toBe('a'); + }); + + it('addShapes appends several shapes in one call', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1), rect('b', 2)]); + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(2); + }); + + it('addShapesAcrossSlices distributes shapes to multiple slices in one update', () => { + useAnnotationStore.getState().addShapesAcrossSlices(SK, { + 0: [rect('a', 1)], + 2: [rect('b', 1), rect('c', 1)], + }); + const byImage = useAnnotationStore.getState().byImage[SK]; + expect(byImage['0']).toHaveLength(1); + expect(byImage['2']).toHaveLength(2); + }); + + it('addShapesAcrossSlices skips slices with an empty shape list and is a no-op if all empty', () => { + useAnnotationStore.getState().addShapesAcrossSlices(SK, { 0: [], 1: [] }); + expect(useAnnotationStore.getState().byImage[SK]).toBeUndefined(); + }); + + it('addShapesAcrossSlices merges into existing slice shapes rather than replacing them', () => { + useAnnotationStore.getState().addShape(SK, 0, rect('a', 1)); + useAnnotationStore.getState().addShapesAcrossSlices(SK, { 0: [rect('b', 1)] }); + expect(useAnnotationStore.getState().byImage[SK]['0'].map((s) => s.id)).toEqual(['a', 'b']); + }); +}); + +describe('removeShape / removeShapes', () => { + beforeEach(() => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1), rect('b', 2), rect('c', 1)]); + }); + + it('removeShape drops one shape by id', () => { + useAnnotationStore.getState().removeShape(SK, 0, 'b'); + expect(useAnnotationStore.getState().byImage[SK]['0'].map((s) => s.id)).toEqual(['a', 'c']); + }); + + it('removeShapes drops several shapes by id in one call', () => { + useAnnotationStore.getState().removeShapes(SK, 0, ['a', 'c']); + expect(useAnnotationStore.getState().byImage[SK]['0'].map((s) => s.id)).toEqual(['b']); + }); +}); + +describe('updateShape', () => { + it('replaces only the targeted shape via the updater fn', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1), rect('b', 2)]); + useAnnotationStore.getState().updateShape(SK, 0, 'a', (sh) => ({ ...sh, classId: 9 } as Shape)); + const shapes = useAnnotationStore.getState().byImage[SK]['0']; + expect(shapes.find((s) => s.id === 'a')?.classId).toBe(9); + expect(shapes.find((s) => s.id === 'b')?.classId).toBe(2); + }); +}); + +describe('setClassForShapes', () => { + it('reassigns only the listed shape ids to the new class', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1), rect('b', 1), rect('c', 1)]); + useAnnotationStore.getState().setClassForShapes(SK, 0, ['a', 'c'], 7); + const shapes = useAnnotationStore.getState().byImage[SK]['0']; + expect(shapes.find((s) => s.id === 'a')?.classId).toBe(7); + expect(shapes.find((s) => s.id === 'b')?.classId).toBe(1); + expect(shapes.find((s) => s.id === 'c')?.classId).toBe(7); + }); +}); + +describe('replaceClassShapesOnSlice', () => { + it('replaces all shapes of a class on the slice, leaving other classes untouched', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1), rect('b', 2)]); + useAnnotationStore.getState().replaceClassShapesOnSlice(SK, 0, 1, [rect('new', 1)]); + const shapes = useAnnotationStore.getState().byImage[SK]['0']; + expect(shapes.map((s) => s.id).sort()).toEqual(['b', 'new']); + }); +}); + +describe('copySliceShapes', () => { + beforeEach(() => { + useAnnotationStore.getState().addShapes(SK, 0, [poly('a', 1), poly('b', 2)]); + }); + + it('clones shapes from the source slice into each target slice with fresh ids', () => { + useAnnotationStore.getState().copySliceShapes(SK, 0, [1, 2]); + const s1 = useAnnotationStore.getState().byImage[SK]['1']; + const s2 = useAnnotationStore.getState().byImage[SK]['2']; + expect(s1).toHaveLength(2); + expect(s2).toHaveLength(2); + expect(s1.map((s) => s.id)).not.toContain('a'); + }); + + it('filters to one class when classId is given', () => { + useAnnotationStore.getState().copySliceShapes(SK, 0, [1], 1); + const s1 = useAnnotationStore.getState().byImage[SK]['1']; + expect(s1).toHaveLength(1); + expect(s1[0].classId).toBe(1); + }); + + it('skips a target slice equal to the source slice', () => { + useAnnotationStore.getState().copySliceShapes(SK, 0, [0, 1]); + // Slice 0 keeps its originals only (no duplicate append onto itself). + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(2); + expect(useAnnotationStore.getState().byImage[SK]['1']).toHaveLength(2); + }); + + it('is a no-op when the source slice has nothing to copy', () => { + useAnnotationStore.getState().copySliceShapes(SK, 5, [6]); + expect(useAnnotationStore.getState().byImage[SK]['6']).toBeUndefined(); + }); + + it('merges copies with shapes already on the target slice', () => { + useAnnotationStore.getState().addShape(SK, 1, rect('existing', 3)); + useAnnotationStore.getState().copySliceShapes(SK, 0, [1]); + expect(useAnnotationStore.getState().byImage[SK]['1']).toHaveLength(3); + }); +}); + +describe('removeShapesByClassId', () => { + it('removes the class across every source and slice, pruning emptied slices/sources', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1), rect('b', 2)]); + useAnnotationStore.getState().addShapes('other::src', 0, [rect('c', 1)]); + useAnnotationStore.getState().removeShapesByClassId(1); + expect(useAnnotationStore.getState().byImage[SK]['0'].map((s) => s.id)).toEqual(['b']); + // The other source had only class-1 shapes, so it's pruned entirely. + expect(useAnnotationStore.getState().byImage['other::src']).toBeUndefined(); + }); + + it('prunes an emptied slice but keeps sibling slices with other classes', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + useAnnotationStore.getState().addShapes(SK, 1, [rect('b', 2)]); + useAnnotationStore.getState().removeShapesByClassId(1); + expect(useAnnotationStore.getState().byImage[SK]['0']).toBeUndefined(); + expect(useAnnotationStore.getState().byImage[SK]['1']).toHaveLength(1); + }); +}); + +describe('removeShapesByClassIdInSource', () => { + it('removes the class only within the given source, leaving other sources alone', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1), rect('b', 2)]); + useAnnotationStore.getState().addShapes('other::src', 0, [rect('c', 1)]); + useAnnotationStore.getState().removeShapesByClassIdInSource(SK, 1); + expect(useAnnotationStore.getState().byImage[SK]['0'].map((s) => s.id)).toEqual(['b']); + expect(useAnnotationStore.getState().byImage['other::src']['0']).toHaveLength(1); + }); + + it('deletes the source entirely when its last class is removed', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + useAnnotationStore.getState().removeShapesByClassIdInSource(SK, 1); + expect(useAnnotationStore.getState().byImage[SK]).toBeUndefined(); + }); + + it('is a no-op for a source with no data', () => { + useAnnotationStore.getState().removeShapesByClassIdInSource('nope', 1); + expect(useAnnotationStore.getState().byImage['nope']).toBeUndefined(); + }); +}); + +describe('removeShapesByOrigin', () => { + it('removes shapes of the given origin, treating missing origin as human', () => { + useAnnotationStore.getState().addShapes(SK, 0, [ + rect('human1', 1), + { ...rect('pred1', 1), origin: 'predicted' } as Shape, + ]); + useAnnotationStore.getState().removeShapesByOrigin(SK, 'predicted'); + expect(useAnnotationStore.getState().byImage[SK]['0'].map((s) => s.id)).toEqual(['human1']); + }); + + it('removes human-origin shapes (undefined origin) when asked', () => { + useAnnotationStore.getState().addShapes(SK, 0, [ + rect('human1', 1), + { ...rect('pred1', 1), origin: 'predicted' } as Shape, + ]); + useAnnotationStore.getState().removeShapesByOrigin(SK, 'human'); + expect(useAnnotationStore.getState().byImage[SK]['0'].map((s) => s.id)).toEqual(['pred1']); + }); + + it('deletes an emptied source and is a no-op for missing sources', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + useAnnotationStore.getState().removeShapesByOrigin(SK, 'human'); + expect(useAnnotationStore.getState().byImage[SK]).toBeUndefined(); + + useAnnotationStore.getState().removeShapesByOrigin('nope', 'human'); + expect(useAnnotationStore.getState().byImage['nope']).toBeUndefined(); + }); +}); + +describe('appendBrushStroke / appendEraseStroke', () => { + it('appendBrushStroke appends a stroke to a brush shape and is a no-op for non-brush shapes', () => { + useAnnotationStore.getState().addShapes(SK, 0, [brush('br', 1), rect('rc', 1)]); + const stroke = { points: [1, 1], radius: 3, mode: 'paint' as const }; + useAnnotationStore.getState().appendBrushStroke(SK, 0, 'br', stroke); + useAnnotationStore.getState().appendBrushStroke(SK, 0, 'rc', stroke); + const shapes = useAnnotationStore.getState().byImage[SK]['0']; + const b = shapes.find((s) => s.id === 'br') as BrushShape; + expect(b.strokes).toHaveLength(1); + const r = shapes.find((s) => s.id === 'rc'); + expect((r as any).strokes).toBeUndefined(); + }); + + it('appendEraseStroke adds an erase-mode brush stroke for brush shapes', () => { + useAnnotationStore.getState().addShapes(SK, 0, [brush('br', 1)]); + useAnnotationStore.getState().appendEraseStroke(SK, 0, 'br', { points: [1, 1], radius: 2 }); + const b = useAnnotationStore.getState().byImage[SK]['0'][0] as BrushShape; + expect(b.strokes).toHaveLength(1); + expect(b.strokes[0].mode).toBe('erase'); + }); + + it('appendEraseStroke adds an `erased` carve-out entry for vector shapes', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('rc', 1)]); + useAnnotationStore.getState().appendEraseStroke(SK, 0, 'rc', { points: [1, 1], radius: 2 }); + const shape = useAnnotationStore.getState().byImage[SK]['0'][0]; + expect(shape.erased).toHaveLength(1); + expect(shape.erased?.[0].radius).toBe(2); + }); +}); + +describe('setShapes', () => { + it('replaces all shapes for the (source, slice)', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + useAnnotationStore.getState().setShapes(SK, 0, [rect('b', 2), rect('c', 2)]); + expect(useAnnotationStore.getState().byImage[SK]['0'].map((s) => s.id)).toEqual(['b', 'c']); + }); +}); + +describe('setSplitForSlice / toggleNegativeSlice', () => { + it('sets the split value for a slice', () => { + useAnnotationStore.getState().setSplitForSlice(SK, 3, 'valid'); + expect(useAnnotationStore.getState().splitBySlice[SK]['3']).toBe('valid'); + }); + + it('toggleNegativeSlice adds then removes a slice from the negative list', () => { + useAnnotationStore.getState().toggleNegativeSlice(SK, 2); + expect(useAnnotationStore.getState().negativeSlices[SK]).toEqual(['2']); + useAnnotationStore.getState().toggleNegativeSlice(SK, 2); + expect(useAnnotationStore.getState().negativeSlices[SK]).toEqual([]); + }); +}); + +describe('loadFromDraft / mergeSourceDraft', () => { + it('loadFromDraft replaces all annotation data wholesale', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + useAnnotationStore.getState().loadFromDraft({ + byImage: { 'new::src': { '0': [rect('x', 1)] } }, + splitBySlice: {}, + negativeSlices: {}, + }); + expect(useAnnotationStore.getState().byImage[SK]).toBeUndefined(); + expect(useAnnotationStore.getState().byImage['new::src']['0']).toHaveLength(1); + }); + + it('mergeSourceDraft merges one source without touching others', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + useAnnotationStore.getState().mergeSourceDraft( + 'other::src', + { '0': [rect('b', 2)] }, + { '0': 'train' }, + ['0'], + ); + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(1); + expect(useAnnotationStore.getState().byImage['other::src']['0']).toHaveLength(1); + expect(useAnnotationStore.getState().splitBySlice['other::src']['0']).toBe('train'); + expect(useAnnotationStore.getState().negativeSlices['other::src']).toEqual(['0']); + }); +}); + +describe('reset', () => { + it('clears shapes, splits, negative slices, and the draft', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + useAnnotationStore.getState().setSplitForSlice(SK, 0, 'test'); + useAnnotationStore.getState().toggleNegativeSlice(SK, 0); + useAnnotationStore.getState().addPolyNode(SK, '0', 1, 1); + useAnnotationStore.getState().reset(); + const s = useAnnotationStore.getState(); + expect(s.byImage).toEqual({}); + expect(s.splitBySlice).toEqual({}); + expect(s.negativeSlices).toEqual({}); + expect(s.draft.tool).toBeNull(); + }); +}); + +describe('touchHistory', () => { + it('creates a shallow-cloned byImage but leaves the data equal', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + const before = useAnnotationStore.getState().byImage; + useAnnotationStore.getState().touchHistory(); + const after = useAnnotationStore.getState().byImage; + expect(after).not.toBe(before); + expect(after).toEqual(before); + }); + + it('records a temporal (undo) entry even though content is unchanged', () => { + useAnnotationStore.temporal.getState().clear(); + expect(useAnnotationStore.temporal.getState().pastStates.length).toBe(0); + useAnnotationStore.getState().touchHistory(); + expect(useAnnotationStore.temporal.getState().pastStates.length).toBeGreaterThan(0); + }); +}); + +describe('draft: addPolyNode / addMagneticNode / clearDraft / commitDraftShapes', () => { + it('addPolyNode starts a fresh draft when context changes and appends within the same context', () => { + useAnnotationStore.getState().addPolyNode(SK, '0', 1, 2); + expect(useAnnotationStore.getState().draft.poly).toEqual([1, 2]); + useAnnotationStore.getState().addPolyNode(SK, '0', 3, 4); + expect(useAnnotationStore.getState().draft.poly).toEqual([1, 2, 3, 4]); + }); + + it('addPolyNode resets the draft when the slice/source context changes', () => { + useAnnotationStore.getState().addPolyNode(SK, '0', 1, 2); + useAnnotationStore.getState().addPolyNode(SK, '1', 5, 6); + expect(useAnnotationStore.getState().draft.poly).toEqual([5, 6]); + expect(useAnnotationStore.getState().draft.sliceKey).toBe('1'); + }); + + it('addMagneticNode accumulates path points for the same context and reseeds otherwise', () => { + useAnnotationStore.getState().addMagneticNode(SK, '0', [1, 1, 2, 2], { x: 1, y: 1 }); + useAnnotationStore.getState().addMagneticNode(SK, '0', [3, 3], { x: 3, y: 3 }); + const draft = useAnnotationStore.getState().draft; + expect(draft.magnetic).toEqual([1, 1, 2, 2, 3, 3]); + expect(draft.magneticSeed).toEqual({ x: 3, y: 3 }); + }); + + it('clearDraft resets to the empty draft', () => { + useAnnotationStore.getState().addPolyNode(SK, '0', 1, 2); + useAnnotationStore.getState().clearDraft(); + expect(useAnnotationStore.getState().draft.tool).toBeNull(); + expect(useAnnotationStore.getState().draft.poly).toEqual([]); + }); + + it('commitDraftShapes writes the slice shapes and clears the draft atomically', () => { + useAnnotationStore.getState().addPolyNode(SK, '0', 1, 2); + useAnnotationStore.getState().commitDraftShapes(SK, 0, [poly('finished', 1)]); + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(1); + expect(useAnnotationStore.getState().draft.tool).toBeNull(); + }); +}); + +describe('undo/redo via zundo temporal middleware', () => { + it('undo restores the prior byImage state and redo re-applies the edit', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(1); + + useAnnotationStore.getState().addShapes(SK, 0, [rect('b', 1)]); + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(2); + + useAnnotationStore.temporal.getState().undo(); + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(1); + expect(useAnnotationStore.getState().byImage[SK]['0'][0].id).toBe('a'); + + useAnnotationStore.temporal.getState().redo(); + expect(useAnnotationStore.getState().byImage[SK]['0']).toHaveLength(2); + }); + + it('undo can walk back past multiple edits to the initial empty state', () => { + useAnnotationStore.getState().addShapes(SK, 0, [rect('a', 1)]); + useAnnotationStore.getState().addShapes(SK, 1, [rect('b', 1)]); + useAnnotationStore.getState().removeShape(SK, 0, 'a'); + + useAnnotationStore.temporal.getState().undo(); + useAnnotationStore.temporal.getState().undo(); + useAnnotationStore.temporal.getState().undo(); + expect(useAnnotationStore.getState().byImage[SK]).toBeUndefined(); + }); +}); diff --git a/frontend/src/stores/annotationStore.ts b/frontend/src/stores/annotationStore.ts index 7673ade..6c1cac5 100644 --- a/frontend/src/stores/annotationStore.ts +++ b/frontend/src/stores/annotationStore.ts @@ -29,12 +29,20 @@ export interface EraseStroke { radius: number; } +/** Provenance: 'predicted' shapes came from committing an iPred run (single-slice + * or volume-apply); undefined/'human' means hand-drawn. Kept forever rather than + * cleared on first edit, so Phase 6's 3D label channel and the Layers panel can + * filter by it at any time. */ +export type ShapeOrigin = 'human' | 'predicted'; + export interface BaseShape { id: string; classId: number; kind: Shape['kind']; /** Optional erase carve-outs (rendered destination-out, subtracted on export). */ erased?: EraseStroke[]; + /** Undefined means human-drawn (backward compatible with every existing shape). */ + origin?: ShapeOrigin; } export interface PolygonShape extends BaseShape { @@ -98,6 +106,10 @@ export interface AnnotationState { addShape: (sourceKey: string, sliceIdx: number, shape: Shape) => void; /** Append several shapes in one update (one undo step) — used by magic-wand. */ addShapes: (sourceKey: string, sliceIdx: number, shapes: Shape[]) => void; + /** Append shapes across MULTIPLE slices in one update (one undo step) — used + * by the iPred volume-apply commit, so accepting a whole-volume prediction + * is one undo, not one per slice. */ + addShapesAcrossSlices: (sourceKey: string, bySlice: Record) => void; removeShape: (sourceKey: string, sliceIdx: number, shapeId: string) => void; /** Remove several shapes in one update (one undo step). */ removeShapes: (sourceKey: string, sliceIdx: number, shapeIds: string[]) => void; @@ -113,6 +125,9 @@ export interface AnnotationState { copySliceShapes: (sourceKey: string, fromSlice: number, toSlices: number[], classId?: number | null) => void; /** Remove every shape with *classId* across all loaded samples (all slices). */ removeShapesByClassId: (classId: number) => void; + /** Remove every shape of *origin* within one sample, across its slices — e.g. + * "reject all predicted" after reviewing a volume-apply commit. One undo step. */ + removeShapesByOrigin: (sourceKey: string, origin: ShapeOrigin) => void; /** Remove every shape with *classId* within a single sample (*sourceKey*), across * its slices. Scoped so deleting a class never touches other samples' annotations. * Tracked by zundo, so Ctrl/Cmd+Z restores the removed regions. */ @@ -187,6 +202,19 @@ export const useAnnotationStore = create()( }; }), + /** Appends shapes across multiple slices in one update (one undo step). */ + addShapesAcrossSlices: (sourceKey, bySlice) => + set((s) => { + const entries = Object.entries(bySlice).filter(([, shapes]) => shapes.length > 0); + if (entries.length === 0) return {}; + const slices = { ...(s.byImage[sourceKey] ?? {}) }; + for (const [sliceIdx, shapes] of entries) { + const sliceKey = String(Number(sliceIdx)); + slices[sliceKey] = [...(slices[sliceKey] ?? []), ...shapes]; + } + return { byImage: { ...s.byImage, [sourceKey]: slices } }; + }), + /** Removes the shape with the given id from the (sourceKey, slice). */ removeShape: (sourceKey, sliceIdx, shapeId) => set((s) => { @@ -326,6 +354,28 @@ export const useAnnotationStore = create()( return { byImage: nextByImage }; }), + /** Removes every shape of *origin* within one sample, pruning emptied slices + * and the source itself; other samples are left untouched. */ + removeShapesByOrigin: (sourceKey, origin) => + set((s) => { + const slices = s.byImage[sourceKey]; + if (!slices) return {}; + const nextSlices: Record = {}; + for (const [sliceKey, shapes] of Object.entries(slices)) { + const filtered = shapes.filter((sh) => (sh.origin ?? 'human') !== origin); + if (filtered.length > 0) { + nextSlices[sliceKey] = filtered; + } + } + const nextByImage = { ...s.byImage }; + if (Object.keys(nextSlices).length > 0) { + nextByImage[sourceKey] = nextSlices; + } else { + delete nextByImage[sourceKey]; + } + return { byImage: nextByImage }; + }), + /** Clone every shape of *fromClassId* into *toClassId* across all slices of * *sourceKey* (fresh ids), merged with existing shapes. One undo step. */ duplicateClassShapes: (sourceKey, fromClassId, toClassId) => diff --git a/frontend/src/stores/clipboardStore.test.ts b/frontend/src/stores/clipboardStore.test.ts new file mode 100644 index 0000000..e05ad09 --- /dev/null +++ b/frontend/src/stores/clipboardStore.test.ts @@ -0,0 +1,30 @@ +import { beforeEach, describe, expect, it } from 'vitest'; +import { useClipboardStore } from './clipboardStore'; + +beforeEach(() => { + useClipboardStore.setState({ shapes: [] }); +}); + +describe('clipboardStore', () => { + it('starts empty', () => { + expect(useClipboardStore.getState().shapes).toEqual([]); + }); + + it('copy stores a deep copy of the given shapes', () => { + const shape = { id: 's1', classId: 1, kind: 'rectangle' as const, x: 0, y: 0, w: 2, h: 2 }; + useClipboardStore.getState().copy([shape]); + const stored = useClipboardStore.getState().shapes; + expect(stored).toEqual([shape]); + expect(stored[0]).not.toBe(shape); // deep copy, not the same reference + + shape.x = 99; + const restored = useClipboardStore.getState().shapes[0]; + expect(restored.kind === 'rectangle' && restored.x).toBe(0); // unaffected by later mutation + }); + + it('clear empties the clipboard', () => { + useClipboardStore.getState().copy([{ id: 's1', classId: 1, kind: 'rectangle', x: 0, y: 0, w: 1, h: 1 }]); + useClipboardStore.getState().clear(); + expect(useClipboardStore.getState().shapes).toEqual([]); + }); +}); diff --git a/frontend/src/stores/connectionStore.test.ts b/frontend/src/stores/connectionStore.test.ts new file mode 100644 index 0000000..8bcbeee --- /dev/null +++ b/frontend/src/stores/connectionStore.test.ts @@ -0,0 +1,56 @@ +import { afterEach, describe, expect, it } from 'vitest'; +import { useConnectionStore } from './connectionStore'; + +const INITIAL = useConnectionStore.getState(); + +afterEach(() => { + useConnectionStore.setState(INITIAL, true); +}); + +describe('connectionStore', () => { + it('starts with no connection and status unknown', () => { + const s = useConnectionStore.getState(); + expect(s.kind).toBeNull(); + expect(s.status).toBe('unknown'); + }); + + it('setConnection resets status to unknown', () => { + useConnectionStore.getState().setStatus('ok'); + useConnectionStore.getState().setConnection({ + kind: 'tiled', + serverUri: 'http://example', + label: 'Example', + sampleCount: 3, + }); + const s = useConnectionStore.getState(); + expect(s.kind).toBe('tiled'); + expect(s.status).toBe('unknown'); + }); + + it('setStatus updates status independently of other fields', () => { + useConnectionStore.getState().setConnection({ + kind: 'tiled', + serverUri: 'http://example', + label: 'Example', + sampleCount: 3, + }); + useConnectionStore.getState().setStatus('error'); + const s = useConnectionStore.getState(); + expect(s.status).toBe('error'); + expect(s.serverUri).toBe('http://example'); + }); + + it('clearConnection resets status to unknown along with everything else', () => { + useConnectionStore.getState().setConnection({ + kind: 'tiled', + serverUri: 'http://example', + label: 'Example', + sampleCount: 3, + }); + useConnectionStore.getState().setStatus('ok'); + useConnectionStore.getState().clearConnection(); + const s = useConnectionStore.getState(); + expect(s.kind).toBeNull(); + expect(s.status).toBe('unknown'); + }); +}); diff --git a/frontend/src/stores/connectionStore.ts b/frontend/src/stores/connectionStore.ts index 56df832..7d2971f 100644 --- a/frontend/src/stores/connectionStore.ts +++ b/frontend/src/stores/connectionStore.ts @@ -27,6 +27,13 @@ export interface ConnectionState { label: string | null; /** Total number of samples reported by /api/connect/summary */ sampleCount: number | null; + /** + * Live Tiled reachability, driven by a periodic health check (see + * useConnectionHealth). 'unknown' until the first check resolves, and + * always 'unknown' for local connections (no network dependency to check). + */ + status: 'unknown' | 'ok' | 'error'; + setStatus: (status: 'unknown' | 'ok' | 'error') => void; setConnection: (payload: { kind: 'tiled' | 'local'; serverUri?: string | null; @@ -49,6 +56,9 @@ export const useConnectionStore = create((set) => ({ localRel: null, label: null, sampleCount: null, + status: 'unknown', + + setStatus: (status) => set({ status }), /** Records the active data-source connection; unspecified fields default to null. */ setConnection: ({ @@ -70,6 +80,7 @@ export const useConnectionStore = create((set) => ({ localRel, label, sampleCount, + status: 'unknown', }), /** Resets all connection fields to null (disconnect). */ @@ -82,6 +93,7 @@ export const useConnectionStore = create((set) => ({ localRoot: null, localRel: null, label: null, + status: 'unknown', sampleCount: null, }), })); diff --git a/frontend/src/stores/datasetStore.ts b/frontend/src/stores/datasetStore.ts index ecdb445..05de0cc 100644 --- a/frontend/src/stores/datasetStore.ts +++ b/frontend/src/stores/datasetStore.ts @@ -10,6 +10,25 @@ export interface ImageMeta { dtype: string; isRgb: boolean; valueRange: [number, number]; + /** The [vmin, vmax] actually used by `/api/image/slice`'s norm="global" + * rendering (a 1st/99th-percentile stretch across the whole volume) — + * distinct from `valueRange` above (slice 0's raw min/max). Needed to + * convert a displayed 0-255 byte value back to a physical intensity, e.g. + * the Sampler-fitted band sent to the 3D viewer. `null` for RGB sources + * (the backend doesn't compute it for those). */ + globalValueRange?: [number, number] | null; + /** Multiscale (Zarr) volumes only — which pyramid level is being displayed. + * `width`/`height`/`nSlices` above always describe the FINEST level, because + * annotations are stored in full-resolution coordinates whichever level is + * open; these describe the image actually drawn underneath them. */ + levelKey?: string | null; + levelIndex?: number | null; + levelCount?: number | null; + levelWidth?: number | null; + levelHeight?: number | null; + levelNSlices?: number | null; + /** Finest-z / level-z. When > 1, only every f-th full-res slice is addressable. */ + zDownsample?: number | null; } export interface RenderOpts { @@ -20,6 +39,21 @@ export interface RenderOpts { cmap: 'gray' | 'viridis'; } +/** + * Denoising applied to the slice before it is normalized for display. + * + * Deliberately NOT part of `RenderOpts`: render options travel into export + * payloads, and a denoise preview must not silently change exported pixels. It + * changes what you see (and what the intensity tools therefore act on) — turning + * it into data is what the "Save denoised copy" bake is for. + */ +export interface DenoiseOpts { + /** A `denoise.ALL_METHODS` entry; 'none' disables it. */ + method: string; + /** 0..1, mapped by the backend onto each method's native parameter. */ + strength: number; +} + export interface DatasetState { /** 'tiled' or 'local' */ kind: string | null; @@ -29,9 +63,11 @@ export interface DatasetState { meta: ImageMeta | null; currentSlice: number; renderOpts: RenderOpts; + denoise: DenoiseOpts; setDataset: (kind: string, source: string, serverUri: string | null, meta: ImageMeta) => void; setSlice: (idx: number) => void; setRenderOpts: (opts: Partial) => void; + setDenoise: (opts: Partial) => void; reset: () => void; } @@ -43,6 +79,8 @@ const DEFAULT_RENDER: RenderOpts = { cmap: 'gray', }; +const DEFAULT_DENOISE: DenoiseOpts = { method: 'none', strength: 0.5 }; + export const useDatasetStore = create((set) => ({ kind: null, source: null, @@ -50,15 +88,23 @@ export const useDatasetStore = create((set) => ({ meta: null, currentSlice: 0, renderOpts: { ...DEFAULT_RENDER }, - /** Activates a new image source and resets the current slice to 0. */ + denoise: { ...DEFAULT_DENOISE }, + /** Activates a new image source and resets the current slice to 0. + * Denoising resets too: its strength is tuned to one volume's noise level and + * carrying it to the next would silently mis-filter it. */ setDataset: (kind, source, serverUri, meta) => - set({ kind, source, serverUri, meta, currentSlice: 0 }), + set({ kind, source, serverUri, meta, currentSlice: 0, denoise: { ...DEFAULT_DENOISE } }), /** Sets the active slice index. */ setSlice: (idx) => set({ currentSlice: idx }), /** Merges partial render options (normalization/scale/percentiles/cmap). */ setRenderOpts: (opts) => set((s) => ({ renderOpts: { ...s.renderOpts, ...opts } })), + /** Merges partial denoise options (method/strength). */ + setDenoise: (opts) => set((s) => ({ denoise: { ...s.denoise, ...opts } })), /** Clears the active dataset and restores default render options. */ reset: () => - set({ kind: null, source: null, serverUri: null, meta: null, currentSlice: 0, renderOpts: { ...DEFAULT_RENDER } }), + set({ + kind: null, source: null, serverUri: null, meta: null, currentSlice: 0, + renderOpts: { ...DEFAULT_RENDER }, denoise: { ...DEFAULT_DENOISE }, + }), })); diff --git a/frontend/src/stores/ipredStore.ts b/frontend/src/stores/ipredStore.ts new file mode 100644 index 0000000..f03ac74 --- /dev/null +++ b/frontend/src/stores/ipredStore.ts @@ -0,0 +1,67 @@ +/** + * ipredStore — preferences + session state for the iPred (interactive + * segmentation) service. Separate from connectionStore (data-source + * connection) and datasetStore (active sample), per this codebase's existing + * store-separation convention. + */ +import { create } from 'zustand'; + +export interface TrainerConfig { + iterations: number; + depth: number; + learning_rate: number; +} + +export const DEFAULT_TRAINER_CONFIG: TrainerConfig = { + iterations: 200, + depth: 6, + learning_rate: 0.1, +}; + +export const DEFAULT_COMPOSITION_ID = 'comp-skimage-slimsam'; +export const DEFAULT_TRAINER_ID = 'catboost'; + +export interface IpredState { + /** Preferred composition document id (modular feature graph). */ + preferredCompositionId: string; + /** Preferred trainer plugin id. */ + preferredTrainerId: string; + preferredTrainerConfig: TrainerConfig; + /** Active ipred session for the currently opened sample. */ + ipredSessionId: string | null; + ipredProjectId: string | null; + setPreferredCompositionId: (id: string) => void; + setPreferredTrainer: (payload: { id?: string; config?: Partial }) => void; + setIpredSession: (payload: { sessionId: string | null; projectId?: string | null }) => void; + reset: () => void; +} + +export const useIpredStore = create((set) => ({ + preferredCompositionId: DEFAULT_COMPOSITION_ID, + preferredTrainerId: DEFAULT_TRAINER_ID, + preferredTrainerConfig: { ...DEFAULT_TRAINER_CONFIG }, + ipredSessionId: null, + ipredProjectId: null, + + setPreferredCompositionId: (id) => set({ preferredCompositionId: id }), + + setPreferredTrainer: ({ id, config }) => + set((s) => ({ + preferredTrainerId: id ?? s.preferredTrainerId, + preferredTrainerConfig: config + ? { ...s.preferredTrainerConfig, ...config } + : s.preferredTrainerConfig, + })), + + setIpredSession: ({ sessionId, projectId = null }) => + set({ ipredSessionId: sessionId, ipredProjectId: projectId }), + + reset: () => + set({ + preferredCompositionId: DEFAULT_COMPOSITION_ID, + preferredTrainerId: DEFAULT_TRAINER_ID, + preferredTrainerConfig: { ...DEFAULT_TRAINER_CONFIG }, + ipredSessionId: null, + ipredProjectId: null, + }), +})); diff --git a/frontend/src/stores/layerVisibilityStore.test.ts b/frontend/src/stores/layerVisibilityStore.test.ts new file mode 100644 index 0000000..edec7db --- /dev/null +++ b/frontend/src/stores/layerVisibilityStore.test.ts @@ -0,0 +1,58 @@ +import { describe, expect, it, beforeEach } from 'vitest'; +import { + isPredictionClassVisible, + isShapeOriginVisible, + useLayerVisibilityStore, +} from '@/stores/layerVisibilityStore'; + +describe('layerVisibilityStore', () => { + beforeEach(() => { + useLayerVisibilityStore.setState({ + groups: { + image: true, + denoise: true, + features: true, + proba: true, + predictions: true, + annotations: true, + manifold: true, + }, + predictionClassVisible: {}, + showPredictionMulti: true, + showPredictionAbstain: true, + annotationOriginVisible: { human: true, predicted: true }, + }); + }); + + it('toggles groups and prediction classes', () => { + const s = useLayerVisibilityStore.getState(); + s.toggleGroup('proba'); + expect(useLayerVisibilityStore.getState().groups.proba).toBe(false); + s.ensurePredictionClasses([1, 2]); + expect(isPredictionClassVisible(useLayerVisibilityStore.getState().predictionClassVisible, 1)).toBe( + true, + ); + s.togglePredictionClass(1); + expect(isPredictionClassVisible(useLayerVisibilityStore.getState().predictionClassVisible, 1)).toBe( + false, + ); + }); + + it('denoise group defaults on and toggles independently', () => { + const s = useLayerVisibilityStore.getState(); + expect(useLayerVisibilityStore.getState().groups.denoise).toBe(true); + s.toggleGroup('denoise'); + expect(useLayerVisibilityStore.getState().groups.denoise).toBe(false); + expect(useLayerVisibilityStore.getState().groups.image).toBe(true); + }); + + it('toggles predicted/human annotation visibility independently', () => { + const s = useLayerVisibilityStore.getState(); + expect(isShapeOriginVisible(useLayerVisibilityStore.getState().annotationOriginVisible, 'predicted')).toBe(true); + expect(isShapeOriginVisible(useLayerVisibilityStore.getState().annotationOriginVisible, undefined)).toBe(true); + s.setAnnotationOriginVisible('predicted', false); + expect(isShapeOriginVisible(useLayerVisibilityStore.getState().annotationOriginVisible, 'predicted')).toBe(false); + // Human-drawn (origin undefined) is unaffected by hiding predicted shapes. + expect(isShapeOriginVisible(useLayerVisibilityStore.getState().annotationOriginVisible, undefined)).toBe(true); + }); +}); diff --git a/frontend/src/stores/layerVisibilityStore.ts b/frontend/src/stores/layerVisibilityStore.ts new file mode 100644 index 0000000..e64e6c4 --- /dev/null +++ b/frontend/src/stores/layerVisibilityStore.ts @@ -0,0 +1,141 @@ +/** + * Canvas layer visibility — function groups + per-class toggles. + * + * Groups control image / denoise / features / probability / predictions / + * annotations / manifold. Per-class keys under predictions and annotations + * (annotations also sync with classStore.isVisible). + */ +import { create } from 'zustand'; + +export type LayerGroupId = + | 'image' + | 'denoise' + | 'features' + | 'proba' + | 'predictions' + | 'annotations' + | 'manifold'; + +export const LAYER_GROUP_META: Record< + LayerGroupId, + { label: string; hint: string } +> = { + image: { label: 'Image', hint: 'Base slice (brightness/levels apply here)' }, + denoise: { label: 'Denoise', hint: 'Server-side denoising of the base slice' }, + features: { label: 'Features', hint: 'Preprocess channel as display base' }, + proba: { label: 'Probability', hint: 'Softmax class heatmap overlay' }, + predictions: { label: 'Predictions', hint: 'Conformal singleton / multi / abstain' }, + annotations: { label: 'Annotations', hint: 'Drawn shapes / scribbles' }, + manifold: { label: 'Suggest', hint: 'Manifold heatmap + markers' }, +}; + +export interface LayerVisibilityState { + groups: Record; + /** Opacity of the probability overlay (0–1). */ + probaOpacity: number; + /** Opacity of the prediction overlay (0–1). */ + predictionsOpacity: number; + /** Per-class visibility for prediction singletons. Missing → visible. */ + predictionClassVisible: Record; + /** Show conformal multi-set hatch. */ + showPredictionMulti: boolean; + /** Show conformal abstain. */ + showPredictionAbstain: boolean; + /** Annotations-layer sub-toggle: show human-drawn vs. iPred-predicted shapes + * independently (see `ShapeOrigin` in annotationStore). Both default visible. */ + annotationOriginVisible: { human: boolean; predicted: boolean }; + setAnnotationOriginVisible: (origin: 'human' | 'predicted', visible: boolean) => void; + setGroup: (id: LayerGroupId, visible: boolean) => void; + toggleGroup: (id: LayerGroupId) => void; + setProbaOpacity: (v: number) => void; + setPredictionsOpacity: (v: number) => void; + setPredictionClassVisible: (classId: number, visible: boolean) => void; + togglePredictionClass: (classId: number) => void; + setShowPredictionMulti: (v: boolean) => void; + setShowPredictionAbstain: (v: boolean) => void; + ensurePredictionClasses: (classIds: number[]) => void; +} + +const DEFAULT_GROUPS: Record = { + image: true, + denoise: true, + features: true, + proba: true, + predictions: true, + annotations: true, + manifold: true, +}; + +export const useLayerVisibilityStore = create((set, get) => ({ + groups: { ...DEFAULT_GROUPS }, + probaOpacity: 0.55, + predictionsOpacity: 0.85, + predictionClassVisible: {}, + showPredictionMulti: true, + showPredictionAbstain: true, + annotationOriginVisible: { human: true, predicted: true }, + + setAnnotationOriginVisible: (origin, visible) => + set((s) => ({ + annotationOriginVisible: { ...s.annotationOriginVisible, [origin]: visible }, + })), + + setGroup: (id, visible) => + set((s) => ({ groups: { ...s.groups, [id]: visible } })), + + toggleGroup: (id) => { + const cur = get().groups[id]; + set((s) => ({ groups: { ...s.groups, [id]: !cur } })); + }, + + setProbaOpacity: (v) => + set({ probaOpacity: Math.min(1, Math.max(0, v)) }), + + setPredictionsOpacity: (v) => + set({ predictionsOpacity: Math.min(1, Math.max(0, v)) }), + + setPredictionClassVisible: (classId, visible) => + set((s) => ({ + predictionClassVisible: { ...s.predictionClassVisible, [classId]: visible }, + })), + + togglePredictionClass: (classId) => { + const cur = get().predictionClassVisible[classId]; + const next = cur === undefined ? false : !cur; + set((s) => ({ + predictionClassVisible: { ...s.predictionClassVisible, [classId]: next }, + })); + }, + + setShowPredictionMulti: (v) => set({ showPredictionMulti: v }), + setShowPredictionAbstain: (v) => set({ showPredictionAbstain: v }), + + ensurePredictionClasses: (classIds) => + set((s) => { + const next = { ...s.predictionClassVisible }; + let changed = false; + for (const id of classIds) { + if (next[id] === undefined) { + next[id] = true; + changed = true; + } + } + return changed ? { predictionClassVisible: next } : s; + }), +})); + +/** True when a prediction class should draw (default visible). */ +export function isPredictionClassVisible( + map: Record, + classId: number, +): boolean { + return map[classId] !== false; +} + +/** True when a shape's origin (human/predicted; undefined = human) should draw. */ +export function isShapeOriginVisible( + visible: { human: boolean; predicted: boolean }, + origin: 'human' | 'predicted' | undefined, +): boolean { + return origin === 'predicted' ? visible.predicted : visible.human; +} diff --git a/frontend/src/stores/predictedRasterStore.test.ts b/frontend/src/stores/predictedRasterStore.test.ts new file mode 100644 index 0000000..7fb1020 --- /dev/null +++ b/frontend/src/stores/predictedRasterStore.test.ts @@ -0,0 +1,56 @@ +import { describe, it, expect, beforeEach } from 'vitest'; +import { usePredictedRasterStore } from './predictedRasterStore'; + +const reset = () => usePredictedRasterStore.setState({ bySource: {} }); + +describe('predictedRasterStore', () => { + beforeEach(reset); + + it('setPointers merges into an existing source without clobbering other slices', () => { + const { setPointers } = usePredictedRasterStore.getState(); + setPointers('sampleA', { '0': { runId: 'run-0', classIds: [1, 2] } }); + setPointers('sampleA', { '1': { runId: 'run-1', classIds: [1, 2] } }); + expect(usePredictedRasterStore.getState().bySource.sampleA).toEqual({ + '0': { runId: 'run-0', classIds: [1, 2] }, + '1': { runId: 'run-1', classIds: [1, 2] }, + }); + }); + + it('setPointers for the same slice overwrites the prior pointer', () => { + const { setPointers } = usePredictedRasterStore.getState(); + setPointers('sampleA', { '0': { runId: 'run-old', classIds: [1] } }); + setPointers('sampleA', { '0': { runId: 'run-new', classIds: [1, 2] } }); + expect(usePredictedRasterStore.getState().bySource.sampleA['0']).toEqual({ + runId: 'run-new', + classIds: [1, 2], + }); + }); + + it('clearSlice removes only the targeted slice', () => { + const { setPointers, clearSlice } = usePredictedRasterStore.getState(); + setPointers('sampleA', { + '0': { runId: 'run-0', classIds: [1] }, + '1': { runId: 'run-1', classIds: [1] }, + }); + clearSlice('sampleA', '0'); + const bySource = usePredictedRasterStore.getState().bySource.sampleA; + expect(bySource['0']).toBeUndefined(); + expect(bySource['1']).toEqual({ runId: 'run-1', classIds: [1] }); + }); + + it('clearSlice on an unknown sample or slice is a no-op', () => { + const { clearSlice } = usePredictedRasterStore.getState(); + expect(() => clearSlice('unknown', '0')).not.toThrow(); + expect(usePredictedRasterStore.getState().bySource).toEqual({}); + }); + + it('clearSource drops every pointer for that sample, leaving others intact', () => { + const { setPointers, clearSource } = usePredictedRasterStore.getState(); + setPointers('sampleA', { '0': { runId: 'run-0', classIds: [1] } }); + setPointers('sampleB', { '0': { runId: 'run-b0', classIds: [1] } }); + clearSource('sampleA'); + const state = usePredictedRasterStore.getState(); + expect(state.bySource.sampleA).toBeUndefined(); + expect(state.bySource.sampleB).toEqual({ '0': { runId: 'run-b0', classIds: [1] } }); + }); +}); diff --git a/frontend/src/stores/predictedRasterStore.ts b/frontend/src/stores/predictedRasterStore.ts new file mode 100644 index 0000000..8631e18 --- /dev/null +++ b/frontend/src/stores/predictedRasterStore.ts @@ -0,0 +1,79 @@ +/** + * predictedRasterStore — lightweight pointers to un-vectorized predicted + * regions from an iPred "Apply across volume" run, keyed by sample + slice. + * + * This is the direct fix for the 297MB-draft incident (139,004 shapes from + * one volume-apply commit, 98.8% predicted-origin): "Commit predicted + * shapes" used to eagerly fetch + vectorize every slice's commit.png into + * real `Shape[]` objects added to `annotationStore` — for a 690-slice + * volume that's hundreds of thousands of polygon objects landing straight + * in the autosaved draft, most of which nobody ever looks at let alone + * edits. + * + * Deliberately a SEPARATE store from `annotationStore`/`useDraftSync`: a + * pointer here is just `{runId, classIds}` (tens of bytes), and — critically + * — is never included in the autosaved draft payload, so committing however + * many slices costs nothing until a specific slice is actually vectorized + * (see `usePixelClassifier.ts`'s live-preview effect, which renders a + * pointer's commit.png directly via the existing Predictions layer, and + * AnnotatePage's "Make this slice editable" action, the only path that ever + * turns a pointer into real `Shape[]`). + * + * Not persisted across page reloads on purpose — same rationale as + * `annotationStore`'s draft being the autosave surface: an un-vectorized + * pointer is recoverable at any time by re-running "Apply across volume" or + * committing again, unlike hand-drawn shapes. + */ +import { create } from 'zustand'; + +export interface PredictedRegionPointer { + /** ipred run id — `ipredRunCommitUrl(runId)` fetches its raw label-map PNG. */ + runId: string; + /** Frontend class ids this run's commit.png pixel values are drawn from + * (matches `labelMapToPolygonShapes`'s `classIds` param exactly — pixel + * value IS the class id, not a 1-based sequential index). */ + classIds: number[]; +} + +export interface PredictedRasterState { + /** sourceKey -> sliceIndex (as string) -> pointer. */ + bySource: Record>; + /** Replace/merge pointers for many slices of one sample at once (a + * "Commit predicted shapes" click). */ + setPointers: (sourceKey: string, pointers: Record) => void; + /** Drop one slice's pointer — called once it's been vectorized into real + * `Shape[]` (or the user wants to discard the prediction for that slice). */ + clearSlice: (sourceKey: string, sliceKey: string) => void; + /** Drop every pointer for a sample — a fresh volume-apply run supersedes + * whatever was committed before. */ + clearSource: (sourceKey: string) => void; +} + +export const usePredictedRasterStore = create((set) => ({ + bySource: {}, + + setPointers: (sourceKey, pointers) => + set((s) => ({ + bySource: { + ...s.bySource, + [sourceKey]: { ...(s.bySource[sourceKey] ?? {}), ...pointers }, + }, + })), + + clearSlice: (sourceKey, sliceKey) => + set((s) => { + const existing = s.bySource[sourceKey]; + if (!existing || !(sliceKey in existing)) return s; + const next = { ...existing }; + delete next[sliceKey]; + return { bySource: { ...s.bySource, [sourceKey]: next } }; + }), + + clearSource: (sourceKey) => + set((s) => { + if (!(sourceKey in s.bySource)) return s; + const next = { ...s.bySource }; + delete next[sourceKey]; + return { bySource: next }; + }), +})); diff --git a/frontend/src/stores/referenceGuideStore.test.ts b/frontend/src/stores/referenceGuideStore.test.ts new file mode 100644 index 0000000..af18bf9 --- /dev/null +++ b/frontend/src/stores/referenceGuideStore.test.ts @@ -0,0 +1,102 @@ +import { beforeEach, describe, expect, it } from 'vitest'; +import { useReferenceGuideStore } from './referenceGuideStore'; + +const initial = useReferenceGuideStore.getState(); + +beforeEach(() => { + useReferenceGuideStore.setState(initial, true); +}); + +const cls = (label: string, overrides: Partial<{ color: string; description: string; exampleCrops: string[] }> = {}) => ({ + label, color: '#000', description: '', exampleCrops: [], ...overrides, +}); + +describe('referenceGuideStore', () => { + it('starts empty', () => { + const s = useReferenceGuideStore.getState(); + expect(s.entries).toEqual([]); + expect(s.notes).toBe(''); + expect(s.loadedFor).toBeNull(); + }); + + it('setGuide replaces the whole guide', () => { + useReferenceGuideStore.getState().setGuide([cls('Cell')], 'my notes', 'local:x.tif'); + const s = useReferenceGuideStore.getState(); + expect(s.entries).toEqual([cls('Cell')]); + expect(s.notes).toBe('my notes'); + expect(s.loadedFor).toBe('local:x.tif'); + }); + + it('addEntry appends a class', () => { + useReferenceGuideStore.getState().addEntry(cls('Pore')); + expect(useReferenceGuideStore.getState().entries).toEqual([cls('Pore')]); + }); + + it('updateEntry merges partial updates by index', () => { + useReferenceGuideStore.getState().setGuide([cls('Pore'), cls('Wall')], '', 'x'); + useReferenceGuideStore.getState().updateEntry(1, { color: '#f00' }); + expect(useReferenceGuideStore.getState().entries[1]).toEqual(cls('Wall', { color: '#f00' })); + expect(useReferenceGuideStore.getState().entries[0]).toEqual(cls('Pore')); + }); + + it('removeEntry removes by index', () => { + useReferenceGuideStore.getState().setGuide([cls('Pore'), cls('Wall')], '', 'x'); + useReferenceGuideStore.getState().removeEntry(0); + expect(useReferenceGuideStore.getState().entries).toEqual([cls('Wall')]); + }); + + it('setNotes updates notes only', () => { + useReferenceGuideStore.getState().setGuide([cls('Pore')], 'old', 'x'); + useReferenceGuideStore.getState().setNotes('new'); + expect(useReferenceGuideStore.getState().notes).toBe('new'); + expect(useReferenceGuideStore.getState().entries).toEqual([cls('Pore')]); + }); + + it('clear resets everything', () => { + useReferenceGuideStore.getState().setGuide([cls('Pore')], 'n', 'x'); + useReferenceGuideStore.getState().clear(); + const s = useReferenceGuideStore.getState(); + expect(s.entries).toEqual([]); + expect(s.notes).toBe(''); + expect(s.loadedFor).toBeNull(); + }); + + describe('applyGenerated', () => { + it('appends newly-discovered classes', () => { + useReferenceGuideStore.getState().applyGenerated([cls('Pore', { color: '#111' })]); + expect(useReferenceGuideStore.getState().entries).toEqual([cls('Pore', { color: '#111' })]); + }); + + it('preserves an existing class description, refreshing color/crops', () => { + useReferenceGuideStore.getState().setGuide( + [cls('Pore', { description: 'lead wrote this', color: '#000' })], '', 'x', + ); + useReferenceGuideStore.getState().applyGenerated([ + cls('Pore', { description: 'auto-generated', color: '#f00', exampleCrops: ['data:img1'] }), + ]); + expect(useReferenceGuideStore.getState().entries).toEqual([ + cls('Pore', { description: 'lead wrote this', color: '#f00', exampleCrops: ['data:img1'] }), + ]); + }); + + it('preserves existing entries not part of the new generation', () => { + useReferenceGuideStore.getState().setGuide([cls('Pore'), cls('Wall')], '', 'x'); + useReferenceGuideStore.getState().applyGenerated([cls('Pore')]); + const labels = useReferenceGuideStore.getState().entries.map((e) => e.label); + expect(labels).toContain('Wall'); + }); + + it('matches existing classes by label case-insensitively (trimmed)', () => { + useReferenceGuideStore.getState().setGuide( + [cls(' pore ', { description: 'kept' })], '', 'x', + ); + useReferenceGuideStore.getState().applyGenerated([cls('Pore', { description: 'new' })]); + expect(useReferenceGuideStore.getState().entries[0].description).toBe('kept'); + }); + + it('uses a default color when neither generated nor existing has one', () => { + useReferenceGuideStore.getState().applyGenerated([{ label: 'Pore', color: '', description: '', exampleCrops: [] }]); + expect(useReferenceGuideStore.getState().entries[0].color).toBe('#1f77b4'); + }); + }); +}); diff --git a/frontend/src/stores/toolStore.ts b/frontend/src/stores/toolStore.ts index e3080f8..c2217b1 100644 --- a/frontend/src/stores/toolStore.ts +++ b/frontend/src/stores/toolStore.ts @@ -3,7 +3,7 @@ */ import { create } from 'zustand'; -export type Tool = 'pan' | 'select' | 'polygon' | 'magnetic' | 'magic' | 'rectangle' | 'ellipse' | 'brush' | 'fill' | 'eraser'; +export type Tool = 'pan' | 'select' | 'polygon' | 'magnetic' | 'magic' | 'rectangle' | 'ellipse' | 'brush' | 'threshold' | 'sampler' | 'fill' | 'eraser'; export type MagicMode = 'contiguous' | 'global'; /** Magic-selection engine: SAM (learned object prior) or the classic wand. */ @@ -21,6 +21,13 @@ export interface ToolState { fillThreshold: number; /** Selected shape ids (multi-select via marquee/shift-click). */ selectedShapeIds: string[]; + /** Threshold brush: paint only where the displayed intensity is in [lo,hi] (0–255). */ + thresholdLo: number; + thresholdHi: number; + /** Show a red overlay of every in-band pixel on the slice (ImageJ-style). */ + thresholdOverlay: boolean; + /** Band width applied when Shift-clicking to sample a pixel's intensity. */ + thresholdSampleWidth: number; /** Magic-wand: similarity tolerance (0–1) and selection mode + denoise. */ magicTolerance: number; magicMode: MagicMode; @@ -59,6 +66,9 @@ export interface ToolState { setBrushSize: (size: number) => void; setFillOpacity: (opacity: number) => void; setFillThreshold: (t: number) => void; + setThresholdBand: (lo: number, hi: number) => void; + setThresholdOverlay: (v: boolean) => void; + setThresholdSampleWidth: (w: number) => void; /** Convenience single-select (clears to [] when null). */ setSelectedShapeId: (id: string | null) => void; setSelectedShapeIds: (ids: string[]) => void; @@ -84,6 +94,10 @@ export const useToolStore = create((set) => ({ brushSize: 10, fillOpacity: 0.5, fillThreshold: 0.1, + thresholdLo: 0, + thresholdHi: 255, + thresholdOverlay: true, + thresholdSampleWidth: 24, selectedShapeIds: [], magicTolerance: 0.08, magicMode: 'contiguous', @@ -108,6 +122,15 @@ export const useToolStore = create((set) => ({ setFillOpacity: (fillOpacity) => set({ fillOpacity }), /** Sets the Fill (paint-bucket) similarity threshold (0–1). */ setFillThreshold: (fillThreshold) => set({ fillThreshold }), + /** Sets the threshold-brush intensity band, clamped to 0–255 and kept ordered. */ + setThresholdBand: (lo, hi) => set({ + thresholdLo: Math.max(0, Math.min(255, Math.round(Math.min(lo, hi)))), + thresholdHi: Math.max(0, Math.min(255, Math.round(Math.max(lo, hi)))), + }), + /** Toggles the red in-band overlay shown while the threshold brush is active. */ + setThresholdOverlay: (thresholdOverlay) => set({ thresholdOverlay }), + /** Sets the band width used when Shift-clicking to sample an intensity. */ + setThresholdSampleWidth: (thresholdSampleWidth) => set({ thresholdSampleWidth }), /** Single-selects a shape, or clears the selection when null. */ setSelectedShapeId: (id) => set({ selectedShapeIds: id ? [id] : [] }), /** Replaces the multi-selection with the given shape ids. */ diff --git a/frontend/src/types/webgpu.d.ts b/frontend/src/types/webgpu.d.ts new file mode 100644 index 0000000..d3a2d6e --- /dev/null +++ b/frontend/src/types/webgpu.d.ts @@ -0,0 +1,8 @@ +// WebGPU global types for the vendored volume renderer. +// +// `@webgpu/types` is not an `@types/*` package, so TypeScript does not pick it up +// automatically. Referencing it here rather than adding a `types` array to +// tsconfig.app.json is deliberate: setting `types` switches off the "include every +// @types package" default, which would silently drop the globals vitest and node +// contribute elsewhere in the app. +/// diff --git a/frontend/tsconfig.app.json b/frontend/tsconfig.app.json index 223ceff..9c49f02 100644 --- a/frontend/tsconfig.app.json +++ b/frontend/tsconfig.app.json @@ -3,7 +3,9 @@ "tsBuildInfoFile": "./node_modules/.tmp/tsconfig.app.tsbuildinfo", "target": "ES2020", "useDefineForClassFields": true, - "lib": ["ES2020", "DOM", "DOM.Iterable"], + // ES2022.Error is declarations-only (it adds `Error`'s `cause` option, which + // the vendored renderer uses). It does not move `target` off ES2020. + "lib": ["ES2020", "ES2022.Error", "DOM", "DOM.Iterable"], "module": "ESNext", "skipLibCheck": true, "moduleResolution": "bundler", @@ -17,8 +19,37 @@ "noUnusedParameters": false, "noFallthroughCasesInSwitch": true, "paths": { - "@/*": ["./src/*"] + "@/*": ["./src/*"], + // The WebGPU volume renderer, pinned as a git submodule (see .gitmodules). + // `@zarrviewer/*` is OUR entry point into it; the `@zarr-viewer/*` and + // `@prism/*` aliases below are ITS OWN internal ones, replicated here + // because tsc resolves imports against this config, not the submodule's. + "@zarrviewer/*": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/*"], + "@zarr-viewer/core": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/core/index.ts"], + "@zarr-viewer/math": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/math/index.ts"], + "@zarr-viewer/scene": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/scene/index.ts"], + "@zarr-viewer/controls": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/controls/index.ts"], + "@zarr-viewer/io": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/io/index.ts"], + "@zarr-viewer/render": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/render/index.ts"], + "@zarr-viewer/fx": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/fx/src/index.ts"], + "@prism/core": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/core/index.ts"], + "@prism/math": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/math/index.ts"], + "@prism/render": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/render/index.ts"], + "@prism/fx": ["./vendor/view_tomography_recon_app/src/zarr-viewer/src/fx/src/index.ts"] } }, - "include": ["src"] + // The submodule is typechecked here rather than as a separate project so a bad + // pointer bump fails `tsc -b` loudly. Its own tsconfig is *stricter* than ours + // (`noUncheckedIndexedAccess`, `verbatimModuleSyntax`), so code written under it + // compiles cleanly under ours. If it ever does not, that is an upstream fix on + // the viewer repo's own branch — do not edit vendored sources to appease tsc. + "include": ["src", "vendor/view_tomography_recon_app/src/zarr-viewer/src"], + "exclude": [ + "vendor/view_tomography_recon_app/src/zarr-viewer/src/fx/test", + // Upstream test fixture predates `RenderProvenance.extendedPreIntegration` + // (added alongside the mip-pyramid/half-res-lighting work) and was never + // updated — a test-only type error with no effect on the exported API we + // actually consume. Excluded rather than patched, per the note above. + "vendor/view_tomography_recon_app/src/zarr-viewer/src/render/accel/test/accel.spec.ts" + ] } diff --git a/frontend/vendor/view_tomography_recon_app b/frontend/vendor/view_tomography_recon_app new file mode 160000 index 0000000..46c1e2a --- /dev/null +++ b/frontend/vendor/view_tomography_recon_app @@ -0,0 +1 @@ +Subproject commit 46c1e2acae30e3abb62402793f2d236330ba5383 diff --git a/frontend/vite.config.ts b/frontend/vite.config.ts index 2c13928..dbeda15 100644 --- a/frontend/vite.config.ts +++ b/frontend/vite.config.ts @@ -15,8 +15,32 @@ function gitCommit(): string { catch { return 'dev'; } } +/** + * VITE_BASE_PATH must be a bare path (e.g. `/bl832/seg_studio/`), never a full + * URL with a scheme/host — the reverse proxy owns the hostname, and can + * differ per environment (staging vs. production) without ever touching this + * build-arg. Fail the build loudly rather than silently baking in a wrong + * `base` that would only surface as a confusing runtime 404. + */ +function basePath(): string { + const raw = process.env.VITE_BASE_PATH; + if (!raw) return '/'; + if (/^[a-z][a-z0-9+.-]*:\/\//i.test(raw) || raw.startsWith('//')) { + throw new Error( + `VITE_BASE_PATH must be a path only (e.g. "/bl832/seg_studio/"), not a full URL — got "${raw}". ` + + 'The reverse proxy owns the hostname; baking one in here would break as soon as it differs (e.g. staging vs. production) or changes.' + ); + } + return raw; +} + export default defineConfig({ plugins: [react(), tsconfigPaths()], + // Read at build time so the same Dockerfile stage can produce either a + // root-hosted image (unset, defaults to '/') or a subpath-hosted one (e.g. + // `/bl832/seg_studio/` behind a stripping reverse proxy) — see main.tsx's + // BrowserRouter basename and config.ts's API_BASE for the other two halves. + base: basePath(), define: { __APP_VERSION__: JSON.stringify(appVersion()), __GIT_COMMIT__: JSON.stringify(gitCommit()), @@ -25,6 +49,10 @@ export default defineConfig({ // (konva, polygon-clipping) land in their own chunks that load with the lazy // Annotate page rather than bloating the initial /connect entry. build: { + // Explicit, not left to Vite's implicit production-mode default — a stray + // future `--mode development` in a build script must not silently ship + // unminified bundles. + minify: 'esbuild', rollupOptions: { output: { manualChunks(id: string) { @@ -58,5 +86,14 @@ export default defineConfig({ globals: true, environment: 'jsdom', setupFiles: ['./src/test/setup.ts'], + // vitest runs in Node, where bare `import 'konva'` resolves via the package's + // "main" field (a canvas-backed Node build requiring the native `canvas` + // module, which isn't installed) instead of "browser" (the jsdom-friendly + // build bundlers use). Force the browser build under test. + alias: [{ find: /^konva$/, replacement: 'konva/lib/index.js' }], + // Scoped to src/ so the vendored renderer's own suite (which needs WebGPU and + // its own runner config) isn't swept into ours by the default glob. Upstream + // tests are upstream's to run. + include: ['src/**/*.{test,spec}.{ts,tsx}'], }, }); diff --git a/ipred/.gitignore b/ipred/.gitignore new file mode 100644 index 0000000..b900aeb --- /dev/null +++ b/ipred/.gitignore @@ -0,0 +1,2 @@ +models/*.pth +models/*.onnx diff --git a/ipred/README.md b/ipred/README.md new file mode 100644 index 0000000..fe0b22d --- /dev/null +++ b/ipred/README.md @@ -0,0 +1,436 @@ +# ipred — architecture and operations guide + +`ipred` is the **iterative prediction** compute service for SAM3 Annotation Studio. It owns feature compositions, preprocessing / feature banks, CatBoost training, Mondrian conformal prediction, and Suggest Labels (manifold). The Vite UI never talks to it directly: the Annotate FastAPI backend (`:8002`) proxies all GUI calls. + +This document describes how the service is laid out, where compute runs, how data reaches it, and how compositions / models are discovered. + +--- + +## 1. Big picture + +```text +┌──────────────┐ HTTP ┌────────────────────┐ HTTP ┌─────────────────┐ +│ Vite UI │ ───────▶ │ Annotate backend │ ───────▶ │ ipred FastAPI │ +│ :5173 │ /api/ipred│ 127.0.0.1:8002 │ IPRED_URL │ 127.0.0.1:8003 │ +└──────────────┘ └────────────────────┘ └────────┬────────┘ + │ + ┌────────────────────────────────────────────────┤ + │ │ + ▼ ▼ + LOCAL_DATA_ROOT/ipred/ Data plane: + catalog.db, projects/*/ • local files under LOCAL_DATA_ROOT + • Tiled over HTTP (:8010) + • optional uploaded array blobs +``` + +| Role | Process | Bind | Responsibility | +|------|---------|------|----------------| +| UI | Vite | `127.0.0.1:5173` | Composition editor, Draw / Train, Layers | +| Annotate | FastAPI | `127.0.0.1:8002` | Tiled/browse/auth; **proxies** `/api/ipred/*` | +| **ipred** | FastAPI (`ipred.api:app`) | `127.0.0.1:8003` | Features, train, infer, manifold | +| Tiled | catalog + API | `127.0.0.1:8010` | Array storage / browse | + +All local binds stay on **loopback** (`127.0.0.1`), never `0.0.0.0`. + +--- + +## 2. Where compute lives + +### Process + +- Package: `ipred/` (editable install into the shared repo `.venv`). +- Entry module: `ipred.api:app` in `ipred/src/ipred/api.py`. +- Launch (via `start_all.sh`): + + ```bash + cd ipred + PYTHONPATH=src uvicorn ipred.api:app --host 127.0.0.1 --port "${IPRED_PORT:-8003}" + ``` + +- The `ipred` console script (`ipred.cli:main`) is an **offline CLI** against the same SQLite/FS layout; it does **not** start uvicorn. + +### Engine filesystem root + +Everything durable for a running engine lives under: + +```text +$LOCAL_DATA_ROOT/ipred/ # default LOCAL_DATA_ROOT=~/data + catalog.db # projects, sessions, banks, models, runs + projects// + features// # FeatureBank blobs + models// # CatBoost + conformal calibration + runs// # proba + conformal maps + arrays// # optional uploaded tensors +``` + +Implemented in `ipred/src/ipred/paths.py`: + +| Helper | Path | +|--------|------| +| `local_data_root()` | `$LOCAL_DATA_ROOT` (default `~/data`) | +| `engine_root()` | `$LOCAL_DATA_ROOT/ipred` | +| `catalog_db_path()` | `…/ipred/catalog.db` | +| `project_blob_dir(id)` | `…/ipred/projects//` | +| `feature_models_root()` | `$LOCAL_DATA_ROOT/.feature_models` (compositions + legacy setups) | + +**Compute affinity:** all CPU/GPU work for encode / CatBoost / PCA / manifold runs **inside the ipred process**. The Annotate backend only forwards HTTP. + +### Startup order (`start_all.sh`) + +1. Shared `.venv` (+ `uv pip install -e ipred`) +2. Load `backend/.env` +3. **Tiled** → wait ready +4. **ipred** → wait `GET http://127.0.0.1:8003/health` +5. **Annotate backend** (receives `IPRED_URL`) → wait `/health` +6. **Frontend** + +On API lifespan, ipred opens the catalog and seeds default feature setups + compositions. + +--- + +## 3. How the UI reaches compute (proxy boundary) + +### Rule + +The browser **never** calls `:8003`. It calls Annotate: + +```text +frontend → ${API_BASE}/api/ipred/... → backend/ipred_client.py → ${IPRED_URL}/... +``` + +- Client library: `frontend/src/lib/ipredApi.ts` +- Server client: `backend/ipred_client.py` (httpx; **no** Python import of the `ipred` package) +- Env: `IPRED_URL` (default `http://127.0.0.1:8003`; legacy alias `CLF_ENGINE_URL`) + +If ipred is down, Annotate returns **503** (`ipred unreachable at …`). + +### Representative route mirror + +| UI / Annotate | ipred | +|---------------|-------| +| `GET /api/ipred/health` | `GET /health` | +| `POST /api/ipred/sessions` | `POST /sessions` | +| `GET /api/ipred/modules` | `GET /modules` | +| `GET/POST /api/ipred/compositions` | `/compositions` | +| `POST /api/ipred/preprocess` | `POST /preprocess` | +| `POST /api/ipred/train` / `infer` / `rethreshold` | same | +| `GET /api/ipred/runs/{id}/proba/{i}.png` | same | +| `POST /api/ipred/runs/{id}/threshold-class` | same | +| `POST /api/ipred/manifold/sample` | same | + +--- + +## 4. Sessions, projects, and identity + +A **project** is a content-addressed identity of *which array source* you are working on. A **session** is a UUID workspace pointer into that project (current feature / model / run ids). + +`POST /sessions` body: + +```json +{ + "kind": "local" | "tiled", + "source": "", + "server_uri": "", + "root": "" +} +``` + +- `project_id` = first 16 hex of SHA1 over `{kind, source, server_uri, root}` +- `session_id` = fresh `uuid4` each open (many sessions can share one project) + +Stored in SQLite `catalog.db`. Blob trees hang off `projects//`. + +Typical open path: Browse / Connect → `useOpenInAnnotate` → `openIpredSession` → `connectionStore.ipredSessionId` + `ipredProjectId`. + +--- + +## 5. Data plane — how pixels reach the engine + +Three mutually exclusive ways to get a 2‑D slice into preprocess: + +### A. Local files (default for local Browse) + +- Read: `array_source.read_slice(kind="local", …)` +- Path: `(root or LOCAL_DATA_ROOT) / source`, with a containment check +- Formats: TIFF (tifffile) or PIL-readable; then z/slice index + +**Requirement:** the ipred host must see that filesystem (same machine, NFS, etc.). + +### B. Tiled over HTTP + +- Read: `tiled.client.from_uri(server_uri or TILED_URI)` + optional `TILED_API_KEY` +- Walk `source` path; pull the slice array over the network + +**Requirement:** network reachability to Tiled. Compute does **not** need the Tiled data directory on disk. + +### C. Uploaded array blobs (compute ≠ data host) + +For when Annotate and ipred do not share `LOCAL_DATA_ROOT`: + +1. `POST /sessions/{id}/arrays` with base64 float32 payload → content-addressed + `$LOCAL_DATA_ROOT/ipred/projects//arrays//{array.npy, meta.json}` +2. `POST /preprocess` with `array_ref` loads that blob + +When `array_ref` is set, **preprocess cache hits are skipped** (identity is upload content + composition, not just slice index on a stable URI). + +### Environment checklist + +| Variable | Purpose | +|----------|---------| +| `LOCAL_DATA_ROOT` | Engine + local file root | +| `TILED_URI` | Default tiled server (`http://127.0.0.1:8010`) | +| `TILED_API_KEY` | Optional tiled auth **only on server side** | +| `IPRED_URL` | Annotate → ipred base URL | + +--- + +## 6. Discovering modules and compositions + +### Module catalog (runtime capability discovery) + +`GET /modules` → `ipred.modules.list_module_catalog()`. + +Registered producers (`ipred/src/ipred/modules/`): + +| Module id | Typical runtime | Output | +|-----------|-----------------|--------| +| `skimage_multiscale` | numpy / skimage | Multi-channel float stack | +| `clahe` | numpy | Optional CLAHE grayscale channel | +| `slimsam` | **onnx** | Dense embedding | +| `tomojepa` | **onnx** if matching `.onnx` present, else **torch** | Dense embedding | +| `pca` | numpy / sklearn | Embedding → PCA channels | + +Each entry reports `id`, `params_schema`, `runtime`, `ready`, and whether it accepts `input_from` / produces channels vs embeddings. + +UI: left column of **CompositionPanel** on the Ipred page. + +### Composition documents (feature graphs) + +A composition is an ordered graph of module instances. The feature bank is the **channel-wise concatenation** of nodes listed in `outputs`. + +Stored under: + +```text +$LOCAL_DATA_ROOT/.feature_models/_compositions//meta.json +``` + +Schema (v1): + +```json +{ + "id": "comp-skimage-mark11", + "kind": "composition", + "name": "…", + "nodes": [ + {"id": "n1", "module": "skimage_multiscale", "params": {…}}, + {"id": "n2", "module": "clahe", "params": {…}}, + {"id": "n3", "module": "tomojepa", "params": {"weights_id": "mark11", "input_size": 512}, "input_from": "n2"}, + {"id": "n4", "module": "pca", "params": {"dims": 64}, "input_from": "n3"} + ], + "outputs": ["n1", "n4"] +} +``` + +- `input_from` wires grayscale (or prior emb) into the next module — this is how “CLAHE → Mark11” is made explicit. +- `content_hash` = short SHA1 of `{kind, nodes, outputs}` — part of the preprocess cache key. +- API: `GET/POST /compositions`, `GET /compositions/{id}`, `POST /compositions/preview` (concat channel labels without running encode). + +**Builtins** (auto-seeded): +`comp-skimage`, `comp-skimage-slimsam`, `comp-slimsam-clahe`, `comp-skimage-mark25`, `comp-mark25-clahe`, `comp-skimage-mark11`, `comp-mark11-clahe`. + +Default preference in the UI (`connectionStore`): + +- `preferredCompositionId = "comp-skimage-slimsam"` +- also mirrored to legacy `preferredFeatureSetupId` for one release + +### Legacy Feature Setups (migration) + +Older “procedure × encoder” combos still exist as disk shelves under `$LOCAL_DATA_ROOT/.feature_models//` (`feature_setups.py`). +`preprocess` accepts either `composition_id` or `feature_setup_id`; legacy ids map via `_LEGACY_SETUP_MAP` (e.g. `default-skimage-slimsam` → `comp-skimage-slimsam`). + +Prefer compositions for new work; setups are a compatibility layer. + +--- + +## 7. Preprocess → FeatureBank + +`POST /preprocess`: + +```json +{ + "session_id": "…", + "composition_id": "comp-skimage-slimsam", + "slice_index": 0, + "array_ref": null +} +``` + +Pipeline: + +1. Resolve composition (or legacy setup → composition). +2. Load slice (local / tiled / array_ref). +3. Cache lookup on `(project_id, setup_id, content_hash, slice_index)` when `array_ref` is null. +4. Else `compose_run.run_composition(doc)` → concat outputs. +5. Write FeatureBank under `projects//features//`: + +| Artifact | Meaning | +|----------|---------| +| `float_stack.npy` | Training features (float16 HxWxC) | +| `uint8_stack.npy` | Display-oriented stack | +| `labels.json` | Channel names (concat order) | +| `channels/NNNN.png` | Browseable per-channel previews | +| `sam_emb.npy` + `sam_meta.json` | Dense emb when an encoder ran (unless fully baked via PCA) | + +UI hook: `useFeatureChannels` (Preprocess stage) calls proxy preprocess and exposes channel navigation. + +--- + +## 8. Train, predict, probability maps + +### Trainer discovery + +`GET /trainers` → plugin list. Today only **`catboost`** (`ipred.trainers.catboost_trainer.CatBoostTrainer`). +UI preference: `connectionStore.preferredTrainerId` (+ depth / trees / LR). + +### Train (`POST /train`) + +- Rasterize sparse Draw shapes → label map. +- Stratified train / calibration split. +- Fit CatBoost → write: + + ```text + projects//models// + model.cbm + meta.json + cal_scores.json + (optional sam_pca.npz) + ``` + +- Register row in catalog `models`; set session `current_model_id`. + +### Infer (`POST /infer`) + +- Full-image `predict_proba` → `proba.npy`. +- Mondrian split-conformal (`conformal.py`) using calibration scores + α → membership / commit / status. +- Run directory: + + ```text + projects//runs// + proba.npy + commit.npy / status.npy / membership.npy + commit.png / status.png + meta.json + ``` + +Status codes: abstain `0`, singleton `1`, multi `2`. + +### Probability overlays + +- `GET /runs/{run_id}/proba/{class_index}.png` — grayscale softmax channel. +- Frontend colorizes with **viridis**, clipping below the per-class threshold (transparent). +- `POST /runs/{run_id}/threshold-class` `{class_id, threshold}` → dense label map for mask-set cache. + +Layers panel keeps Image / Features / Probability / Predictions / Annotations independently toggleable. + +### Two “model” concepts (do not confuse) + +| Store | Path | Owner | Purpose | +|-------|------|-------|---------| +| **ipred project models** | `$LOCAL_DATA_ROOT/ipred/projects//models/` | ipred | Session train artifacts | +| **Annotate clf shelf** | `$LOCAL_DATA_ROOT/.clf_models/` | Annotate (`backend/clf_shelf.py`) | Named reusable classifiers from Connect/Browse scaffold UI | + +The shelf is separately flagged `# REMOVE THIS AND USE YOUR OWN STUFF` in the UI. + +--- + +## 9. Encoder backends (ONNX vs torch) + +| Encoder | Preferred | Fallback / notes | +|---------|-----------|------------------| +| **SlimSAM** | ONNX (`FEATURE_ENCODER_ONNX` or repo SlimSAM `vision_encoder.onnx`) | No torch path | +| **TomoJEPA Mark25/11** | ONNX if `TOMOJEPA[_11]_ONNX` or `ipred/models/tomojepa{25,11}.onnx` **and** spatial size matches requested `input_size` | Else torch `.pth` (`TOMOJEPA[_11]_WEIGHTS` or `ipred/models/*.pth`; needs `ipred[torch]`) | + +Weights / ONNX are **gitignored**. Export: + +```bash +python -m ipred.scripts.export_tomojepa_onnx --input-size 512 +``` + +`GET /modules` reports `runtime` and `ready` so the Composition UI can grey out broken modules. + +--- + +## 10. Manifold / Suggest Labels + +Lightweight placement suggestions on an existing FeatureBank: + +1. `POST /manifold/sample` — PCA-whitened variance windows + exclusion; optional ROI from selected shapes. +2. `GET /manifold/{sample_id}/heatmap.png` — residual interestingness. + +Implementation: `manifold.py` + `manifold_jobs.py`; results TTL-cached in process. Wired from Preprocess via `ManifoldSuggestPanel` / `useFeatureManifold` (still through `/api/ipred/…`). + +--- + +## 11. Frontend map + +| Surface | Path | Talks to | +|---------|------|----------| +| Ipred hub page | `frontend/src/app/pages/IpredPage.tsx` | Composition + trainer prefs | +| Composition window | `frontend/src/components/CompositionPanel/` | `/modules`, `/compositions` | +| Prefer composition | `connectionStore.preferredCompositionId` | Used by preprocess / train hooks | +| Preprocess / Draw / Train | `AnnotateWorkspace` stages | features, classifier, layers | +| Layers | `LayersPanel` + `layerVisibilityStore` | Canvas only (no ipred) | +| API wrapper | `frontend/src/lib/ipredApi.ts` | Annotate proxy only | + +--- + +## 12. Package map (`ipred/src/ipred/`) + +| Module | Responsibility | +|--------|----------------| +| `api.py` | FastAPI routes | +| `cli.py` | Offline JSON CLI | +| `catalog.py` | SQLite schema + CRUD | +| `paths.py` | Roots under `LOCAL_DATA_ROOT` | +| `array_source.py` | local / tiled slice load | +| `array_blobs.py` | uploaded content-addressed arrays | +| `compositions.py` | Composition docs + migration | +| `compose_run.py` | Execute module graph | +| `modules/*` | Feature module registry | +| `preprocess.py` | Cache + FeatureBank write | +| `feature_setups.py` | Legacy setup shelf | +| `features.py` | Skimage / PNG helpers | +| `sam_embed.py` / `tomojepa_*.py` | Encoders | +| `train_infer.py` | Train / infer / rethreshold / proba | +| `trainers/*` | Trainer plugins | +| `conformal.py` | Mondrian conformal maps | +| `labels.py` | Shapes → label raster | +| `manifold*.py` | Suggest Labels | +| `cache.py` | In-memory TTL (manifold) | + +Tests live in `ipred/tests/` (pytest). Optional deps: base wheel includes **onnxruntime**; torch/timm via `ipred[torch]`. + +--- + +## 13. Mental model (one paragraph) + +**ipred is a loopback compute worker** whose durable state sits under `$LOCAL_DATA_ROOT/ipred`. The UI discovers **modules** (`GET /modules`) and builds **compositions** (graphs). Preprocess runs that graph against a session’s data identity (local path, Tiled URI, or uploaded blob) and caches a **FeatureBank**. Train writes a **project model**; infer writes a **run** with softmax + conformal maps. Annotate is only a security/data façade: it never exposes Tiled keys or `:8003` to the browser, and optional array upload lets compute run without sharing the data filesystem. + +--- + +## 14. Quick ops + +```bash +# health +curl -s http://127.0.0.1:8003/health + +# modules +curl -s http://127.0.0.1:8003/modules | python -m json.tool + +# full stack +./start_all.sh +# ipred: http://127.0.0.1:8003 +# annotate: http://127.0.0.1:8002 (proxies /api/ipred) +``` + +Set secrets and URLs in `backend/.env` (see `backend/.env.example`): `LOCAL_DATA_ROOT`, `IPRED_URL`, `TILED_*`, `TOMOJEPA_*`, `TOMOJEPA_ONNX`, etc. diff --git a/ipred/models/README.md b/ipred/models/README.md new file mode 100644 index 0000000..8b81c25 --- /dev/null +++ b/ipred/models/README.md @@ -0,0 +1,35 @@ +# ipred/models + +Local model weights for iPred's feature modules. Gitignored — nothing here ships +with the repo. + +## SlimSAM (ONNX) + +Auto-vendored by the Magic tool's fetch script — no separate step needed here. +`slimsam_mod.py` resolves it from `frontend/public/models/slimsam-77-uniform/onnx/vision_encoder.onnx`: + +``` +node frontend/scripts/fetch-sam-model.mjs +``` + +## TomoJEPA (Mark25 / Mark11) + +**Not fetched automatically — no public source exists.** These are private +trained checkpoints (~365 MB each), not published on HuggingFace or anywhere +else; `github.com/phzwart/tomojepa` is the training toolkit only and ships no +checkpoint files. + +To enable TomoJEPA-based compositions, place the checkpoint(s) here: + +``` +ipred/models/tomojepa25.pth # Mark25 +ipred/models/tomojepa11.pth # Mark11 +``` + +or point at them via `TOMOJEPA_WEIGHTS` / `TOMOJEPA11_WEIGHTS` env vars. +ONNX exports (`tomojepa25.onnx` / `tomojepa11.onnx`, or `TOMOJEPA_ONNX` / +`TOMOJEPA11_ONNX`) are preferred when present — see +`python -m ipred.scripts.export_tomojepa_onnx`. + +Until a checkpoint is present, `GET /modules` reports `tomojepa.ready == false` +and the frontend greys out compositions that depend on it. diff --git a/ipred/pyproject.toml b/ipred/pyproject.toml new file mode 100644 index 0000000..f452ed6 --- /dev/null +++ b/ipred/pyproject.toml @@ -0,0 +1,38 @@ +[build-system] +requires = ["setuptools>=68"] +build-backend = "setuptools.build_meta" + +[project] +name = "ipred" +version = "0.1.0" +description = "Generic iterative prediction backend (features, train plugins, conformal products)" +requires-python = ">=3.11" +dependencies = [ + "fastapi>=0.115", + "uvicorn[standard]>=0.30", + "numpy>=1.26", + "pillow>=10.3", + "python-dotenv>=1.0", + "scikit-image>=0.22", + "scikit-learn>=1.4", + "catboost>=1.2", + "onnxruntime>=1.17", + "tifffile>=2024.0", + "httpx>=0.27", + "tiled[client]>=0.1", +] + +[project.optional-dependencies] +test = ["pytest>=8", "pytest-asyncio>=0.23", "httpx>=0.27"] +torch = ["torch>=2.2", "timm>=1.0", "onnx>=1.16"] + +[project.scripts] +ipred = "ipred.cli:main" + +[tool.setuptools.packages.find] +where = ["src"] + +[tool.pytest.ini_options] +asyncio_mode = "auto" +testpaths = ["tests"] +pythonpath = ["src"] diff --git a/ipred/src/ipred/__init__.py b/ipred/src/ipred/__init__.py new file mode 100644 index 0000000..8fe81e6 --- /dev/null +++ b/ipred/src/ipred/__init__.py @@ -0,0 +1,3 @@ +"""ipred — generic iterative prediction backend.""" + +__version__ = "0.1.0" diff --git a/ipred/src/ipred/api.py b/ipred/src/ipred/api.py new file mode 100644 index 0000000..785cf0c --- /dev/null +++ b/ipred/src/ipred/api.py @@ -0,0 +1,560 @@ +"""FastAPI app for the ipred iterative prediction backend.""" + +from __future__ import annotations + +import logging +import shutil +from contextlib import asynccontextmanager +from pathlib import Path +from typing import Any, Optional + +from fastapi import FastAPI, HTTPException +from fastapi.responses import FileResponse, Response +from pydantic import BaseModel, Field + +from ipred import compositions, feature_setups, manifold_jobs, preprocess, train_infer +from ipred.catalog import Catalog +from ipred.modules import list_module_catalog +from ipred.trainers import list_trainers + +logger = logging.getLogger(__name__) + +_catalog: Catalog | None = None + + +def get_catalog() -> Catalog: + """Return process-wide catalog.""" + global _catalog + if _catalog is None: + _catalog = Catalog() + feature_setups.ensure_default_setups(_catalog) + compositions.ensure_default_compositions(_catalog) + return _catalog + + +@asynccontextmanager +async def lifespan(_app: FastAPI): + """Ensure defaults on startup.""" + get_catalog() + yield + + +app = FastAPI(title="ipred", version="0.1.0", lifespan=lifespan) + + +class ProjectIdentity(BaseModel): + """Source identity for a project / session.""" + + kind: str + source: str + server_uri: Optional[str] = None + root: Optional[str] = None + + +class PreprocessBody(BaseModel): + """Featurize request.""" + + session_id: str + feature_setup_id: Optional[str] = None + composition_id: Optional[str] = None + slice_index: int = 0 + array_ref: Optional[str] = None + + +class CompositionUpsertBody(BaseModel): + """Create/update a composition document.""" + + name: str + nodes: list[dict[str, Any]] + outputs: list[str] + composition_id: Optional[str] = None + builtin: bool = False + + +class ArrayUploadBody(BaseModel): + """Upload a float array (base64 raw float32) for remote preprocess.""" + + session_id: str + shape: list[int] + dtype: str = "float32" + data_b64: str + array_ref: Optional[str] = None + + +class TrainBody(BaseModel): + """Train request.""" + + session_id: str + shapes: list[dict[str, Any]] + feature_id: Optional[str] = None + trainer_id: str = "catboost" + config: dict[str, Any] = Field(default_factory=dict) + + +class TrainMultiBody(BaseModel): + """Multi-slice train request — pools labeled pixels across slices.""" + + session_id: str + slices: dict[str, list[dict[str, Any]]] + feature_ids: dict[str, str] + trainer_id: str = "catboost" + config: dict[str, Any] = Field(default_factory=dict) + + +class InferBody(BaseModel): + """Infer request.""" + + session_id: str + model_id: Optional[str] = None + feature_id: Optional[str] = None + alpha: float = 0.05 + store_probabilities: bool = True + + +class RethresholdBody(BaseModel): + """Rethreshold request.""" + + session_id: str + alpha: float + run_id: Optional[str] = None + + +class SetupUpsertBody(BaseModel): + """Create/update a Feature Setup.""" + + name: str + kind: str + procedure_id: Optional[str] = None + params: Optional[dict[str, Any]] = None + encoder_setup_id: Optional[str] = None + weights_path: Optional[str] = None + weights_format: Optional[str] = None + inference: Optional[dict[str, Any]] = None + setup_id: Optional[str] = None + + +class ManifoldSampleBody(BaseModel): + """Suggest Labels sample request.""" + + feature_id: str + k: int = 24 + box_size: Optional[int] = None + stride: Optional[int] = None + pca_dims: int = 16 + shapes: Optional[list[dict[str, Any]]] = None + + +@app.get("/health") +def health() -> dict[str, str]: + """Liveness probe.""" + return {"status": "ok", "service": "ipred"} + + +@app.post("/sessions") +def open_session(body: ProjectIdentity) -> dict[str, Any]: + """Open a session for a project (creates project if needed).""" + cat = get_catalog() + session = cat.open_session( + kind=body.kind, + source=body.source, + server_uri=body.server_uri, + root=body.root, + ) + return { + "session_id": session.session_id, + "project_id": session.project_id, + "current_feature_id": session.current_feature_id, + "current_model_id": session.current_model_id, + "current_run_id": session.current_run_id, + } + + +@app.get("/sessions/{session_id}") +def get_session(session_id: str) -> dict[str, Any]: + """Fetch session state.""" + session = get_catalog().get_session(session_id) + if session is None: + raise HTTPException(404, "session not found") + project = get_catalog().get_project(session.project_id) + return { + "session_id": session.session_id, + "project_id": session.project_id, + "project": { + "kind": project.kind if project else None, + "source": project.source if project else None, + "server_uri": project.server_uri if project else None, + "root": project.root if project else None, + }, + "current_feature_id": session.current_feature_id, + "current_model_id": session.current_model_id, + "current_run_id": session.current_run_id, + } + + +@app.get("/setups") +def api_list_setups() -> dict[str, Any]: + """List Feature Setups (legacy) + ensure compositions exist.""" + feature_setups.ensure_default_setups(get_catalog()) + compositions.ensure_default_compositions(get_catalog()) + return {"setups": feature_setups.list_setups()} + + +@app.get("/modules") +def api_list_modules() -> dict[str, Any]: + """List composable feature modules.""" + return {"modules": list_module_catalog()} + + +@app.get("/compositions") +def api_list_compositions() -> dict[str, Any]: + """List composition documents.""" + compositions.ensure_default_compositions(get_catalog()) + return {"compositions": compositions.list_compositions()} + + +@app.get("/compositions/{composition_id}") +def api_get_composition(composition_id: str) -> dict[str, Any]: + """Get one composition (+ concat preview labels).""" + try: + doc = compositions.resolve_composition(composition_id) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + doc = dict(doc) + doc["preview_labels"] = compositions.preview_concat_labels(doc) + return doc + + +@app.post("/compositions") +def api_upsert_composition(body: CompositionUpsertBody) -> dict[str, Any]: + """Create or update a composition.""" + try: + return compositions.save_composition( + name=body.name, + nodes=body.nodes, + outputs=body.outputs, + composition_id=body.composition_id, + builtin=body.builtin, + catalog=get_catalog(), + ) + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + + +@app.post("/compositions/preview") +def api_preview_composition(body: CompositionUpsertBody) -> dict[str, Any]: + """Return concat labels without saving.""" + try: + labels = compositions.preview_concat_labels( + {"nodes": body.nodes, "outputs": body.outputs} + ) + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + return {"preview_labels": labels} + + +@app.post("/sessions/{session_id}/arrays") +def api_upload_array(session_id: str, body: ArrayUploadBody) -> dict[str, Any]: + """Upload a content-addressed array blob for remote preprocess.""" + import base64 + + import numpy as np + + from ipred import array_blobs + + session = get_catalog().get_session(session_id) + if session is None: + raise HTTPException(404, "session not found") + if body.session_id and body.session_id != session_id: + raise HTTPException(422, "session_id mismatch") + try: + raw = base64.b64decode(body.data_b64) + shape = tuple(int(x) for x in body.shape) + arr = np.frombuffer(raw, dtype=np.dtype(body.dtype)).reshape(shape) + except Exception as exc: + raise HTTPException(422, f"invalid array payload: {exc}") from exc + return array_blobs.save_array_blob( + session.project_id, + arr, + array_ref=body.array_ref, + ) + + +@app.get("/setups/{setup_id}") +def api_get_setup(setup_id: str) -> dict[str, Any]: + """Get one Feature Setup.""" + try: + return feature_setups.resolve_setup(setup_id) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + + +@app.post("/setups") +def api_upsert_setup(body: SetupUpsertBody) -> dict[str, Any]: + """Create or update a Feature Setup.""" + try: + return feature_setups.save_setup( + name=body.name, + kind=body.kind, + procedure_id=body.procedure_id, + params=body.params, + encoder_setup_id=body.encoder_setup_id, + weights_path=body.weights_path, + weights_format=body.weights_format, + inference=body.inference, + setup_id=body.setup_id, + catalog=get_catalog(), + ) + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + + +@app.get("/trainers") +def api_list_trainers() -> dict[str, Any]: + """List trainer plugin ids.""" + return {"trainers": list_trainers()} + + +@app.post("/preprocess") +def api_preprocess(body: PreprocessBody) -> dict[str, Any]: + """Cache-aware featurize (composition_id preferred; setup_id still accepted).""" + try: + return preprocess.run_preprocess( + get_catalog(), + session_id=body.session_id, + feature_setup_id=body.feature_setup_id, + composition_id=body.composition_id, + slice_index=body.slice_index, + array_ref=body.array_ref, + ) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + except FileNotFoundError as exc: + raise HTTPException(404, str(exc)) from exc + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + except Exception as exc: + logger.exception("preprocess failed") + raise HTTPException(500, f"preprocess failed: {exc}") from exc + + +@app.get("/features/{feature_id}/channels/{index}") +def api_feature_channel(feature_id: str, index: int) -> Response: + """Return a channel PNG for a feature bank.""" + row = get_catalog().get_feature_bank(feature_id) + if row is None: + raise HTTPException(404, "feature bank not found") + try: + path = preprocess.channel_png_path(row["project_id"], feature_id, index) + except ValueError as exc: + logger.warning("rejected invalid channel path request: %s", exc) + raise HTTPException(404, "channel not found") from exc + if not path.is_file(): + raise HTTPException(404, "channel not found") + return FileResponse(path, media_type="image/png") + + +@app.delete("/features/{feature_id}") +def api_delete_feature_bank(feature_id: str) -> dict[str, Any]: + """Delete a feature bank (DB row + its on-disk blob dir). + + For batch-apply jobs (Phase 4.5's "Apply across volume") to release each + slice's feature bank immediately after it's used — see + `Catalog.delete_feature_bank`'s doc for why this exists. Deleting an + already-gone or unknown feature_id is a no-op, not an error: cleanup + calls are best-effort and must never fail the job that triggered them. + """ + blob_dir = get_catalog().delete_feature_bank(feature_id) + if blob_dir: + shutil.rmtree(blob_dir, ignore_errors=True) + return {"deleted": blob_dir is not None} + + +@app.post("/train") +def api_train(body: TrainBody) -> dict[str, Any]: + """Train via trainer plugin.""" + try: + return train_infer.run_train( + get_catalog(), + session_id=body.session_id, + shapes=body.shapes, + feature_id=body.feature_id, + trainer_id=body.trainer_id, + config=body.config, + ) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + except Exception as exc: + logger.exception("train failed") + raise HTTPException(500, f"train failed: {exc}") from exc + + +@app.post("/train/multi") +def api_train_multi(body: TrainMultiBody) -> dict[str, Any]: + """Train one model pooling labeled pixels across multiple slices.""" + try: + return train_infer.run_train_multi_slice( + get_catalog(), + session_id=body.session_id, + per_slice_shapes={int(k): v for k, v in body.slices.items()}, + feature_ids={int(k): v for k, v in body.feature_ids.items()}, + trainer_id=body.trainer_id, + config=body.config, + ) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + except Exception as exc: + logger.exception("multi-slice train failed") + raise HTTPException(500, f"multi-slice train failed: {exc}") from exc + + +@app.post("/infer") +def api_infer(body: InferBody) -> dict[str, Any]: + """Infer + conformal products.""" + try: + return train_infer.run_infer( + get_catalog(), + session_id=body.session_id, + model_id=body.model_id, + feature_id=body.feature_id, + alpha=body.alpha, + store_probabilities=body.store_probabilities, + ) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + except Exception as exc: + logger.exception("infer failed") + raise HTTPException(500, f"infer failed: {exc}") from exc + + +@app.post("/rethreshold") +def api_rethreshold(body: RethresholdBody) -> dict[str, Any]: + """Rethreshold from cached proba.""" + try: + return train_infer.run_rethreshold( + get_catalog(), + session_id=body.session_id, + run_id=body.run_id, + alpha=body.alpha, + ) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + except Exception as exc: + logger.exception("rethreshold failed") + raise HTTPException(500, f"rethreshold failed: {exc}") from exc + + +@app.get("/runs/{run_id}/commit.png") +def api_run_commit(run_id: str) -> Response: + """Commit map PNG.""" + run = get_catalog().get_run(run_id) + if run is None: + raise HTTPException(404, "run not found") + path = Path(run["blob_dir"]) / "commit.png" + if not path.is_file(): + raise HTTPException(404, "commit.png missing") + return FileResponse(path, media_type="image/png") + + +@app.get("/runs/{run_id}/status.png") +def api_run_status(run_id: str) -> Response: + """Status map PNG.""" + run = get_catalog().get_run(run_id) + if run is None: + raise HTTPException(404, "run not found") + path = Path(run["blob_dir"]) / "status.png" + if not path.is_file(): + raise HTTPException(404, "status.png missing") + return FileResponse(path, media_type="image/png") + + +@app.get("/runs/{run_id}/proba/{class_index}.png") +def api_run_proba_channel(run_id: str, class_index: int) -> Response: + """Softmax probability heatmap PNG for one class column (0-based index).""" + try: + data = train_infer.proba_heatmap_png(get_catalog(), run_id, class_index) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + except FileNotFoundError as exc: + raise HTTPException(404, str(exc)) from exc + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + return Response(content=data, media_type="image/png") + + +class ThresholdClassBody(BaseModel): + class_id: int + threshold: float = 0.5 + + +@app.post("/runs/{run_id}/threshold-class") +def api_threshold_class(run_id: str, body: ThresholdClassBody) -> dict[str, Any]: + """Threshold one softmax class into a dense label map (for mask-set caches).""" + try: + return train_infer.threshold_class_label_map( + get_catalog(), + run_id, + class_id=body.class_id, + threshold=body.threshold, + ) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + except FileNotFoundError as exc: + raise HTTPException(404, str(exc)) from exc + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + + +@app.get("/runs/{run_id}") +def api_get_run(run_id: str) -> dict[str, Any]: + """Run metadata.""" + run = get_catalog().get_run(run_id) + if run is None: + raise HTTPException(404, "run not found") + import json + + meta = json.loads(run["meta_json"]) + return {**meta, "blob_dir": run["blob_dir"], "alpha": run["alpha"]} + + +@app.post("/manifold/sample") +def api_manifold_sample(body: ManifoldSampleBody) -> dict[str, Any]: + """Suggest Labels: variance boxes + residual heatmap on a feature bank.""" + try: + return manifold_jobs.run_manifold_sample( + get_catalog(), + feature_id=body.feature_id, + k=body.k, + box_size=body.box_size, + stride=body.stride, + pca_dims=body.pca_dims, + shapes=body.shapes, + ) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + except ValueError as exc: + raise HTTPException(422, str(exc)) from exc + except Exception as exc: + logger.exception("manifold sample failed") + raise HTTPException(500, f"manifold sample failed: {exc}") from exc + + +@app.get("/manifold/{sample_id}/heatmap.png") +def api_manifold_heatmap(sample_id: str) -> Response: + """Residual heatmap PNG for a manifold sample.""" + try: + png = manifold_jobs.heatmap_png(sample_id) + except KeyError as exc: + raise HTTPException(404, str(exc)) from exc + return Response( + content=png, + media_type="image/png", + headers={"Cache-Control": "private, max-age=300"}, + ) diff --git a/ipred/src/ipred/array_blobs.py b/ipred/src/ipred/array_blobs.py new file mode 100644 index 0000000..04ceb93 --- /dev/null +++ b/ipred/src/ipred/array_blobs.py @@ -0,0 +1,56 @@ +"""Content-addressed array blobs for remote compute (data ≠ compute host).""" + +from __future__ import annotations + +import hashlib +import json +import uuid +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +import numpy as np + +from ipred.paths import project_blob_dir + + +def arrays_dir(project_id: str) -> Path: + """Blob directory for uploaded arrays.""" + path = project_blob_dir(project_id) / "arrays" + path.mkdir(parents=True, exist_ok=True) + return path + + +def save_array_blob( + project_id: str, + arr: np.ndarray, + *, + array_ref: str | None = None, + meta: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Persist ``arr`` and return ``{array_ref, shape, dtype, ...}``.""" + data = np.asarray(arr) + raw = data.astype(np.float32, copy=False).tobytes() + digest = hashlib.sha1(raw).hexdigest()[:16] + ref = array_ref or f"{digest}_{uuid.uuid4().hex[:8]}" + dest = arrays_dir(project_id) / ref + dest.mkdir(parents=True, exist_ok=True) + np.save(dest / "array.npy", data) + record = { + "array_ref": ref, + "shape": list(data.shape), + "dtype": str(data.dtype), + "sha1": digest, + "created_at": datetime.now(timezone.utc).isoformat(), + **(meta or {}), + } + (dest / "meta.json").write_text(json.dumps(record, indent=2), encoding="utf-8") + return record + + +def load_array_blob(project_id: str, array_ref: str) -> np.ndarray: + """Load a previously uploaded array.""" + path = arrays_dir(project_id) / array_ref / "array.npy" + if not path.is_file(): + raise FileNotFoundError(f"array_ref not found: {array_ref}") + return np.load(path).astype(np.float32) diff --git a/ipred/src/ipred/array_source.py b/ipred/src/ipred/array_source.py new file mode 100644 index 0000000..02ac0d1 --- /dev/null +++ b/ipred/src/ipred/array_source.py @@ -0,0 +1,183 @@ +"""Minimal local / Tiled slice reader (no annotate-backend imports).""" + +from __future__ import annotations + +import os +from pathlib import Path +from typing import Any + +import numpy as np + + +def _local_root(root: str | None) -> Path: + if root: + return Path(root).expanduser().resolve() + return Path(os.getenv("LOCAL_DATA_ROOT", "~/data")).expanduser().resolve() + + +def read_slice( + *, + kind: str, + source: str, + slice_index: int = 0, + server_uri: str | None = None, + root: str | None = None, +) -> np.ndarray: + """Load one 2D/ HxWxC slice array.""" + if kind == "local": + return _read_local(source, slice_index=slice_index, root=root) + if kind == "tiled": + return _read_tiled( + source, slice_index=slice_index, server_uri=server_uri + ) + raise ValueError(f"unknown kind {kind!r}") + + +def _read_local( + source: str, + *, + slice_index: int, + root: str | None, +) -> np.ndarray: + base = _local_root(root) + path = (base / source).resolve() + if not str(path).startswith(str(base)): + raise PermissionError("local path escapes LOCAL_DATA_ROOT") + if not path.is_file(): + raise FileNotFoundError(str(path)) + suffix = path.suffix.lower() + if suffix in {".tif", ".tiff"}: + import tifffile + + # Read exactly one PAGE, not the whole stack (`tifffile.imread` decodes + # every page up front). A batch-apply job over N slices previously + # decoded the entire stack N times — an O(N^2) cost that dominated + # "predict across all slices" wall-clock time on larger volumes. + with tifffile.TiffFile(str(path)) as tif: + n_pages = len(tif.pages) + if n_pages <= 1: + return _index_slice(np.asarray(tif.asarray()), slice_index) + idx = int(slice_index) if 0 <= int(slice_index) < n_pages else 0 + page = np.asarray(tif.pages[idx].asarray()) + return _index_slice(page, 0) + from PIL import Image as PILImage + + arr = np.asarray(PILImage.open(path)) + return _index_slice(np.asarray(arr), slice_index) + + +def _is_container_node(node: Any) -> bool: + sf = getattr(node, "structure_family", None) + return str(getattr(sf, "value", sf)) == "container" + + +def _descend_to_array(node: Any, max_depth: int = 8) -> Any: + """Resolve a Tiled container to the single array/stack it represents. + + Handles two shapes seen in this repo's Tiled catalogs: + + * An OME-NGFF multiscale pyramid (``scale0/image``, ``scale1/image``, …) — + descends to the finest (``scale0``) level's array. Without this, a + 5-level pyramid container's 5 children get handed to ``np.asarray``, + which yields a 1-D array of length 5 (one object per child) instead of + image data, tripping the "unsupported array shape" check below. + * A wrapper container whose only/first child is itself a container (e.g. + drag-and-drop ingest nesting) — descends until reaching an array or a + container of arrays (slice stack). + """ + if not _is_container_node(node): + return node + + try: + scale_keys = sorted( + (k for k in node if str(k).startswith("scale")), + key=lambda k: int("".join(ch for ch in str(k) if ch.isdigit()) or 0), + ) + except Exception: # noqa: BLE001 — not enumerable → not a pyramid + scale_keys = [] + if scale_keys: + level = node[scale_keys[0]] + if _is_container_node(level): + try: + level = level[next(iter(level))] + except StopIteration: + pass + return level + + depth = 0 + while _is_container_node(node) and depth < max_depth: + try: + first = next(iter(node)) + except StopIteration: + return node # empty container — nothing to descend into + child = node[first] + if not _is_container_node(child): + return node # container of arrays → treat as a slice stack + node = child + depth += 1 + return node + + +def _read_tiled( + source: str, + *, + slice_index: int, + server_uri: str | None, +) -> np.ndarray: + from tiled.client import from_uri + + uri = (server_uri or os.getenv("TILED_URI") or "http://127.0.0.1:8010").rstrip( + "/" + ) + api_key = os.getenv("TILED_API_KEY") + client: Any = from_uri(uri, api_key=api_key) if api_key else from_uri(uri) + node: Any = client + for part in source.strip("/").split("/"): + if not part: + continue + node = node[part] + node = _descend_to_array(node) + if _is_container_node(node): + # Container of per-slice arrays (not a bare NHW/NHWC array node). + keys = sorted(node) + node = node[keys[int(slice_index)]] + return _index_slice(np.asarray(node), 0) + + # A bare NHW/NHWC array node: index it BEFORE calling np.asarray, so only + # the requested slice is fetched/decoded over the wire — `np.asarray(node)` + # first (the previous behavior) pulls the ENTIRE stack for every single + # slice request, turning an N-slice batch-apply job into an O(N^2) read. + # `.shape` is cheap metadata (no data fetch), so it's safe to inspect + # before deciding whether/how to index — mirrors `_index_slice`'s own + # HWC-vs-NHW heuristic, applied to shape metadata instead of a realized + # array. Matches the working pattern in `backend/arrays.py: read_slice`. + shape = tuple(int(s) for s in node.shape) + if len(shape) <= 2: + return np.asarray(node) + if len(shape) == 3 and shape[-1] in (1, 3, 4) and ( + shape[0] > 8 or shape[0] == shape[1] + ): + return np.asarray(node) # HWC single image — nothing to index, it's the whole thing + # Genuine NHW/NHWC stack: node[idx] is Tiled's own lazy single-slice + # fetch — the whole reason this branch exists — so the result is already + # exactly one 2-D/HWC frame with no further slicing needed. + return np.asarray(node[int(slice_index)]) + + +def _index_slice(arr: np.ndarray, slice_index: int) -> np.ndarray: + """Pick HW or HWC frame from NHW / NHWC / HW / HWC.""" + a = np.asarray(arr) + if a.ndim == 2: + return a + if a.ndim == 3: + # HWC vs NHW + if a.shape[-1] in (1, 3, 4) and a.shape[0] > 8: + return a # HWC single + if a.shape[-1] in (1, 3, 4) and a.shape[0] <= 8: + # ambiguous small N — treat as HWC if last dim channel-like and square-ish + if a.shape[0] == a.shape[1]: + return a + return a[int(slice_index)] + if a.ndim == 4: + return a[int(slice_index)] + raise ValueError(f"unsupported array shape {a.shape}") diff --git a/ipred/src/ipred/cache.py b/ipred/src/ipred/cache.py new file mode 100644 index 0000000..3f2de80 --- /dev/null +++ b/ipred/src/ipred/cache.py @@ -0,0 +1,40 @@ +"""Tiny TTL cache for short-lived ipred artifacts (manifold heatmaps).""" + +from __future__ import annotations + +import time +from threading import Lock +from typing import Any, Generic, TypeVar + +T = TypeVar("T") + + +class TTLCache(Generic[T]): + """Thread-safe TTL + max-entries cache.""" + + def __init__(self, *, ttl_seconds: float, max_entries: int) -> None: + self.ttl = float(ttl_seconds) + self.max_entries = int(max_entries) + self._data: dict[str, tuple[float, T]] = {} + self._lock = Lock() + + def get(self, key: str) -> T | None: + now = time.monotonic() + with self._lock: + item = self._data.get(key) + if item is None: + return None + ts, val = item + if now - ts > self.ttl: + del self._data[key] + return None + return val + + def set(self, key: str, value: T) -> None: + now = time.monotonic() + with self._lock: + self._data[key] = (now, value) + if len(self._data) > self.max_entries: + # Drop oldest + oldest = min(self._data.items(), key=lambda kv: kv[1][0])[0] + del self._data[oldest] diff --git a/ipred/src/ipred/catalog.py b/ipred/src/ipred/catalog.py new file mode 100644 index 0000000..4506386 --- /dev/null +++ b/ipred/src/ipred/catalog.py @@ -0,0 +1,470 @@ +"""SQLite catalog for projects, sessions, feature banks, models, and runs.""" + +from __future__ import annotations + +import hashlib +import json +import sqlite3 +import uuid +from contextlib import contextmanager +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Iterator + +from ipred.paths import catalog_db_path, project_blob_dir + + +def _utc_now() -> str: + return datetime.now(timezone.utc).isoformat() + + +def project_id_for( + *, + kind: str, + source: str, + server_uri: str | None = None, + root: str | None = None, +) -> str: + """Stable project id from source identity.""" + payload = json.dumps( + { + "kind": kind, + "source": source, + "server_uri": server_uri or "", + "root": root or "", + }, + sort_keys=True, + separators=(",", ":"), + ) + return hashlib.sha1(payload.encode("utf-8")).hexdigest()[:16] + + +@dataclass(frozen=True) +class ProjectRow: + """Project catalog row.""" + + project_id: str + kind: str + source: str + server_uri: str | None + root: str | None + created_at: str + + +@dataclass(frozen=True) +class SessionRow: + """Session catalog row.""" + + session_id: str + project_id: str + current_feature_id: str | None + current_model_id: str | None + current_run_id: str | None + created_at: str + updated_at: str + + +class Catalog: + """SQLite-backed engine catalog.""" + + def __init__(self, db_path: Path | None = None) -> None: + self.db_path = Path(db_path) if db_path else catalog_db_path() + self.db_path.parent.mkdir(parents=True, exist_ok=True) + self._init_schema() + + @contextmanager + def connect(self) -> Iterator[sqlite3.Connection]: + """Yield a connection with row factory.""" + conn = sqlite3.connect(str(self.db_path)) + conn.row_factory = sqlite3.Row + try: + yield conn + conn.commit() + except Exception: + conn.rollback() + raise + finally: + conn.close() + + def _init_schema(self) -> None: + with self.connect() as conn: + conn.executescript( + """ + CREATE TABLE IF NOT EXISTS projects ( + project_id TEXT PRIMARY KEY, + kind TEXT NOT NULL, + source TEXT NOT NULL, + server_uri TEXT, + root TEXT, + created_at TEXT NOT NULL + ); + + CREATE TABLE IF NOT EXISTS sessions ( + session_id TEXT PRIMARY KEY, + project_id TEXT NOT NULL REFERENCES projects(project_id), + current_feature_id TEXT, + current_model_id TEXT, + current_run_id TEXT, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL + ); + + CREATE TABLE IF NOT EXISTS feature_setups ( + setup_id TEXT PRIMARY KEY, + name TEXT NOT NULL, + kind TEXT NOT NULL, + content_hash TEXT NOT NULL, + meta_json TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL + ); + + CREATE TABLE IF NOT EXISTS feature_banks ( + feature_id TEXT PRIMARY KEY, + project_id TEXT NOT NULL, + setup_id TEXT NOT NULL, + content_hash TEXT NOT NULL, + slice_index INTEGER NOT NULL, + n_channels INTEGER NOT NULL, + height INTEGER NOT NULL, + width INTEGER NOT NULL, + blob_dir TEXT NOT NULL, + setup_snapshot TEXT NOT NULL, + status TEXT NOT NULL, + created_at TEXT NOT NULL, + UNIQUE(project_id, setup_id, content_hash, slice_index) + ); + + CREATE TABLE IF NOT EXISTS models ( + model_id TEXT PRIMARY KEY, + project_id TEXT NOT NULL, + feature_id TEXT NOT NULL, + trainer_id TEXT NOT NULL, + blob_dir TEXT NOT NULL, + meta_json TEXT NOT NULL, + created_at TEXT NOT NULL + ); + + CREATE TABLE IF NOT EXISTS runs ( + run_id TEXT PRIMARY KEY, + project_id TEXT NOT NULL, + model_id TEXT NOT NULL, + feature_id TEXT NOT NULL, + alpha REAL NOT NULL, + blob_dir TEXT NOT NULL, + meta_json TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL + ); + """ + ) + + def ensure_project( + self, + *, + kind: str, + source: str, + server_uri: str | None = None, + root: str | None = None, + ) -> ProjectRow: + """Create or return the project for this source identity.""" + pid = project_id_for( + kind=kind, source=source, server_uri=server_uri, root=root + ) + now = _utc_now() + with self.connect() as conn: + row = conn.execute( + "SELECT * FROM projects WHERE project_id = ?", (pid,) + ).fetchone() + if row is None: + conn.execute( + """ + INSERT INTO projects + (project_id, kind, source, server_uri, root, created_at) + VALUES (?, ?, ?, ?, ?, ?) + """, + (pid, kind, source, server_uri, root, now), + ) + project_blob_dir(pid) + row = conn.execute( + "SELECT * FROM projects WHERE project_id = ?", (pid,) + ).fetchone() + return self._project_from_row(row) + + def open_session( + self, + *, + kind: str, + source: str, + server_uri: str | None = None, + root: str | None = None, + ) -> SessionRow: + """Ensure project exists and open a new session on it.""" + project = self.ensure_project( + kind=kind, source=source, server_uri=server_uri, root=root + ) + sid = uuid.uuid4().hex + now = _utc_now() + with self.connect() as conn: + conn.execute( + """ + INSERT INTO sessions + (session_id, project_id, current_feature_id, current_model_id, + current_run_id, created_at, updated_at) + VALUES (?, ?, NULL, NULL, NULL, ?, ?) + """, + (sid, project.project_id, now, now), + ) + row = conn.execute( + "SELECT * FROM sessions WHERE session_id = ?", (sid,) + ).fetchone() + return self._session_from_row(row) + + def get_session(self, session_id: str) -> SessionRow | None: + """Look up a session.""" + with self.connect() as conn: + row = conn.execute( + "SELECT * FROM sessions WHERE session_id = ?", (session_id,) + ).fetchone() + return self._session_from_row(row) if row else None + + def get_project(self, project_id: str) -> ProjectRow | None: + """Look up a project.""" + with self.connect() as conn: + row = conn.execute( + "SELECT * FROM projects WHERE project_id = ?", (project_id,) + ).fetchone() + return self._project_from_row(row) if row else None + + def set_session_currents( + self, + session_id: str, + *, + feature_id: str | None = None, + model_id: str | None = None, + run_id: str | None = None, + ) -> SessionRow: + """Update session current_* pointers (only provided fields).""" + session = self.get_session(session_id) + if session is None: + raise KeyError(f"unknown session {session_id}") + feat = feature_id if feature_id is not None else session.current_feature_id + model = model_id if model_id is not None else session.current_model_id + run = run_id if run_id is not None else session.current_run_id + now = _utc_now() + with self.connect() as conn: + conn.execute( + """ + UPDATE sessions + SET current_feature_id = ?, current_model_id = ?, + current_run_id = ?, updated_at = ? + WHERE session_id = ? + """, + (feat, model, run, now, session_id), + ) + row = conn.execute( + "SELECT * FROM sessions WHERE session_id = ?", (session_id,) + ).fetchone() + return self._session_from_row(row) + + def upsert_feature_setup( + self, + *, + setup_id: str, + name: str, + kind: str, + content_hash: str, + meta: dict[str, Any], + ) -> None: + """Insert or update a feature setup index row.""" + now = _utc_now() + payload = json.dumps(meta, sort_keys=True) + with self.connect() as conn: + existing = conn.execute( + "SELECT setup_id FROM feature_setups WHERE setup_id = ?", + (setup_id,), + ).fetchone() + if existing is None: + conn.execute( + """ + INSERT INTO feature_setups + (setup_id, name, kind, content_hash, meta_json, + created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?) + """, + (setup_id, name, kind, content_hash, payload, now, now), + ) + else: + conn.execute( + """ + UPDATE feature_setups + SET name = ?, kind = ?, content_hash = ?, meta_json = ?, + updated_at = ? + WHERE setup_id = ? + """, + (name, kind, content_hash, payload, now, setup_id), + ) + + def find_feature_bank( + self, + *, + project_id: str, + setup_id: str, + content_hash: str, + slice_index: int, + ) -> dict[str, Any] | None: + """Return a ready feature-bank row dict or None.""" + with self.connect() as conn: + row = conn.execute( + """ + SELECT * FROM feature_banks + WHERE project_id = ? AND setup_id = ? AND content_hash = ? + AND slice_index = ? AND status = 'ready' + """, + (project_id, setup_id, content_hash, slice_index), + ).fetchone() + return dict(row) if row else None + + def insert_feature_bank(self, record: dict[str, Any]) -> None: + """Insert a feature-bank row.""" + with self.connect() as conn: + conn.execute( + """ + INSERT INTO feature_banks + (feature_id, project_id, setup_id, content_hash, slice_index, + n_channels, height, width, blob_dir, setup_snapshot, status, + created_at) + VALUES + (:feature_id, :project_id, :setup_id, :content_hash, + :slice_index, :n_channels, :height, :width, :blob_dir, + :setup_snapshot, :status, :created_at) + """, + record, + ) + + def get_feature_bank(self, feature_id: str) -> dict[str, Any] | None: + """Look up a feature bank by id.""" + with self.connect() as conn: + row = conn.execute( + "SELECT * FROM feature_banks WHERE feature_id = ?", + (feature_id,), + ).fetchone() + return dict(row) if row else None + + def delete_feature_bank(self, feature_id: str) -> str | None: + """Delete a feature-bank row and return its `blob_dir` for the caller + to remove from disk (this method does not touch the filesystem — it + only owns the DB row, matching every other method here). + + Feature banks bake 32-64 full-resolution float32 PCA channels + (~1 GB per 2000x2000 slice, see `sam_embed.emb_grid_to_pca_channels`) + and were never evicted: a single interactive slice benefits from the + cache (re-predicting after a param tweak), but a volume-wide batch + job (Phase 4.5's "Apply across volume") touches every slice exactly + once, so there is nothing later in that job to reuse the cache for — + it just accumulates forever. 690 slices at ~1 GB each is enough to + fill a laptop's disk mid-job. Batch-apply calls this right after each + slice's inference completes; ordinary single-slice interactive use + (train/preprocess/infer outside a batch job) is untouched. + + Returns None if the feature_id doesn't exist (already deleted, or a + bad id) — the caller should treat that as a no-op, not an error. + """ + with self.connect() as conn: + row = conn.execute( + "SELECT blob_dir FROM feature_banks WHERE feature_id = ?", + (feature_id,), + ).fetchone() + if row is None: + return None + conn.execute("DELETE FROM feature_banks WHERE feature_id = ?", (feature_id,)) + return row["blob_dir"] + + def insert_model(self, record: dict[str, Any]) -> None: + """Insert a model row.""" + with self.connect() as conn: + conn.execute( + """ + INSERT INTO models + (model_id, project_id, feature_id, trainer_id, blob_dir, + meta_json, created_at) + VALUES + (:model_id, :project_id, :feature_id, :trainer_id, :blob_dir, + :meta_json, :created_at) + """, + record, + ) + + def get_model(self, model_id: str) -> dict[str, Any] | None: + """Look up a model row.""" + with self.connect() as conn: + row = conn.execute( + "SELECT * FROM models WHERE model_id = ?", (model_id,) + ).fetchone() + return dict(row) if row else None + + def insert_run(self, record: dict[str, Any]) -> None: + """Insert a run row.""" + with self.connect() as conn: + conn.execute( + """ + INSERT INTO runs + (run_id, project_id, model_id, feature_id, alpha, blob_dir, + meta_json, created_at, updated_at) + VALUES + (:run_id, :project_id, :model_id, :feature_id, :alpha, + :blob_dir, :meta_json, :created_at, :updated_at) + """, + record, + ) + + def update_run( + self, + run_id: str, + *, + alpha: float, + meta_json: str, + ) -> None: + """Update run alpha/meta after rethreshold.""" + now = _utc_now() + with self.connect() as conn: + conn.execute( + """ + UPDATE runs + SET alpha = ?, meta_json = ?, updated_at = ? + WHERE run_id = ? + """, + (alpha, meta_json, now, run_id), + ) + + def get_run(self, run_id: str) -> dict[str, Any] | None: + """Look up a run row.""" + with self.connect() as conn: + row = conn.execute( + "SELECT * FROM runs WHERE run_id = ?", (run_id,) + ).fetchone() + return dict(row) if row else None + + @staticmethod + def _project_from_row(row: sqlite3.Row) -> ProjectRow: + return ProjectRow( + project_id=row["project_id"], + kind=row["kind"], + source=row["source"], + server_uri=row["server_uri"], + root=row["root"], + created_at=row["created_at"], + ) + + @staticmethod + def _session_from_row(row: sqlite3.Row) -> SessionRow: + return SessionRow( + session_id=row["session_id"], + project_id=row["project_id"], + current_feature_id=row["current_feature_id"], + current_model_id=row["current_model_id"], + current_run_id=row["current_run_id"], + created_at=row["created_at"], + updated_at=row["updated_at"], + ) diff --git a/ipred/src/ipred/cli.py b/ipred/src/ipred/cli.py new file mode 100644 index 0000000..bf20b8d --- /dev/null +++ b/ipred/src/ipred/cli.py @@ -0,0 +1,188 @@ +"""CLI entrypoint ``ipred`` with --project session sugar.""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path +from typing import Any + +from ipred import feature_setups, preprocess, train_infer +from ipred.catalog import Catalog +from ipred.trainers import list_trainers + + +def _catalog() -> Catalog: + cat = Catalog() + feature_setups.ensure_default_setups(cat) + return cat + + +def _resolve_session(args: argparse.Namespace, cat: Catalog) -> str: + """Return session_id from --session or --project sugar.""" + if getattr(args, "session", None): + return str(args.session) + if getattr(args, "kind", None) and getattr(args, "source", None): + session = cat.open_session( + kind=args.kind, + source=args.source, + server_uri=getattr(args, "server_uri", None), + root=getattr(args, "root", None), + ) + print(json.dumps({"session_id": session.session_id, "project_id": session.project_id})) + return session.session_id + raise SystemExit("provide --session or --kind/--source (--project sugar)") + + +def _add_project_sugar(parser: argparse.ArgumentParser) -> None: + parser.add_argument("--session", default=None, help="Existing session id") + parser.add_argument("--kind", default=None, help="Project kind (local|tiled)") + parser.add_argument("--source", default=None, help="Project source path") + parser.add_argument("--server-uri", dest="server_uri", default=None) + parser.add_argument("--root", default=None, help="Local root override") + + +def main(argv: list[str] | None = None) -> None: + """CLI main.""" + parser = argparse.ArgumentParser( + prog="ipred", + description="Iterative prediction backend CLI", + ) + sub = parser.add_subparsers(dest="cmd", required=True) + + p_session = sub.add_parser("session", help="Session commands") + sess_sub = p_session.add_subparsers(dest="session_cmd", required=True) + p_open = sess_sub.add_parser("open", help="Open a session for a project") + p_open.add_argument("--kind", required=True) + p_open.add_argument("--source", required=True) + p_open.add_argument("--server-uri", dest="server_uri", default=None) + p_open.add_argument("--root", default=None) + + p_setup = sub.add_parser("setup", help="Feature Setup commands") + setup_sub = p_setup.add_subparsers(dest="setup_cmd", required=True) + setup_sub.add_parser("list", help="List setups") + p_show = setup_sub.add_parser("show", help="Show one setup") + p_show.add_argument("setup_id") + + p_pre = sub.add_parser("preprocess", help="Featurize (cache-aware)") + _add_project_sugar(p_pre) + p_pre.add_argument("--setup", required=True, help="Feature setup id") + p_pre.add_argument("--slice", type=int, default=0) + + p_train = sub.add_parser("train", help="Train a model") + _add_project_sugar(p_train) + p_train.add_argument("--labels", required=True, help="JSON file of shapes") + p_train.add_argument("--trainer", default="catboost") + p_train.add_argument("--feature-id", dest="feature_id", default=None) + p_train.add_argument("--config", default=None, help="JSON config file or string") + + p_inf = sub.add_parser("infer", help="Infer + conformal maps") + _add_project_sugar(p_inf) + p_inf.add_argument("--alpha", type=float, default=0.05) + p_inf.add_argument("--model-id", dest="model_id", default=None) + p_inf.add_argument("--feature-id", dest="feature_id", default=None) + + p_rt = sub.add_parser("rethreshold", help="Rethreshold from cached proba") + _add_project_sugar(p_rt) + p_rt.add_argument("--alpha", type=float, required=True) + p_rt.add_argument("--run-id", dest="run_id", default=None) + + sub.add_parser("trainers", help="List trainer plugins") + + args = parser.parse_args(argv) + cat = _catalog() + + if args.cmd == "session" and args.session_cmd == "open": + session = cat.open_session( + kind=args.kind, + source=args.source, + server_uri=args.server_uri, + root=args.root, + ) + _print( + { + "session_id": session.session_id, + "project_id": session.project_id, + } + ) + return + + if args.cmd == "setup": + if args.setup_cmd == "list": + _print({"setups": feature_setups.list_setups()}) + return + if args.setup_cmd == "show": + _print(feature_setups.resolve_setup(args.setup_id)) + return + + if args.cmd == "trainers": + _print({"trainers": list_trainers()}) + return + + if args.cmd == "preprocess": + sid = _resolve_session(args, cat) + out = preprocess.run_preprocess( + cat, + session_id=sid, + feature_setup_id=args.setup, + slice_index=args.slice, + ) + _print(out) + return + + if args.cmd == "train": + sid = _resolve_session(args, cat) + shapes = json.loads(Path(args.labels).read_text(encoding="utf-8")) + if isinstance(shapes, dict) and "shapes" in shapes: + shapes = shapes["shapes"] + config: dict[str, Any] = {} + if args.config: + cfg_path = Path(args.config) + if cfg_path.is_file(): + config = json.loads(cfg_path.read_text(encoding="utf-8")) + else: + config = json.loads(args.config) + out = train_infer.run_train( + cat, + session_id=sid, + shapes=shapes, + feature_id=args.feature_id, + trainer_id=args.trainer, + config=config, + ) + _print(out) + return + + if args.cmd == "infer": + sid = _resolve_session(args, cat) + out = train_infer.run_infer( + cat, + session_id=sid, + model_id=args.model_id, + feature_id=args.feature_id, + alpha=args.alpha, + ) + _print(out) + return + + if args.cmd == "rethreshold": + sid = _resolve_session(args, cat) + out = train_infer.run_rethreshold( + cat, + session_id=sid, + run_id=args.run_id, + alpha=args.alpha, + ) + _print(out) + return + + parser.error(f"unhandled command {args.cmd}") + + +def _print(obj: Any) -> None: + print(json.dumps(obj, indent=2, default=str)) + + +if __name__ == "__main__": + main(sys.argv[1:]) diff --git a/ipred/src/ipred/compose_run.py b/ipred/src/ipred/compose_run.py new file mode 100644 index 0000000..43c6081 --- /dev/null +++ b/ipred/src/ipred/compose_run.py @@ -0,0 +1,84 @@ +"""Execute a composition document on a raw image array.""" + +from __future__ import annotations + +from typing import Any + +import numpy as np + +from ipred import features +from ipred.compositions import validate_composition +from ipred.modules import get_module +from ipred.modules.base import ChannelBlock, ModuleContext +from ipred.modules.pca_mod import uint8_from_float + + +def run_composition( + arr: np.ndarray, + doc: dict[str, Any], +) -> tuple[ + np.ndarray, + np.ndarray, + list[str], + np.ndarray | None, + dict[str, Any] | None, +]: + """Run composition nodes; concatenate ``outputs`` channel blocks. + + Returns: + uint8_stack, float_stack, labels, sam_emb, sam_meta + (emb/meta from the last PCA or encoder that produced them). + """ + validate_composition(doc) + gray = features.to_grayscale(arr) + ctx = ModuleContext(raw=arr, gray=gray, params={}) + by_id = {n["id"]: n for n in doc["nodes"]} + + # Topological-ish: run in document order (authors must order deps first) + for node in doc["nodes"]: + nid = node["id"] + mid = str(node["module"]) + mod = get_module(mid) + params = dict(node.get("params") or {}) + input_from = node.get("input_from") + input_image = None + if input_from: + upstream = ctx.node_outputs.get(input_from) + if upstream is None: + raise ValueError( + f"node {nid}: input_from {input_from} not run yet" + ) + if upstream.image_2d is not None: + input_image = upstream.image_2d + elif upstream.float_stack is not None and upstream.float_stack.shape[-1] >= 1: + input_image = upstream.float_stack[..., 0] + params["_input_from"] = input_from + ctx.params = params + ctx.input_image = input_image + block = mod.run(ctx) + ctx.node_outputs[nid] = block + + float_parts: list[np.ndarray] = [] + labels: list[str] = [] + sam_emb: np.ndarray | None = None + sam_meta: dict[str, Any] | None = None + for oid in doc.get("outputs") or []: + block = ctx.node_outputs[oid] + if block.float_stack is not None and block.float_stack.size: + float_parts.append(np.asarray(block.float_stack, dtype=np.float32)) + labels.extend(block.labels) + if block.emb is not None: + sam_emb = block.emb + sam_meta = block.emb_meta + # Also pick up emb from non-output encoder nodes if PCA was in outputs + if sam_emb is None: + for block in ctx.node_outputs.values(): + if block.emb is not None: + sam_emb = block.emb + sam_meta = block.emb_meta + + if not float_parts: + raise ValueError("composition produced no channel outputs") + float_stack = np.concatenate(float_parts, axis=-1) + uint8_stack = uint8_from_float(float_stack) + return uint8_stack, float_stack, labels, sam_emb, sam_meta diff --git a/ipred/src/ipred/compositions.py b/ipred/src/ipred/compositions.py new file mode 100644 index 0000000..a3b6d70 --- /dev/null +++ b/ipred/src/ipred/compositions.py @@ -0,0 +1,407 @@ +"""Composition documents — ordered module graphs for feature banks.""" + +from __future__ import annotations + +import hashlib +import json +import logging +import uuid +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from ipred.catalog import Catalog +from ipred.modules import get_module, list_module_catalog +from ipred.paths import feature_models_root + +logger = logging.getLogger(__name__) + +KIND_COMPOSITION = "composition" + + +def compositions_root() -> Path: + """Directory for composition meta.json files.""" + root = feature_models_root() / "_compositions" + root.mkdir(parents=True, exist_ok=True) + return root + + +def content_hash(doc: dict[str, Any]) -> str: + """Stable hash of composition structure (excludes name/timestamps).""" + payload = { + "kind": KIND_COMPOSITION, + "nodes": doc.get("nodes"), + "outputs": doc.get("outputs"), + } + raw = json.dumps(payload, sort_keys=True, separators=(",", ":")) + return hashlib.sha1(raw.encode("utf-8")).hexdigest()[:16] + + +def _dir(comp_id: str) -> Path: + return compositions_root() / comp_id + + +def load_composition(comp_id: str) -> dict[str, Any] | None: + """Load composition meta or None.""" + path = _dir(comp_id) / "meta.json" + if not path.is_file(): + return None + return json.loads(path.read_text(encoding="utf-8")) + + +def list_compositions() -> list[dict[str, Any]]: + """List all compositions on disk.""" + root = compositions_root() + out: list[dict[str, Any]] = [] + for child in sorted(root.iterdir()) if root.is_dir() else []: + meta_path = child / "meta.json" + if not meta_path.is_file(): + continue + try: + out.append(json.loads(meta_path.read_text(encoding="utf-8"))) + except (OSError, json.JSONDecodeError): + continue + return out + + +def save_composition( + *, + name: str, + nodes: list[dict[str, Any]], + outputs: list[str], + composition_id: str | None = None, + builtin: bool = False, + catalog: Catalog | None = None, +) -> dict[str, Any]: + """Create or overwrite a composition document.""" + validate_composition({"nodes": nodes, "outputs": outputs}) + cid = composition_id or uuid.uuid4().hex + now = datetime.now(timezone.utc).isoformat() + existing = load_composition(cid) + created = existing.get("created_at") if existing else now + meta: dict[str, Any] = { + "id": cid, + "name": (name or cid).strip() or cid, + "kind": KIND_COMPOSITION, + "builtin": bool(builtin), + "nodes": list(nodes), + "outputs": list(outputs), + "created_at": created, + "updated_at": now, + } + meta["content_hash"] = content_hash(meta) + dest = _dir(cid) + dest.mkdir(parents=True, exist_ok=True) + (dest / "meta.json").write_text(json.dumps(meta, indent=2), encoding="utf-8") + if catalog is not None: + catalog.upsert_feature_setup( + setup_id=cid, + name=meta["name"], + kind=KIND_COMPOSITION, + content_hash=meta["content_hash"], + meta=meta, + ) + return meta + + +def resolve_composition(comp_id: str) -> dict[str, Any]: + """Load composition or raise KeyError.""" + meta = load_composition(comp_id) + if meta is None: + raise KeyError(f"unknown composition {comp_id}") + return meta + + +def validate_composition(doc: dict[str, Any]) -> None: + """Raise ValueError if nodes/outputs are invalid.""" + nodes = doc.get("nodes") or [] + outputs = doc.get("outputs") or [] + if not isinstance(nodes, list) or not nodes: + raise ValueError("composition requires at least one node") + ids = [n.get("id") for n in nodes] + if len(ids) != len(set(ids)): + raise ValueError("duplicate node ids") + for n in nodes: + mid = n.get("module") + if not mid: + raise ValueError(f"node {n.get('id')} missing module") + try: + get_module(str(mid)) + except KeyError as exc: + raise ValueError(str(exc)) from exc + src = n.get("input_from") + if src is not None and src not in ids: + raise ValueError(f"node {n.get('id')} input_from unknown: {src}") + for oid in outputs: + if oid not in ids: + raise ValueError(f"output node unknown: {oid}") + + +def preview_concat_labels(doc: dict[str, Any]) -> list[str]: + """Labels the bank would concatenate for ``outputs`` order.""" + validate_composition(doc) + by_id = {n["id"]: n for n in doc["nodes"]} + labels: list[str] = [] + for oid in doc.get("outputs") or []: + node = by_id[oid] + mod = get_module(str(node["module"])) + params = dict(node.get("params") or {}) + labels.extend(mod.preview_labels(params)) + return labels + + +def setup_id_to_composition_id(setup_id: str) -> str | None: + """Map legacy procedure setup ids → builtin composition ids.""" + return _LEGACY_SETUP_MAP.get(setup_id) + + +_LEGACY_SETUP_MAP: dict[str, str] = { + "default-skimage": "comp-skimage", + "default-skimage-slimsam": "comp-skimage-slimsam", + "default-slimsam-clahe": "comp-slimsam-clahe", + "default-skimage-mark25": "comp-skimage-mark25", + "default-mark25-clahe": "comp-mark25-clahe", + "default-skimage-mark11": "comp-skimage-mark11", + "default-mark11-clahe": "comp-mark11-clahe", + # weights-only → CLAHE companion compositions + "default-mark25": "comp-mark25-clahe", + "default-mark11": "comp-mark11-clahe", + "default-slimsam": "comp-skimage-slimsam", +} + + +def _skimage_params() -> dict[str, Any]: + return { + "sigma_min": 1.0, + "sigma_max": 8.0, + "intensity": True, + "edges": True, + "texture": True, + "clahe": True, + } + + +def ensure_default_compositions(catalog: Catalog | None = None) -> list[dict[str, Any]]: + """Create builtin compositions matching legacy procedure setups.""" + specs: list[tuple[str, str, list[dict[str, Any]], list[str]]] = [ + ( + "comp-skimage", + "Skimage multiscale", + [{"id": "n1", "module": "skimage_multiscale", "params": _skimage_params()}], + ["n1"], + ), + ( + "comp-skimage-slimsam", + "Skimage + SlimSAM", + [ + {"id": "n1", "module": "skimage_multiscale", "params": _skimage_params()}, + {"id": "n2", "module": "slimsam", "params": {}}, + { + "id": "n3", + "module": "pca", + "params": {"dims": 32}, + "input_from": "n2", + }, + ], + ["n1", "n3"], + ), + ( + "comp-slimsam-clahe", + "SlimSAM + CLAHE", + [ + { + "id": "n1", + "module": "clahe", + "params": { + "clahe": True, + "clip_limit": 0.01, + "include_in_bank": True, + }, + }, + { + "id": "n2", + "module": "slimsam", + "params": {}, + "input_from": "n1", + }, + { + "id": "n3", + "module": "pca", + "params": {"dims": 64}, + "input_from": "n2", + }, + ], + ["n1", "n3"], + ), + ( + "comp-skimage-mark25", + "Skimage + Mark25", + [ + { + "id": "n1", + "module": "skimage_multiscale", + "params": { + **_skimage_params(), + }, + }, + { + "id": "n2", + "module": "tomojepa", + "params": { + "weights_id": "mark25", + "input_size": 512, + "resize": True, + }, + }, + { + "id": "n3", + "module": "pca", + "params": {"dims": 64}, + "input_from": "n2", + }, + ], + ["n1", "n3"], + ), + ( + "comp-mark25-clahe", + "Mark25 + CLAHE", + [ + { + "id": "n1", + "module": "clahe", + "params": { + "clahe": True, + "clip_limit": 0.01, + "include_in_bank": True, + }, + }, + { + "id": "n2", + "module": "tomojepa", + "params": { + "weights_id": "mark25", + "input_size": 512, + "resize": True, + }, + "input_from": "n1", + }, + { + "id": "n3", + "module": "pca", + "params": {"dims": 64}, + "input_from": "n2", + }, + ], + ["n1", "n3"], + ), + ( + "comp-skimage-mark11", + "Skimage + Mark11", + [ + {"id": "n1", "module": "skimage_multiscale", "params": _skimage_params()}, + { + "id": "n2", + "module": "tomojepa", + "params": { + "weights_id": "mark11", + "input_size": 512, + "resize": True, + }, + }, + { + "id": "n3", + "module": "pca", + "params": {"dims": 64}, + "input_from": "n2", + }, + ], + ["n1", "n3"], + ), + ( + "comp-mark11-clahe", + "Mark11 + CLAHE", + [ + { + "id": "n1", + "module": "clahe", + "params": { + "clahe": True, + "clip_limit": 0.01, + "include_in_bank": True, + }, + }, + { + "id": "n2", + "module": "tomojepa", + "params": { + "weights_id": "mark11", + "input_size": 512, + "resize": True, + }, + "input_from": "n1", + }, + { + "id": "n3", + "module": "pca", + "params": {"dims": 64}, + "input_from": "n2", + }, + ], + ["n1", "n3"], + ), + ] + created: list[dict[str, Any]] = [] + for cid, name, nodes, outputs in specs: + existing = load_composition(cid) + if existing is None: + created.append( + save_composition( + name=name, + nodes=nodes, + outputs=outputs, + composition_id=cid, + builtin=True, + catalog=catalog, + ) + ) + else: + if catalog is not None: + catalog.upsert_feature_setup( + setup_id=cid, + name=existing["name"], + kind=KIND_COMPOSITION, + content_hash=existing.get("content_hash") + or content_hash(existing), + meta=existing, + ) + created.append(existing) + return created + + +def resolve_preprocess_id(setup_or_comp_id: str) -> str: + """Normalize legacy setup id or composition id to a composition id.""" + if load_composition(setup_or_comp_id) is not None: + return setup_or_comp_id + mapped = setup_id_to_composition_id(setup_or_comp_id) + if mapped and load_composition(mapped) is not None: + return mapped + # compositions might not be seeded yet + if mapped: + return mapped + raise KeyError(f"unknown composition or setup {setup_or_comp_id}") + + +__all__ = [ + "KIND_COMPOSITION", + "content_hash", + "ensure_default_compositions", + "list_compositions", + "list_module_catalog", + "load_composition", + "preview_concat_labels", + "resolve_composition", + "resolve_preprocess_id", + "save_composition", + "setup_id_to_composition_id", + "validate_composition", +] diff --git a/ipred/src/ipred/conformal.py b/ipred/src/ipred/conformal.py new file mode 100644 index 0000000..fb282ea --- /dev/null +++ b/ipred/src/ipred/conformal.py @@ -0,0 +1,95 @@ +"""Mondrian split-conformal thresholds and maps from cached proba.""" + +from __future__ import annotations + +from typing import Any + +import numpy as np + +STATUS_ABSTAIN = 0 +STATUS_SINGLETON = 1 +STATUS_MULTI = 2 + + +def conformal_quantile(scores: np.ndarray, alpha: float) -> float: + """Finite-sample split-conformal quantile.""" + s = np.sort(np.asarray(scores, dtype=np.float64).ravel()) + n = s.size + if n == 0: + return 1.0 + if not 0.0 < alpha < 1.0: + raise ValueError("alpha must be in (0, 1)") + k = int(np.ceil((n + 1) * (1.0 - alpha))) + k = min(max(k, 1), n) + return float(s[k - 1]) + + +def mondrian_thresholds( + cal_scores_by_class: dict[int, np.ndarray], + alpha: float, +) -> dict[int, float]: + """Per-class conformal thresholds q_y(alpha).""" + return { + int(c): conformal_quantile(scores, alpha) + for c, scores in cal_scores_by_class.items() + } + + +def maps_from_proba( + proba: np.ndarray, + class_ids: list[int], + q_by_class: dict[int, float], +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Build commit, status, and packed membership from HxWxK proba. + + Returns: + commit (HxW uint8), status (HxW uint8), membership_bits (HxW uint32) + where bit i means class_ids[i] is in the prediction set. + """ + if proba.ndim != 3: + raise ValueError("proba must be HxWxK") + h, w, k = proba.shape + if k != len(class_ids): + raise ValueError("proba K must match class_ids length") + thr_prob = {cid: 1.0 - q for cid, q in q_by_class.items()} + in_set = np.zeros((h, w, k), dtype=bool) + for j, cid in enumerate(class_ids): + t = thr_prob.get(int(cid)) + if t is None: + continue + in_set[..., j] = proba[..., j] >= t + + set_sizes = in_set.sum(axis=2) + status = np.full((h, w), STATUS_ABSTAIN, dtype=np.uint8) + status[set_sizes == 1] = STATUS_SINGLETON + status[set_sizes > 1] = STATUS_MULTI + + commit = np.zeros((h, w), dtype=np.uint8) + singleton = set_sizes == 1 + if np.any(singleton): + cols = np.argmax(in_set[singleton], axis=1) + classes = np.asarray(class_ids, dtype=np.uint8) + commit[singleton] = classes[cols] + + membership = np.zeros((h, w), dtype=np.uint32) + for j in range(min(k, 32)): + membership |= in_set[..., j].astype(np.uint32) << j + + return commit, status, membership + + +def counts_from_status(status: np.ndarray) -> dict[str, int]: + """Count abstain / singleton / multi pixels.""" + return { + "singleton": int(np.sum(status == STATUS_SINGLETON)), + "multi": int(np.sum(status == STATUS_MULTI)), + "abstain": int(np.sum(status == STATUS_ABSTAIN)), + } + + +def cal_scores_from_json(raw: dict[str, Any]) -> dict[int, np.ndarray]: + """Parse cal_scores.json into numpy arrays.""" + out: dict[int, np.ndarray] = {} + for k, v in raw.items(): + out[int(k)] = np.asarray(v, dtype=np.float64) + return out diff --git a/ipred/src/ipred/feature_setups.py b/ipred/src/ipred/feature_setups.py new file mode 100644 index 0000000..365b02c --- /dev/null +++ b/ipred/src/ipred/feature_setups.py @@ -0,0 +1,416 @@ +"""Feature Setup shelf — procedure or weights configurations.""" + +from __future__ import annotations + +import hashlib +import json +import logging +import os +import uuid +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +from ipred.catalog import Catalog +from ipred.paths import feature_models_root + +logger = logging.getLogger(__name__) + +PROCEDURE_SKIMAGE = "skimage_multiscale_v1" +PROCEDURE_CLAHE_ENCODER = "clahe_encoder_v1" +WEIGHTS_ONNX_VISION = "onnx_vision_encoder_v1" +WEIGHTS_TOMOJEPA_MARK25 = "torch_tomojepa_mark25_v1" +WEIGHTS_TOMOJEPA_MARK11 = "torch_tomojepa_mark11_v1" +TOMOJEPA_WEIGHTS_FORMATS = frozenset( + {WEIGHTS_TOMOJEPA_MARK25, WEIGHTS_TOMOJEPA_MARK11} +) + +DEFAULT_SKIMAGE_PARAMS: dict[str, Any] = { + "sigma_min": 1.0, + "sigma_max": 8.0, + "intensity": True, + "edges": True, + "texture": True, + "clahe": True, +} + +DEFAULT_CLAHE_ENCODER_PARAMS: dict[str, Any] = { + "clahe": True, + "clip_limit": 0.01, + "kernel_size": None, + "resize": True, + "input_size": 512, + "pca_dims": 64, +} + +DEFAULT_SKIMAGE_MARK25_PARAMS: dict[str, Any] = { + **DEFAULT_SKIMAGE_PARAMS, + "resize": True, + "input_size": 512, + "pca_dims": 64, +} + + +def content_hash(meta: dict[str, Any]) -> str: + """Stable hash of setup config (excludes timestamps / name display).""" + payload = { + "id": meta.get("id"), + "kind": meta.get("kind"), + "procedure_id": meta.get("procedure_id"), + "params": meta.get("params"), + "encoder_setup_id": meta.get("encoder_setup_id"), + "weights_path": meta.get("weights_path"), + "weights_format": meta.get("weights_format"), + "inference": meta.get("inference"), + } + raw = json.dumps(payload, sort_keys=True, separators=(",", ":")) + return hashlib.sha1(raw.encode("utf-8")).hexdigest()[:16] + + +def _dir(setup_id: str) -> Path: + return feature_models_root() / setup_id + + +def _write_meta(meta: dict[str, Any]) -> Path: + dest = _dir(meta["id"]) + dest.mkdir(parents=True, exist_ok=True) + path = dest / "meta.json" + path.write_text(json.dumps(meta, indent=2), encoding="utf-8") + return path + + +def load_setup(setup_id: str) -> dict[str, Any] | None: + """Load a setup meta.json or None.""" + path = _dir(setup_id) / "meta.json" + if not path.is_file(): + return None + return json.loads(path.read_text(encoding="utf-8")) + + +def list_setups() -> list[dict[str, Any]]: + """List all Feature Setups on disk.""" + root = feature_models_root() + out: list[dict[str, Any]] = [] + for child in sorted(root.iterdir()) if root.is_dir() else []: + meta_path = child / "meta.json" + if not meta_path.is_file(): + continue + try: + meta = json.loads(meta_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError): + continue + out.append(meta) + return out + + +def save_setup( + *, + name: str, + kind: str, + procedure_id: str | None = None, + params: dict[str, Any] | None = None, + encoder_setup_id: str | None = None, + weights_path: str | None = None, + weights_format: str | None = None, + inference: dict[str, Any] | None = None, + setup_id: str | None = None, + builtin: bool = False, + catalog: Catalog | None = None, +) -> dict[str, Any]: + """Create or overwrite a Feature Setup.""" + if kind not in ("procedure", "weights"): + raise ValueError("kind must be 'procedure' or 'weights'") + sid = setup_id or uuid.uuid4().hex + now = datetime.now(timezone.utc).isoformat() + existing = load_setup(sid) + created = existing.get("created_at") if existing else now + meta: dict[str, Any] = { + "id": sid, + "name": (name or sid).strip() or sid, + "kind": kind, + "builtin": bool(builtin), + "created_at": created, + "updated_at": now, + } + if kind == "procedure": + if not procedure_id: + raise ValueError("procedure setups require procedure_id") + meta["procedure_id"] = procedure_id + if params is not None: + meta["params"] = dict(params) + elif procedure_id == PROCEDURE_CLAHE_ENCODER: + meta["params"] = dict(DEFAULT_CLAHE_ENCODER_PARAMS) + else: + meta["params"] = dict(DEFAULT_SKIMAGE_PARAMS) + if encoder_setup_id: + meta["encoder_setup_id"] = encoder_setup_id + else: + if not weights_path: + raise ValueError("weights setups require weights_path") + meta["weights_path"] = weights_path + meta["weights_format"] = weights_format or WEIGHTS_ONNX_VISION + meta["inference"] = dict(inference or {}) + + meta["content_hash"] = content_hash(meta) + _write_meta(meta) + if catalog is not None: + catalog.upsert_feature_setup( + setup_id=sid, + name=meta["name"], + kind=kind, + content_hash=meta["content_hash"], + meta=meta, + ) + return meta + + +def resolve_slimsam_weights_path() -> str | None: + """Locate SlimSAM ONNX; honors FEATURE_ENCODER_ONNX.""" + env = os.getenv("FEATURE_ENCODER_ONNX") + if env: + p = Path(env).expanduser().resolve() + if p.is_file(): + return str(p) + # Prefer repo-adjacent locations without importing annotate backend. + here = Path(__file__).resolve() + repo = here.parents[3] # .../repo/ipred/src/ipred → repo + candidates = [ + repo / "backend" / "models" / "slimsam-77-uniform" / "onnx" / "vision_encoder.onnx", + repo / "frontend" / "public" / "models" / "slimsam-77-uniform" / "onnx" / "vision_encoder.onnx", + ] + for c in candidates: + if c.is_file(): + return str(c.resolve()) + return None + + +def resolve_tomojepa_weights_path( + filename: str = "tomojepa25.pth", + *, + env_var: str | None = "TOMOJEPA_WEIGHTS", +) -> str | None: + """Locate a TomoJEPA ``.pth``; honors ``env_var`` when set.""" + if env_var: + env = os.getenv(env_var) + if env: + p = Path(env).expanduser().resolve() + if p.is_file(): + return str(p) + here = Path(__file__).resolve() + candidates = [ + here.parents[2] / "models" / filename, + here.parents[3] / "ipred" / "models" / filename, + ] + for c in candidates: + if c.is_file(): + return str(c.resolve()) + return None + + +def _upsert_existing(meta: dict[str, Any], catalog: Catalog | None) -> None: + if catalog is None: + return + catalog.upsert_feature_setup( + setup_id=meta["id"], + name=meta["name"], + kind=meta["kind"], + content_hash=meta.get("content_hash") or content_hash(meta), + meta=meta, + ) + + +def _ensure_procedure_with_param_backfill( + *, + setup_id: str, + name: str, + procedure_id: str, + params: dict[str, Any], + encoder_setup_id: str | None, + catalog: Catalog | None, +) -> dict[str, Any]: + """Create procedure setup or backfill missing default params.""" + existing = load_setup(setup_id) + if existing is None: + return save_setup( + setup_id=setup_id, + name=name, + kind="procedure", + procedure_id=procedure_id, + params=dict(params), + encoder_setup_id=encoder_setup_id, + builtin=True, + catalog=catalog, + ) + cur = dict(existing.get("params") or {}) + changed = False + for key, val in params.items(): + if key not in cur: + cur[key] = val + changed = True + if changed: + return save_setup( + setup_id=setup_id, + name=existing.get("name") or name, + kind="procedure", + procedure_id=procedure_id, + params=cur, + encoder_setup_id=existing.get("encoder_setup_id") or encoder_setup_id, + builtin=True, + catalog=catalog, + ) + _upsert_existing(existing, catalog) + return existing + + +def _ensure_tomojepa_variant( + catalog: Catalog | None, + *, + weights_id: str, + weights_name: str, + combo_id: str, + combo_name: str, + clahe_id: str, + clahe_name: str, + filename: str, + env_var: str, + weights_format: str, +) -> list[dict[str, Any]]: + """Weights + skimage combo + CLAHE procedure for one TomoJEPA checkpoint.""" + out: list[dict[str, Any]] = [] + tomo_path = resolve_tomojepa_weights_path(filename, env_var=env_var) or "" + weights = load_setup(weights_id) + if weights is None: + weights = save_setup( + setup_id=weights_id, + name=weights_name, + kind="weights", + weights_path=tomo_path or "(missing)", + weights_format=weights_format, + inference={"input_size": 512, "pca_dims": 32}, + builtin=True, + catalog=catalog, + ) + else: + _upsert_existing(weights, catalog) + out.append(weights) + + out.append( + _ensure_procedure_with_param_backfill( + setup_id=combo_id, + name=combo_name, + procedure_id=PROCEDURE_SKIMAGE, + params=dict(DEFAULT_SKIMAGE_MARK25_PARAMS), + encoder_setup_id=weights_id, + catalog=catalog, + ) + ) + out.append( + _ensure_procedure_with_param_backfill( + setup_id=clahe_id, + name=clahe_name, + procedure_id=PROCEDURE_CLAHE_ENCODER, + params=dict(DEFAULT_CLAHE_ENCODER_PARAMS), + encoder_setup_id=weights_id, + catalog=catalog, + ) + ) + return out + + +def ensure_default_setups(catalog: Catalog | None = None) -> list[dict[str, Any]]: + """Create built-in setups if missing; return all defaults.""" + created: list[dict[str, Any]] = [] + sk = load_setup("default-skimage") + if sk is None: + sk = save_setup( + setup_id="default-skimage", + name="Default skimage multiscale", + kind="procedure", + procedure_id=PROCEDURE_SKIMAGE, + params=dict(DEFAULT_SKIMAGE_PARAMS), + builtin=True, + catalog=catalog, + ) + else: + _upsert_existing(sk, catalog) + created.append(sk) + + weights_path = resolve_slimsam_weights_path() or "" + slim = load_setup("default-slimsam") + if slim is None: + slim = save_setup( + setup_id="default-slimsam", + name="Default SlimSAM encoder", + kind="weights", + weights_path=weights_path or "(missing)", + weights_format=WEIGHTS_ONNX_VISION, + inference={"input_size": 1024, "pca_dims": 32}, + builtin=True, + catalog=catalog, + ) + else: + _upsert_existing(slim, catalog) + created.append(slim) + + combo = load_setup("default-skimage-slimsam") + if combo is None: + combo = save_setup( + setup_id="default-skimage-slimsam", + name="Default skimage + SlimSAM", + kind="procedure", + procedure_id=PROCEDURE_SKIMAGE, + params=dict(DEFAULT_SKIMAGE_PARAMS), + encoder_setup_id="default-slimsam", + builtin=True, + catalog=catalog, + ) + else: + _upsert_existing(combo, catalog) + created.append(combo) + + created.append( + _ensure_procedure_with_param_backfill( + setup_id="default-slimsam-clahe", + name="Default SlimSAM + CLAHE", + procedure_id=PROCEDURE_CLAHE_ENCODER, + params=dict(DEFAULT_CLAHE_ENCODER_PARAMS), + encoder_setup_id="default-slimsam", + catalog=catalog, + ) + ) + + for variant in ( + { + "weights_id": "default-mark25", + "weights_name": "Default Mark25 TomoJEPA encoder", + "combo_id": "default-skimage-mark25", + "combo_name": "Default skimage + Mark25", + "clahe_id": "default-mark25-clahe", + "clahe_name": "Default Mark25 + CLAHE", + "filename": "tomojepa25.pth", + "env_var": "TOMOJEPA_WEIGHTS", + "weights_format": WEIGHTS_TOMOJEPA_MARK25, + }, + { + "weights_id": "default-mark11", + "weights_name": "Default Mark11 TomoJEPA encoder", + "combo_id": "default-skimage-mark11", + "combo_name": "Default skimage + Mark11", + "clahe_id": "default-mark11-clahe", + "clahe_name": "Default Mark11 + CLAHE", + "filename": "tomojepa11.pth", + "env_var": "TOMOJEPA11_WEIGHTS", + "weights_format": WEIGHTS_TOMOJEPA_MARK11, + }, + ): + created.extend(_ensure_tomojepa_variant(catalog, **variant)) + + return created + + +def resolve_setup(setup_id: str) -> dict[str, Any]: + """Load setup or raise KeyError.""" + meta = load_setup(setup_id) + if meta is None: + raise KeyError(f"unknown feature setup {setup_id}") + return meta diff --git a/ipred/src/ipred/features.py b/ipred/src/ipred/features.py new file mode 100644 index 0000000..faf185b --- /dev/null +++ b/ipred/src/ipred/features.py @@ -0,0 +1,180 @@ +"""Multiscale feature computation (ported; no annotate-backend imports).""" + +from __future__ import annotations + +from io import BytesIO + +import numpy as np +from PIL import Image as PILImage +from skimage.exposure import equalize_adapthist +from skimage.feature import multiscale_basic_features + + +def list_sigmas(sigma_min: float, sigma_max: float, num_sigma: int | None = None) -> np.ndarray: + """Return the σ grid used by multiscale_basic_features.""" + if sigma_min <= 0 or sigma_max < sigma_min: + raise ValueError(f"invalid sigma range [{sigma_min}, {sigma_max}]") + if num_sigma is None: + num_sigma = int(np.log2(sigma_max) - np.log2(sigma_min) + 1) + return np.logspace( + np.log2(sigma_min), + np.log2(sigma_max), + num=int(num_sigma), + base=2, + endpoint=True, + ) + + +def _fmt_sigma(sigma: float) -> str: + if abs(sigma - round(sigma)) < 1e-6: + return str(int(round(sigma))) + return f"{sigma:g}" + + +def feature_channel_labels( + *, + sigma_min: float = 1.0, + sigma_max: float = 8.0, + intensity: bool = True, + edges: bool = True, + texture: bool = True, + num_sigma: int | None = None, +) -> list[str]: + """Human-readable labels matching skimage per-σ order.""" + if not any((intensity, edges, texture)): + raise ValueError("at least one of intensity, edges, texture must be True") + labels: list[str] = [] + for sigma in list_sigmas(sigma_min, sigma_max, num_sigma): + s = _fmt_sigma(float(sigma)) + if intensity: + labels.append(f"intensity σ={s}") + if edges: + labels.append(f"edges σ={s}") + if texture: + labels.append(f"texture λ− σ={s}") + labels.append(f"texture λ+ σ={s}") + return labels + + +def to_grayscale(arr: np.ndarray) -> np.ndarray: + """Convert a slice to float grayscale in [0, 1].""" + a = np.asarray(arr) + if a.ndim == 3 and a.shape[-1] in (3, 4): + rgb = a[..., :3].astype(np.float64) + gray = 0.299 * rgb[..., 0] + 0.587 * rgb[..., 1] + 0.114 * rgb[..., 2] + elif a.ndim == 2: + gray = a.astype(np.float64) + else: + raise ValueError(f"unsupported array shape {a.shape}") + finite = gray[np.isfinite(gray)] + if finite.size == 0: + return np.zeros_like(gray, dtype=np.float64) + lo, hi = float(np.min(finite)), float(np.max(finite)) + if hi <= lo: + return np.zeros_like(gray, dtype=np.float64) + return np.clip((gray - lo) / (hi - lo), 0.0, 1.0) + + +def _float_feats_to_uint8(feats: np.ndarray, *, clahe: bool) -> np.ndarray: + """Convert float HxWxC to display uint8.""" + h, w, c = feats.shape + out = np.empty((h, w, c), dtype=np.uint8) + for i in range(c): + ch = feats[..., i].astype(np.float64) + lo, hi = float(np.nanmin(ch)), float(np.nanmax(ch)) + norm = (ch - lo) / (hi - lo) if hi > lo else np.zeros_like(ch) + if clahe: + norm = equalize_adapthist(np.clip(norm, 0.0, 1.0), clip_limit=0.01) + out[..., i] = np.clip(np.round(norm * 255.0), 0, 255).astype(np.uint8) + return out + + +def apply_clahe( + gray: np.ndarray, + *, + clip_limit: float = 0.01, + kernel_size: int | None = None, +) -> np.ndarray: + """CLAHE on a [0, 1] grayscale image; returns float32 in [0, 1].""" + g = np.clip(np.asarray(gray, dtype=np.float64), 0.0, 1.0) + kwargs: dict = {"clip_limit": float(clip_limit)} + if kernel_size is not None and int(kernel_size) > 0: + kwargs["kernel_size"] = int(kernel_size) + out = equalize_adapthist(g, **kwargs) + return np.asarray(out, dtype=np.float32) + + +def compute_clahe_stack( + gray: np.ndarray, + *, + clip_limit: float = 0.01, + kernel_size: int | None = None, + apply: bool = True, +) -> tuple[np.ndarray, np.ndarray, list[str]]: + """Return single-channel ``(uint8_stack, float_stack, labels)`` with optional CLAHE. + + Args: + gray: Grayscale float image in ``[0, 1]``. + clip_limit: skimage ``equalize_adapthist`` clip limit. + kernel_size: Optional CLAHE tile size; ``None`` uses skimage default. + apply: When False, store min-max gray only (no CLAHE). + """ + g = np.asarray(gray, dtype=np.float32) + if apply: + ch = apply_clahe(g, clip_limit=clip_limit, kernel_size=kernel_size) + label = "clahe" + else: + ch = np.clip(g, 0.0, 1.0).astype(np.float32) + label = "intensity" + float_stack = ch[..., None] + uint8_stack = np.clip(np.round(float_stack * 255.0), 0, 255).astype(np.uint8) + return uint8_stack, float_stack, [label] + + +def compute_feature_stacks( + gray: np.ndarray, + *, + sigma_min: float = 1.0, + sigma_max: float = 8.0, + intensity: bool = True, + edges: bool = True, + texture: bool = True, + clahe: bool = True, + num_sigma: int | None = None, +) -> tuple[np.ndarray, np.ndarray, list[str]]: + """Return ``(uint8_stack, float_stack, labels)``.""" + labels = feature_channel_labels( + sigma_min=sigma_min, + sigma_max=sigma_max, + intensity=intensity, + edges=edges, + texture=texture, + num_sigma=num_sigma, + ) + feats = multiscale_basic_features( + np.asarray(gray, dtype=np.float32), + intensity=intensity, + edges=edges, + texture=texture, + sigma_min=sigma_min, + sigma_max=sigma_max, + num_sigma=num_sigma, + workers=1, + ) + if feats.shape[-1] != len(labels): + raise RuntimeError( + f"feature count mismatch: got {feats.shape[-1]}, expected {len(labels)}" + ) + float_stack = np.nan_to_num(feats.astype(np.float32), nan=0.0, posinf=0.0, neginf=0.0) + uint8_stack = _float_feats_to_uint8(float_stack, clahe=clahe) + return uint8_stack, float_stack, labels + + +def encode_channel_png(stack: np.ndarray, index: int) -> bytes: + """Encode one feature channel as grayscale PNG.""" + if index < 0 or index >= stack.shape[-1]: + raise IndexError(f"channel index {index} out of range") + img = PILImage.fromarray(stack[..., index], mode="L") + buf = BytesIO() + img.save(buf, format="PNG") + return buf.getvalue() diff --git a/ipred/src/ipred/labels.py b/ipred/src/ipred/labels.py new file mode 100644 index 0000000..474d737 --- /dev/null +++ b/ipred/src/ipred/labels.py @@ -0,0 +1,150 @@ +"""Rasterize annotation shapes to a sparse label map (self-contained).""" + +from __future__ import annotations + +from typing import Any + +import numpy as np +from PIL import Image as PILImage +from PIL import ImageDraw + + +def build_label_map(shapes: list[dict[str, Any]], height: int, width: int) -> np.ndarray: + """Compose uint8 label map (0 = unlabeled); last shape wins.""" + labels = np.zeros((height, width), dtype=np.uint8) + for shape in shapes: + class_id = int(shape.get("classId") or shape.get("class_id") or 0) + if class_id <= 0 or class_id > 255: + continue + mask = shape_to_mask(shape, height, width) + labels[mask] = class_id + return labels + + +def shape_to_mask(shape: dict[str, Any], h: int, w: int) -> np.ndarray: + """Rasterize one shape to boolean mask.""" + kind = shape.get("kind") + if kind == "rectangle": + return _rect_mask( + float(shape["x"]), + float(shape["y"]), + float(shape["w"]), + float(shape["h"]), + h, + w, + ) + if kind == "ellipse": + return _ellipse_mask( + float(shape["cx"]), + float(shape["cy"]), + float(shape["rx"]), + float(shape["ry"]), + h, + w, + ) + if kind == "polygon": + return _polygon_mask(shape.get("points") or [], h, w, shape.get("holes")) + if kind == "brush": + return _brush_mask(shape.get("strokes") or [], h, w) + raise ValueError(f"Unknown shape kind: {kind!r}") + + +def _rect_mask(x: float, y: float, ww: float, hh: float, h: int, w: int) -> np.ndarray: + x0 = max(0, int(np.floor(min(x, x + ww)))) + x1 = min(w, int(np.ceil(max(x, x + ww)))) + y0 = max(0, int(np.floor(min(y, y + hh)))) + y1 = min(h, int(np.ceil(max(y, y + hh)))) + mask = np.zeros((h, w), dtype=bool) + if x1 > x0 and y1 > y0: + mask[y0:y1, x0:x1] = True + return mask + + +def _ellipse_mask( + cx: float, cy: float, rx: float, ry: float, h: int, w: int +) -> np.ndarray: + yy, xx = np.ogrid[:h, :w] + rx = max(rx, 1e-6) + ry = max(ry, 1e-6) + return ((xx - cx) / rx) ** 2 + ((yy - cy) / ry) ** 2 <= 1.0 + + +def _xy_pairs(points: list[Any]) -> list[tuple[float, float]]: + """Normalize studio point payloads to ``(x, y)`` pairs. + + Studio wire format for polygons/brushes is a flat ``[x0, y0, x1, y1, ...]`` + float list (same as coco_export). Nested ``[[x,y], ...]`` and ``{x,y}`` + dicts are also accepted. + """ + if not points: + return [] + if all(isinstance(p, (int, float)) for p in points): + if len(points) < 2: + return [] + return [ + (float(points[i]), float(points[i + 1])) + for i in range(0, len(points) - 1, 2) + ] + xy: list[tuple[float, float]] = [] + for p in points: + if isinstance(p, (list, tuple)) and len(p) >= 2: + xy.append((float(p[0]), float(p[1]))) + elif isinstance(p, dict): + xy.append((float(p["x"]), float(p["y"]))) + return xy + + +def _polygon_mask( + points: list[Any], h: int, w: int, holes: list[Any] | None = None +) -> np.ndarray: + """Rasterize the outer ring, then carve out each hole ring. + + Holes come from the studio's clip-to-other-classes tool (see + ``clipToClasses.ts``), which represents "this region minus already-labeled + neighbor classes" as a polygon with holes rather than reshaping the outer + ring. Ignoring them would paint the full outer extent — including pixels + that visually belong to other classes on the canvas. + """ + xy = _xy_pairs(points) + if len(xy) < 3: + return np.zeros((h, w), dtype=bool) + img = PILImage.new("L", (w, h), 0) + draw = ImageDraw.Draw(img) + draw.polygon(xy, outline=1, fill=1) + for hole in holes or []: + hxy = _xy_pairs(hole) + if len(hxy) >= 3: + draw.polygon(hxy, outline=0, fill=0) + return np.asarray(img, dtype=np.uint8) > 0 + + +def _brush_mask(strokes: list[dict[str, Any]], h: int, w: int) -> np.ndarray: + """Paint OR / erase AND-NOT along stroke points.""" + mask = np.zeros((h, w), dtype=bool) + for stroke in strokes: + pts = stroke.get("points") or [] + xy = _xy_pairs(pts) + if not xy: + continue + # Studio stamps use radius directly; legacy payloads may send diameter as size. + if stroke.get("radius") is not None: + radius = max(1, int(round(float(stroke["radius"])))) + else: + radius = max(1, int(round(float(stroke.get("size") or 3) / 2))) + mode = str(stroke.get("mode") or stroke.get("type") or "paint") + erase = mode in ("erase", "eraser") + layer = PILImage.new("L", (w, h), 0) + ld = ImageDraw.Draw(layer) + if len(xy) == 1: + x, y = xy[0] + ld.ellipse((x - radius, y - radius, x + radius, y + radius), fill=1) + else: + ld.line(xy, fill=1, width=max(1, radius * 2)) + for x, y in xy: + ld.ellipse((x - radius, y - radius, x + radius, y + radius), fill=1) + layer_m = np.asarray(layer, dtype=np.uint8) > 0 + if erase: + mask &= ~layer_m + else: + mask |= layer_m + return mask diff --git a/ipred/src/ipred/manifold.py b/ipred/src/ipred/manifold.py new file mode 100644 index 0000000..56cafdd --- /dev/null +++ b/ipred/src/ipred/manifold.py @@ -0,0 +1,359 @@ +"""Greedy feature-variance box sampling for annotation guidance. + +Sliding windows scored by whitened feature variance; picks suppress nearby +(spatial Chebyshev exclusion ≥ box side) and feature-similar boxes. +""" + +from __future__ import annotations + +import logging +import math +import uuid +from dataclasses import dataclass +from io import BytesIO +from typing import Any + +import numpy as np +from PIL import Image as PILImage +from sklearn.decomposition import PCA +from sklearn.preprocessing import StandardScaler + +from ipred.cache import TTLCache + +logger = logging.getLogger(__name__) + +_sample_cache: TTLCache = TTLCache(ttl_seconds=600.0, max_entries=32) +_COS_HARD = 0.95 +_COS_ALPHA = 2.0 + + +@dataclass(frozen=True) +class ManifoldSample: + """Inducing boxes + residual interestingness heatmap.""" + + job_id: str + points: list[dict[str, Any]] + heatmap: np.ndarray # HxW float32 in [0, 1] + meta: dict[str, Any] + + +@dataclass(frozen=True) +class CachedManifoldSample: + """Server-side manifold sample for PNG GET.""" + + sample_id: str + heatmap_png: bytes + points: list[dict[str, Any]] + meta: dict[str, Any] + + +def derived_exclusion_radius(h: int, w: int, k: int) -> float: + """Disk radius so roughly K disks pack the image plane.""" + k = max(1, int(k)) + area = float(h * w) + r = math.sqrt(area / (math.pi * k)) + r_max = min(h, w) / 4.0 + return float(max(8.0, min(r, r_max))) + + +def sample_inducing_points( + float_stack: np.ndarray, + *, + feature_id: str, + k: int = 24, + box_size: int | None = None, + stride: int | None = None, + pca_dims: int = 16, + seed: int = 0, + mask: np.ndarray | None = None, +) -> ManifoldSample: + """Greedy variance boxes with spatial + feature-space exclusion. + + Args: + float_stack: Feature bank HxWxC. + feature_id: Owning feature bank id (stored on the sample). + k: Number of inducing boxes (clamped to [2, 128]). + box_size: Full square side length in image pixels. When omitted, derived + from the K-based exclusion radius (``2 * round(r)``). + stride: Optional hop override; default ``max(4, box_half)``. + pca_dims: PCA dimensionality after standardization. + seed: RNG for PCA solver stability. + mask: Optional HxW placement mask. When set, only windows whose full + box lies inside the mask are eligible. + + Returns: + ManifoldSample with box centers, boxes, radius, and residual heatmap. + """ + feats = np.asarray(float_stack, dtype=np.float32) + if feats.ndim != 3: + raise ValueError(f"float_stack must be HxWxC, got {feats.shape}") + h, w, c = feats.shape + if c < 1: + raise ValueError("empty feature stack") + + place_mask: np.ndarray | None = None + mask_pixels = 0 + if mask is not None: + m = np.asarray(mask) + if m.shape != (h, w): + raise ValueError(f"mask shape {m.shape} does not match image {(h, w)}") + place_mask = m.astype(bool, copy=False) + mask_pixels = int(place_mask.sum()) + if mask_pixels < 1: + raise ValueError("placement mask is empty") + + k = int(max(2, min(128, k))) + r_pack = derived_exclusion_radius(h, w, k) + if box_size is None: + side = int(max(8, 2 * round(r_pack))) + else: + side = int(box_size) + # Clamp: at least 8px side, at most half the short image edge + side = int(max(8, min(side, min(h, w) // 2 * 2))) # even-ish via floor + if side % 2 == 1: + side += 1 + b = max(4, side // 2) # box half-size + # Spatial exclusion in Chebyshev (L∞) distance so axis-aligned boxes do + # not overlap: centers must be ≥ side+1 apart (1px gap for stroke). + # Also respect K-packing floor. + r = float(max(r_pack, side + 1)) + hop = int(stride) if stride is not None else int(max(4, b)) + + # Coarse feature grid for PCA fit + window stats + grid_stride = max(2, hop // 2) + gy = np.arange(0, h, grid_stride, dtype=np.int32) + gx = np.arange(0, w, grid_stride, dtype=np.int32) + yy, xx = np.meshgrid(gy, gx, indexing="ij") + flat_y = yy.ravel() + flat_x = xx.ravel() + X = feats[flat_y, flat_x, :].astype(np.float64) + X = np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0) + n = X.shape[0] + if n < 4: + raise ValueError("too few pixels to score boxes") + + scaler = StandardScaler() + Xs = scaler.fit_transform(X) + d = int(max(1, min(pca_dims, c, n - 1, 16))) + pca = PCA(n_components=d, random_state=seed) + Z = pca.fit_transform(Xs) + explained = float(np.sum(pca.explained_variance_ratio_)) + + # Map (y,x) -> Z row via nearest grid sample (for window aggregation) + # Store Z as grid gh×gw×d + gh, gw = len(gy), len(gx) + Z_grid = Z.reshape(gh, gw, d) + + # Sliding window centers on hop grid + cy = np.arange(b, h - b + 1, hop, dtype=np.int32) + cx = np.arange(b, w - b + 1, hop, dtype=np.int32) + if cy.size == 0: + cy = np.array([h // 2], dtype=np.int32) + if cx.size == 0: + cx = np.array([w // 2], dtype=np.int32) + win_yy, win_xx = np.meshgrid(cy, cx, indexing="ij") + centers_y = win_yy.ravel() + centers_x = win_xx.ravel() + n_win = int(centers_y.size) + + scores = np.zeros(n_win, dtype=np.float64) + means = np.zeros((n_win, d), dtype=np.float64) + + for i in range(n_win): + y0 = int(centers_y[i] - b) + y1 = int(centers_y[i] + b) + x0 = int(centers_x[i] - b) + x1 = int(centers_x[i] + b) + # Grid indices covering the box + gi0 = max(0, (y0 + grid_stride - 1) // grid_stride) + gi1 = min(gh, (y1 // grid_stride) + 1) + gj0 = max(0, (x0 + grid_stride - 1) // grid_stride) + gj1 = min(gw, (x1 // grid_stride) + 1) + patch = Z_grid[gi0:gi1, gj0:gj1, :].reshape(-1, d) + if patch.shape[0] < 2: + continue + means[i] = patch.mean(axis=0) + # tr(Cov) = sum of feature variances + scores[i] = float(np.var(patch, axis=0).sum()) + + residual = scores.copy() + n_windows_in_mask = n_win + if place_mask is not None: + for i in range(n_win): + cx_i = int(centers_x[i]) + cy_i = int(centers_y[i]) + x0 = max(0, cx_i - b) + y0 = max(0, cy_i - b) + x1 = min(w, cx_i + b) + y1 = min(h, cy_i + b) + if x1 <= x0 or y1 <= y0 or not bool(place_mask[y0:y1, x0:x1].all()): + residual[i] = 0.0 + scores[i] = 0.0 + n_windows_in_mask = int(np.count_nonzero(residual > 0)) + if n_windows_in_mask < 1: + raise ValueError( + "no candidate boxes fit fully inside the placement mask; " + "draw a larger mask or reduce box size" + ) + + points: list[dict[str, Any]] = [] + + for pick_i in range(k): + if not np.any(residual > 0): + break + i_star = int(np.argmax(residual)) + score = float(residual[i_star]) + if score <= 0: + break + cx_i = int(centers_x[i_star]) + cy_i = int(centers_y[i_star]) + mu = means[i_star] + mu_norm = float(np.linalg.norm(mu)) + x0 = max(0, cx_i - b) + y0 = max(0, cy_i - b) + # Exclusive max so Konva/CSS width = x1 - x0 equals 2*b when unclipped + x1 = min(w, cx_i + b) + y1 = min(h, cy_i + b) + points.append( + { + "x": cx_i, + "y": cy_i, + "cluster": pick_i, + "score": score, + "radius": float(r), + "box_size": int(2 * b), + "box": {"x0": x0, "y0": y0, "x1": x1, "y1": y1}, + } + ) + residual[i_star] = 0.0 + + # Spatial (Chebyshev) + feature suppress. Exclude with ≥ side so + # axis-aligned boxes of width `side` do not share interior area. + for j in range(n_win): + if residual[j] <= 0: + continue + dx = abs(float(centers_x[j] - cx_i)) + dy = abs(float(centers_y[j] - cy_i)) + if max(dx, dy) < r: + residual[j] = 0.0 + continue + mj = means[j] + mj_norm = float(np.linalg.norm(mj)) + if mu_norm < 1e-12 or mj_norm < 1e-12: + continue + cos = float(np.dot(mu, mj) / (mu_norm * mj_norm)) + cos = max(0.0, cos) + if cos > _COS_HARD: + residual[j] = 0.0 + else: + residual[j] *= 1.0 - (cos ** _COS_ALPHA) + + # Hard guarantee: drop any pick whose AABB still overlaps an earlier one + # (protects against hop / clip / float edge cases). + points = _dedupe_overlapping_boxes(points) + + heatmap = _residual_heatmap( + h, w, centers_y, centers_x, residual, hop=hop, half=b + ) + + meta = { + "k": k, + "n_picked": len(points), + "n_windows": n_win, + "n_windows_in_mask": n_windows_in_mask, + "n_subsample": int(n), + "pca_dims": d, + "explained_variance": explained, + "stride": hop, + "radius": float(r), + "box_half": b, + "box_size": int(2 * b), + "mask_pixels": mask_pixels, + "has_mask": place_mask is not None, + } + return ManifoldSample(job_id=feature_id, points=points, heatmap=heatmap, meta=meta) + + +def _boxes_overlap(a: dict[str, int], b: dict[str, int]) -> bool: + """True if two exclusive-end AABBs share interior area.""" + return a["x0"] < b["x1"] and b["x0"] < a["x1"] and a["y0"] < b["y1"] and b["y0"] < a["y1"] + + +def _dedupe_overlapping_boxes(points: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Keep greedy order; drop later picks whose box overlaps an earlier one.""" + kept: list[dict[str, Any]] = [] + for p in points: + box = p.get("box") + if not isinstance(box, dict): + kept.append(p) + continue + if any(_boxes_overlap(box, k["box"]) for k in kept if isinstance(k.get("box"), dict)): + continue + kept.append(p) + for i, p in enumerate(kept): + p["cluster"] = i + return kept + + +def _residual_heatmap( + h: int, + w: int, + centers_y: np.ndarray, + centers_x: np.ndarray, + residual: np.ndarray, + *, + hop: int, + half: int, +) -> np.ndarray: + """Paint residual scores onto a coarse grid and upsample to HxW.""" + out = np.zeros((h, w), dtype=np.float32) + max_s = float(np.max(residual)) if residual.size else 0.0 + if max_s < 1e-12: + return out + for i in range(residual.size): + if residual[i] <= 0: + continue + val = float(residual[i] / max_s) + cy = int(centers_y[i]) + cx = int(centers_x[i]) + y0 = max(0, cy - half) + y1 = min(h, cy + half) + x0 = max(0, cx - half) + x1 = min(w, cx + half) + # Take max so overlapping windows keep strongest residual + patch = out[y0:y1, x0:x1] + np.maximum(patch, val, out=patch) + return out + + +def encode_heatmap_png(heatmap: np.ndarray) -> bytes: + """Encode HxW float [0,1] residual map as grayscale PNG (0–255).""" + u8 = np.clip(np.asarray(heatmap, dtype=np.float32) * 255.0, 0, 255).astype(np.uint8) + img = PILImage.fromarray(u8, mode="L") + buf = BytesIO() + img.save(buf, format="PNG") + return buf.getvalue() + + +def store_sample(result: ManifoldSample) -> CachedManifoldSample: + """Cache heatmap PNG + points; return handle for GET.""" + sample_id = uuid.uuid4().hex + cached = CachedManifoldSample( + sample_id=sample_id, + heatmap_png=encode_heatmap_png(result.heatmap), + points=list(result.points), + meta={ + "sample_id": sample_id, + "job_id": result.job_id, + **result.meta, + "points": list(result.points), + }, + ) + _sample_cache.set(sample_id, cached) + return cached + + +def get_sample(sample_id: str) -> CachedManifoldSample | None: + """Look up a cached manifold sample.""" + s = _sample_cache.get(sample_id) + return s if isinstance(s, CachedManifoldSample) else None diff --git a/ipred/src/ipred/manifold_jobs.py b/ipred/src/ipred/manifold_jobs.py new file mode 100644 index 0000000..a806386 --- /dev/null +++ b/ipred/src/ipred/manifold_jobs.py @@ -0,0 +1,74 @@ +"""Orchestrate manifold suggest on persisted feature banks.""" + +from __future__ import annotations + +from typing import Any + +import numpy as np + +from ipred import manifold +from ipred.catalog import Catalog +from ipred.labels import shape_to_mask +from ipred.preprocess import load_feature_bank_arrays + + +def run_manifold_sample( + catalog: Catalog, + *, + feature_id: str, + k: int = 24, + box_size: int | None = None, + stride: int | None = None, + pca_dims: int = 16, + shapes: list[dict[str, Any]] | None = None, +) -> dict[str, Any]: + """Sample inducing boxes on a feature bank; return API response dict.""" + row = catalog.get_feature_bank(feature_id) + if row is None: + raise KeyError(f"unknown feature bank {feature_id}") + bank = load_feature_bank_arrays(row["blob_dir"]) + float_stack = bank["float_stack"] + h, w = float_stack.shape[:2] + + place_mask = None + if shapes: + place_mask = np.zeros((h, w), dtype=bool) + for shape in shapes: + place_mask |= shape_to_mask(shape, h, w) + if not bool(place_mask.any()): + raise ValueError("placement mask is empty") + + result = manifold.sample_inducing_points( + float_stack, + feature_id=feature_id, + k=k, + box_size=box_size, + stride=stride, + pca_dims=pca_dims, + mask=place_mask, + ) + cached = manifold.store_sample(result) + return { + "sample_id": cached.sample_id, + "feature_id": feature_id, + "points": cached.points, + "k": result.meta["k"], + "n_picked": result.meta["n_picked"], + "n_subsample": result.meta["n_subsample"], + "pca_dims": result.meta["pca_dims"], + "explained_variance": result.meta["explained_variance"], + "stride": result.meta["stride"], + "radius": result.meta["radius"], + "box_size": result.meta["box_size"], + "mask_pixels": result.meta.get("mask_pixels", 0), + "n_windows_in_mask": result.meta.get("n_windows_in_mask"), + "has_mask": result.meta.get("has_mask", False), + } + + +def heatmap_png(sample_id: str) -> bytes: + """Return cached heatmap PNG bytes.""" + cached = manifold.get_sample(sample_id) + if cached is None: + raise KeyError(f"unknown manifold sample {sample_id}") + return cached.heatmap_png diff --git a/ipred/src/ipred/modules/__init__.py b/ipred/src/ipred/modules/__init__.py new file mode 100644 index 0000000..78937eb --- /dev/null +++ b/ipred/src/ipred/modules/__init__.py @@ -0,0 +1,73 @@ +"""Module registry — discover and list FeatureModules.""" + +from __future__ import annotations + +from typing import Any + +from ipred.modules.base import FeatureModule, ModuleMeta +from ipred.modules.clahe_mod import ClaheModule +from ipred.modules.pca_mod import PcaModule +from ipred.modules.skimage_mod import SkimageMultiscaleModule +from ipred.modules.slimsam_mod import SlimSamModule +from ipred.modules.tomojepa_mod import TomoJepaModule + +_REGISTRY: dict[str, FeatureModule] | None = None + + +def _build_registry() -> dict[str, FeatureModule]: + mods: list[FeatureModule] = [ + SkimageMultiscaleModule(), + ClaheModule(), + SlimSamModule(), + TomoJepaModule(), + PcaModule(), + ] + return {m.meta.id: m for m in mods} + + +def get_registry() -> dict[str, FeatureModule]: + """Return module id → instance map.""" + global _REGISTRY + if _REGISTRY is None: + _REGISTRY = _build_registry() + return _REGISTRY + + +def get_module(module_id: str) -> FeatureModule: + """Lookup module or raise KeyError.""" + reg = get_registry() + if module_id not in reg: + raise KeyError(f"unknown module {module_id}") + return reg[module_id] + + +def list_module_catalog() -> list[dict[str, Any]]: + """JSON-serializable catalog for GET /modules.""" + out: list[dict[str, Any]] = [] + for mod in get_registry().values(): + meta: ModuleMeta = mod.meta + # Prefer ONNX in catalog only when graph matches default input_size + runtime = meta.runtime + if meta.id == "tomojepa": + from ipred import tomojepa_onnx + + onnx = TomoJepaModule()._onnx_path({"weights_id": "mark25"}) + if ( + onnx is not None + and tomojepa_onnx.onnx_matches_input_size(onnx, 512) + ): + runtime = "onnx" + out.append( + { + "id": meta.id, + "name": meta.name, + "description": meta.description, + "runtime": runtime, + "ready": bool(mod.ready()), + "accepts_input_from": meta.accepts_input_from, + "produces_channels": meta.produces_channels, + "produces_embedding": meta.produces_embedding, + "params_schema": meta.params_schema, + } + ) + return out diff --git a/ipred/src/ipred/modules/base.py b/ipred/src/ipred/modules/base.py new file mode 100644 index 0000000..7e1e6e7 --- /dev/null +++ b/ipred/src/ipred/modules/base.py @@ -0,0 +1,65 @@ +"""Feature module protocol — typed channel producers for compositions.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Protocol + +import numpy as np + + +@dataclass +class ModuleMeta: + """Catalog entry for a feature module.""" + + id: str + name: str + description: str + runtime: str # onnx | torch | numpy + params_schema: dict[str, Any] = field(default_factory=dict) + accepts_input_from: bool = False + produces_channels: bool = True + produces_embedding: bool = False + + +@dataclass +class ChannelBlock: + """One module's contribution to the feature bank.""" + + float_stack: np.ndarray | None = None # HxWxC + labels: list[str] = field(default_factory=list) + emb: np.ndarray | None = None # Hp×Wp×D dense grid + emb_meta: dict[str, Any] | None = None + # Optional 2-D image for downstream encoder nodes (HxW float) + image_2d: np.ndarray | None = None + + +@dataclass +class ModuleContext: + """Shared state while running a composition on one slice.""" + + raw: np.ndarray + gray: np.ndarray + params: dict[str, Any] + # Outputs of nodes already run, keyed by node id + node_outputs: dict[str, ChannelBlock] = field(default_factory=dict) + # Resolved input image for modules that accept input_from + input_image: np.ndarray | None = None + + +class FeatureModule(Protocol): + """Runnable feature module.""" + + meta: ModuleMeta + + def ready(self) -> bool: + """True when dependencies/weights are available.""" + ... + + def run(self, ctx: ModuleContext) -> ChannelBlock: + """Produce channels and/or an embedding.""" + ... + + def preview_labels(self, params: dict[str, Any]) -> list[str]: + """Labels that would appear in the bank for these params (no compute).""" + ... diff --git a/ipred/src/ipred/modules/clahe_mod.py b/ipred/src/ipred/modules/clahe_mod.py new file mode 100644 index 0000000..cc823d4 --- /dev/null +++ b/ipred/src/ipred/modules/clahe_mod.py @@ -0,0 +1,55 @@ +"""CLAHE / intensity image module (feeds encoders or contributes a channel).""" + +from __future__ import annotations + +from typing import Any + +from ipred import features +from ipred.modules.base import ChannelBlock, ModuleContext, ModuleMeta + + +class ClaheModule: + """Single-channel CLAHE (or raw intensity) plus image_2d for encoders.""" + + meta = ModuleMeta( + id="clahe", + name="CLAHE", + description="Adaptive histogram equalization → one channel + encoder input", + runtime="numpy", + accepts_input_from=False, + produces_channels=True, + params_schema={ + "clahe": {"type": "boolean", "default": True}, + "clip_limit": {"type": "number", "default": 0.01}, + "kernel_size": {"type": ["integer", "null"], "default": None}, + "include_in_bank": {"type": "boolean", "default": True}, + }, + ) + + def ready(self) -> bool: + return True + + def preview_labels(self, params: dict[str, Any]) -> list[str]: + if not bool(params.get("include_in_bank", True)): + return [] + return ["clahe"] if bool(params.get("clahe", True)) else ["intensity"] + + def run(self, ctx: ModuleContext) -> ChannelBlock: + p = ctx.params + clahe_on = bool(p.get("clahe", True)) + clip_limit = float(p.get("clip_limit", 0.01)) + ks = p.get("kernel_size", None) + kernel_size = int(ks) if ks not in (None, "", False) else None + _uint8, float_stack, labels = features.compute_clahe_stack( + ctx.gray, + clip_limit=clip_limit, + kernel_size=kernel_size, + apply=clahe_on, + ) + image_2d = float_stack[..., 0] + include = bool(p.get("include_in_bank", True)) + return ChannelBlock( + float_stack=float_stack if include else None, + labels=labels if include else [], + image_2d=image_2d, + ) diff --git a/ipred/src/ipred/modules/pca_mod.py b/ipred/src/ipred/modules/pca_mod.py new file mode 100644 index 0000000..611e8fa --- /dev/null +++ b/ipred/src/ipred/modules/pca_mod.py @@ -0,0 +1,78 @@ +"""PCA reduce dense embeddings into bank channels.""" + +from __future__ import annotations + +from typing import Any + +import numpy as np + +from ipred import features, sam_embed +from ipred.modules.base import ChannelBlock, ModuleContext, ModuleMeta + + +class PcaModule: + """Bake encoder embedding grid → upsampled PCA float channels.""" + + meta = ModuleMeta( + id="pca", + name="PCA reduce", + description="PCA on dense embedding → heatmaps in feature bank", + runtime="numpy", + accepts_input_from=True, + produces_channels=True, + produces_embedding=False, + params_schema={ + "dims": {"type": "integer", "default": 64}, + }, + ) + + def ready(self) -> bool: + return True + + def preview_labels(self, params: dict[str, Any]) -> list[str]: + dims = max(1, int(params.get("dims", 64))) + return [f"pca{i}" for i in range(dims)] + + def run(self, ctx: ModuleContext) -> ChannelBlock: + src_id = None + # Prefer explicitly linked node; else last embedding in context + emb = None + emb_meta: dict[str, Any] | None = None + # input_from resolved upstream sets nothing on ctx for emb — + # look up from node_outputs via params _input_from injected by runner + input_from = ctx.params.get("_input_from") + if input_from and input_from in ctx.node_outputs: + block = ctx.node_outputs[input_from] + emb, emb_meta = block.emb, block.emb_meta + src_id = input_from + if emb is None: + for nid, block in reversed(list(ctx.node_outputs.items())): + if block.emb is not None: + emb, emb_meta = block.emb, block.emb_meta + src_id = nid + break + if emb is None or emb_meta is None: + raise ValueError("pca module requires an upstream embedding node") + dims = max(1, int(ctx.params.get("dims", 64))) + h, w = int(ctx.gray.shape[0]), int(ctx.gray.shape[1]) + pca_float, pca_labels, info = sam_embed.emb_grid_to_pca_channels( + emb, + out_h=h, + out_w=w, + n_components=dims, + ) + meta = dict(emb_meta) + meta.update(info) + meta["pca_from_node"] = src_id + # Keep emb for bank persistence (sam_emb.npy) + return ChannelBlock( + float_stack=pca_float.astype(np.float32), + labels=pca_labels, + emb=emb, + emb_meta=meta, + ) + + +def uint8_from_float(float_stack: np.ndarray) -> np.ndarray: + """Display uint8 for a float HxWxC bank.""" + return features._float_feats_to_uint8(float_stack, clahe=False) diff --git a/ipred/src/ipred/modules/skimage_mod.py b/ipred/src/ipred/modules/skimage_mod.py new file mode 100644 index 0000000..7fa696a --- /dev/null +++ b/ipred/src/ipred/modules/skimage_mod.py @@ -0,0 +1,54 @@ +"""Skimage multiscale feature module.""" + +from __future__ import annotations + +from typing import Any + +from ipred import features +from ipred.modules.base import ChannelBlock, ModuleContext, ModuleMeta + + +class SkimageMultiscaleModule: + """Produce multiscale intensity/edges/texture channels.""" + + meta = ModuleMeta( + id="skimage_multiscale", + name="Skimage multiscale", + description="ilastik-style multiscale basic features", + runtime="numpy", + accepts_input_from=False, + produces_channels=True, + params_schema={ + "sigma_min": {"type": "number", "default": 1.0}, + "sigma_max": {"type": "number", "default": 8.0}, + "intensity": {"type": "boolean", "default": True}, + "edges": {"type": "boolean", "default": True}, + "texture": {"type": "boolean", "default": True}, + "clahe": {"type": "boolean", "default": True}, + }, + ) + + def ready(self) -> bool: + return True + + def preview_labels(self, params: dict[str, Any]) -> list[str]: + return features.feature_channel_labels( + sigma_min=float(params.get("sigma_min", 1.0)), + sigma_max=float(params.get("sigma_max", 8.0)), + intensity=bool(params.get("intensity", True)), + edges=bool(params.get("edges", True)), + texture=bool(params.get("texture", True)), + ) + + def run(self, ctx: ModuleContext) -> ChannelBlock: + p = ctx.params + _uint8, float_stack, labels = features.compute_feature_stacks( + ctx.gray, + sigma_min=float(p.get("sigma_min", 1.0)), + sigma_max=float(p.get("sigma_max", 8.0)), + intensity=bool(p.get("intensity", True)), + edges=bool(p.get("edges", True)), + texture=bool(p.get("texture", True)), + clahe=bool(p.get("clahe", True)), + ) + return ChannelBlock(float_stack=float_stack, labels=labels) diff --git a/ipred/src/ipred/modules/slimsam_mod.py b/ipred/src/ipred/modules/slimsam_mod.py new file mode 100644 index 0000000..5fe9ff2 --- /dev/null +++ b/ipred/src/ipred/modules/slimsam_mod.py @@ -0,0 +1,50 @@ +"""SlimSAM ONNX encoder module.""" + +from __future__ import annotations + +from typing import Any + +from ipred import feature_setups, sam_embed +from ipred.modules.base import ChannelBlock, ModuleContext, ModuleMeta + + +class SlimSamModule: + """Dense SlimSAM vision-encoder embeddings (ONNX Runtime).""" + + meta = ModuleMeta( + id="slimsam", + name="SlimSAM", + description="SlimSAM vision encoder (ONNX)", + runtime="onnx", + accepts_input_from=True, + produces_channels=False, + produces_embedding=True, + params_schema={ + "weights_path": {"type": "string", "default": None}, + }, + ) + + def ready(self) -> bool: + path = feature_setups.resolve_slimsam_weights_path() + return sam_embed.encoder_available(path) + + def preview_labels(self, params: dict[str, Any]) -> list[str]: + del params + return [] # PCA module consumes embedding + + def run(self, ctx: ModuleContext) -> ChannelBlock: + src = ctx.input_image if ctx.input_image is not None else ctx.raw + wpath = ctx.params.get("weights_path") or feature_setups.resolve_slimsam_weights_path() + if not sam_embed.encoder_available(wpath): + raise ValueError("SlimSAM ONNX weights not available") + emb, orig_hw, reshaped_hw = sam_embed.encode_image_embeddings( + src, weights_path=wpath + ) + meta = { + "orig_hw": list(orig_hw), + "reshaped_hw": list(reshaped_hw), + "weights_path": wpath, + "encoder": "slimsam", + "weights_format": feature_setups.WEIGHTS_ONNX_VISION, + } + return ChannelBlock(emb=emb, emb_meta=meta) diff --git a/ipred/src/ipred/modules/tomojepa_mod.py b/ipred/src/ipred/modules/tomojepa_mod.py new file mode 100644 index 0000000..9d9362a --- /dev/null +++ b/ipred/src/ipred/modules/tomojepa_mod.py @@ -0,0 +1,154 @@ +"""TomoJEPA / Mark encoder module (torch; ONNX when available).""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +from ipred import feature_setups, tomojepa_embed +from ipred.modules.base import ChannelBlock, ModuleContext, ModuleMeta + + +def _resolve_weights(params: dict[str, Any]) -> str | None: + """Resolve Mark25/11 weights from params.""" + explicit = params.get("weights_path") + if explicit and str(explicit) != "(missing)": + p = Path(str(explicit)).expanduser() + if p.is_file(): + return str(p.resolve()) + weights_id = str(params.get("weights_id") or "mark25").lower() + if weights_id in ("mark11", "tomojepa11", "11"): + return feature_setups.resolve_tomojepa_weights_path( + "tomojepa11.pth", env_var="TOMOJEPA11_WEIGHTS" + ) + return feature_setups.resolve_tomojepa_weights_path( + "tomojepa25.pth", env_var="TOMOJEPA_WEIGHTS" + ) + + +def _encoder_name(params: dict[str, Any], path: str | None) -> str: + wid = str(params.get("weights_id") or "").lower() + if wid in ("mark11", "tomojepa11", "11"): + return "mark11" + if path and "tomojepa11" in Path(path).name: + return "mark11" + return "mark25" + + +def _weights_format(name: str) -> str: + if name == "mark11": + return feature_setups.WEIGHTS_TOMOJEPA_MARK11 + return feature_setups.WEIGHTS_TOMOJEPA_MARK25 + + +class TomoJepaModule: + """Dense TomoJEPA embeddings (torch; prefers ONNX when present).""" + + meta = ModuleMeta( + id="tomojepa", + name="TomoJEPA", + description="Mark25/Mark11 dense projector (torch or ONNX)", + runtime="torch", + accepts_input_from=True, + produces_channels=False, + produces_embedding=True, + params_schema={ + "weights_id": {"type": "string", "default": "mark25", "enum": ["mark25", "mark11"]}, + "weights_path": {"type": "string", "default": None}, + "input_size": {"type": "integer", "default": 512}, + "resize": {"type": "boolean", "default": True}, + }, + ) + + def ready(self) -> bool: + if self._onnx_path({"weights_id": "mark25"}) or self._onnx_path( + {"weights_id": "mark11"} + ): + return True + for fname, env in ( + ("tomojepa25.pth", "TOMOJEPA_WEIGHTS"), + ("tomojepa11.pth", "TOMOJEPA11_WEIGHTS"), + ): + p = feature_setups.resolve_tomojepa_weights_path(fname, env_var=env) + if p and tomojepa_embed.encoder_available(p): + return True + return False + + def preview_labels(self, params: dict[str, Any]) -> list[str]: + del params + return [] + + def _onnx_path(self, params: dict[str, Any]) -> Path | None: + import os + + explicit = params.get("onnx_path") + if explicit and Path(str(explicit)).expanduser().is_file(): + return Path(str(explicit)).expanduser().resolve() + wid = str(params.get("weights_id") or "mark25").lower() + fname = "tomojepa11.onnx" if wid in ("mark11", "11") else "tomojepa25.onnx" + env = "TOMOJEPA11_ONNX" if "11" in fname else "TOMOJEPA_ONNX" + env_p = os.getenv(env) + if env_p and Path(env_p).expanduser().is_file(): + return Path(env_p).expanduser().resolve() + # modules/tomojepa_mod.py → …/ipred/src/ipred/modules → parents[3]=ipred/ + here = Path(__file__).resolve() + for c in ( + here.parents[3] / "models" / fname, + here.parents[4] / "ipred" / "models" / fname, + ): + if c.is_file(): + return c.resolve() + return None + + def run(self, ctx: ModuleContext) -> ChannelBlock: + src = ctx.input_image if ctx.input_image is not None else ctx.gray + p = ctx.params + resize = bool(p.get("resize", True)) + input_size = int(p.get("input_size", 512)) + onnx_path = self._onnx_path(p) + use_onnx = False + if onnx_path is not None: + from ipred import tomojepa_onnx + + use_onnx = tomojepa_onnx.onnx_matches_input_size( + onnx_path, input_size + ) + if use_onnx and onnx_path is not None: + from ipred import tomojepa_onnx + + emb, orig_hw, reshaped_hw = tomojepa_onnx.encode_dense_embeddings( + src, + weights_path=str(onnx_path), + input_size=input_size, + resize=resize, + ) + runtime = "onnx" + path_s = str(onnx_path) + else: + wpath = _resolve_weights(p) + if not tomojepa_embed.encoder_available(wpath): + raise ValueError( + "TomoJEPA weights not available (torch/.pth or ONNX)" + ) + emb, orig_hw, reshaped_hw = tomojepa_embed.encode_dense_embeddings( + src, + weights_path=wpath, + input_size=input_size, + resize=resize, + ) + runtime = "torch" + path_s = wpath + enc_name = _encoder_name(p, path_s) + meta = { + "orig_hw": list(orig_hw), + "reshaped_hw": list(reshaped_hw), + "weights_path": path_s, + "encoder": enc_name, + "input_size": input_size, + "resize": resize, + "weights_format": _weights_format(enc_name), + "runtime": runtime, + "model_input_range": [-1.0, 1.0], + "intensity_norm": "minmax_01_then_linear_m11", + } + return ChannelBlock(emb=emb, emb_meta=meta) diff --git a/ipred/src/ipred/paths.py b/ipred/src/ipred/paths.py new file mode 100644 index 0000000..4a3a071 --- /dev/null +++ b/ipred/src/ipred/paths.py @@ -0,0 +1,37 @@ +"""Filesystem roots for ipred under LOCAL_DATA_ROOT.""" + +from __future__ import annotations + +import os +from pathlib import Path + + +def local_data_root() -> Path: + """Return resolved ``LOCAL_DATA_ROOT`` (default ``~/data``).""" + return Path(os.getenv("LOCAL_DATA_ROOT", "~/data")).expanduser().resolve() + + +def engine_root() -> Path: + """Return ``$LOCAL_DATA_ROOT/ipred``.""" + root = local_data_root() / "ipred" + root.mkdir(parents=True, exist_ok=True) + return root + + +def catalog_db_path() -> Path: + """SQLite catalog path.""" + return engine_root() / "catalog.db" + + +def feature_models_root() -> Path: + """Feature Setup shelf directory.""" + root = local_data_root() / ".feature_models" + root.mkdir(parents=True, exist_ok=True) + return root + + +def project_blob_dir(project_id: str) -> Path: + """Blob root for one project.""" + path = engine_root() / "projects" / project_id + path.mkdir(parents=True, exist_ok=True) + return path diff --git a/ipred/src/ipred/preprocess.py b/ipred/src/ipred/preprocess.py new file mode 100644 index 0000000..ff937c9 --- /dev/null +++ b/ipred/src/ipred/preprocess.py @@ -0,0 +1,227 @@ +"""Cache-aware featurize for a session (composition-first).""" + +from __future__ import annotations + +import json +import logging +import re +import uuid +from datetime import datetime, timezone +from pathlib import Path +from typing import Any + +import numpy as np + +from ipred import array_source, compositions, features +from ipred.catalog import Catalog +from ipred.compose_run import run_composition +from ipred.paths import project_blob_dir + +logger = logging.getLogger(__name__) + + +def _utc_now() -> str: + return datetime.now(timezone.utc).isoformat() + + +def run_preprocess( + catalog: Catalog, + *, + session_id: str, + feature_setup_id: str | None = None, + composition_id: str | None = None, + slice_index: int = 0, + array_ref: str | None = None, +) -> dict[str, Any]: + """Featurize with cache hit/miss; update session current_feature_id. + + Prefers ``composition_id``; ``feature_setup_id`` is mapped via legacy table. + Optional ``array_ref`` selects a session-uploaded blob instead of path/Tiled. + """ + session = catalog.get_session(session_id) + if session is None: + raise KeyError(f"unknown session {session_id}") + project = catalog.get_project(session.project_id) + if project is None: + raise KeyError(f"unknown project {session.project_id}") + + compositions.ensure_default_compositions(catalog) + raw_id = composition_id or feature_setup_id + if not raw_id: + raise ValueError("composition_id or feature_setup_id required") + try: + cid = compositions.resolve_preprocess_id(raw_id) + except KeyError: + # Ensure defaults then retry + compositions.ensure_default_compositions(catalog) + cid = compositions.resolve_preprocess_id(raw_id) + + doc = compositions.resolve_composition(cid) + chash = doc.get("content_hash") or compositions.content_hash(doc) + + hit = catalog.find_feature_bank( + project_id=project.project_id, + setup_id=cid, + content_hash=chash, + slice_index=slice_index, + ) + if hit is not None and array_ref is None: + catalog.set_session_currents(session_id, feature_id=hit["feature_id"]) + return _bank_response(hit, cache_hit=True) + + if array_ref: + arr = _load_array_ref(project.project_id, array_ref) + else: + arr = array_source.read_slice( + kind=project.kind, + source=project.source, + slice_index=slice_index, + server_uri=project.server_uri, + root=project.root, + ) + uint8_stack, float_stack, labels, sam_emb, sam_meta = run_composition( + arr, doc + ) + feature_id = uuid.uuid4().hex + blob = project_blob_dir(project.project_id) / "features" / feature_id + blob.mkdir(parents=True, exist_ok=True) + np.save(blob / "float_stack.npy", float_stack.astype(np.float16)) + np.save(blob / "uint8_stack.npy", uint8_stack) + (blob / "labels.json").write_text( + json.dumps(labels, indent=2), encoding="utf-8" + ) + if sam_emb is not None: + np.save(blob / "sam_emb.npy", sam_emb.astype(np.float32)) + (blob / "sam_meta.json").write_text( + json.dumps(sam_meta), encoding="utf-8" + ) + channels_dir = blob / "channels" + channels_dir.mkdir(exist_ok=True) + for i in range(uint8_stack.shape[-1]): + png = features.encode_channel_png(uint8_stack, i) + (channels_dir / f"{i:04d}.png").write_bytes(png) + + h, w, c = float_stack.shape + snapshot = dict(doc) + record = { + "feature_id": feature_id, + "project_id": project.project_id, + "setup_id": cid, + "content_hash": chash, + "slice_index": int(slice_index), + "n_channels": int(c), + "height": int(h), + "width": int(w), + "blob_dir": str(blob), + "setup_snapshot": json.dumps(snapshot), + "status": "ready", + "created_at": _utc_now(), + } + catalog.insert_feature_bank(record) + catalog.set_session_currents(session_id, feature_id=feature_id) + return _bank_response(record, cache_hit=False, labels=labels) + + +def _load_array_ref(project_id: str, array_ref: str) -> np.ndarray: + """Load a content-addressed slice uploaded via the data plane.""" + from ipred.array_blobs import load_array_blob + + return load_array_blob(project_id, array_ref) + + +def load_feature_bank_arrays(blob_dir: str | Path) -> dict[str, Any]: + """Load persisted feature arrays from a bank directory.""" + blob = Path(blob_dir) + float_stack = np.load(blob / "float_stack.npy").astype(np.float32) + uint8_stack = np.load(blob / "uint8_stack.npy") + labels = json.loads((blob / "labels.json").read_text(encoding="utf-8")) + sam_emb = None + sam_meta = None + if (blob / "sam_emb.npy").is_file(): + sam_emb = np.load(blob / "sam_emb.npy") + if (blob / "sam_meta.json").is_file(): + sam_meta = json.loads( + (blob / "sam_meta.json").read_text(encoding="utf-8") + ) + return { + "float_stack": float_stack, + "uint8_stack": uint8_stack, + "labels": labels, + "sam_emb": sam_emb, + "sam_meta": sam_meta, + } + + +_FEATURE_ID_RE = re.compile(r"^[0-9a-f]{32}$") # exactly uuid.uuid4().hex's shape + + +def channel_png_path(project_id: str, feature_id: str, index: int) -> Path: + """Path to a cached channel PNG. + + Deliberately does NOT take the catalog's stored ``blob_dir`` string and + build a path from it — however safe that value always happens to be in + practice (a server-generated ``uuid.uuid4().hex`` under ``engine_root()``, + see ``run_preprocess`` above), a static analyzer has no way to know that, + and treats any DB value reached via a user-supplied ``feature_id`` lookup + as still tainted (CodeQL py/path-injection). Instead, validate the + *request's own* ``feature_id`` against the exact fixed shape it is always + generated in, before it ever touches a path, and rebuild the directory + fresh from known-safe components (``project_blob_dir`` already + constrains ``project_id`` the same way via its own hex-digest shape). A + request whose ``feature_id`` doesn't match a real feature bank's id can + never reach the filesystem layer at all. + """ + if not _FEATURE_ID_RE.match(feature_id): + raise ValueError(f"invalid feature_id: {feature_id!r}") + return project_blob_dir(project_id) / "features" / feature_id / "channels" / f"{index:04d}.png" + + +def _bank_response( + record: dict[str, Any], + *, + cache_hit: bool, + labels: list[str] | None = None, +) -> dict[str, Any]: + blob = Path(record["blob_dir"]) + if labels is None and (blob / "labels.json").is_file(): + labels = json.loads((blob / "labels.json").read_text(encoding="utf-8")) + return { + "feature_id": record["feature_id"], + "project_id": record["project_id"], + "setup_id": record["setup_id"], + "composition_id": record["setup_id"], + "slice_index": record["slice_index"], + "n_channels": record["n_channels"], + "height": record["height"], + "width": record["width"], + "labels": labels or [], + "cache_hit": cache_hit, + "blob_dir": record["blob_dir"], + } + + +# Keep soft import for tests that still poke at legacy helpers +def _compute_from_setup( + arr: np.ndarray, + setup: dict[str, Any], + snapshot: dict[str, Any], +) -> tuple[ + np.ndarray, + np.ndarray, + list[str], + np.ndarray | None, + dict[str, Any] | None, +]: + """Legacy bridge: convert procedure setup → composition run via resolve.""" + del snapshot + sid = setup.get("id") or "" + compositions.ensure_default_compositions() + try: + cid = compositions.resolve_preprocess_id(sid) + doc = compositions.resolve_composition(cid) + return run_composition(arr, doc) + except KeyError as exc: + raise ValueError( + "weights/procedure setups must map to a composition; " + f"unknown {sid}" + ) from exc diff --git a/ipred/src/ipred/sam_embed.py b/ipred/src/ipred/sam_embed.py new file mode 100644 index 0000000..7e1c446 --- /dev/null +++ b/ipred/src/ipred/sam_embed.py @@ -0,0 +1,219 @@ +"""SlimSAM ONNX embeddings for feature banks (self-contained).""" + +from __future__ import annotations + +import logging +import os +import threading +from functools import lru_cache +from pathlib import Path +from typing import Any + +import numpy as np +from PIL import Image as PILImage + +logger = logging.getLogger(__name__) + +_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32).reshape(3, 1, 1) +_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32).reshape(3, 1, 1) +_LONG_EDGE = 1024 +_PAD = 1024 + +_session_lock = threading.Lock() + + +def resolve_encoder_path(explicit: str | None = None) -> Path | None: + """Resolve ONNX path from explicit path, env, or repo defaults.""" + if explicit and explicit != "(missing)": + p = Path(explicit).expanduser().resolve() + if p.is_file(): + return p + env = os.getenv("FEATURE_ENCODER_ONNX") + if env: + p = Path(env).expanduser().resolve() + if p.is_file(): + return p + here = Path(__file__).resolve() + repo = here.parents[3] + for c in ( + repo / "backend" / "models" / "slimsam-77-uniform" / "onnx" / "vision_encoder.onnx", + repo / "frontend" / "public" / "models" / "slimsam-77-uniform" / "onnx" / "vision_encoder.onnx", + ): + if c.is_file(): + return c.resolve() + return None + + +def encoder_available(weights_path: str | None = None) -> bool: + """True when an ONNX encoder file is present.""" + return resolve_encoder_path(weights_path) is not None + + +@lru_cache(maxsize=4) +def _session_for(path_str: str): + import onnxruntime as ort + + return ort.InferenceSession(path_str, providers=["CPUExecutionProvider"]) + + +def rgb_uint8_from_array(arr: np.ndarray) -> np.ndarray: + """Convert slice to HxWx3 uint8 RGB.""" + a = np.asarray(arr) + if a.ndim == 3 and a.shape[-1] in (3, 4): + rgb = a[..., :3].astype(np.float64) + elif a.ndim == 2: + g = a.astype(np.float64) + rgb = np.stack([g, g, g], axis=-1) + else: + raise ValueError(f"unsupported array shape {a.shape}") + finite = rgb[np.isfinite(rgb)] + if finite.size == 0: + return np.zeros((*rgb.shape[:2], 3), dtype=np.uint8) + lo, hi = float(np.min(finite)), float(np.max(finite)) + if hi <= lo: + return np.zeros((*rgb.shape[:2], 3), dtype=np.uint8) + scaled = (rgb - lo) / (hi - lo) * 255.0 + return np.clip(scaled, 0, 255).astype(np.uint8) + + +def encode_image_embeddings( + arr: np.ndarray, + *, + weights_path: str | None = None, +) -> tuple[np.ndarray, tuple[int, int], tuple[int, int]]: + """Return ``(emb eh×ew×C, orig_hw, reshaped_hw)``.""" + path = resolve_encoder_path(weights_path) + if path is None: + raise FileNotFoundError("SlimSAM vision_encoder.onnx not found") + + rgb = rgb_uint8_from_array(arr) + oh, ow = int(rgb.shape[0]), int(rgb.shape[1]) + scale = _LONG_EDGE / max(oh, ow) + rh = max(1, int(round(oh * scale))) + rw = max(1, int(round(ow * scale))) + resized = np.asarray( + PILImage.fromarray(rgb).resize((rw, rh), PILImage.BILINEAR), + dtype=np.uint8, + ) + canvas = np.zeros((_PAD, _PAD, 3), dtype=np.uint8) + canvas[:rh, :rw] = resized + x = canvas.astype(np.float32) / 255.0 + x = np.transpose(x, (2, 0, 1)) + x = (x - _MEAN) / _STD + x = np.expand_dims(x, 0) + + with _session_lock: + sess = _session_for(str(path)) + inp = sess.get_inputs()[0].name + out = sess.run(None, {inp: x})[0] + + # NCHW → HWC + emb = np.transpose(np.asarray(out[0], dtype=np.float32), (1, 2, 0)) + return emb, (oh, ow), (rh, rw) + + +def bilinear_sample_emb( + emb: np.ndarray, + ys: np.ndarray, + xs: np.ndarray, + *, + orig_h: int, + orig_w: int, + reshaped_h: int, + reshaped_w: int, +) -> np.ndarray: + """Sample embedding at pixel coords (orig image space).""" + eh, ew, c = emb.shape + # Map orig → reshaped → emb grid + scale_y = reshaped_h / max(orig_h, 1) + scale_x = reshaped_w / max(orig_w, 1) + fy = (ys.astype(np.float64) * scale_y) * (eh / max(reshaped_h, 1)) + fx = (xs.astype(np.float64) * scale_x) * (ew / max(reshaped_w, 1)) + fy = np.clip(fy, 0, eh - 1.001) + fx = np.clip(fx, 0, ew - 1.001) + y0 = np.floor(fy).astype(np.int64) + x0 = np.floor(fx).astype(np.int64) + y1 = np.minimum(y0 + 1, eh - 1) + x1 = np.minimum(x0 + 1, ew - 1) + wy = fy - y0 + wx = fx - x0 + ia = emb[y0, x0] + ib = emb[y0, x1] + ic = emb[y1, x0] + id_ = emb[y1, x1] + wa = ((1 - wy) * (1 - wx))[:, None] + wb = ((1 - wy) * wx)[:, None] + wc = (wy * (1 - wx))[:, None] + wd = (wy * wx)[:, None] + return (wa * ia + wb * ib + wc * ic + wd * id_).astype(np.float32) + + +def fit_pca(x: np.ndarray, n_components: int = 32) -> tuple[np.ndarray, np.ndarray]: + """Fit PCA; return mean and components (n_comp × d).""" + from sklearn.decomposition import PCA + + n = min(n_components, x.shape[0], x.shape[1]) + pca = PCA(n_components=n, svd_solver="randomized", random_state=0) + pca.fit(x) + return pca.mean_.astype(np.float32), pca.components_.astype(np.float32) + + +def transform_pca( + x: np.ndarray, + mean: np.ndarray, + components: np.ndarray, +) -> np.ndarray: + """Project rows with fitted PCA.""" + return ((x - mean) @ components.T).astype(np.float32) + + +def emb_grid_to_pca_channels( + emb: np.ndarray, + *, + out_h: int, + out_w: int, + n_components: int = 64, + max_fit_samples: int = 50_000, +) -> tuple[np.ndarray, list[str], dict[str, Any]]: + """PCA on dense emb tokens, upsample to ``out_h×out_w×K`` float channels. + + Args: + emb: Encoder embedding grid ``eh×ew×C``. + out_h / out_w: Full-resolution image size for heatmaps / float_stack. + n_components: Target PCA dims (capped by tokens and channels). + max_fit_samples: Subsample cap when fitting PCA on large grids. + + Returns: + ``(float_hwk, labels, info)`` with labels ``pca0…pca{K-1}``. + """ + from skimage.transform import resize + + eh, ew, c = emb.shape + flat = np.asarray(emb, dtype=np.float32).reshape(-1, c) + n_tok = flat.shape[0] + k_target = max(1, int(n_components)) + if n_tok > max_fit_samples: + rng = np.random.default_rng(0) + idx = rng.choice(n_tok, size=max_fit_samples, replace=False) + fit_x = flat[idx] + else: + fit_x = flat + mean, comp = fit_pca(fit_x, n_components=k_target) + projected = transform_pca(flat, mean, comp) + k = projected.shape[1] + grid = projected.reshape(eh, ew, k) + up = resize( + grid, + (int(out_h), int(out_w), k), + order=1, + mode="edge", + anti_aliasing=False, + preserve_range=True, + ).astype(np.float32) + labels = [f"pca{i}" for i in range(k)] + info = { + "pca_dims": int(k), + "emb_shape": [int(eh), int(ew), int(c)], + "baked_into_float_stack": True, + } + return up, labels, info diff --git a/ipred/src/ipred/scripts/__init__.py b/ipred/src/ipred/scripts/__init__.py new file mode 100644 index 0000000..1a9ec9c --- /dev/null +++ b/ipred/src/ipred/scripts/__init__.py @@ -0,0 +1 @@ +"""Export scripts package.""" diff --git a/ipred/src/ipred/scripts/export_tomojepa_onnx.py b/ipred/src/ipred/scripts/export_tomojepa_onnx.py new file mode 100644 index 0000000..291bbb5 --- /dev/null +++ b/ipred/src/ipred/scripts/export_tomojepa_onnx.py @@ -0,0 +1,62 @@ +#!/usr/bin/env python3 +"""Export TomoJEPA Mark25/Mark11 checkpoints to ONNX. + +Usage: + python -m ipred.scripts.export_tomojepa_onnx + python -m ipred.scripts.export_tomojepa_onnx --weights ipred/models/tomojepa11.pth +""" + +from __future__ import annotations + +import argparse +import logging +import sys +from pathlib import Path + +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger("export_tomojepa_onnx") + + +def main(argv: list[str] | None = None) -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--weights", + action="append", + default=None, + help="Path to .pth (repeatable). Default: both tomojepa25/11 if present.", + ) + parser.add_argument("--input-size", type=int, default=512) + parser.add_argument( + "--out-dir", + type=Path, + default=None, + help="Output directory (default: same as weights)", + ) + args = parser.parse_args(argv) + + from ipred.tomojepa_onnx import export_tomojepa_onnx + + here = Path(__file__).resolve() + models = here.parents[2] / "models" + weights = args.weights + if not weights: + weights = [] + for name in ("tomojepa25.pth", "tomojepa11.pth"): + p = models / name + if p.is_file(): + weights.append(str(p)) + if not weights: + logger.error("no weights found") + return 1 + + for w in weights: + pth = Path(w) + out_dir = args.out_dir or pth.parent + onnx_path = out_dir / (pth.stem + ".onnx") + export_tomojepa_onnx(pth, onnx_path, input_size=args.input_size) + logger.info("wrote %s", onnx_path) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/ipred/src/ipred/tomojepa_embed.py b/ipred/src/ipred/tomojepa_embed.py new file mode 100644 index 0000000..8551728 --- /dev/null +++ b/ipred/src/ipred/tomojepa_embed.py @@ -0,0 +1,198 @@ +"""Mark25 / TomoJEPA dense embeddings for feature banks.""" + +from __future__ import annotations + +import logging +import os +import threading +from functools import lru_cache +from pathlib import Path +from typing import Any + +import numpy as np +from PIL import Image as PILImage + +logger = logging.getLogger(__name__) + +_model_lock = threading.Lock() +_DEFAULT_INPUT_SIZE = 512 +_PATCH = 16 + + +def torch_available() -> bool: + """True when torch and timm can be imported.""" + try: + import torch # noqa: F401 + import timm # noqa: F401 + except ImportError: + return False + return True + + +def resolve_weights_path(explicit: str | None = None) -> Path | None: + """Resolve TomoJEPA weights from explicit path, env, or repo default.""" + if explicit and explicit != "(missing)": + p = Path(explicit).expanduser().resolve() + if p.is_file(): + return p + env = os.getenv("TOMOJEPA_WEIGHTS") + if env: + p = Path(env).expanduser().resolve() + if p.is_file(): + return p + here = Path(__file__).resolve() + # .../repo/ipred/src/ipred → repo/ipred/models + candidates = [ + here.parents[2] / "models" / "tomojepa25.pth", + here.parents[3] / "ipred" / "models" / "tomojepa25.pth", + ] + for c in candidates: + if c.is_file(): + return c.resolve() + return None + + +def encoder_available(weights_path: str | None = None) -> bool: + """True when torch/timm are importable and weights exist.""" + if not torch_available(): + return False + return resolve_weights_path(weights_path) is not None + + +def load_checkpoint_state(path: Path | str) -> dict[str, Any]: + """Load ``ckpt['net']`` with optional ``module.`` prefix stripped.""" + import torch + + ckpt = torch.load(str(path), map_location="cpu", weights_only=False) + if not isinstance(ckpt, dict) or "net" not in ckpt: + raise ValueError(f"expected TomoJEPA ckpt with 'net' key: {path}") + return {k.replace("module.", ""): v for k, v in ckpt["net"].items()} + + +def build_encoder(**kwargs: Any): + """Build uninitialized DINOv3ViTEncoder (requires torch/timm).""" + from ipred.tomojepa_encoder import DINOv3ViTEncoder + + return DINOv3ViTEncoder(**kwargs) + + +@lru_cache(maxsize=4) +def _cached_model(path_str: str): + import torch + + from ipred.tomojepa_encoder import DINOv3ViTEncoder + + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + state = load_checkpoint_state(path_str) + # Mark25 dense proj is 64-D; Mark11 is 256-D — read from checkpoint. + proj_w = state.get("proj.8.weight") + if proj_w is None: + raise ValueError(f"checkpoint missing proj.8.weight: {path_str}") + proj_dim = int(proj_w.shape[0]) + net = DINOv3ViTEncoder( + proj_dim=proj_dim, img_size=_DEFAULT_INPUT_SIZE, in_chans=1, pretrained=False + ) + net.load_state_dict(state, strict=True) + net.eval() + net.to(device) + return net, device + + +def _pad_to_multiple(h: int, w: int, multiple: int = _PATCH) -> tuple[int, int]: + ph = (multiple - (h % multiple)) % multiple + pw = (multiple - (w % multiple)) % multiple + return h + ph, w + pw + + +def _gray_unit_interval(arr: np.ndarray) -> np.ndarray: + a = np.asarray(arr) + if a.ndim == 3 and a.shape[-1] in (3, 4): + rgb = a[..., :3].astype(np.float64) + gray = 0.299 * rgb[..., 0] + 0.587 * rgb[..., 1] + 0.114 * rgb[..., 2] + elif a.ndim == 2: + gray = a.astype(np.float64) + else: + raise ValueError(f"unsupported array shape {a.shape}") + finite = gray[np.isfinite(gray)] + if finite.size == 0: + return np.zeros_like(gray, dtype=np.float32) + lo, hi = float(np.min(finite)), float(np.max(finite)) + if hi <= lo: + return np.zeros_like(gray, dtype=np.float32) + return ((gray - lo) / (hi - lo)).astype(np.float32) + + +def to_minus_one_one(gray01: np.ndarray) -> np.ndarray: + """Linear map from ``[0, 1]`` to ``[-1, 1]`` (Mark25 / TomoJEPA convention).""" + return (np.asarray(gray01, dtype=np.float32) * 2.0 - 1.0).astype(np.float32) + + +def encode_dense_embeddings( + arr: np.ndarray, + *, + weights_path: str | None = None, + input_size: int = _DEFAULT_INPUT_SIZE, + resize: bool = True, +) -> tuple[np.ndarray, tuple[int, int], tuple[int, int]]: + """Return ``(emb Hp×Wp×64, orig_hw, reshaped_hw)``. + + When ``resize`` is True (default), the grayscale slice is resized to a + square ``input_size×input_size``. When False, native resolution is kept and + only padded to a multiple of 16 (Mark25 ``dynamic_img_size``). + + Intensity pipeline: min-max to ``[0, 1]`` (CLAHE already in that range when + passed from ``clahe_encoder_v1``), then linear shift to ``[-1, 1]`` before + the network. + """ + import torch + + path = resolve_weights_path(weights_path) + if path is None: + raise FileNotFoundError("TomoJEPA weights (.pth) not found") + if not torch_available(): + raise ImportError( + "TomoJEPA requires torch and timm — install with: pip install -e '.[torch]'" + ) + + gray = _gray_unit_interval(arr) + oh, ow = int(gray.shape[0]), int(gray.shape[1]) + if resize: + size = max(int(input_size), _PATCH) + resized = np.asarray( + PILImage.fromarray(gray, mode="F").resize( + (size, size), PILImage.BILINEAR + ), + dtype=np.float32, + ) + rh, rw = size, size + else: + resized = gray.astype(np.float32, copy=True) + rh, rw = oh, ow + + pad_h, pad_w = _pad_to_multiple(rh, rw, _PATCH) + if pad_h != rh or pad_w != rw: + canvas = np.zeros((pad_h, pad_w), dtype=np.float32) + canvas[:rh, :rw] = resized + resized = canvas + else: + pad_h, pad_w = rh, rw + + # TomoJEPA expects intensities in [-1, 1] after CLAHE / unit-interval prep. + model_in = to_minus_one_one(resized) + + x = torch.from_numpy(np.array(model_in, dtype=np.float32, copy=True)[None, None, None]) + with _model_lock: + net, device = _cached_model(str(path)) + x = x.to(device) + with torch.inference_mode(): + _glob, dense, _feats = net(x) + dense_np = dense.detach().cpu().numpy()[0, 0] # [L, 64] + + grid_h = pad_h // _PATCH + grid_w = pad_w // _PATCH + if dense_np.shape[0] != grid_h * grid_w: + side = int(round(dense_np.shape[0] ** 0.5)) + grid_h = grid_w = side + emb = dense_np.reshape(grid_h, grid_w, dense_np.shape[-1]).astype(np.float32) + # reshaped_hw is the coordinate frame covered by the emb grid (incl. pad). + return emb, (oh, ow), (pad_h, pad_w) diff --git a/ipred/src/ipred/tomojepa_encoder.py b/ipred/src/ipred/tomojepa_encoder.py new file mode 100644 index 0000000..44b2c29 --- /dev/null +++ b/ipred/src/ipred/tomojepa_encoder.py @@ -0,0 +1,107 @@ +"""Mark25 DINOv3ViTEncoder — ViT-S/16 + dense/global projectors. + +Architecture reconstructed to match TomoJEPA checkpoints (``ckpt['net']``): +timm ``vit_small_patch16_dinov3`` backbone, BatchNorm dense MLP → ``proj_dim`` +(Mark25: 64-D, Mark11: 256-D), LayerNorm global MLP → 16-D. +""" + +from __future__ import annotations + +from typing import Any + +import torch +import torch.nn as nn + + +def _dense_mlp(in_dim: int, hidden: int, out_dim: int) -> nn.Sequential: + return nn.Sequential( + nn.Linear(in_dim, hidden), + nn.BatchNorm1d(hidden), + nn.GELU(), + nn.Dropout(0.0), + nn.Linear(hidden, hidden), + nn.BatchNorm1d(hidden), + nn.GELU(), + nn.Dropout(0.0), + nn.Linear(hidden, out_dim), + nn.Identity(), + ) + + +def _global_mlp(in_dim: int, hidden: int, out_dim: int) -> nn.Sequential: + return nn.Sequential( + nn.Linear(in_dim, hidden), + nn.LayerNorm(hidden), + nn.GELU(), + nn.Dropout(0.0), + nn.Linear(hidden, hidden), + nn.LayerNorm(hidden), + nn.GELU(), + nn.Dropout(0.0), + nn.Linear(hidden, out_dim), + nn.Identity(), + ) + + +class DINOv3ViTEncoder(nn.Module): + """Mark25 TomoJEPA encoder producing global and dense projections.""" + + def __init__( + self, + proj_dim: int = 64, + img_size: int = 512, + in_chans: int = 1, + pretrained: bool = False, + *, + embed_dim: int = 384, + global_dim: int = 16, + hidden_dim: int = 2048, + ) -> None: + super().__init__() + del pretrained # checkpoint is loaded separately; never ImageNet init + import timm + + self.backbone = timm.create_model( + "vit_small_patch16_dinov3", + pretrained=False, + in_chans=in_chans, + img_size=img_size, + num_classes=0, + dynamic_img_size=True, + ) + self.proj = _dense_mlp(embed_dim, hidden_dim, proj_dim) + self.global_proj = _global_mlp(embed_dim, hidden_dim, global_dim) + self.num_prefix_tokens = int(getattr(self.backbone, "num_prefix_tokens", 5)) + self.proj_dim = int(proj_dim) + self.patch_size = int(getattr(self.backbone, "patch_size", 16) or 16) + + def forward( + self, x: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Run encoder on multi-view batch. + + Args: + x: Float tensor ``[B, C, V, H, W]`` (typically C=1). + + Returns: + ``(global_proj, dense_proj, feats)`` where dense is + ``[B, V, L, proj_dim]`` and global is ``[B, V, global_dim]``. + """ + if x.ndim != 5: + raise ValueError(f"expected [B,C,V,H,W], got shape {tuple(x.shape)}") + b, c, v, h, w = x.shape + flat = x.permute(0, 2, 1, 3, 4).reshape(b * v, c, h, w) + feats = self.backbone.forward_features(flat) + cls = feats[:, 0] + patches = feats[:, self.num_prefix_tokens :] + bv, length, dim = patches.shape + dense = self.proj(patches.reshape(bv * length, dim)).reshape( + b, v, length, -1 + ) + glob = self.global_proj(cls).reshape(b, v, -1) + return glob, dense, feats + + +def build_encoder(**kwargs: Any) -> DINOv3ViTEncoder: + """Construct an uninitialized Mark25 encoder.""" + return DINOv3ViTEncoder(**kwargs) diff --git a/ipred/src/ipred/tomojepa_onnx.py b/ipred/src/ipred/tomojepa_onnx.py new file mode 100644 index 0000000..cfbf7c7 --- /dev/null +++ b/ipred/src/ipred/tomojepa_onnx.py @@ -0,0 +1,208 @@ +"""TomoJEPA ONNX Runtime backend (optional; falls back to torch).""" + +from __future__ import annotations + +import logging +import threading +from functools import lru_cache +from pathlib import Path +from typing import Any + +import numpy as np +from PIL import Image as PILImage + +from ipred.tomojepa_embed import _gray_unit_interval, to_minus_one_one + +logger = logging.getLogger(__name__) + +_session_lock = threading.Lock() +_PATCH = 16 +_DEFAULT_INPUT_SIZE = 512 + + +def onnx_available() -> bool: + """True when onnxruntime can be imported.""" + try: + import onnxruntime # noqa: F401 + except ImportError: + return False + return True + + +@lru_cache(maxsize=4) +def _session_for(path_str: str): + import onnxruntime as ort + + return ort.InferenceSession(path_str, providers=["CPUExecutionProvider"]) + + +def fixed_spatial_size(weights_path: str | Path) -> int | None: + """Return fixed square H=W from ONNX input, or None when dynamic / unknown.""" + if not onnx_available(): + return None + path = Path(weights_path) + if not path.is_file(): + return None + with _session_lock: + sess = _session_for(str(path.resolve())) + shape = sess.get_inputs()[0].shape + if not shape: + return None + # [N,C,D,H,W] or [N,C,H,W] + dims = list(shape) + if len(dims) >= 2: + h, w = dims[-2], dims[-1] + if isinstance(h, int) and isinstance(w, int) and h == w and h > 0: + return int(h) + return None + + +def onnx_matches_input_size( + weights_path: str | Path, input_size: int +) -> bool: + """True when ONNX is dynamic or fixed spatial size equals ``input_size``.""" + fixed = fixed_spatial_size(weights_path) + if fixed is None: + return True + return fixed == int(input_size) + + +def _pad_to_multiple(h: int, w: int, multiple: int = _PATCH) -> tuple[int, int]: + ph = (multiple - (h % multiple)) % multiple + pw = (multiple - (w % multiple)) % multiple + return h + ph, w + pw + + +def encode_dense_embeddings( + arr: np.ndarray, + *, + weights_path: str, + input_size: int = _DEFAULT_INPUT_SIZE, + resize: bool = True, +) -> tuple[np.ndarray, tuple[int, int], tuple[int, int]]: + """Run exported TomoJEPA ONNX; return ``(emb Hp×Wp×D, orig_hw, reshaped_hw)``. + + Expects ONNX input name ``input`` with shape ``[1,1,1,H,W]`` or ``[1,1,H,W]`` + float32 in ``[-1, 1]``, and output ``dense`` as ``[1,1,L,D]`` or ``[1,L,D]``. + + When the graph has a fixed spatial size, that size is used (must match + ``input_size`` when ``resize`` is True, else ValueError). + """ + if not onnx_available(): + raise ImportError("onnxruntime required for TomoJEPA ONNX") + path = Path(weights_path) + if not path.is_file(): + raise FileNotFoundError(str(path)) + + fixed = fixed_spatial_size(path) + effective_size = int(input_size) + if fixed is not None: + if resize and fixed != effective_size: + raise ValueError( + f"ONNX fixed size {fixed} != input_size {effective_size}; " + "use torch fallback or re-export ONNX" + ) + effective_size = fixed + + gray = _gray_unit_interval(arr) + oh, ow = int(gray.shape[0]), int(gray.shape[1]) + if resize: + size = max(effective_size, _PATCH) + resized = np.asarray( + PILImage.fromarray(gray, mode="F").resize( + (size, size), PILImage.BILINEAR + ), + dtype=np.float32, + ) + rh, rw = size, size + else: + resized = gray.astype(np.float32, copy=True) + rh, rw = oh, ow + if fixed is not None and (rh != fixed or rw != fixed): + raise ValueError( + f"array {rh}x{rw} does not match ONNX fixed size {fixed}" + ) + + pad_h, pad_w = _pad_to_multiple(rh, rw, _PATCH) + if pad_h != rh or pad_w != rw: + canvas = np.zeros((pad_h, pad_w), dtype=np.float32) + canvas[:rh, :rw] = resized + resized = canvas + else: + pad_h, pad_w = rh, rw + + model_in = to_minus_one_one(resized) + x5 = np.array(model_in, dtype=np.float32, copy=True)[None, None, None] + x4 = np.array(model_in, dtype=np.float32, copy=True)[None, None] + + with _session_lock: + sess = _session_for(str(path.resolve())) + in_meta = sess.get_inputs()[0] + in_name = in_meta.name + # Choose rank matching exported model + shape = in_meta.shape + rank = len(shape) if shape else 5 + feed = {in_name: x5 if rank >= 5 else x4} + outs = sess.run(None, feed) + dense = np.asarray(outs[0], dtype=np.float32) + # Squeeze to [L, D] + while dense.ndim > 2: + if dense.shape[0] == 1: + dense = dense[0] + else: + break + if dense.ndim != 2: + raise ValueError(f"unexpected ONNX dense shape {dense.shape}") + + grid_h = pad_h // _PATCH + grid_w = pad_w // _PATCH + if dense.shape[0] != grid_h * grid_w: + side = int(round(dense.shape[0] ** 0.5)) + grid_h = grid_w = side + emb = dense.reshape(grid_h, grid_w, dense.shape[-1]).astype(np.float32) + return emb, (oh, ow), (pad_h, pad_w) + + +def export_tomojepa_onnx( + pth_path: str | Path, + onnx_path: str | Path, + *, + input_size: int = 512, + opset: int = 17, +) -> Path: + """Export ``DINOv3ViTEncoder`` checkpoint to ONNX (fixed square size).""" + import torch + + from ipred.tomojepa_embed import build_encoder, load_checkpoint_state + + pth = Path(pth_path) + out = Path(onnx_path) + out.parent.mkdir(parents=True, exist_ok=True) + state = load_checkpoint_state(pth) + proj_dim = int(state["proj.8.weight"].shape[0]) + net = build_encoder(proj_dim=proj_dim, img_size=input_size, in_chans=1) + net.load_state_dict(state, strict=True) + net.eval() + + class _DenseOnly(torch.nn.Module): + def __init__(self, enc: Any) -> None: + super().__init__() + self.enc = enc + + def forward(self, x: torch.Tensor) -> torch.Tensor: + _g, dense, _f = self.enc(x) + return dense + + wrapper = _DenseOnly(net) + dummy = torch.zeros(1, 1, 1, input_size, input_size, dtype=torch.float32) + torch.onnx.export( + wrapper, + dummy, + str(out), + input_names=["input"], + output_names=["dense"], + opset_version=opset, + dynamo=False, + ) + logger.info("exported TomoJEPA ONNX → %s (proj_dim=%s)", out, proj_dim) + return out.resolve() diff --git a/ipred/src/ipred/train_infer.py b/ipred/src/ipred/train_infer.py new file mode 100644 index 0000000..23acad0 --- /dev/null +++ b/ipred/src/ipred/train_infer.py @@ -0,0 +1,778 @@ +"""Train / infer / rethreshold orchestration.""" + +from __future__ import annotations + +import json +import logging +import uuid +from datetime import datetime, timezone +from functools import lru_cache +from io import BytesIO +from pathlib import Path +from typing import Any + +import numpy as np +from PIL import Image as PILImage + +from ipred import conformal, sam_embed +from ipred.catalog import Catalog +from ipred.labels import build_label_map +from ipred.paths import project_blob_dir +from ipred.preprocess import load_feature_bank_arrays +from ipred.trainers import get_trainer + +logger = logging.getLogger(__name__) + + +def _utc_now() -> str: + return datetime.now(timezone.utc).isoformat() + + +def _shape_bbox(shape: dict[str, Any]) -> tuple[float, float, float, float] | None: + """Rough (x0, y0, x1, y1) extent of one shape, for mismatch diagnostics.""" + kind = shape.get("kind") + try: + if kind == "rectangle": + x, y, ww, hh = float(shape["x"]), float(shape["y"]), float(shape["w"]), float(shape["h"]) + return min(x, x + ww), min(y, y + hh), max(x, x + ww), max(y, y + hh) + if kind == "ellipse": + cx, cy, rx, ry = (float(shape["cx"]), float(shape["cy"]), float(shape["rx"]), float(shape["ry"])) + return cx - rx, cy - ry, cx + rx, cy + ry + if kind == "polygon": + pts = shape.get("points") or [] + xs = pts[0::2] + ys = pts[1::2] + if not xs or not ys: + return None + return min(xs), min(ys), max(xs), max(ys) + if kind == "brush": + xs: list[float] = [] + ys: list[float] = [] + for stroke in shape.get("strokes") or []: + pts = stroke.get("points") or [] + xs.extend(pts[0::2]) + ys.extend(pts[1::2]) + if not xs or not ys: + return None + return min(xs), min(ys), max(xs), max(ys) + except (KeyError, TypeError, ValueError): + return None + return None + + +def _shape_class_summary(shapes: list[dict[str, Any]]) -> str: + """Per-class shape count + combined bbox, for error messages when + rasterization yields fewer classes than the caller annotated.""" + by_class: dict[int, list[tuple[float, float, float, float]]] = {} + for shape in shapes: + cid = int(shape.get("classId") or shape.get("class_id") or 0) + bbox = _shape_bbox(shape) + by_class.setdefault(cid, []) + if bbox is not None: + by_class[cid].append(bbox) + parts = [] + for cid, boxes in sorted(by_class.items()): + if boxes: + x0 = min(b[0] for b in boxes) + y0 = min(b[1] for b in boxes) + x1 = max(b[2] for b in boxes) + y1 = max(b[3] for b in boxes) + parts.append(f"class {cid}: {len(boxes)} shape(s) spanning ({x0:.0f},{y0:.0f})-({x1:.0f},{y1:.0f})") + else: + parts.append(f"class {cid}: shape(s) with unparsable extent") + return "; ".join(parts) if parts else "no shapes" + + +def stratified_train_cal_split( + ys: np.ndarray, + *, + train_frac: float = 0.8, + rng: np.random.Generator, +) -> tuple[np.ndarray, np.ndarray]: + """Stratified indices into train / calibration.""" + if not 0.0 < train_frac < 1.0: + raise ValueError("train_frac must be in (0, 1)") + train_parts: list[np.ndarray] = [] + cal_parts: list[np.ndarray] = [] + for cls in np.unique(ys): + idx = np.flatnonzero(ys == cls) + rng.shuffle(idx) + n = idx.size + if n == 1: + train_parts.append(idx) + continue + n_train = max(1, min(n - 1, int(round(n * train_frac)))) + train_parts.append(idx[:n_train]) + cal_parts.append(idx[n_train:]) + train_idx = np.concatenate(train_parts) if train_parts else np.array([], dtype=np.int64) + cal_idx = np.concatenate(cal_parts) if cal_parts else np.array([], dtype=np.int64) + if cal_idx.size == 0: + raise ValueError("calibration split empty — need more labeled pixels per class") + return train_idx, cal_idx + + +def _stratified_sample( + ys: np.ndarray, max_samples: int, rng: np.random.Generator +) -> np.ndarray: + n = ys.shape[0] + if n <= max_samples: + return np.arange(n) + classes, counts = np.unique(ys, return_counts=True) + alloc = np.maximum(1, np.floor(counts / n * max_samples).astype(int)) + while alloc.sum() < max_samples: + alloc[int(np.argmax(counts - alloc))] += 1 + while alloc.sum() > max_samples: + i = int(np.argmax(alloc)) + if alloc[i] > 1: + alloc[i] -= 1 + else: + break + chosen: list[np.ndarray] = [] + for cls, k in zip(classes, alloc): + idx = np.flatnonzero(ys == cls) + take = min(int(k), idx.size) + chosen.append(rng.choice(idx, size=take, replace=False)) + return np.concatenate(chosen) if chosen else np.arange(0) + + +def _raw_pixel_features( + bank: dict[str, Any], + yy: np.ndarray, + xx: np.ndarray, +) -> tuple[np.ndarray, list[str], np.ndarray | None]: + """Skimage feature columns + raw (pre-PCA) SAM embedding columns, if any. + + Split out of ``_pixel_features`` so multi-slice training can pool raw SAM + pixels across slices and fit ONE PCA over the pooled set — fitting PCA + per-slice would make the reduced columns incomparable slice-to-slice. + """ + float_stack = bank["float_stack"] + h, w, _ = float_stack.shape + ys_i = np.clip(yy.astype(np.int64), 0, h - 1) + xs_i = np.clip(xx.astype(np.int64), 0, w - 1) + x_sk = float_stack[ys_i, xs_i].astype(np.float32) + labels = list(bank["labels"]) + sam_emb = bank.get("sam_emb") + sam_meta = bank.get("sam_meta") + # Preprocess may already bake encoder PCA into float_stack channels. + if sam_meta and sam_meta.get("baked_into_float_stack"): + return x_sk, labels, None + if sam_emb is None or sam_meta is None: + return x_sk, labels, None + oh, ow = sam_meta["orig_hw"] + rh, rw = sam_meta["reshaped_hw"] + x_sam = sam_embed.bilinear_sample_emb( + sam_emb, + ys_i.astype(np.float64), + xs_i.astype(np.float64), + orig_h=int(oh), + orig_w=int(ow), + reshaped_h=int(rh), + reshaped_w=int(rw), + ) + return x_sk, labels, x_sam + + +def _pixel_features( + bank: dict[str, Any], + yy: np.ndarray, + xx: np.ndarray, + *, + sam_pca_mean: np.ndarray | None = None, + sam_pca_components: np.ndarray | None = None, + fit_pca: bool = False, + sam_pca_dims: int = 32, +) -> tuple[np.ndarray, list[str], np.ndarray | None, np.ndarray | None]: + x_sk, labels, x_sam = _raw_pixel_features(bank, yy, xx) + if x_sam is None: + return x_sk, labels, None, None + mean, comp = sam_pca_mean, sam_pca_components + if fit_pca: + mean, comp = sam_embed.fit_pca(x_sam, n_components=sam_pca_dims) + if mean is not None and comp is not None: + x_sam = sam_embed.transform_pca(x_sam, mean, comp) + x = np.concatenate([x_sk, x_sam], axis=1) + feat_labels = labels + [f"sam{i}" for i in range(x_sam.shape[1])] + return x, feat_labels, mean, comp + + +def run_train( + catalog: Catalog, + *, + session_id: str, + shapes: list[dict[str, Any]], + feature_id: str | None = None, + trainer_id: str = "catboost", + config: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Train a plugin model on session feature bank + shapes.""" + session = catalog.get_session(session_id) + if session is None: + raise KeyError(f"unknown session {session_id}") + fid = feature_id or session.current_feature_id + if not fid: + raise ValueError("no feature_id; run preprocess first") + bank_row = catalog.get_feature_bank(fid) + if bank_row is None: + raise KeyError(f"unknown feature bank {fid}") + bank = load_feature_bank_arrays(bank_row["blob_dir"]) + h, w, _ = bank["float_stack"].shape + label_map = build_label_map(shapes, h, w) + if not np.any(label_map > 0): + raise ValueError( + f"no labeled pixels — annotate at least one shape " + f"({_shape_class_summary(shapes)} vs. feature bank {w}x{h})" + ) + + yy, xx = np.nonzero(label_map) + y_all = label_map[yy, xx].astype(np.int32) + if len(np.unique(y_all)) < 2: + raise ValueError( + "need at least two classes with labeled pixels — got " + f"{sorted(int(c) for c in np.unique(y_all))} after rasterizing to the feature " + f"bank ({w}x{h}); shapes sent: {_shape_class_summary(shapes)}. If the shape " + "classes/extents don't match, the feature bank may be at a different " + "resolution than the annotated slice." + ) + + cfg = dict(config or {}) + rng = np.random.default_rng(int(cfg.get("random_seed", 0))) + max_samples = int(cfg.get("max_samples", 200_000)) + train_frac = float(cfg.get("train_frac", 0.8)) + cap_idx = _stratified_sample(y_all, max_samples, rng) + yy_c, xx_c, y_c = yy[cap_idx], xx[cap_idx], y_all[cap_idx] + train_idx, cal_idx = stratified_train_cal_split(y_c, train_frac=train_frac, rng=rng) + + uses_sam = ( + bank.get("sam_emb") is not None + and not (bank.get("sam_meta") or {}).get("baked_into_float_stack") + ) + x_tr, feat_labels, pca_mean, pca_comp = _pixel_features( + bank, + yy_c[train_idx], + xx_c[train_idx], + fit_pca=uses_sam, + sam_pca_dims=int(cfg.get("sam_pca_dims", 32)), + ) + x_cal, _, _, _ = _pixel_features( + bank, + yy_c[cal_idx], + xx_c[cal_idx], + sam_pca_mean=pca_mean, + sam_pca_components=pca_comp, + ) + + trainer = get_trainer(trainer_id) + arts = trainer.train( + x_tr, + y_c[train_idx], + x_cal, + y_c[cal_idx], + feature_labels=feat_labels, + config=cfg, + ) + return _save_trained_model( + catalog, + session=session, + session_id=session_id, + fid=fid, + trainer_id=trainer_id, + arts=arts, + uses_sam=uses_sam, + train_frac=train_frac, + pca_mean=pca_mean, + pca_comp=pca_comp, + ) + + +def _save_trained_model( + catalog: Catalog, + *, + session: Any, + session_id: str, + fid: str, + trainer_id: str, + arts: Any, + uses_sam: bool, + train_frac: float, + pca_mean: np.ndarray | None, + pca_comp: np.ndarray | None, + extra_meta: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Persist a trained ``ModelArtifacts`` to blob storage + the catalog. + + Shared tail of ``run_train`` and ``run_train_multi_slice`` — the only + difference between single- and multi-slice training is how ``arts`` (and + the pooled SAM PCA, if any) got built; saving/cataloging is identical. + """ + arts.params["uses_sam"] = uses_sam + arts.params["train_frac"] = train_frac + + model_id = uuid.uuid4().hex + blob = project_blob_dir(session.project_id) / "models" / model_id + blob.mkdir(parents=True, exist_ok=True) + trainer = get_trainer(trainer_id) + trainer.save(arts.model_handle, blob) + feature_importances = list(arts.extras.get("feature_importances") or []) + meta = { + "model_id": model_id, + "feature_id": fid, + "trainer_id": trainer_id, + "class_ids": arts.class_ids, + "feature_labels": arts.feature_labels, + "cal_scores_by_class": arts.cal_scores_by_class, + "train_accuracy": arts.train_accuracy, + "n_train": arts.n_train, + "n_cal": arts.n_cal, + "n_samples": arts.n_samples, + "params": arts.params, + "uses_sam": uses_sam, + "feature_importances": feature_importances, + **(extra_meta or {}), + } + if pca_mean is not None and pca_comp is not None: + np.savez(blob / "sam_pca.npz", mean=pca_mean, components=pca_comp) + meta["has_sam_pca"] = True + (blob / "meta.json").write_text(json.dumps(meta, indent=2), encoding="utf-8") + (blob / "cal_scores.json").write_text( + json.dumps(arts.cal_scores_by_class), encoding="utf-8" + ) + + catalog.insert_model( + { + "model_id": model_id, + "project_id": session.project_id, + "feature_id": fid, + "trainer_id": trainer_id, + "blob_dir": str(blob), + "meta_json": json.dumps(meta), + "created_at": _utc_now(), + } + ) + catalog.set_session_currents(session_id, model_id=model_id) + return { + "model_id": model_id, + "feature_id": fid, + "trainer_id": trainer_id, + "class_ids": arts.class_ids, + "train_accuracy": arts.train_accuracy, + "n_train": arts.n_train, + "n_cal": arts.n_cal, + "n_samples": arts.n_samples, + "params": arts.params, + "feature_importances": feature_importances, + **(extra_meta or {}), + } + + +def run_train_multi_slice( + catalog: Catalog, + *, + session_id: str, + per_slice_shapes: dict[int, list[dict[str, Any]]], + feature_ids: dict[int, str], + trainer_id: str = "catboost", + config: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Train one model pooling labeled pixels across multiple slices. + + Each slice keeps its own feature bank (its own composition run against + that slice), but SAM PCA — when in play — is fit ONCE on pixels pooled + across every slice's training split, so the reduced columns mean the same + thing regardless of which slice a pixel came from. A per-slice PCA fit + (naively looping ``run_train``) would make them incomparable. + """ + session = catalog.get_session(session_id) + if session is None: + raise KeyError(f"unknown session {session_id}") + if not per_slice_shapes: + raise ValueError("no slices with shapes to train on") + + cfg = dict(config or {}) + rng = np.random.default_rng(int(cfg.get("random_seed", 0))) + max_samples = int(cfg.get("max_samples", 200_000)) + train_frac = float(cfg.get("train_frac", 0.8)) + per_slice_cap = max(1, max_samples // max(1, len(per_slice_shapes))) + + banks: dict[int, dict[str, Any]] = {} + train_parts: list[tuple[int, np.ndarray, np.ndarray, np.ndarray]] = [] + cal_parts: list[tuple[int, np.ndarray, np.ndarray, np.ndarray]] = [] + uses_sam = False + summaries: list[str] = [] + + for slice_index, shapes in sorted(per_slice_shapes.items()): + fid = feature_ids.get(slice_index) + if not fid: + raise ValueError(f"no feature_id for slice {slice_index}") + bank_row = catalog.get_feature_bank(fid) + if bank_row is None: + raise KeyError(f"unknown feature bank {fid}") + bank = load_feature_bank_arrays(bank_row["blob_dir"]) + banks[slice_index] = bank + h, w, _ = bank["float_stack"].shape + label_map = build_label_map(shapes, h, w) + if not np.any(label_map > 0): + summaries.append(f"slice {slice_index}: no labeled pixels") + continue + yy, xx = np.nonzero(label_map) + y_all = label_map[yy, xx].astype(np.int32) + cap_idx = _stratified_sample(y_all, per_slice_cap, rng) + yy_c, xx_c, y_c = yy[cap_idx], xx[cap_idx], y_all[cap_idx] + train_idx, cal_idx = stratified_train_cal_split(y_c, train_frac=train_frac, rng=rng) + train_parts.append((slice_index, yy_c[train_idx], xx_c[train_idx], y_c[train_idx])) + cal_parts.append((slice_index, yy_c[cal_idx], xx_c[cal_idx], y_c[cal_idx])) + uses_sam = uses_sam or ( + bank.get("sam_emb") is not None + and not (bank.get("sam_meta") or {}).get("baked_into_float_stack") + ) + summaries.append( + f"slice {slice_index}: {y_all.size} labeled px, " + f"classes {sorted(int(c) for c in np.unique(y_all))}" + ) + + if not train_parts: + raise ValueError( + "no labeled pixels across the given slices — " + "; ".join(summaries) + ) + + all_y_train = np.concatenate([p[3] for p in train_parts]) + all_y_cal = ( + np.concatenate([p[3] for p in cal_parts]) + if cal_parts + else np.array([], dtype=np.int32) + ) + if len(np.unique(np.concatenate([all_y_train, all_y_cal]))) < 2: + raise ValueError( + "need at least two classes with labeled pixels across the selected " + "slices — " + "; ".join(summaries) + ) + + # Raw (pre-PCA) extraction per slice, pooled before any PCA fit. + x_sk_train_parts, x_sam_train_parts = [], [] + x_sk_cal_parts, x_sam_cal_parts = [], [] + feat_labels: list[str] | None = None + for slice_index, yy, xx, _y in train_parts: + x_sk, labels, x_sam = _raw_pixel_features(banks[slice_index], yy, xx) + feat_labels = feat_labels or labels + x_sk_train_parts.append(x_sk) + if x_sam is not None: + x_sam_train_parts.append(x_sam) + for slice_index, yy, xx, _y in cal_parts: + x_sk, _labels, x_sam = _raw_pixel_features(banks[slice_index], yy, xx) + x_sk_cal_parts.append(x_sk) + if x_sam is not None: + x_sam_cal_parts.append(x_sam) + + x_sk_train = np.concatenate(x_sk_train_parts, axis=0) + x_sk_cal = ( + np.concatenate(x_sk_cal_parts, axis=0) + if x_sk_cal_parts + else np.empty((0, x_sk_train.shape[1]), dtype=np.float32) + ) + + pca_mean = pca_comp = None + if uses_sam and x_sam_train_parts: + x_sam_train_raw = np.concatenate(x_sam_train_parts, axis=0) + x_sam_cal_raw = ( + np.concatenate(x_sam_cal_parts, axis=0) + if x_sam_cal_parts + else np.empty((0, x_sam_train_raw.shape[1]), dtype=x_sam_train_raw.dtype) + ) + pca_mean, pca_comp = sam_embed.fit_pca( + x_sam_train_raw, n_components=int(cfg.get("sam_pca_dims", 32)) + ) + x_sam_train = sam_embed.transform_pca(x_sam_train_raw, pca_mean, pca_comp) + x_sam_cal = ( + sam_embed.transform_pca(x_sam_cal_raw, pca_mean, pca_comp) + if x_sam_cal_raw.shape[0] + else np.empty((0, pca_comp.shape[0]), dtype=np.float32) + ) + x_tr = np.concatenate([x_sk_train, x_sam_train], axis=1) + x_cal = ( + np.concatenate([x_sk_cal, x_sam_cal], axis=1) + if x_sk_cal.shape[0] + else np.empty((0, x_tr.shape[1]), dtype=np.float32) + ) + feat_labels = (feat_labels or []) + [f"sam{i}" for i in range(x_sam_train.shape[1])] + else: + x_tr = x_sk_train + x_cal = x_sk_cal + + trainer = get_trainer(trainer_id) + arts = trainer.train( + x_tr, all_y_train, x_cal, all_y_cal, feature_labels=feat_labels or [], config=cfg + ) + trained_slices = sorted(slice_index for slice_index, *_ in train_parts) + # The default feature_id is only an infer-time fallback when a caller doesn't + # pass one explicitly — real infer calls always target one specific slice. + default_feature_id = feature_ids[trained_slices[0]] + return _save_trained_model( + catalog, + session=session, + session_id=session_id, + fid=default_feature_id, + trainer_id=trainer_id, + arts=arts, + uses_sam=uses_sam, + train_frac=train_frac, + pca_mean=pca_mean, + pca_comp=pca_comp, + extra_meta={"trained_slice_indices": trained_slices}, + ) + + +@lru_cache(maxsize=4) +def _load_model_cached(trainer_id: str, blob_model_str: str) -> tuple[Any, Any, np.ndarray | None, np.ndarray | None]: + """Deserialize a trained model (+ its SAM PCA, if any) once and reuse it. + + Model blob directories are never mutated in place — every training run + gets a fresh `uuid.uuid4()` model_id/blob_dir (see `run_train`/ + `_save_trained_model`) — so caching by `(trainer_id, blob_dir)` can never + serve stale weights. Without this, a volume-wide "apply across all + slices" job (Phase 4.5) reloaded the CatBoost model and PCA from disk on + every single slice instead of once per job. + """ + blob_model = Path(blob_model_str) + trainer = get_trainer(trainer_id) + handle = trainer.load(blob_model) + pca_mean = pca_comp = None + pca_path = blob_model / "sam_pca.npz" + if pca_path.is_file(): + z = np.load(pca_path) + pca_mean, pca_comp = z["mean"], z["components"] + return trainer, handle, pca_mean, pca_comp + + +def run_infer( + catalog: Catalog, + *, + session_id: str, + model_id: str | None = None, + feature_id: str | None = None, + alpha: float = 0.05, + row_chunk: int = 128, + store_probabilities: bool = True, +) -> dict[str, Any]: + """Full-image predict_proba + conformal maps; persist float16 proba. + + `store_probabilities=False` skips writing `proba.npy` to disk (the math + is unchanged — `proba` is still computed in memory to derive `commit`/ + `status`/`membership`). For a volume-wide batch-apply job (Phase 4.5), + `proba.npy` is the single largest thing a run writes (a full H×W×K + float16 array) and is never read back by that flow — only `commit.png` + is (see `AnnotatePage.tsx`'s `handleCommitVolumeApply`). Leave this + `True` (the default) for interactive single-slice infer calls: it's what + `run_rethreshold`/`threshold_class_map` read back to let a user adjust + alpha or view a per-class heatmap after the fact — pointing either of + those at a run saved with `store_probabilities=False` raises + `FileNotFoundError`, by design. + """ + session = catalog.get_session(session_id) + if session is None: + raise KeyError(f"unknown session {session_id}") + mid = model_id or session.current_model_id + fid = feature_id or session.current_feature_id + if not mid or not fid: + raise ValueError("model_id and feature_id required (train/preprocess first)") + model_row = catalog.get_model(mid) + bank_row = catalog.get_feature_bank(fid) + if model_row is None or bank_row is None: + raise KeyError("model or feature bank missing") + + blob_model = Path(model_row["blob_dir"]) + meta = json.loads((blob_model / "meta.json").read_text(encoding="utf-8")) + trainer, handle, pca_mean, pca_comp = _load_model_cached(model_row["trainer_id"], str(blob_model)) + bank = load_feature_bank_arrays(bank_row["blob_dir"]) + + class_ids = [int(c) for c in meta["class_ids"]] + h, w, _ = bank["float_stack"].shape + k = len(class_ids) + proba = np.empty((h, w, k), dtype=np.float16) + + for y0 in range(0, h, row_chunk): + y1 = min(h, y0 + row_chunk) + yy = np.repeat(np.arange(y0, y1), w) + xx = np.tile(np.arange(w), y1 - y0) + x, _, _, _ = _pixel_features( + bank, yy, xx, sam_pca_mean=pca_mean, sam_pca_components=pca_comp + ) + block = trainer.predict_proba(handle, x).astype(np.float32) + # Align columns to meta class_ids order + model_classes = [int(c) for c in np.asarray(handle.classes_).tolist()] + if model_classes != class_ids: + remap = np.zeros_like(block) + for j, cid in enumerate(class_ids): + if cid in model_classes: + remap[:, j] = block[:, model_classes.index(cid)] + block = remap + proba[y0:y1] = block.reshape(y1 - y0, w, k).astype(np.float16) + + cal = conformal.cal_scores_from_json(meta["cal_scores_by_class"]) + q_by_class = conformal.mondrian_thresholds(cal, alpha) + commit, status, membership = conformal.maps_from_proba( + proba.astype(np.float32), class_ids, q_by_class + ) + counts = conformal.counts_from_status(status) + + run_id = uuid.uuid4().hex + blob = project_blob_dir(session.project_id) / "runs" / run_id + blob.mkdir(parents=True, exist_ok=True) + if store_probabilities: + np.save(blob / "proba.npy", proba) + np.save(blob / "commit.npy", commit) + np.save(blob / "status.npy", status) + np.save(blob / "membership.npy", membership) + _write_label_png(blob / "commit.png", commit) + _write_label_png(blob / "status.png", status) + + run_meta = { + "run_id": run_id, + "model_id": mid, + "feature_id": fid, + "alpha": alpha, + "class_ids": class_ids, + "q_by_class": {str(k_): v for k_, v in q_by_class.items()}, + "counts": counts, + } + (blob / "meta.json").write_text(json.dumps(run_meta, indent=2), encoding="utf-8") + now = _utc_now() + catalog.insert_run( + { + "run_id": run_id, + "project_id": session.project_id, + "model_id": mid, + "feature_id": fid, + "alpha": float(alpha), + "blob_dir": str(blob), + "meta_json": json.dumps(run_meta), + "created_at": now, + "updated_at": now, + } + ) + catalog.set_session_currents(session_id, run_id=run_id) + return {**run_meta, "blob_dir": str(blob)} + + +def run_rethreshold( + catalog: Catalog, + *, + session_id: str, + run_id: str | None = None, + alpha: float, +) -> dict[str, Any]: + """Recompute membership/commit/status from cached proba (no re-infer).""" + session = catalog.get_session(session_id) + if session is None: + raise KeyError(f"unknown session {session_id}") + rid = run_id or session.current_run_id + if not rid: + raise ValueError("run_id required") + run = catalog.get_run(rid) + if run is None: + raise KeyError(f"unknown run {rid}") + blob = Path(run["blob_dir"]) + proba = np.load(blob / "proba.npy").astype(np.float32) + run_meta = json.loads((blob / "meta.json").read_text(encoding="utf-8")) + model_row = catalog.get_model(run["model_id"]) + if model_row is None: + raise KeyError("model missing for run") + model_meta = json.loads( + Path(model_row["blob_dir"], "meta.json").read_text(encoding="utf-8") + ) + cal = conformal.cal_scores_from_json(model_meta["cal_scores_by_class"]) + class_ids = [int(c) for c in run_meta["class_ids"]] + q_by_class = conformal.mondrian_thresholds(cal, alpha) + commit, status, membership = conformal.maps_from_proba(proba, class_ids, q_by_class) + counts = conformal.counts_from_status(status) + np.save(blob / "commit.npy", commit) + np.save(blob / "status.npy", status) + np.save(blob / "membership.npy", membership) + _write_label_png(blob / "commit.png", commit) + _write_label_png(blob / "status.png", status) + run_meta.update( + { + "alpha": alpha, + "q_by_class": {str(k): v for k, v in q_by_class.items()}, + "counts": counts, + } + ) + (blob / "meta.json").write_text(json.dumps(run_meta, indent=2), encoding="utf-8") + catalog.update_run(rid, alpha=float(alpha), meta_json=json.dumps(run_meta)) + catalog.set_session_currents(session_id, run_id=rid) + return {**run_meta, "blob_dir": str(blob)} + + +def _write_label_png(path: Path, arr: np.ndarray) -> None: + img = PILImage.fromarray(arr.astype(np.uint8), mode="L") + buf = BytesIO() + img.save(buf, format="PNG") + path.write_bytes(buf.getvalue()) + + +def _run_blob(catalog: Catalog, run_id: str) -> tuple[dict[str, Any], Path, dict[str, Any]]: + run = catalog.get_run(run_id) + if run is None: + raise KeyError(f"unknown run {run_id}") + blob = Path(run["blob_dir"]) + meta = json.loads(run["meta_json"]) + return run, blob, meta + + +def proba_heatmap_png(catalog: Catalog, run_id: str, class_index: int) -> bytes: + """Return grayscale PNG of softmax channel ``class_index`` (0…K-1).""" + _run, blob, meta = _run_blob(catalog, run_id) + proba_path = blob / "proba.npy" + if not proba_path.is_file(): + raise FileNotFoundError("proba.npy missing — run Predict first") + proba = np.load(proba_path) + if proba.ndim != 3: + raise ValueError("proba must be HxWxK") + k = int(proba.shape[2]) + if not 0 <= int(class_index) < k: + raise ValueError(f"class_index {class_index} out of range 0..{k - 1}") + ch = np.clip(proba[:, :, int(class_index)].astype(np.float32), 0.0, 1.0) + uint8 = np.clip(np.round(ch * 255.0), 0, 255).astype(np.uint8) + buf = BytesIO() + PILImage.fromarray(uint8, mode="L").save(buf, format="PNG") + return buf.getvalue() + + +def threshold_class_label_map( + catalog: Catalog, + run_id: str, + *, + class_id: int, + threshold: float, +) -> dict[str, Any]: + """Binary class map from softmax: ``class_id`` where ``p >= threshold``, else 0. + + Returns width/height and base64-encoded uint8 label map for mask-set caches. + """ + import base64 + + _run, blob, meta = _run_blob(catalog, run_id) + class_ids = [int(c) for c in meta.get("class_ids") or []] + if int(class_id) not in class_ids: + raise ValueError(f"class_id {class_id} not in run class_ids {class_ids}") + j = class_ids.index(int(class_id)) + proba_path = blob / "proba.npy" + if not proba_path.is_file(): + raise FileNotFoundError("proba.npy missing — run Predict first") + proba = np.load(proba_path).astype(np.float32) + t = float(threshold) + if not 0.0 <= t <= 1.0: + raise ValueError("threshold must be in [0, 1]") + h, w, _ = proba.shape + labels = np.zeros((h, w), dtype=np.uint8) + labels[proba[:, :, j] >= t] = np.uint8(int(class_id)) + raw = labels.reshape(-1).tobytes() + return { + "run_id": run_id, + "class_id": int(class_id), + "class_index": int(j), + "threshold": t, + "width": int(w), + "height": int(h), + "n_positive": int(np.count_nonzero(labels)), + "label_map_b64": base64.b64encode(raw).decode("ascii"), + } diff --git a/ipred/src/ipred/trainers/__init__.py b/ipred/src/ipred/trainers/__init__.py new file mode 100644 index 0000000..1a360ad --- /dev/null +++ b/ipred/src/ipred/trainers/__init__.py @@ -0,0 +1,25 @@ +"""Trainer plugin registry.""" + +from __future__ import annotations + +from ipred.trainers.base import ModelArtifacts, TrainerPlugin +from ipred.trainers.catboost_trainer import CatBoostTrainer + +_REGISTRY: dict[str, TrainerPlugin] = { + CatBoostTrainer.id: CatBoostTrainer(), +} + + +def get_trainer(trainer_id: str) -> TrainerPlugin: + """Return a registered trainer or raise KeyError.""" + try: + return _REGISTRY[trainer_id] + except KeyError as exc: + raise KeyError( + f"unknown trainer {trainer_id!r}; available: {sorted(_REGISTRY)}" + ) from exc + + +def list_trainers() -> list[str]: + """Return registered trainer ids.""" + return sorted(_REGISTRY) diff --git a/ipred/src/ipred/trainers/base.py b/ipred/src/ipred/trainers/base.py new file mode 100644 index 0000000..3edd6d7 --- /dev/null +++ b/ipred/src/ipred/trainers/base.py @@ -0,0 +1,53 @@ +"""Trainer plugin interface.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Protocol + +import numpy as np + + +@dataclass +class ModelArtifacts: + """Artifacts produced by ``train``.""" + + class_ids: list[int] + feature_labels: list[str] + cal_scores_by_class: dict[int, list[float]] + train_accuracy: float + n_train: int + n_cal: int + n_samples: int + params: dict[str, Any] + extras: dict[str, Any] = field(default_factory=dict) + # In-memory handle for immediate predict_proba (plugin-specific) + model_handle: Any = None + + +class TrainerPlugin(Protocol): + """Train / predict_proba contract for pixel classifiers.""" + + id: str + + def train( + self, + x_train: np.ndarray, + y_train: np.ndarray, + x_cal: np.ndarray, + y_cal: np.ndarray, + *, + feature_labels: list[str], + config: dict[str, Any], + ) -> ModelArtifacts: + """Fit model and compute Mondrian calibration scores.""" + + def predict_proba(self, model_handle: Any, x: np.ndarray) -> np.ndarray: + """Return NxK probabilities aligned with training class order.""" + + def save(self, model_handle: Any, dest_dir: Path) -> None: + """Persist model weights into *dest_dir*.""" + + def load(self, dest_dir: Path) -> Any: + """Load model weights from *dest_dir*.""" diff --git a/ipred/src/ipred/trainers/catboost_trainer.py b/ipred/src/ipred/trainers/catboost_trainer.py new file mode 100644 index 0000000..5bf01fc --- /dev/null +++ b/ipred/src/ipred/trainers/catboost_trainer.py @@ -0,0 +1,107 @@ +"""CatBoost TrainerPlugin.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any + +import numpy as np +from catboost import CatBoostClassifier + +from ipred.trainers.base import ModelArtifacts + + +class CatBoostTrainer: + """CatBoost multiclass pixel trainer with Mondrian cal scores.""" + + id = "catboost" + + def train( + self, + x_train: np.ndarray, + y_train: np.ndarray, + x_cal: np.ndarray, + y_cal: np.ndarray, + *, + feature_labels: list[str], + config: dict[str, Any], + ) -> ModelArtifacts: + """Fit CatBoost and collect per-class nonconformity scores.""" + iterations = int(config.get("iterations", 200)) + depth = int(config.get("depth", 6)) + learning_rate = float(config.get("learning_rate", 0.1)) + random_seed = int(config.get("random_seed", 0)) + params = { + "iterations": iterations, + "depth": depth, + "learning_rate": learning_rate, + "loss_function": "MultiClass", + "random_seed": random_seed, + } + clf = CatBoostClassifier( + iterations=iterations, + depth=depth, + learning_rate=learning_rate, + loss_function="MultiClass", + verbose=False, + allow_writing_files=False, + random_seed=random_seed, + thread_count=-1, + ) + clf.fit(x_train, y_train) + pred = np.asarray(clf.predict(x_train)).reshape(-1).astype(np.int32) + train_accuracy = float(np.mean(pred == y_train)) + + proba_cal = np.asarray(clf.predict_proba(x_cal), dtype=np.float64) + classes = [int(c) for c in np.asarray(clf.classes_).tolist()] + class_to_col = {c: i for i, c in enumerate(classes)} + cal_scores: dict[int, list[float]] = {} + for cid in classes: + mask = y_cal == cid + if not np.any(mask): + continue + col = class_to_col[cid] + scores = 1.0 - proba_cal[mask, col] + cal_scores[cid] = np.sort(scores.astype(np.float64)).tolist() + if not cal_scores: + raise ValueError("no calibration scores") + + labels = list(feature_labels) + raw_imp = np.asarray(clf.get_feature_importance(), dtype=np.float64) + if len(labels) != raw_imp.size: + labels = [f"f{i}" for i in range(raw_imp.size)] + order = np.argsort(-raw_imp) + feature_importances = [ + {"label": labels[int(i)], "importance": float(raw_imp[int(i)])} for i in order + ] + + return ModelArtifacts( + class_ids=classes, + feature_labels=labels, + cal_scores_by_class=cal_scores, + train_accuracy=train_accuracy, + n_train=int(y_train.shape[0]), + n_cal=int(y_cal.shape[0]), + n_samples=int(y_train.shape[0] + y_cal.shape[0]), + params=params, + extras={"feature_importances": feature_importances}, + model_handle=clf, + ) + + def predict_proba(self, model_handle: Any, x: np.ndarray) -> np.ndarray: + """Return NxK float64 probabilities.""" + proba = np.asarray(model_handle.predict_proba(x), dtype=np.float64) + if proba.ndim == 1: + proba = np.stack([1.0 - proba, proba], axis=1) + return proba + + def save(self, model_handle: Any, dest_dir: Path) -> None: + """Write ``model.cbm``.""" + dest_dir.mkdir(parents=True, exist_ok=True) + model_handle.save_model(str(dest_dir / "model.cbm")) + + def load(self, dest_dir: Path) -> Any: + """Load ``model.cbm``.""" + clf = CatBoostClassifier() + clf.load_model(str(dest_dir / "model.cbm")) + return clf diff --git a/ipred/tests/test_api_health.py b/ipred/tests/test_api_health.py new file mode 100644 index 0000000..29680bc --- /dev/null +++ b/ipred/tests/test_api_health.py @@ -0,0 +1,51 @@ +"""API health + session open.""" + +from __future__ import annotations + +import pytest +from fastapi.testclient import TestClient + +from ipred import api + + +@pytest.fixture() +def client(tmp_path, monkeypatch: pytest.MonkeyPatch) -> TestClient: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + api._catalog = None + with TestClient(api.app) as c: + yield c + + +def test_health(client: TestClient) -> None: + r = client.get("/health") + assert r.status_code == 200 + assert r.json()["status"] == "ok" + + +def test_open_session_and_list_setups(client: TestClient) -> None: + r = client.post( + "/sessions", + json={"kind": "local", "source": "x.png", "root": "/tmp"}, + ) + assert r.status_code == 200 + body = r.json() + assert "session_id" in body + assert "project_id" in body + + r2 = client.get("/setups") + assert r2.status_code == 200 + ids = {s["id"] for s in r2.json()["setups"]} + assert "default-skimage" in ids + assert "default-mark25" in ids + assert "default-skimage-mark25" in ids + assert "default-mark25-clahe" in ids + assert "default-mark11" in ids + assert "default-skimage-mark11" in ids + assert "default-mark11-clahe" in ids + assert "default-slimsam-clahe" in ids + + r3 = client.get("/compositions") + assert r3.status_code == 200 + cids = {c["id"] for c in r3.json()["compositions"]} + assert "comp-skimage" in cids + assert "comp-skimage-mark11" in cids diff --git a/ipred/tests/test_api_train_multi.py b/ipred/tests/test_api_train_multi.py new file mode 100644 index 0000000..c04a73d --- /dev/null +++ b/ipred/tests/test_api_train_multi.py @@ -0,0 +1,62 @@ +"""POST /train/multi over HTTP.""" + +from __future__ import annotations + +import numpy as np +import pytest +import tifffile +from fastapi.testclient import TestClient + +from ipred import api + + +@pytest.fixture() +def client(tmp_path, monkeypatch: pytest.MonkeyPatch) -> TestClient: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + api._catalog = None + with TestClient(api.app) as c: + yield c + + +def test_train_multi_over_http(client: TestClient, tmp_path) -> None: + stack = np.zeros((2, 48, 48), dtype=np.uint8) + stack[:, :, 24:] = 200 + tifffile.imwrite(tmp_path / "stack.tif", stack, photometric="minisblack") + + r = client.post( + "/sessions", json={"kind": "local", "source": "stack.tif", "root": str(tmp_path)} + ) + assert r.status_code == 200 + session_id = r.json()["session_id"] + + feature_ids: dict[str, str] = {} + for slice_index in (0, 1): + r = client.post( + "/preprocess", + json={ + "session_id": session_id, + "feature_setup_id": "default-skimage", + "slice_index": slice_index, + }, + ) + assert r.status_code == 200, r.text + feature_ids[str(slice_index)] = r.json()["feature_id"] + + shapes = [ + {"kind": "rectangle", "classId": 1, "x": 2, "y": 2, "w": 18, "h": 40}, + {"kind": "rectangle", "classId": 2, "x": 28, "y": 2, "w": 18, "h": 40}, + ] + r = client.post( + "/train/multi", + json={ + "session_id": session_id, + "slices": {"0": shapes, "1": shapes}, + "feature_ids": feature_ids, + "trainer_id": "catboost", + "config": {"iterations": 20, "depth": 4, "random_seed": 0}, + }, + ) + assert r.status_code == 200, r.text + body = r.json() + assert body["class_ids"] == [1, 2] + assert body["trained_slice_indices"] == [0, 1] diff --git a/ipred/tests/test_array_blobs.py b/ipred/tests/test_array_blobs.py new file mode 100644 index 0000000..f590091 --- /dev/null +++ b/ipred/tests/test_array_blobs.py @@ -0,0 +1,77 @@ +"""Array blob data plane for remote preprocess.""" + +from __future__ import annotations + +import base64 + +import numpy as np +import pytest +from fastapi.testclient import TestClient + +from ipred import api, array_blobs, compositions + + +@pytest.fixture() +def client(tmp_path, monkeypatch: pytest.MonkeyPatch) -> TestClient: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + api._catalog = None + with TestClient(api.app) as c: + yield c + + +def test_save_load_array_blob(tmp_path, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + arr = np.arange(12, dtype=np.float32).reshape(3, 4) + rec = array_blobs.save_array_blob("proj1", arr) + out = array_blobs.load_array_blob("proj1", rec["array_ref"]) + np.testing.assert_array_equal(out, arr) + + +def test_modules_and_compositions_api(client: TestClient) -> None: + r = client.get("/modules") + assert r.status_code == 200 + ids = {m["id"] for m in r.json()["modules"]} + assert "tomojepa" in ids + + r2 = client.get("/compositions") + assert r2.status_code == 200 + cids = {c["id"] for c in r2.json()["compositions"]} + assert "comp-skimage" in cids + + r3 = client.get("/compositions/comp-skimage") + assert r3.status_code == 200 + assert "preview_labels" in r3.json() + + +def test_upload_array_and_preprocess(client: TestClient) -> None: + sess = client.post( + "/sessions", + json={"kind": "local", "source": "x.png", "root": "/tmp"}, + ).json() + arr = np.linspace(0, 1, 16 * 16, dtype=np.float32).reshape(16, 16) + raw = arr.astype(np.float32).tobytes() + up = client.post( + f"/sessions/{sess['session_id']}/arrays", + json={ + "session_id": sess["session_id"], + "shape": [16, 16], + "dtype": "float32", + "data_b64": base64.b64encode(raw).decode("ascii"), + }, + ) + assert up.status_code == 200 + ref = up.json()["array_ref"] + compositions.ensure_default_compositions(api.get_catalog()) + prep = client.post( + "/preprocess", + json={ + "session_id": sess["session_id"], + "composition_id": "comp-skimage", + "array_ref": ref, + "slice_index": 0, + }, + ) + assert prep.status_code == 200, prep.text + body = prep.json() + assert body["n_channels"] > 0 + assert body["composition_id"] == "comp-skimage" diff --git a/ipred/tests/test_array_source.py b/ipred/tests/test_array_source.py new file mode 100644 index 0000000..4fe3a1a --- /dev/null +++ b/ipred/tests/test_array_source.py @@ -0,0 +1,143 @@ +import numpy as np + +from ipred import array_source + + +class _FakeArray: + def __init__(self, shape): + self._data = np.zeros(shape, dtype="float32") + self.structure_family = "array" + self.shape = shape + # Spies so tests can assert on WHICH access pattern was actually used — + # the whole point of the O(N^2) regression fix is that a single-slice + # read must index this node, not realize the whole thing via __array__. + self.array_calls = 0 + self.getitem_calls: list[int] = [] + + def __array__(self, dtype=None): + self.array_calls += 1 + return self._data + + def __getitem__(self, idx): + self.getitem_calls.append(idx) + return self._data[idx] + + +class _FakeContainer(dict): + structure_family = "container" + + +def _fake_pyramid(): + """Container mimicking a registered multiscale volume: scaleN -> {image: array}.""" + return _FakeContainer( + { + "scale0": _FakeContainer({"image": _FakeArray((9, 32, 32))}), + "scale1": _FakeContainer({"image": _FakeArray((5, 16, 16))}), + "scale2": _FakeContainer({"image": _FakeArray((3, 8, 8))}), + "scale3": _FakeContainer({"image": _FakeArray((2, 4, 4))}), + "scale4": _FakeContainer({"image": _FakeArray((1, 2, 2))}), + } + ) + + +class TestDescendToArray: + def test_descends_multiscale_pyramid_to_finest_level(self) -> None: + node = array_source._descend_to_array(_fake_pyramid()) + assert not array_source._is_container_node(node) + assert np.asarray(node).shape == (9, 32, 32) + + def test_descends_wrapper_container_to_array_stack(self) -> None: + wrapper = _FakeContainer({"vol": _FakeContainer({"a": _FakeArray((4, 4))})}) + node = array_source._descend_to_array(wrapper) + assert array_source._is_container_node(node) # container of arrays == stack + + def test_leaves_bare_array_untouched(self) -> None: + arr = _FakeArray((4, 4)) + assert array_source._descend_to_array(arr) is arr + + def test_five_level_pyramid_does_not_collapse_to_length_5_array(self) -> None: + # Regression: a naive np.asarray(container) on an un-descended 5-child + # pyramid container previously yielded shape (5,), tripping + # "unsupported array shape" in _index_slice during train/infer. + node = array_source._descend_to_array(_fake_pyramid()) + data = np.asarray(node) + assert data.shape != (5,) + sliced = array_source._index_slice(data, 0) + assert sliced.shape == (32, 32) + + +class TestReadTiledLazySlice: + """Regression: a batch-apply job over N slices used to call np.asarray(node) + (materializing the WHOLE stack) on every single-slice request — an O(N^2) + read that dominated "predict across all slices" wall-clock time. A real + single-slice read must index the node BEFORE realizing it.""" + + def _patch_from_uri(self, monkeypatch, client): + import tiled.client + + monkeypatch.setattr(tiled.client, "from_uri", lambda *a, **k: client) + + def test_bare_nhw_stack_indexes_without_materializing_whole_array(self, monkeypatch) -> None: + arr = _FakeArray((50, 16, 16)) + client = _FakeContainer({"ds": arr}) + self._patch_from_uri(monkeypatch, client) + + out = array_source.read_slice(kind="tiled", source="ds", slice_index=7, server_uri="http://x") + + assert out.shape == (16, 16) + assert arr.getitem_calls == [7] + assert arr.array_calls == 0 # the whole-stack path must never fire + + def test_container_of_per_slice_arrays_only_realizes_the_requested_one(self, monkeypatch) -> None: + wanted = _FakeArray((16, 16)) + other = _FakeArray((16, 16)) + client = _FakeContainer({"ds": _FakeContainer({"0": other, "1": wanted})}) + self._patch_from_uri(monkeypatch, client) + + out = array_source.read_slice(kind="tiled", source="ds", slice_index=1, server_uri="http://x") + + assert out.shape == (16, 16) + assert wanted.array_calls == 1 + assert other.array_calls == 0 # the sibling slice must never be touched + + def test_hwc_single_image_is_returned_whole(self, monkeypatch) -> None: + arr = _FakeArray((64, 64, 3)) + client = _FakeContainer({"ds": arr}) + self._patch_from_uri(monkeypatch, client) + + out = array_source.read_slice(kind="tiled", source="ds", slice_index=0, server_uri="http://x") + + assert out.shape == (64, 64, 3) + assert arr.getitem_calls == [] # a single HWC image has nothing to index into + + +class TestReadLocalTiffPage: + """Regression: `_read_local` used to call `tifffile.imread` (decodes every + page) for a single-slice request.""" + + def test_reads_one_page_of_a_multipage_tiff(self, tmp_path, monkeypatch) -> None: + import tifffile + + path = tmp_path / "stack.tif" + pages = np.stack([np.full((8, 8), i, dtype=np.uint8) for i in range(5)]) + tifffile.imwrite(str(path), pages) + + # `tifffile.imread` decodes the whole stack up front — assert the fix + # never calls it for a multi-page file. + monkeypatch.setattr( + tifffile, "imread", lambda *a, **k: (_ for _ in ()).throw(AssertionError("full-stack imread called")) + ) + + out = array_source._read_local("stack.tif", slice_index=3, root=str(tmp_path)) + assert out.shape == (8, 8) + assert int(out[0, 0]) == 3 + + def test_single_page_tiff_still_works(self, tmp_path) -> None: + import tifffile + + path = tmp_path / "single.tif" + tifffile.imwrite(str(path), np.full((8, 8), 9, dtype=np.uint8)) + + out = array_source._read_local("single.tif", slice_index=0, root=str(tmp_path)) + assert out.shape == (8, 8) + assert int(out[0, 0]) == 9 diff --git a/ipred/tests/test_catalog_feature_banks.py b/ipred/tests/test_catalog_feature_banks.py new file mode 100644 index 0000000..7165624 --- /dev/null +++ b/ipred/tests/test_catalog_feature_banks.py @@ -0,0 +1,56 @@ +"""Catalog.delete_feature_bank — cleanup for batch-apply jobs (see +backend/ipred_batch_jobs.py's volume-apply job, which relies on this to avoid +accumulating a ~1 GB feature bank per slice with no eviction).""" +from __future__ import annotations + +import pytest + +from ipred.catalog import Catalog + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def _insert_bank(catalog: Catalog, feature_id: str, blob_dir: str, slice_index: int = 0) -> None: + catalog.insert_feature_bank( + { + "feature_id": feature_id, + "project_id": "proj", + "setup_id": "setup", + "content_hash": "hash", + "slice_index": slice_index, + "n_channels": 3, + "height": 8, + "width": 8, + "blob_dir": blob_dir, + "setup_snapshot": "{}", + "status": "ready", + "created_at": "2026-01-01T00:00:00+00:00", + } + ) + + +def test_delete_feature_bank_removes_the_row_and_returns_its_blob_dir(catalog: Catalog) -> None: + _insert_bank(catalog, "fid-1", "/blobs/fid-1") + + blob_dir = catalog.delete_feature_bank("fid-1") + + assert blob_dir == "/blobs/fid-1" + assert catalog.get_feature_bank("fid-1") is None + + +def test_delete_feature_bank_is_a_noop_for_an_unknown_id(catalog: Catalog) -> None: + assert catalog.delete_feature_bank("does-not-exist") is None + + +def test_delete_feature_bank_does_not_touch_other_rows(catalog: Catalog) -> None: + _insert_bank(catalog, "fid-1", "/blobs/fid-1", slice_index=0) + _insert_bank(catalog, "fid-2", "/blobs/fid-2", slice_index=1) + + catalog.delete_feature_bank("fid-1") + + assert catalog.get_feature_bank("fid-1") is None + assert catalog.get_feature_bank("fid-2") is not None diff --git a/ipred/tests/test_catalog_sessions.py b/ipred/tests/test_catalog_sessions.py new file mode 100644 index 0000000..33e3882 --- /dev/null +++ b/ipred/tests/test_catalog_sessions.py @@ -0,0 +1,39 @@ +"""Catalog project/session persistence.""" + +from __future__ import annotations + +import os + +import pytest + +from ipred.catalog import Catalog, project_id_for + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def test_project_id_stable() -> None: + a = project_id_for(kind="local", source="a.tif", root="/data") + b = project_id_for(kind="local", source="a.tif", root="/data") + c = project_id_for(kind="local", source="b.tif", root="/data") + assert a == b + assert a != c + + +def test_open_session_creates_project(catalog: Catalog) -> None: + s1 = catalog.open_session(kind="local", source="img.tif", root=str(os.environ["LOCAL_DATA_ROOT"])) + s2 = catalog.open_session(kind="local", source="img.tif", root=str(os.environ["LOCAL_DATA_ROOT"])) + assert s1.project_id == s2.project_id + assert s1.session_id != s2.session_id + got = catalog.get_session(s1.session_id) + assert got is not None + assert got.current_feature_id is None + + +def test_set_session_currents(catalog: Catalog) -> None: + s = catalog.open_session(kind="local", source="x.png") + updated = catalog.set_session_currents(s.session_id, feature_id="abc") + assert updated.current_feature_id == "abc" diff --git a/ipred/tests/test_channel_png_path.py b/ipred/tests/test_channel_png_path.py new file mode 100644 index 0000000..3a2b36b --- /dev/null +++ b/ipred/tests/test_channel_png_path.py @@ -0,0 +1,40 @@ +"""channel_png_path validates feature_id against its known-safe shape before +ever constructing a filesystem path from it (CodeQL py/path-injection).""" + +from __future__ import annotations + +import uuid + +import pytest + +from ipred import preprocess + + +@pytest.fixture() +def local_data_root(tmp_path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return tmp_path + + +def test_valid_feature_id_resolves_under_project_blob_dir(local_data_root) -> None: + from ipred.paths import project_blob_dir + + feature_id = uuid.uuid4().hex + path = preprocess.channel_png_path("abc123", feature_id, 0) + assert path == project_blob_dir("abc123") / "features" / feature_id / "channels" / "0000.png" + + +def test_traversal_payload_as_feature_id_is_rejected(local_data_root) -> None: + with pytest.raises(ValueError, match="invalid feature_id"): + preprocess.channel_png_path("abc123", "../../../../etc/passwd", 0) + + +def test_wrong_length_feature_id_is_rejected(local_data_root) -> None: + with pytest.raises(ValueError, match="invalid feature_id"): + preprocess.channel_png_path("abc123", "not-a-real-uuid", 0) + + +def test_uppercase_hex_feature_id_is_rejected(local_data_root) -> None: + # uuid.uuid4().hex is always lowercase — anything else can't be a real one. + with pytest.raises(ValueError, match="invalid feature_id"): + preprocess.channel_png_path("abc123", uuid.uuid4().hex.upper(), 0) diff --git a/ipred/tests/test_compositions.py b/ipred/tests/test_compositions.py new file mode 100644 index 0000000..66ba5ee --- /dev/null +++ b/ipred/tests/test_compositions.py @@ -0,0 +1,80 @@ +"""Composition validation, builtins, and legacy setup mapping.""" + +from __future__ import annotations + +import numpy as np +import pytest + +from ipred import compositions, compose_run, feature_setups +from ipred.catalog import Catalog +from ipred.modules import list_module_catalog + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def test_module_catalog_lists_core_modules() -> None: + ids = {m["id"] for m in list_module_catalog()} + assert ids >= { + "skimage_multiscale", + "clahe", + "slimsam", + "tomojepa", + "pca", + } + + +def test_ensure_default_compositions(catalog: Catalog) -> None: + comps = compositions.ensure_default_compositions(catalog) + ids = {c["id"] for c in comps} + assert "comp-skimage" in ids + assert "comp-skimage-mark11" in ids + assert "comp-mark25-clahe" in ids + assert "comp-slimsam-clahe" in ids + + +def test_legacy_setup_maps_to_composition(catalog: Catalog) -> None: + compositions.ensure_default_compositions(catalog) + assert ( + compositions.resolve_preprocess_id("default-skimage-mark25") + == "comp-skimage-mark25" + ) + assert compositions.resolve_preprocess_id("default-mark11") == "comp-mark11-clahe" + + +def test_preview_concat_skimage(catalog: Catalog) -> None: + compositions.ensure_default_compositions(catalog) + doc = compositions.resolve_composition("comp-skimage") + labels = compositions.preview_concat_labels(doc) + assert any("intensity" in lab for lab in labels) + + +def test_run_composition_skimage_only(catalog: Catalog) -> None: + compositions.ensure_default_compositions(catalog) + doc = compositions.resolve_composition("comp-skimage") + arr = np.linspace(0, 1, 32 * 32, dtype=np.float32).reshape(32, 32) + uint8, floats, labels, emb, meta = compose_run.run_composition(arr, doc) + assert emb is None + assert floats.ndim == 3 + assert uint8.shape == floats.shape + assert len(labels) == floats.shape[-1] + + +def test_content_hash_stable() -> None: + doc = { + "nodes": [{"id": "n1", "module": "skimage_multiscale", "params": {}}], + "outputs": ["n1"], + } + compositions.validate_composition(doc) + h1 = compositions.content_hash(doc) + h2 = compositions.content_hash({**doc, "name": "x"}) + assert h1 == h2 + + +def test_feature_setups_still_seed(catalog: Catalog) -> None: + feature_setups.ensure_default_setups(catalog) + compositions.ensure_default_compositions(catalog) + assert feature_setups.load_setup("default-skimage") is not None diff --git a/ipred/tests/test_emb_pca_channels.py b/ipred/tests/test_emb_pca_channels.py new file mode 100644 index 0000000..ecbc79f --- /dev/null +++ b/ipred/tests/test_emb_pca_channels.py @@ -0,0 +1,30 @@ +"""Encoder PCA → full-res heatmap channels.""" + +from __future__ import annotations + +import numpy as np + +from ipred import sam_embed + + +def test_emb_grid_to_pca_channels_shape() -> None: + rng = np.random.default_rng(0) + emb = rng.normal(size=(8, 8, 16)).astype(np.float32) + up, labels, info = sam_embed.emb_grid_to_pca_channels( + emb, out_h=32, out_w=40, n_components=8 + ) + assert up.shape == (32, 40, 8) + assert labels == [f"pca{i}" for i in range(8)] + assert info["pca_dims"] == 8 + assert info["baked_into_float_stack"] is True + + +def test_emb_grid_pca_caps_at_rank() -> None: + rng = np.random.default_rng(1) + emb = rng.normal(size=(4, 4, 5)).astype(np.float32) + up, labels, info = sam_embed.emb_grid_to_pca_channels( + emb, out_h=16, out_w=16, n_components=64 + ) + assert up.shape[-1] == 5 + assert len(labels) == 5 + assert info["pca_dims"] == 5 diff --git a/ipred/tests/test_feature_setups.py b/ipred/tests/test_feature_setups.py new file mode 100644 index 0000000..19ea35f --- /dev/null +++ b/ipred/tests/test_feature_setups.py @@ -0,0 +1,58 @@ +"""Feature Setup shelf defaults.""" + +from __future__ import annotations + +import pytest + +from ipred.catalog import Catalog +from ipred import feature_setups + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def test_ensure_defaults(catalog: Catalog) -> None: + defaults = feature_setups.ensure_default_setups(catalog) + assert len(defaults) == 10 + ids = {d["id"] for d in defaults} + assert "default-skimage" in ids + assert "default-slimsam" in ids + assert "default-skimage-slimsam" in ids + assert "default-slimsam-clahe" in ids + assert "default-mark25" in ids + assert "default-skimage-mark25" in ids + assert "default-mark25-clahe" in ids + assert "default-mark11" in ids + assert "default-skimage-mark11" in ids + assert "default-mark11-clahe" in ids + sk = feature_setups.resolve_setup("default-skimage") + assert sk["kind"] == "procedure" + assert sk["procedure_id"] == feature_setups.PROCEDURE_SKIMAGE + mark = feature_setups.resolve_setup("default-mark25") + assert mark["kind"] == "weights" + assert mark["weights_format"] == feature_setups.WEIGHTS_TOMOJEPA_MARK25 + assert mark["inference"]["input_size"] == 512 + combo = feature_setups.resolve_setup("default-skimage-mark25") + assert combo["encoder_setup_id"] == "default-mark25" + assert combo["params"]["input_size"] == 512 + assert combo["params"]["resize"] is True + assert combo["params"]["pca_dims"] == 64 + clahe = feature_setups.resolve_setup("default-mark25-clahe") + assert clahe["procedure_id"] == feature_setups.PROCEDURE_CLAHE_ENCODER + assert clahe["encoder_setup_id"] == "default-mark25" + assert clahe["params"]["resize"] is True + assert clahe["params"]["input_size"] == 512 + assert clahe["params"]["clahe"] is True + assert clahe["params"]["pca_dims"] == 64 + mark11 = feature_setups.resolve_setup("default-mark11") + assert mark11["weights_format"] == feature_setups.WEIGHTS_TOMOJEPA_MARK11 + combo11 = feature_setups.resolve_setup("default-skimage-mark11") + assert combo11["encoder_setup_id"] == "default-mark11" + clahe11 = feature_setups.resolve_setup("default-mark11-clahe") + assert clahe11["encoder_setup_id"] == "default-mark11" + slim_clahe = feature_setups.resolve_setup("default-slimsam-clahe") + assert slim_clahe["procedure_id"] == feature_setups.PROCEDURE_CLAHE_ENCODER + assert slim_clahe["encoder_setup_id"] == "default-slimsam" diff --git a/ipred/tests/test_infer_store_probabilities.py b/ipred/tests/test_infer_store_probabilities.py new file mode 100644 index 0000000..75d46e4 --- /dev/null +++ b/ipred/tests/test_infer_store_probabilities.py @@ -0,0 +1,72 @@ +"""run_infer(store_probabilities=False) — the option volume-apply batch jobs +use (backend/ipred_batch_jobs.py) to skip writing the largest thing a run +saves, since that flow only ever reads commit.png back.""" +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import pytest +from PIL import Image as PILImage + +from ipred import feature_setups, preprocess, train_infer +from ipred.catalog import Catalog + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def _trained_session(catalog: Catalog, tmp_path) -> str: + feature_setups.ensure_default_setups(catalog) + img = np.zeros((48, 48), dtype=np.uint8) + img[:, 24:] = 220 + PILImage.fromarray(img, mode="L").save(tmp_path / "blob.png") + + session = catalog.open_session(kind="local", source="blob.png", root=str(tmp_path)) + preprocess.run_preprocess(catalog, session_id=session.session_id, feature_setup_id="default-skimage") + shapes = [ + {"kind": "rectangle", "classId": 1, "x": 2, "y": 2, "w": 18, "h": 40}, + {"kind": "rectangle", "classId": 2, "x": 28, "y": 2, "w": 18, "h": 40}, + ] + train_infer.run_train( + catalog, session_id=session.session_id, shapes=shapes, + trainer_id="catboost", config={"iterations": 20, "depth": 4, "random_seed": 0}, + ) + return session.session_id + + +def test_store_probabilities_false_skips_proba_npy(catalog: Catalog, tmp_path) -> None: + session_id = _trained_session(catalog, tmp_path) + + run = train_infer.run_infer(catalog, session_id=session_id, alpha=0.2, store_probabilities=False) + + blob = Path(run["blob_dir"]) + assert not (blob / "proba.npy").exists() + # Everything volume-apply's commit step actually needs must still exist. + assert (blob / "commit.png").is_file() + assert (blob / "commit.npy").exists() + assert (blob / "status.npy").exists() + assert (blob / "membership.npy").exists() + + +def test_store_probabilities_true_is_still_the_default(catalog: Catalog, tmp_path) -> None: + session_id = _trained_session(catalog, tmp_path) + + run = train_infer.run_infer(catalog, session_id=session_id, alpha=0.2) + + assert (Path(run["blob_dir"]) / "proba.npy").is_file() + + +def test_store_probabilities_false_produces_the_same_commit_map(catalog: Catalog, tmp_path) -> None: + """The math must be identical either way — only persistence changes.""" + session_id = _trained_session(catalog, tmp_path) + + with_proba = train_infer.run_infer(catalog, session_id=session_id, alpha=0.2, store_probabilities=True) + without_proba = train_infer.run_infer(catalog, session_id=session_id, alpha=0.2, store_probabilities=False) + + a = np.load(Path(with_proba["blob_dir"]) / "commit.npy") + b = np.load(Path(without_proba["blob_dir"]) / "commit.npy") + assert np.array_equal(a, b) diff --git a/ipred/tests/test_labels.py b/ipred/tests/test_labels.py new file mode 100644 index 0000000..c412864 --- /dev/null +++ b/ipred/tests/test_labels.py @@ -0,0 +1,62 @@ +"""Label rasterization accepts studio shape wire formats.""" + +from __future__ import annotations + +import numpy as np + +from ipred.labels import build_label_map, shape_to_mask + + +def test_polygon_flat_xy_list() -> None: + """Studio polygons store points as flat [x0,y0,x1,y1,...].""" + shape = { + "kind": "polygon", + "classId": 6, + "points": [10.0, 10.0, 40.0, 10.0, 40.0, 40.0, 10.0, 40.0], + } + mask = shape_to_mask(shape, 64, 64) + assert int(mask.sum()) > 100 + + labels = build_label_map([shape], 64, 64) + assert int((labels == 6).sum()) == int(mask.sum()) + + +def test_polygon_nested_pairs_still_work() -> None: + shape = { + "kind": "polygon", + "classId": 1, + "points": [[10, 10], [40, 10], [40, 40], [10, 40]], + } + assert shape_to_mask(shape, 64, 64).sum() > 100 + + +def test_polygon_holes_are_excluded() -> None: + """A background polygon clipped around another class (holes) must not + paint over that class's region — see clipToClasses.ts on the frontend.""" + outer = [0.0, 0.0, 64.0, 0.0, 64.0, 64.0, 0.0, 64.0] + hole = [20.0, 20.0, 40.0, 20.0, 40.0, 40.0, 20.0, 40.0] + shape = {"kind": "polygon", "classId": 1, "points": outer, "holes": [hole]} + mask = shape_to_mask(shape, 64, 64) + solid = shape_to_mask({"kind": "polygon", "classId": 1, "points": outer}, 64, 64) + assert int(mask.sum()) < int(solid.sum()) + assert not mask[30, 30] + + other = {"kind": "polygon", "classId": 2, "points": hole} + labels = build_label_map([shape, other], 64, 64) + assert labels[30, 30] == 2 + assert labels[5, 5] == 1 + + +def test_brush_flat_points_and_radius() -> None: + """Studio brush strokes use flat points + radius (not nested {x,y}).""" + shape = { + "kind": "brush", + "classId": 2, + "strokes": [ + {"points": [8.0, 8.0, 20.0, 20.0, 32.0, 8.0], "radius": 4.0, "mode": "paint"}, + ], + } + mask = shape_to_mask(shape, 48, 48) + assert int(mask.sum()) > 20 + labels = build_label_map([shape], 48, 48) + assert np.any(labels == 2) diff --git a/ipred/tests/test_manifold.py b/ipred/tests/test_manifold.py new file mode 100644 index 0000000..e32b9d6 --- /dev/null +++ b/ipred/tests/test_manifold.py @@ -0,0 +1,42 @@ +"""Manifold Suggest Labels on persisted ipred feature banks.""" + +from __future__ import annotations + +import numpy as np +from PIL import Image as PILImage +import pytest + +from ipred.catalog import Catalog +from ipred import feature_setups, manifold_jobs, preprocess + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def test_manifold_sample_on_feature_bank(catalog: Catalog, tmp_path) -> None: + feature_setups.ensure_default_setups(catalog) + img = np.zeros((64, 64), dtype=np.uint8) + img[:, 32:] = 200 + img[16:48, 16:48] = 120 + PILImage.fromarray(img, mode="L").save(tmp_path / "m.png") + + session = catalog.open_session(kind="local", source="m.png", root=str(tmp_path)) + bank = preprocess.run_preprocess( + catalog, + session_id=session.session_id, + feature_setup_id="default-skimage", + ) + out = manifold_jobs.run_manifold_sample( + catalog, + feature_id=bank["feature_id"], + k=8, + box_size=16, + ) + assert out["sample_id"] + assert out["n_picked"] >= 1 + assert len(out["points"]) >= 1 + png = manifold_jobs.heatmap_png(out["sample_id"]) + assert png[:8] == b"\x89PNG\r\n\x1a\n" diff --git a/ipred/tests/test_preprocess_cache.py b/ipred/tests/test_preprocess_cache.py new file mode 100644 index 0000000..b98250e --- /dev/null +++ b/ipred/tests/test_preprocess_cache.py @@ -0,0 +1,116 @@ +"""Cache-aware preprocess.""" + +from __future__ import annotations + +from pathlib import Path + +import numpy as np +from PIL import Image as PILImage +import pytest + +from ipred.catalog import Catalog +from ipred import feature_setups, preprocess, tomojepa_embed + +WEIGHTS = Path(__file__).resolve().parents[1] / "models" / "tomojepa25.pth" + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def test_preprocess_cache_hit(catalog: Catalog, tmp_path) -> None: + feature_setups.ensure_default_setups(catalog) + # Write a local PNG under LOCAL_DATA_ROOT + img = np.zeros((32, 32), dtype=np.uint8) + img[:, 16:] = 200 + rel = "sample.png" + PILImage.fromarray(img, mode="L").save(tmp_path / rel) + + session = catalog.open_session(kind="local", source=rel, root=str(tmp_path)) + first = preprocess.run_preprocess( + catalog, + session_id=session.session_id, + feature_setup_id="default-skimage", + slice_index=0, + ) + assert first["cache_hit"] is False + assert first["n_channels"] > 0 + + second = preprocess.run_preprocess( + catalog, + session_id=session.session_id, + feature_setup_id="default-skimage", + slice_index=0, + ) + assert second["cache_hit"] is True + assert second["feature_id"] == first["feature_id"] + + +@pytest.mark.skipif( + not WEIGHTS.is_file() or not tomojepa_embed.encoder_available(str(WEIGHTS)), + reason="tomojepa25.pth or torch/timm not available", +) +def test_preprocess_mark25_writes_dense_emb(catalog: Catalog, tmp_path) -> None: + feature_setups.ensure_default_setups(catalog) + img = np.linspace(0, 255, 64 * 48, dtype=np.float32).reshape(64, 48).astype( + np.uint8 + ) + rel = "tomo.png" + PILImage.fromarray(img, mode="L").save(tmp_path / rel) + + session = catalog.open_session(kind="local", source=rel, root=str(tmp_path)) + out = preprocess.run_preprocess( + catalog, + session_id=session.session_id, + feature_setup_id="default-skimage-mark25", + slice_index=0, + ) + assert out["cache_hit"] is False + bank = preprocess.load_feature_bank_arrays(out["blob_dir"]) + assert bank["sam_emb"] is not None + assert bank["sam_emb"].shape == (32, 32, 64) + assert bank["sam_meta"]["encoder"] == "mark25" + assert bank["sam_meta"]["input_size"] == 512 + + +@pytest.mark.skipif( + not WEIGHTS.is_file() or not tomojepa_embed.encoder_available(str(WEIGHTS)), + reason="tomojepa25.pth or torch/timm not available", +) +def test_preprocess_mark25_clahe_no_skimage(catalog: Catalog, tmp_path) -> None: + feature_setups.ensure_default_setups(catalog) + img = np.linspace(0, 255, 80 * 64, dtype=np.float32).reshape(80, 64).astype( + np.uint8 + ) + rel = "tomo_clahe.png" + PILImage.fromarray(img, mode="L").save(tmp_path / rel) + + session = catalog.open_session(kind="local", source=rel, root=str(tmp_path)) + out = preprocess.run_preprocess( + catalog, + session_id=session.session_id, + feature_setup_id="default-mark25-clahe", + slice_index=0, + ) + assert out["cache_hit"] is False + assert out["n_channels"] == 1 + 64 + bank = preprocess.load_feature_bank_arrays(out["blob_dir"]) + assert bank["labels"][0] == "clahe" + assert bank["labels"][1:5] == ["pca0", "pca1", "pca2", "pca3"] + assert bank["float_stack"].shape[-1] == 65 + assert bank["float_stack"].shape[:2] == (80, 64) + assert bank["sam_emb"] is not None + assert bank["sam_emb"].shape == (32, 32, 64) + assert bank["sam_meta"]["encoder"] == "mark25" + assert bank["sam_meta"]["resize"] is True + assert bank["sam_meta"]["input_size"] == 512 + assert bank["sam_meta"]["baked_into_float_stack"] is True + assert bank["sam_meta"]["pca_dims"] == 64 + # Channel PNGs written for heatmap browsing + from pathlib import Path + + ch_dir = Path(out["blob_dir"]) / "channels" + assert (ch_dir / "0000.png").is_file() + assert (ch_dir / "0064.png").is_file() diff --git a/ipred/tests/test_proba_channels.py b/ipred/tests/test_proba_channels.py new file mode 100644 index 0000000..b2dccf6 --- /dev/null +++ b/ipred/tests/test_proba_channels.py @@ -0,0 +1,97 @@ +"""Softmax proba heatmaps + per-class threshold maps.""" + +from __future__ import annotations + +import base64 + +import numpy as np +import pytest + +from ipred.catalog import Catalog +from ipred import feature_setups, preprocess, train_infer +from ipred.paths import project_blob_dir + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def _fake_run(catalog: Catalog, tmp_path) -> str: + """Insert a minimal run with synthetic HxWxK softmax proba.""" + feature_setups.ensure_default_setups(catalog) + session = catalog.open_session(kind="local", source="x.png", root=str(tmp_path)) + project_id = session.project_id + run_id = "runproba01" + blob = project_blob_dir(project_id) / "runs" / run_id + blob.mkdir(parents=True, exist_ok=True) + # Two classes; left half prefers class 1, right prefers class 2 + h, w, k = 8, 10, 2 + proba = np.zeros((h, w, k), dtype=np.float16) + proba[:, :5, 0] = 0.8 + proba[:, :5, 1] = 0.2 + proba[:, 5:, 0] = 0.3 + proba[:, 5:, 1] = 0.7 + np.save(blob / "proba.npy", proba) + meta = { + "run_id": run_id, + "model_id": "m", + "feature_id": "f", + "alpha": 0.05, + "class_ids": [1, 2], + "counts": {"singleton": 0, "multi": 0, "abstain": 0}, + } + import json + from datetime import datetime, timezone + + (blob / "meta.json").write_text(json.dumps(meta), encoding="utf-8") + now = datetime.now(timezone.utc).isoformat() + catalog.insert_run( + { + "run_id": run_id, + "project_id": project_id, + "model_id": "m", + "feature_id": "f", + "alpha": 0.05, + "blob_dir": str(blob), + "meta_json": json.dumps(meta), + "created_at": now, + "updated_at": now, + } + ) + return run_id + + +def test_proba_heatmap_png(catalog: Catalog, tmp_path) -> None: + run_id = _fake_run(catalog, tmp_path) + png = train_infer.proba_heatmap_png(catalog, run_id, 0) + assert png[:8] == b"\x89PNG\r\n\x1a\n" + png1 = train_infer.proba_heatmap_png(catalog, run_id, 1) + assert png1[:8] == b"\x89PNG\r\n\x1a\n" + with pytest.raises(ValueError): + train_infer.proba_heatmap_png(catalog, run_id, 9) + + +def test_threshold_class_label_map(catalog: Catalog, tmp_path) -> None: + run_id = _fake_run(catalog, tmp_path) + out = train_infer.threshold_class_label_map( + catalog, run_id, class_id=1, threshold=0.5 + ) + assert out["width"] == 10 + assert out["height"] == 8 + assert out["class_id"] == 1 + raw = base64.b64decode(out["label_map_b64"]) + labels = np.frombuffer(raw, dtype=np.uint8).reshape(8, 10) + assert np.all(labels[:, :5] == 1) + assert np.all(labels[:, 5:] == 0) + assert out["n_positive"] == 8 * 5 + + out2 = train_infer.threshold_class_label_map( + catalog, run_id, class_id=2, threshold=0.6 + ) + labels2 = np.frombuffer( + base64.b64decode(out2["label_map_b64"]), dtype=np.uint8 + ).reshape(8, 10) + assert np.all(labels2[:, 5:] == 2) + assert np.all(labels2[:, :5] == 0) diff --git a/ipred/tests/test_tomojepa_embed.py b/ipred/tests/test_tomojepa_embed.py new file mode 100644 index 0000000..2968e91 --- /dev/null +++ b/ipred/tests/test_tomojepa_embed.py @@ -0,0 +1,129 @@ +"""Mark25 / TomoJEPA dense embedding encode.""" + +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import pytest + +from ipred import tomojepa_embed + +WEIGHTS = Path(__file__).resolve().parents[1] / "models" / "tomojepa25.pth" +WEIGHTS11 = Path(__file__).resolve().parents[1] / "models" / "tomojepa11.pth" +needs_weights = pytest.mark.skipif( + not WEIGHTS.is_file() or not tomojepa_embed.torch_available(), + reason="tomojepa25.pth or torch/timm not available", +) +needs_mark11 = pytest.mark.skipif( + not WEIGHTS11.is_file() or not tomojepa_embed.torch_available(), + reason="tomojepa11.pth or torch/timm not available", +) + + +def test_resolve_weights_path_default(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("TOMOJEPA_WEIGHTS", raising=False) + resolved = tomojepa_embed.resolve_weights_path(None) + if WEIGHTS.is_file(): + assert resolved is not None + assert resolved.name == "tomojepa25.pth" + else: + assert resolved is None + + +def test_resolve_weights_path_env(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + fake = tmp_path / "custom.pth" + fake.write_bytes(b"x") + monkeypatch.setenv("TOMOJEPA_WEIGHTS", str(fake)) + assert tomojepa_embed.resolve_weights_path(None) == fake.resolve() + + +def test_to_minus_one_one_linear() -> None: + x = np.array([[0.0, 0.5, 1.0]], dtype=np.float32) + y = tomojepa_embed.to_minus_one_one(x) + np.testing.assert_allclose(y, [[-1.0, 0.0, 1.0]], atol=1e-6) + + +@needs_weights +def test_encode_feeds_model_in_minus_one_one(monkeypatch: pytest.MonkeyPatch) -> None: + """All Mark25 paths must pass intensities in [-1, 1] into the network.""" + captured: list[float] = [] + + real_cached = tomojepa_embed._cached_model + + def _wrap(path_str: str): + net, device = real_cached(path_str) + orig_forward = net.forward + + def forward(x): # noqa: ANN001 + t = x.detach().cpu().numpy() + captured.append(float(t.min())) + captured.append(float(t.max())) + return orig_forward(x) + + net.forward = forward # type: ignore[method-assign] + return net, device + + monkeypatch.setattr(tomojepa_embed, "_cached_model", _wrap) + arr = np.linspace(10, 200, 64 * 48, dtype=np.float32).reshape(64, 48) + tomojepa_embed.encode_dense_embeddings( + arr, weights_path=str(WEIGHTS), input_size=128 + ) + assert captured, "model forward was not called" + assert min(captured) >= -1.0 - 1e-5 + assert max(captured) <= 1.0 + 1e-5 + # After min-max + linear map, extremes should reach near ±1. + assert min(captured) < -0.9 + assert max(captured) > 0.9 + + +@needs_weights +def test_load_state_dict_strict() -> None: + net = tomojepa_embed.build_encoder() + state = tomojepa_embed.load_checkpoint_state(WEIGHTS) + net.load_state_dict(state, strict=True) + + +@needs_weights +def test_encode_dense_shape_512() -> None: + arr = np.linspace(0, 1, 400 * 300, dtype=np.float32).reshape(400, 300) + emb, orig_hw, reshaped_hw = tomojepa_embed.encode_dense_embeddings( + arr, weights_path=str(WEIGHTS), input_size=512 + ) + assert orig_hw == (400, 300) + assert reshaped_hw == (512, 512) + assert emb.shape == (32, 32, 64) + assert emb.dtype == np.float32 + + +@needs_weights +def test_encode_dense_honors_input_size() -> None: + arr = np.random.default_rng(0).random((128, 96), dtype=np.float32) + emb, _, reshaped = tomojepa_embed.encode_dense_embeddings( + arr, weights_path=str(WEIGHTS), input_size=256 + ) + assert reshaped == (256, 256) + assert emb.shape == (16, 16, 64) + + +@needs_weights +def test_encode_native_no_resize() -> None: + arr = np.random.default_rng(1).random((100, 80), dtype=np.float32) + emb, orig_hw, reshaped = tomojepa_embed.encode_dense_embeddings( + arr, weights_path=str(WEIGHTS), resize=False + ) + assert orig_hw == (100, 80) + # padded to multiple of 16 + assert reshaped == (112, 80) + assert emb.shape == (7, 5, 64) + + +@needs_mark11 +def test_encode_mark11_dense_is_256d() -> None: + """Mark11 dense projector is 256-D (Mark25 is 64-D).""" + tomojepa_embed._cached_model.cache_clear() + arr = np.linspace(0, 1, 64 * 48, dtype=np.float32).reshape(64, 48) + emb, _, _ = tomojepa_embed.encode_dense_embeddings( + arr, weights_path=str(WEIGHTS11), input_size=128 + ) + assert emb.shape == (8, 8, 256) diff --git a/ipred/tests/test_tomojepa_onnx.py b/ipred/tests/test_tomojepa_onnx.py new file mode 100644 index 0000000..5e69528 --- /dev/null +++ b/ipred/tests/test_tomojepa_onnx.py @@ -0,0 +1,27 @@ +"""TomoJEPA ONNX encode smoke (skips if onnx file missing).""" + +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import pytest + +from ipred import tomojepa_onnx + +ONNX25 = Path(__file__).resolve().parents[1] / "models" / "tomojepa25.onnx" +needs_onnx = pytest.mark.skipif( + not ONNX25.is_file() or not tomojepa_onnx.onnx_available(), + reason="tomojepa25.onnx not available", +) + + +@needs_onnx +def test_onnx_encode_dense_shape() -> None: + arr = np.linspace(0, 1, 64 * 48, dtype=np.float32).reshape(64, 48) + emb, orig, reshaped = tomojepa_onnx.encode_dense_embeddings( + arr, weights_path=str(ONNX25), input_size=128 + ) + assert orig == (64, 48) + assert reshaped == (128, 128) + assert emb.shape == (8, 8, 64) diff --git a/ipred/tests/test_train_multi_slice.py b/ipred/tests/test_train_multi_slice.py new file mode 100644 index 0000000..2682825 --- /dev/null +++ b/ipred/tests/test_train_multi_slice.py @@ -0,0 +1,101 @@ +"""Multi-slice training pools labeled pixels across slices into one model.""" + +from __future__ import annotations + +import numpy as np +import pytest +import tifffile + +from ipred import feature_setups, preprocess, train_infer +from ipred.catalog import Catalog + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def _make_stack(tmp_path, n=3, size=48): + stack = np.zeros((n, size, size), dtype=np.uint8) + for i in range(n): + stack[i, :, size // 2:] = 200 + tifffile.imwrite(tmp_path / "stack.tif", stack, photometric="minisblack") + return "stack.tif" + + +def test_train_multi_slice_pools_across_slices(catalog: Catalog, tmp_path) -> None: + feature_setups.ensure_default_setups(catalog) + source = _make_stack(tmp_path, n=3, size=48) + session = catalog.open_session(kind="local", source=source, root=str(tmp_path)) + + per_slice_shapes: dict[int, list[dict]] = {} + feature_ids: dict[int, str] = {} + for slice_index in (0, 1, 2): + bank = preprocess.run_preprocess( + catalog, + session_id=session.session_id, + feature_setup_id="default-skimage", + slice_index=slice_index, + ) + feature_ids[slice_index] = bank["feature_id"] + per_slice_shapes[slice_index] = [ + {"kind": "rectangle", "classId": 1, "x": 2, "y": 2, "w": 18, "h": 40}, + {"kind": "rectangle", "classId": 2, "x": 28, "y": 2, "w": 18, "h": 40}, + ] + + trained = train_infer.run_train_multi_slice( + catalog, + session_id=session.session_id, + per_slice_shapes=per_slice_shapes, + feature_ids=feature_ids, + trainer_id="catboost", + config={"iterations": 20, "depth": 4, "random_seed": 0}, + ) + assert trained["model_id"] + assert trained["class_ids"] == [1, 2] + assert trained["trained_slice_indices"] == [0, 1, 2] + # Pooled across 3 slices should see roughly 3x the pixels of one slice alone. + assert trained["n_samples"] > 1000 + + # The model is usable for ordinary single-slice infer against any one slice. + run = train_infer.run_infer( + catalog, + session_id=session.session_id, + model_id=trained["model_id"], + feature_id=feature_ids[1], + alpha=0.2, + ) + assert run["run_id"] + assert run["class_ids"] == [1, 2] + + +def test_train_multi_slice_requires_two_classes_pooled(catalog: Catalog, tmp_path) -> None: + feature_setups.ensure_default_setups(catalog) + source = _make_stack(tmp_path, n=2, size=48) + session = catalog.open_session(kind="local", source=source, root=str(tmp_path)) + + per_slice_shapes: dict[int, list[dict]] = {} + feature_ids: dict[int, str] = {} + for slice_index in (0, 1): + bank = preprocess.run_preprocess( + catalog, + session_id=session.session_id, + feature_setup_id="default-skimage", + slice_index=slice_index, + ) + feature_ids[slice_index] = bank["feature_id"] + # Only class 1 labeled on every slice — never two classes, even pooled. + per_slice_shapes[slice_index] = [ + {"kind": "rectangle", "classId": 1, "x": 2, "y": 2, "w": 18, "h": 40}, + ] + + with pytest.raises(ValueError, match="need at least two classes"): + train_infer.run_train_multi_slice( + catalog, + session_id=session.session_id, + per_slice_shapes=per_slice_shapes, + feature_ids=feature_ids, + trainer_id="catboost", + config={}, + ) diff --git a/ipred/tests/test_train_rethreshold.py b/ipred/tests/test_train_rethreshold.py new file mode 100644 index 0000000..2d896b9 --- /dev/null +++ b/ipred/tests/test_train_rethreshold.py @@ -0,0 +1,71 @@ +"""Train → infer → rethreshold without recompute.""" + +from __future__ import annotations + +import numpy as np +from PIL import Image as PILImage +import pytest + +from ipred.catalog import Catalog +from ipred import feature_setups, preprocess, train_infer + + +@pytest.fixture() +def catalog(tmp_path, monkeypatch: pytest.MonkeyPatch) -> Catalog: + monkeypatch.setenv("LOCAL_DATA_ROOT", str(tmp_path)) + return Catalog(tmp_path / "ipred" / "catalog.db") + + +def test_train_infer_rethreshold(catalog: Catalog, tmp_path) -> None: + feature_setups.ensure_default_setups(catalog) + img = np.zeros((48, 48), dtype=np.uint8) + img[:, 24:] = 220 + PILImage.fromarray(img, mode="L").save(tmp_path / "blob.png") + + session = catalog.open_session( + kind="local", source="blob.png", root=str(tmp_path) + ) + preprocess.run_preprocess( + catalog, + session_id=session.session_id, + feature_setup_id="default-skimage", + ) + shapes = [ + {"kind": "rectangle", "classId": 1, "x": 2, "y": 2, "w": 18, "h": 40}, + {"kind": "rectangle", "classId": 2, "x": 28, "y": 2, "w": 18, "h": 40}, + ] + trained = train_infer.run_train( + catalog, + session_id=session.session_id, + shapes=shapes, + trainer_id="catboost", + config={"iterations": 20, "depth": 4, "random_seed": 0}, + ) + assert trained["model_id"] + imps = trained.get("feature_importances") or [] + assert imps, "CatBoost train must return ranked feature_importances" + assert "label" in imps[0] and "importance" in imps[0] + assert imps[0]["importance"] >= imps[-1]["importance"] + + run = train_infer.run_infer( + catalog, + session_id=session.session_id, + alpha=0.2, + ) + assert run["run_id"] + blob = run["blob_dir"] + assert (tmp_path / "ipred").exists() or True + from pathlib import Path + + proba_path = Path(blob) / "proba.npy" + assert proba_path.is_file() + proba_before = np.load(proba_path).copy() + + rerun = train_infer.run_rethreshold( + catalog, + session_id=session.session_id, + alpha=0.05, + ) + assert rerun["alpha"] == 0.05 + proba_after = np.load(proba_path) + np.testing.assert_array_equal(proba_before, proba_after) diff --git a/mkdocs.yml b/mkdocs.yml index bcd7bfb..580a45c 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -88,8 +88,11 @@ nav: - 3. Annotate: guide/annotate.md - 4. Annotation guide: guide/reference-guide.md - 5. Export & download: guide/export.md + - 6. Train a deep model: guide/train.md + - 7. 3D volume view: guide/volume.md - Reference: - Software architecture: reference/architecture.md + - Production deployment: reference/deployment.md - Keyboard shortcuts: reference/shortcuts.md - Troubleshooting: reference/troubleshooting.md diff --git a/start_all.sh b/start_all.sh index 3dc6dbd..454e413 100755 --- a/start_all.sh +++ b/start_all.sh @@ -21,6 +21,7 @@ TILED_CONFIG="$SCRIPT_DIR/tiled/config.yml" MKDOCS_CONFIG="$SCRIPT_DIR/mkdocs.yml" TILED_PORT="${TILED_PORT:-8010}" BACKEND_PORT="${BACKEND_PORT:-8002}" +IPRED_PORT="${IPRED_PORT:-8003}" FRONTEND_PORT="${FRONTEND_PORT:-5173}" DOCS_PORT="${DOCS_PORT:-8000}" # PROD=1 (or SERVE_MODE=prod): build the optimized SPA and have the backend serve @@ -30,6 +31,7 @@ STATIC_DIR="$BACKEND_DIR/static" RUN_DIR="$SCRIPT_DIR/.run" TILED_PID_FILE="$RUN_DIR/tiled.pid" BACKEND_PID_FILE="$RUN_DIR/backend.pid" +IPRED_PID_FILE="$RUN_DIR/ipred.pid" FRONTEND_PID_FILE="$RUN_DIR/frontend.pid" DOCS_PID_FILE="$RUN_DIR/docs.pid" ENV_DIR="" @@ -126,6 +128,7 @@ cleanup_managed_processes() { stop_managed_process "$DOCS_PID_FILE" "docs" stop_managed_process "$FRONTEND_PID_FILE" "frontend" stop_managed_process "$BACKEND_PID_FILE" "backend" + stop_managed_process "$IPRED_PID_FILE" "iPred" stop_managed_process "$TILED_PID_FILE" "Tiled" } @@ -222,6 +225,7 @@ reclaim_orphaned_repo_ports() { fi stop_repo_listener_on_port "$FRONTEND_PORT" "frontend" "$FRONTEND_DIR" "vite" "" stop_repo_listener_on_port "$BACKEND_PORT" "backend" "$BACKEND_DIR" "annotation_server:app" "uvicorn" + stop_repo_listener_on_port "$IPRED_PORT" "iPred" "$SCRIPT_DIR/ipred" "ipred.api:app" "uvicorn" stop_repo_listener_on_port "$DOCS_PORT" "docs" "$SCRIPT_DIR" "mkdocs" "" # Repo-scoped: only reclaim OUR own stale Tiled (its command line contains this # repo's config path). A foreign Tiled on the port is left alone — we coexist by @@ -283,6 +287,64 @@ ensure_backend_env() { fi } +ensure_ml_env() { + # Best-effort + backgrounded-in-spirit: never blocks or fails startup — the + # Train tab (and denoise_bake.py's "model" denoise method) already degrade + # to a clear "unavailable" state via train_common.torch_available() / + # dlsia_available() when these aren't installed, so skipping this is always + # a safe default, not a broken one. + local want_install="${INSTALL_ML:-}" + # Default to installing on Apple Silicon Mac (torch gets MPS acceleration + # there); other platforms opt in explicitly, since a CUDA/CPU torch wheel is + # a multi-GB download the user may not want on every fresh machine. + if [ -z "$want_install" ] && [ "$(uname -s)" = "Darwin" ] && [ "$(uname -m)" = "arm64" ]; then + want_install=1 + fi + if [ "$want_install" != "1" ]; then + echo -e "${YELLOW} Skipping ML deps (Train tab) — set INSTALL_ML=1 to install torch/dlsia here.${NC}" + return 0 + fi + + if "$PYTHON" -c "import torch" >/dev/null 2>&1; then + echo -e "${GREEN} torch already installed — Train tab available.${NC}" + else + echo -e "${CYAN}==> Installing torch for the Train tab (can take a few minutes on first run)...${NC}" + if uv pip install --python "$PYTHON" "torch>=2.4"; then + echo -e "${GREEN} torch installed.${NC}" + else + echo -e "${YELLOW} torch install failed — the Train tab will report 'unavailable'. Retry manually:${NC}" + echo -e "${YELLOW} uv pip install --python \"$PYTHON\" \"torch>=2.4\"${NC}" + fi + fi + + if "$PYTHON" -c "import dlsia" >/dev/null 2>&1; then + echo -e "${GREEN} dlsia already installed — TUNet model family available.${NC}" + else + echo -e "${CYAN}==> Installing dlsia (+ qlty) for the Train tab's TUNet model family...${NC}" + if uv pip install --python "$PYTHON" "dlsia>=0.3" "qlty>=1.5"; then + echo -e "${GREEN} dlsia installed.${NC}" + else + echo -e "${YELLOW} dlsia install failed — the TUNet model family will report 'unavailable'.${NC}" + fi + fi + + # Some ops (e.g. certain interpolate modes) aren't implemented on MPS yet; + # fall back to CPU for just that op instead of erroring. + export PYTORCH_ENABLE_MPS_FALLBACK=1 +} + +ensure_ipred_env() { + # NOTE: probe "ipred.api", not bare "ipred" — the repo's top-level ipred/ + # directory (sibling of ipred/src/) is itself an importable namespace package + # from cwd, so "import ipred" can silently succeed even when the real + # editable install (ipred/src/ipred) was never pip-installed. + if ! "$PYTHON" -c "import ipred.api" >/dev/null 2>&1; then + echo -e "${YELLOW} Installing iPred (interactive segmentation service) via uv...${NC}" + uv pip install --python "$PYTHON" -e "$SCRIPT_DIR/ipred" >/dev/null 2>&1 || return 1 + fi + "$PYTHON" -c "import ipred.api" >/dev/null 2>&1 +} + ensure_frontend_runtime() { if command -v npm >/dev/null 2>&1 && can_run_npm "$(command -v npm)"; then NPM_CMD=("$(command -v npm)") @@ -325,10 +387,11 @@ tiled_cmd() { cleanup() { echo "" echo -e "${YELLOW}Shutting down...${NC}" - kill "$TILED_PID" "$BACKEND_PID" "$FRONTEND_PID" ${DOCS_PID:+"$DOCS_PID"} 2>/dev/null || true - wait "$TILED_PID" "$BACKEND_PID" "$FRONTEND_PID" ${DOCS_PID:+"$DOCS_PID"} 2>/dev/null || true + kill "$TILED_PID" "$BACKEND_PID" "$FRONTEND_PID" ${IPRED_PID:+"$IPRED_PID"} ${DOCS_PID:+"$DOCS_PID"} 2>/dev/null || true + wait "$TILED_PID" "$BACKEND_PID" "$FRONTEND_PID" ${IPRED_PID:+"$IPRED_PID"} ${DOCS_PID:+"$DOCS_PID"} 2>/dev/null || true cleanup_pid_file "$TILED_PID_FILE" cleanup_pid_file "$BACKEND_PID_FILE" + cleanup_pid_file "$IPRED_PID_FILE" cleanup_pid_file "$FRONTEND_PID_FILE" cleanup_pid_file "$DOCS_PID_FILE" echo -e "${GREEN}Done.${NC}" @@ -337,6 +400,7 @@ cleanup() { trap cleanup SIGINT SIGTERM ensure_backend_env +ensure_ml_env ensure_frontend_runtime cleanup_managed_processes reclaim_orphaned_repo_ports @@ -359,6 +423,15 @@ if [ "$BACKEND_PORT" != "$_orig_backend_port" ]; then fi export API_PROXY_TARGET="http://127.0.0.1:${BACKEND_PORT}" +# iPred: fall back to the next free port if busy. Exported as IPRED_URL before +# the backend starts, so ipred_client.py picks up the chosen port. +_orig_ipred_port="$IPRED_PORT" +IPRED_PORT="$(pick_free_port "$IPRED_PORT" "iPred")" +if [ "$IPRED_PORT" != "$_orig_ipred_port" ]; then + echo -e "${YELLOW} iPred port ${_orig_ipred_port} is in use — using ${IPRED_PORT} instead.${NC}" +fi +export IPRED_URL="http://127.0.0.1:${IPRED_PORT}" + # Frontend: fall back to the next free port if busy (Vite serves on --port below). _orig_frontend_port="$FRONTEND_PORT" FRONTEND_PORT="$(pick_free_port "$FRONTEND_PORT" "Frontend")" @@ -468,6 +541,11 @@ ensure_sam_model # resolved server-side by backend/tiled_config.py. Keys are never sent to the # frontend. (Port must match backend/tiled_config.py default: 8010.) # --------------------------------------------------------------------------- +# Generated 3-D volume pyramids (backend/tiff_stack_source.py) are served by +# Tiled in place, from a path listed in tiled/config.yml's readable_storage. +# Created BEFORE Tiled starts, since that is when readable_storage is resolved. +mkdir -p "$SCRIPT_DIR/.tiled/volumes" + echo -e "${CYAN}==> Starting Tiled (port ${TILED_PORT})...${NC}" # Repair catalog asset paths in case the repo was moved or cloned to a new location. @@ -528,6 +606,23 @@ if [ "$TILED_READY" != 1 ]; then exit 1 fi +# --------------------------------------------------------------------------- +# Git submodules: the WebGPU volume renderer that powers the 3D tab lives in +# frontend/vendor/view_tomography_recon_app. A clone without --recurse-submodules +# leaves it empty, which otherwise shows up as an unresolved import halfway +# through a Vite build. Initialise it here so that never happens. +# --------------------------------------------------------------------------- +ZARR_VIEWER_ENTRY="$FRONTEND_DIR/vendor/view_tomography_recon_app/src/zarr-viewer/src/ome-zarr-viewer.ts" +if [ ! -f "$ZARR_VIEWER_ENTRY" ]; then + echo -e "${CYAN}==> Initialising git submodules (3D volume renderer)...${NC}" + if git -C "$SCRIPT_DIR" submodule update --init --recursive; then + echo -e "${GREEN} Submodules ready.${NC}" + else + echo -e "${YELLOW} Could not initialise submodules — the 3D tab will not build.${NC}" + echo -e "${YELLOW} Fix with: git submodule update --init --recursive${NC}" + fi +fi + # --------------------------------------------------------------------------- # Production SPA build (PROD=1): build the optimized frontend and stage it in # backend/static/ BEFORE the backend starts — the SPA mount is decided at import @@ -551,6 +646,41 @@ else rm -rf "$STATIC_DIR" # ensure the backend serves API-only in dev fi +# --------------------------------------------------------------------------- +# iPred — standalone interactive-segmentation service (CatBoost + conformal +# prediction on composable feature banks). Optional: the Assist/Predict stages +# show a clear "not running" state when this is down, and everything else in +# the app works regardless (backend/ipred_routes.py returns 503 rather than +# erroring). Startup failures here are warnings, never fatal. +# --------------------------------------------------------------------------- +IPRED_PID="" +if ensure_ipred_env; then + echo -e "${CYAN}==> Starting iPred (port ${IPRED_PORT})...${NC}" + (cd "$SCRIPT_DIR" && "$ENV_DIR/bin/uvicorn" ipred.api:app --host 127.0.0.1 --port "$IPRED_PORT") & + IPRED_PID=$! + echo "$IPRED_PID" > "$IPRED_PID_FILE" + echo -e "${GREEN} iPred PID: $IPRED_PID${NC}" + + echo -e "${CYAN} Waiting for iPred...${NC}" + IPRED_READY=0 + for i in $(seq 1 20); do + if curl -sf "http://127.0.0.1:${IPRED_PORT}/health" >/dev/null 2>&1; then + echo -e "${GREEN} iPred ready at http://127.0.0.1:${IPRED_PORT}${NC}" + IPRED_READY=1 + break + fi + if ! kill -0 "$IPRED_PID" 2>/dev/null; then + break + fi + sleep 0.5 + done + if [ "$IPRED_READY" != 1 ]; then + echo -e "${YELLOW} iPred did not come up in time — Assist/Predict will show a 'not running' state until it does.${NC}" + fi +else + echo -e "${YELLOW}==> Skipping iPred (install failed) — Assist/Predict will show a 'not running' state.${NC}" +fi + # --------------------------------------------------------------------------- # Backend # --------------------------------------------------------------------------- @@ -632,6 +762,11 @@ else echo -e "${GREEN} Frontend : http://127.0.0.1:${FRONTEND_PORT}${NC}" fi echo -e "${GREEN} Backend : http://127.0.0.1:${BACKEND_PORT}${NC}" +if [ "$IPRED_READY" = "1" ]; then + echo -e "${GREEN} iPred : http://127.0.0.1:${IPRED_PORT}${NC}" +else + echo -e "${YELLOW} iPred : not running (Assist/Predict disabled)${NC}" +fi if [ -n "$DOCS_PID" ]; then echo -e "${GREEN} Docs : http://127.0.0.1:${DOCS_PORT}${NC}" fi diff --git a/tiled/config.docker.yml b/tiled/config.docker.yml new file mode 100644 index 0000000..be21762 --- /dev/null +++ b/tiled/config.docker.yml @@ -0,0 +1,34 @@ +# Portable counterpart to config.yml, for the app-full Docker image +# (Dockerfile's app-full stage / docker-entrypoint-full.sh) — config.yml +# hardcodes a local dev machine's absolute data path +# (readable_storage: /Users/.../data), which has no meaning inside a +# container. Everything here lives under /data, the same volume mount point +# the backend's LOCAL_DATA_ROOT already uses (see docker-compose.full.yml) — +# one volume to persist for the whole container's state. +authentication: + allow_anonymous_access: true + # No hardcoded key here either — docker-entrypoint-full.sh generates one + # (or uses TILED_API_KEY if you set it) and passes it via `tiled serve + # --api-key`, exactly like start_all.sh does for local dev. Anonymous + # access is read-only; the key authorizes writes. + +allow_origins: + - "*" + +trees: + - path: / + tree: tiled.catalog:from_uri + args: + uri: "sqlite+aiosqlite:////data/.tiled/catalog.db" + writable_storage: "/data/.tiled/data" + readable_storage: + # Bind-mount your own source datasets here (see docker-compose.full.yml + # / docker-compose.local.yml's volumes). Uploads/ingest also land here. + # Must match the bind-mount target in both compose files exactly — + # Tiled refuses to register an external asset outside this list. + - "/data/processed" + # Generated 3-D volume pyramids (backend/tiff_stack_source.py) and + # mask pyramids (backend/mask_pyramid.py) — must stay in step with + # VOLUME_CACHE_DIR, set by docker-entrypoint-full.sh. + - "/data/.tiled/volumes" + init_if_not_exists: true diff --git a/tiled/config.yml b/tiled/config.yml index 418b86a..81f4049 100644 --- a/tiled/config.yml +++ b/tiled/config.yml @@ -16,5 +16,10 @@ trees: writable_storage: ".tiled/data" readable_storage: - "/Users/david/Documents/data" + # Generated 3-D volume pyramids (backend/tiff_stack_source.py). Tiled + # serves these in place, so the path must be readable here or the node + # registers with no children and every chunk request 500s. + # Keep in step with VOLUME_CACHE_DIR. + - "./.tiled/volumes" init_if_not_exists: true