diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 56ddeb7..95f3723 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -142,15 +142,50 @@ jobs: exit 1 fi done - sha256sum -- *.exe *.dmg *.AppImage > SHA256SUMS.txt - sha256sum --check SHA256SUMS.txt + python <<'PY' + from hashlib import sha256 + from pathlib import Path + + files = [] + for extension in ("exe", "dmg", "AppImage"): + files.extend(sorted(Path(".").glob(f"*.{extension}"))) + + with Path("SHA256SUMS.txt").open("w", encoding="utf-8") as manifest: + for path in files: + manifest.write(f"{path.name}\t{sha256(path.read_bytes()).hexdigest()}\n") + + entries = {} + with Path("SHA256SUMS.txt").open("r", encoding="utf-8") as manifest: + for line in manifest: + name, digest = line.rstrip("\n").split("\t", 1) + entries[name] = digest + + for path in files: + digest = sha256(path.read_bytes()).hexdigest() + if entries.get(path.name) != digest: + raise SystemExit(f"Checksum mismatch for {path.name}") + PY - name: Publish complete prerelease working-directory: artifacts shell: bash run: | - state=$(gh api --paginate "repos/$GH_REPO/releases?per_page=100" \ - --jq '.[] | select(.tag_name == env.RELEASE_TAG) | .draft') + owner=${GH_REPO%/*} + repo=${GH_REPO#*/} + state=$(gh api graphql \ + -F owner="$owner" \ + -F name="$repo" \ + -F tag="$RELEASE_TAG" \ + -f query=' + query($owner: String!, $name: String!, $tag: String!) { + repository(owner: $owner, name: $name) { + release(tagName: $tag) { + isDraft + } + } + } + ' \ + --jq '.data.repository.release | if . == null then "" else .isDraft end') if [[ "$state" == "false" ]]; then echo "Release $RELEASE_TAG is already published; leaving its assets unchanged." exit 0 diff --git a/README.md b/README.md index ccb8a31..b212570 100644 --- a/README.md +++ b/README.md @@ -28,6 +28,8 @@ Local Forge is a fast, local-first desktop workbench for open models. Chat with - Persistent sessions, image attachments, search, pinning, and deletion - Local Diffusers and `.safetensors` / `.ckpt` model registration - Cancellable image generation with live step progress and Library import +- Local image descriptions and generation prompts through Ollama vision models +- Optional Face Fix, tiled 2x/4x UltraSharp upscaling, and saved NSFW region masks - Cancellable QLoRA training with live progress and standard PEFT output - Screenshot and reference-image imports shared by Studio and Library - Live RAM and NVIDIA telemetry, with unavailable data left unknown @@ -47,4 +49,4 @@ npm run dev Opening Vite directly uses a browser preview adapter. Use the Electron window for filesystem imports, MCP servers, encrypted credential persistence, and real hardware data. -See the [runtime and project guide](docs/RUNTIME_AND_PROJECT_GUIDE.md) for image and training setup, MCP servers, validation, releases, and project status. +See the [runtime and project guide](docs/RUNTIME_AND_PROJECT_GUIDE.md) for image and training setup, MCP servers, validation, releases, and project status. See the [contribution guide](CONTRIBUTING.md), [security policy](SECURITY.md), [asset notes](docs/THIRD_PARTY_ASSETS.md), and [MIT License](LICENSE) for repository policies and project details. diff --git a/docs/RUNTIME_AND_PROJECT_GUIDE.md b/docs/RUNTIME_AND_PROJECT_GUIDE.md index 2e6ef18..c7208c5 100644 --- a/docs/RUNTIME_AND_PROJECT_GUIDE.md +++ b/docs/RUNTIME_AND_PROJECT_GUIDE.md @@ -11,7 +11,9 @@ python -m pip install -r runtime/requirements-training.txt Select that interpreter under **Settings > Runtimes > Python runtime**. Local Forge never installs Python packages or downloads model weights automatically. -Studio executes local Diffusers directories containing `model_index.json`. Registered `.safetensors` and `.ckpt` files remain visible for inventory but must be converted to Diffusers format before generation. Completed PNGs are added to Library and stored under Local Forge's user-data `outputs/images` directory. +Studio executes local Diffusers directories containing `model_index.json`. Registered `.safetensors` and `.ckpt` files remain visible for inventory but must be converted to Diffusers format before generation. Completed PNGs are added to Library and stored under Local Forge's user-data `outputs/images` directory. Optional NSFW segmentation writes one binary mask per detected region beside the completed image. + +The Studio **Describe image** and **Create prompt** actions use the Ollama model currently selected in Workbench. That model must advertise Ollama's `vision` capability. Local Forge does not pre-classify or block adult images; description quality and model-level refusals depend on the selected vision model. Tune requires a local Hugging Face Transformers model directory containing `config.json`; Ollama tags and GGUF files are inference artifacts and are not offered as trainable base models. Datasets may be JSON, JSONL, or CSV and should contain either a `text` field or chat `messages`. Completed adapters use standard PEFT format under `outputs/adapters`. @@ -35,17 +37,17 @@ npm run build ## Automated releases -Every push to `main` (including a merged pull request) runs tests, lint, and the application build. After those checks pass, GitHub Actions packages that exact commit for Windows x64 (`.exe`), macOS universal (`.dmg`, Intel and Apple Silicon), and Linux x64 (`.AppImage`). +Every push to `main` (including a merged pull request) runs tests, lint, and the application build. After those checks pass, GitHub Actions packages that exact commit for Windows x64 (`.exe`), one universal macOS `.dmg` that runs on both Intel and Apple Silicon, and Linux x64 (`.AppImage`). Download installers from [GitHub Releases](https://github.com/Azayzel/local-forge/releases). Each main build is an unsigned prerelease, tagged `v-main.` (for example, `v0.1.0-main.42`). The installer version matches the tag without its `v` prefix. These alpha builds are not marked as the latest stable release. -All three installers and `SHA256SUMS.txt` are uploaded before the release is published. To verify a download, compare its SHA-256 hash with the matching entry in that file: +All three installers and `SHA256SUMS.txt` are uploaded before the release is published. That file lists one tab-separated `filenamesha256` entry per installer. To verify a download, compare its SHA-256 hash with the matching entry in that file: ```powershell Get-FileHash .\Local-Forge-0.1.0-main.42-win-x64.exe -Algorithm SHA256 ``` -Use your downloaded installer's actual filename. On Linux, download all three installers and the checksum file into one directory and run `sha256sum --check SHA256SUMS.txt`; on macOS use `shasum -a 256 -c SHA256SUMS.txt`. +Use your downloaded installer's actual filename. On Linux, run `sha256sum ./Local-Forge-0.1.0-main.42-linux-x64.AppImage`; on macOS run `shasum -a 256 ./Local-Forge-0.1.0-main.42-mac-universal.dmg`; in either case, compare the printed hash with the matching entry in `SHA256SUMS.txt`. Windows and macOS may show security warnings: these builds are not yet signed or notarized. See [release maintenance](../CONTRIBUTING.md#release-maintenance) for setup, versioning, and retries. diff --git a/docs/THIRD_PARTY_ASSETS.md b/docs/THIRD_PARTY_ASSETS.md index 22e9c4a..54bd192 100644 --- a/docs/THIRD_PARTY_ASSETS.md +++ b/docs/THIRD_PARTY_ASSETS.md @@ -17,6 +17,20 @@ The files are `1200 x 1800`, `1200 x 800`, `1200 x 800`, and `1200 x 1800` respe Images imported by a user are copied into that user's private Local Forge data directory. They are never included in the source tree or project releases. +## Optional enhancement models + +Local Forge does not bundle enhancement weights. When a user chooses Install in Studio, the app downloads an immutable, checksum-verified revision to that user's private Local Forge data directory: + +| Purpose | Repository | File | SHA-256 | +| --- | --- | --- | --- | +| 4x image upscaling | [lokCX/4x-Ultrasharp](https://huggingface.co/lokCX/4x-Ultrasharp) | `4x-UltraSharp.pth` | `a5812231fc936b42af08a5edba784195495d303d5b3248c24489ef0c4021fe01` | +| Face detection | [Bingsu/adetailer](https://huggingface.co/Bingsu/adetailer) | `face_yolov8n.pt` | `70b640f8f60b1cf0dcc72f30caf3da9495eb2fb6509da48c53374ad6806e6a9c` | +| Breast segmentation | [NSFW-API/NSFW_Segmentation](https://huggingface.co/NSFW-API/NSFW_Segmentation) | `nsfw-seg-breast-x.pt` | `5ad882ddaf149873be131943b373da9f15a0603c91508e2291ece83d729c8ecc` | +| Penis segmentation | [NSFW-API/NSFW_Segmentation](https://huggingface.co/NSFW-API/NSFW_Segmentation) | `nsfw-seg-penis-x.pt` | `49d9fc8ee67d3bdee44e46bf75aeb76058f5cd074bac027b55ff071b63a32e21` | +| Vagina segmentation | [NSFW-API/NSFW_Segmentation](https://huggingface.co/NSFW-API/NSFW_Segmentation) | `nsfw-seg-vagina-x.pt` | `f8349ae348f5cf041809c38bf38d2c80a0f3dccdb1375b97971358cde742197f` | + +These files remain subject to the terms published by their respective authors and are not covered by Local Forge's MIT code license. + ## Fonts and icons - Manrope and Space Grotesk are bundled through Fontsource packages and distributed under the SIL Open Font License 1.1. diff --git a/electron/chat.test.ts b/electron/chat.test.ts index 5eb4d32..372a48b 100644 --- a/electron/chat.test.ts +++ b/electron/chat.test.ts @@ -7,7 +7,7 @@ import type { ChatStreamEvent, McpServerConfig, } from "../src/types"; -import { runMcpChat } from "./chat"; +import { runMcpChat, warmChatModel } from "./chat"; import { closeMcpConnections } from "./mcp"; const fixture: McpServerConfig = { @@ -26,6 +26,41 @@ const fixture: McpServerConfig = { afterAll(() => closeMcpConnections()); describe("MCP chat orchestration", () => { + it("preloads the selected model with the chat keep-alive", async () => { + let requestPath = ""; + let requestBody: Record = {}; + const server = createServer(async (incoming, response) => { + requestPath = incoming.url ?? ""; + const chunks: Buffer[] = []; + for await (const chunk of incoming) chunks.push(Buffer.from(chunk)); + requestBody = JSON.parse( + Buffer.concat(chunks).toString("utf8"), + ) as Record; + response.writeHead(200, { "content-type": "application/json" }); + response.end(JSON.stringify({ done: true })); + }); + await new Promise((resolve) => + server.listen(0, "127.0.0.1", resolve), + ); + const address = server.address() as AddressInfo; + + try { + await warmChatModel(`http://127.0.0.1:${address.port}`, "fixture-model"); + } finally { + await new Promise((resolve, reject) => + server.close((error) => (error ? reject(error) : resolve())), + ); + } + + expect(requestPath).toBe("/api/generate"); + expect(requestBody).toMatchObject({ + model: "fixture-model", + prompt: "", + stream: false, + keep_alive: "30m", + }); + }); + it("approves and executes a tool before continuing the Ollama chat", async () => { const requests: Array> = []; const server = createServer(async (incoming, response) => { @@ -105,6 +140,8 @@ describe("MCP chat orchestration", () => { } expect(requests).toHaveLength(2); + expect(requests[0].think).toBe(false); + expect(requests[0].keep_alive).toBe("30m"); expect(requests[0].tools).toEqual( expect.arrayContaining([ expect.objectContaining({ diff --git a/electron/chat.ts b/electron/chat.ts index e6d9fa1..eee0a26 100644 --- a/electron/chat.ts +++ b/electron/chat.ts @@ -7,6 +7,7 @@ import type { import { callMcpTool, listEnabledMcpTools, type McpToolBinding } from "./mcp"; const MAX_TOOL_ROUNDS = 8; +const MODEL_KEEP_ALIVE = "30m"; interface OllamaToolCall { function: { @@ -40,18 +41,40 @@ interface RoundResult { evalDurationNs: number; } -function runtimeEndpoint(baseUrl: string): string { +function runtimeEndpoint(baseUrl: string, route = "/api/chat"): string { const url = new URL(baseUrl); const localHosts = new Set(["localhost", "127.0.0.1", "::1", "[::1]"]); if (url.protocol !== "http:" || !localHosts.has(url.hostname)) { throw new Error("Local Forge only connects to runtimes on this machine."); } - url.pathname = "/api/chat"; + url.pathname = route; url.search = ""; url.hash = ""; return url.toString(); } +export async function warmChatModel( + baseUrl: string, + model: string, +): Promise { + if (!model.trim()) return; + const response = await fetch(runtimeEndpoint(baseUrl, "/api/generate"), { + method: "POST", + headers: { "content-type": "application/json" }, + body: JSON.stringify({ + model, + prompt: "", + stream: false, + keep_alive: MODEL_KEEP_ALIVE, + }), + }); + if (!response.ok) { + const detail = await response.text(); + throw new Error(detail || `Runtime returned ${response.status}.`); + } + await response.text(); +} + async function readJsonLines( response: Response, onValue: (value: Record) => void, @@ -135,6 +158,8 @@ async function runRound( model: options.request.model, messages, options: options.request.options, + think: false, + keep_alive: MODEL_KEEP_ALIVE, tools: bindings.length > 0 ? toolDefinitions(bindings) : undefined, stream: true, }), diff --git a/electron/content-security-policy.test.ts b/electron/content-security-policy.test.ts new file mode 100644 index 0000000..99a7d37 --- /dev/null +++ b/electron/content-security-policy.test.ts @@ -0,0 +1,26 @@ +import { readFile } from "node:fs/promises"; +import path from "node:path"; +import { describe, expect, it } from "vitest"; + +describe("content security policy", () => { + it("allows both Local Forge image protocols", async () => { + const html = await readFile(path.join(process.cwd(), "index.html"), "utf8"); + const page = new DOMParser().parseFromString(html, "text/html"); + const policy = page + .querySelector('meta[http-equiv="Content-Security-Policy"]') + ?.getAttribute("content"); + const imageSources = policy + ?.split(";") + .find((directive) => directive.trim().startsWith("img-src ")) + ?.trim() + .split(/\s+/) + .slice(1); + + expect(imageSources).toEqual( + expect.arrayContaining([ + "local-forge-attachment:", + "local-forge-output:", + ]), + ); + }); +}); diff --git a/electron/enhancement-models.test.ts b/electron/enhancement-models.test.ts new file mode 100644 index 0000000..86e61b4 --- /dev/null +++ b/electron/enhancement-models.test.ts @@ -0,0 +1,121 @@ +import { promises as fs } from "node:fs"; +import { createHash } from "node:crypto"; +import os from "node:os"; +import path from "node:path"; +import { afterEach, describe, expect, it } from "vitest"; +import { + discoverEnhancementModels, + downloadTrustedAsset, + type TrustedAsset, +} from "./enhancement-models"; + +const temporaryDirectories: string[] = []; + +async function temporaryDirectory(): Promise { + const directory = await fs.mkdtemp( + path.join(os.tmpdir(), "local-forge-enhancements-"), + ); + temporaryDirectories.push(directory); + return directory; +} + +afterEach(async () => { + await Promise.all( + temporaryDirectories + .splice(0) + .map((directory) => fs.rm(directory, { recursive: true, force: true })), + ); +}); + +describe("discoverEnhancementModels", () => { + it("discovers the Lavely enhancement model layout", async () => { + const root = await temporaryDirectory(); + const upscaler = path.join(root, "upscalers", "4x-UltraSharp.pth"); + const faceDetector = path.join(root, "face_detector", "face_yolov8n.pt"); + const segmenterDirectory = path.join(root, "nsfw_segmentation"); + await fs.mkdir(path.dirname(upscaler), { recursive: true }); + await fs.mkdir(path.dirname(faceDetector), { recursive: true }); + await fs.mkdir(segmenterDirectory, { recursive: true }); + await fs.writeFile(upscaler, "upscaler"); + await fs.writeFile(faceDetector, "detector"); + await Promise.all( + ["breast", "penis", "vagina"].map((region) => + fs.writeFile( + path.join(segmenterDirectory, `nsfw-seg-${region}-x.pt`), + region, + ), + ), + ); + + await expect(discoverEnhancementModels([root])).resolves.toEqual({ + upscalerModelPath: upscaler, + faceDetectorModelPath: faceDetector, + nsfwSegmenterModelPath: segmenterDirectory, + }); + }); + + it("ignores missing and incomplete model files", async () => { + const root = await temporaryDirectory(); + const upscaler = path.join(root, "upscalers", "4x-UltraSharp.pth"); + const segmenterDirectory = path.join(root, "nsfw_segmentation"); + await fs.mkdir(path.dirname(upscaler), { recursive: true }); + await fs.mkdir(segmenterDirectory, { recursive: true }); + await fs.writeFile(upscaler, ""); + await fs.writeFile( + path.join(segmenterDirectory, "nsfw-seg-breast-x.pt"), + "breast", + ); + await fs.writeFile( + path.join(segmenterDirectory, "nsfw-seg-penis-x.pt"), + "penis", + ); + + await expect(discoverEnhancementModels([root])).resolves.toEqual({ + upscalerModelPath: "", + faceDetectorModelPath: "", + nsfwSegmenterModelPath: "", + }); + }); +}); + +describe("downloadTrustedAsset", () => { + const payload = Buffer.from("trusted model payload"); + const asset: TrustedAsset = { + directory: "upscalers", + filename: "model.pth", + url: "https://models.example/model.pth", + bytes: payload.byteLength, + sha256: createHash("sha256").update(payload).digest("hex"), + }; + + it("promotes a verified download to its final path", async () => { + const root = await temporaryDirectory(); + const target = await downloadTrustedAsset( + root, + asset, + async () => new Response(payload), + ); + + await expect(fs.readFile(target)).resolves.toEqual(payload); + await expect(fs.stat(`${target}.part`)).rejects.toMatchObject({ + code: "ENOENT", + }); + }); + + it("rejects a corrupt download without leaving model files", async () => { + const root = await temporaryDirectory(); + const target = path.join(root, asset.directory, asset.filename); + + await expect( + downloadTrustedAsset( + root, + asset, + async () => new Response("not the expected model"), + ), + ).rejects.toThrow("integrity check"); + await expect(fs.stat(target)).rejects.toMatchObject({ code: "ENOENT" }); + await expect(fs.stat(`${target}.part`)).rejects.toMatchObject({ + code: "ENOENT", + }); + }); +}); diff --git a/electron/enhancement-models.ts b/electron/enhancement-models.ts new file mode 100644 index 0000000..123885c --- /dev/null +++ b/electron/enhancement-models.ts @@ -0,0 +1,199 @@ +import { createHash } from "node:crypto"; +import { promises as fs } from "node:fs"; +import path from "node:path"; +import type { + EnhancementInstallResult, + EnhancementModelKind, + EnhancementModelPaths, +} from "../src/types"; + +export interface TrustedAsset { + directory: string; + filename: string; + url: string; + bytes: number; + sha256: string; +} + +const TRUSTED_ASSETS: Record = { + upscaler: [ + { + directory: "upscalers", + filename: "4x-UltraSharp.pth", + url: "https://huggingface.co/lokCX/4x-Ultrasharp/resolve/1856559b50de25116a7c07261177dd128f1f5664/4x-UltraSharp.pth", + bytes: 66_961_958, + sha256: + "a5812231fc936b42af08a5edba784195495d303d5b3248c24489ef0c4021fe01", + }, + ], + faceDetector: [ + { + directory: "face_detector", + filename: "face_yolov8n.pt", + url: "https://huggingface.co/Bingsu/adetailer/resolve/53cc19de382014514d9d4038601d261a7faa9b7b/face_yolov8n.pt", + bytes: 6_230_011, + sha256: + "70b640f8f60b1cf0dcc72f30caf3da9495eb2fb6509da48c53374ad6806e6a9c", + }, + ], + nsfwSegmenter: [ + { + directory: "nsfw_segmentation", + filename: "nsfw-seg-breast-x.pt", + url: "https://huggingface.co/NSFW-API/NSFW_Segmentation/resolve/818199566a6e3aa0acec2dd1f98bc98c9b99750d/nsfw-seg-breast-x.pt", + bytes: 124_800_929, + sha256: + "5ad882ddaf149873be131943b373da9f15a0603c91508e2291ece83d729c8ecc", + }, + { + directory: "nsfw_segmentation", + filename: "nsfw-seg-penis-x.pt", + url: "https://huggingface.co/NSFW-API/NSFW_Segmentation/resolve/818199566a6e3aa0acec2dd1f98bc98c9b99750d/nsfw-seg-penis-x.pt", + bytes: 124_844_641, + sha256: + "49d9fc8ee67d3bdee44e46bf75aeb76058f5cd074bac027b55ff071b63a32e21", + }, + { + directory: "nsfw_segmentation", + filename: "nsfw-seg-vagina-x.pt", + url: "https://huggingface.co/NSFW-API/NSFW_Segmentation/resolve/818199566a6e3aa0acec2dd1f98bc98c9b99750d/nsfw-seg-vagina-x.pt", + bytes: 124_817_249, + sha256: + "f8349ae348f5cf041809c38bf38d2c80a0f3dccdb1375b97971358cde742197f", + }, + ], +}; + +const UPSCALER_FILES = [ + ["upscalers", "4x-UltraSharp.pth"], + ["4x-UltraSharp.pth"], +]; +const FACE_DETECTOR_FILES = [ + ["face_detector", "face_yolov8n.pt"], + ["face_yolov8n.pt"], +]; +const NSFW_SEGMENTER_FILES = [ + "nsfw-seg-breast-x.pt", + "nsfw-seg-penis-x.pt", + "nsfw-seg-vagina-x.pt", +]; + +async function firstUsableFile( + roots: string[], + candidates: string[][], +): Promise { + for (const root of roots) { + for (const candidate of candidates) { + const filePath = path.resolve(root, ...candidate); + try { + const stat = await fs.stat(filePath); + if (stat.isFile() && stat.size > 0) return filePath; + } catch { + // Continue through known locations. + } + } + } + return ""; +} + +async function firstUsableSegmenterDirectory(roots: string[]): Promise { + for (const root of roots) { + const directory = path.resolve(root, "nsfw_segmentation"); + const filesAreUsable = await Promise.all( + NSFW_SEGMENTER_FILES.map(async (filename) => { + try { + const stat = await fs.stat(path.join(directory, filename)); + return stat.isFile() && stat.size > 0; + } catch { + return false; + } + }), + ); + if (filesAreUsable.every(Boolean)) return directory; + } + return ""; +} + +export async function discoverEnhancementModels( + modelRoots: string[], +): Promise { + const [upscalerModelPath, faceDetectorModelPath, nsfwSegmenterModelPath] = + await Promise.all([ + firstUsableFile(modelRoots, UPSCALER_FILES), + firstUsableFile(modelRoots, FACE_DETECTOR_FILES), + firstUsableSegmenterDirectory(modelRoots), + ]); + return { + upscalerModelPath, + faceDetectorModelPath, + nsfwSegmenterModelPath, + }; +} + +function validPayload(payload: Buffer, asset: TrustedAsset): boolean { + return ( + payload.byteLength === asset.bytes && + createHash("sha256").update(payload).digest("hex") === asset.sha256 + ); +} + +export async function downloadTrustedAsset( + modelsRoot: string, + asset: TrustedAsset, + fetchAsset: (url: string) => Promise, +): Promise { + const targetDirectory = path.resolve(modelsRoot, asset.directory); + const target = path.join(targetDirectory, asset.filename); + const temporary = `${target}.part`; + + try { + const existing = await fs.readFile(target); + if (validPayload(existing, asset)) return target; + } catch (error) { + if ((error as NodeJS.ErrnoException).code !== "ENOENT") throw error; + } + + const response = await fetchAsset(asset.url); + if (!response.ok) { + throw new Error( + `Model download failed with HTTP ${response.status} ${response.statusText}.`, + ); + } + const payload = Buffer.from(await response.arrayBuffer()); + if (!validPayload(payload, asset)) { + throw new Error("Downloaded model failed its integrity check."); + } + + await fs.mkdir(targetDirectory, { recursive: true }); + try { + await fs.writeFile(temporary, payload); + await fs.rm(target, { force: true }); + await fs.rename(temporary, target); + } catch (error) { + await fs.rm(temporary, { force: true }); + throw error; + } + return target; +} + +export async function installEnhancementModel( + modelsRoot: string, + kind: EnhancementModelKind, + fetchAsset: (url: string) => Promise, +): Promise { + const assets = TRUSTED_ASSETS[kind]; + if (!assets) throw new Error("Unknown enhancement model."); + const installedPaths: string[] = []; + for (const asset of assets) { + installedPaths.push( + await downloadTrustedAsset(modelsRoot, asset, fetchAsset), + ); + } + return { + kind, + path: + kind === "nsfwSegmenter" + ? path.dirname(installedPaths[0]) + : installedPaths[0], + }; +} diff --git a/electron/jobs.test.ts b/electron/jobs.test.ts index cdf135d..79c8d36 100644 --- a/electron/jobs.test.ts +++ b/electron/jobs.test.ts @@ -51,10 +51,20 @@ describe("LocalJobManager", () => { const modelPath = path.join(root, "image-model"); const upscalerPath = path.join(root, "4x-UltraSharp.pth"); const faceDetectorPath = path.join(root, "face-detector.pt"); + const segmenterPath = path.join(root, "nsfw-segmentation"); await fs.mkdir(modelPath); + await fs.mkdir(segmenterPath); await fs.writeFile(path.join(modelPath, "model_index.json"), "{}"); await fs.writeFile(upscalerPath, "fixture"); await fs.writeFile(faceDetectorPath, "fixture"); + await Promise.all( + ["breast", "penis", "vagina"].map((region) => + fs.writeFile( + path.join(segmenterPath, `nsfw-seg-${region}-x.pt`), + "fixture", + ), + ), + ); const events: JobEvent[] = []; const manager = new LocalJobManager({ runtimeDirectory: () => root, @@ -87,6 +97,8 @@ describe("LocalJobManager", () => { upscale: true, upscaleFactor: 4, upscalerModelPath: upscalerPath, + nsfwSegmentation: true, + nsfwSegmenterModelPath: segmenterPath, }, (event) => events.push(event), ); @@ -107,7 +119,7 @@ describe("LocalJobManager", () => { root, "outputs", "images", - "fixture-preview-1.png", + "preview-fixture-1.png", ), step: 1, total: 2, @@ -134,6 +146,12 @@ describe("LocalJobManager", () => { height: 512, }), ); + await expect( + fs.stat(path.join(root, "outputs", "images", "preview-fixture-1.png")), + ).rejects.toMatchObject({ code: "ENOENT" }); + await expect( + fs.stat(path.join(root, "outputs", "images", "fixture.png")), + ).resolves.toMatchObject({ size: 7 }); await expect( fs.readFile(path.join(root, "outputs", "images", "request.json"), "utf8"), ).resolves.toEqual( @@ -151,10 +169,130 @@ describe("LocalJobManager", () => { upscale: true, upscale_factor: 4, upscaler_model: upscalerPath, + nsfw_segmentation: true, + nsfw_segmenter_model_dir: segmenterPath, }), ); }); + it("purges stale previews without deleting final images", async () => { + const root = await temporaryDirectory(); + const outputDirectory = path.join(root, "outputs", "images"); + const preview = path.join(outputDirectory, "preview-stale-job-001.png"); + const final = path.join(outputDirectory, "final-image.png"); + await fs.mkdir(outputDirectory, { recursive: true }); + await fs.writeFile(preview, "preview"); + await fs.writeFile(final, "final"); + const manager = new LocalJobManager({ + runtimeDirectory: () => root, + outputDirectory: () => path.join(root, "outputs"), + }); + + await manager.cleanupPreviews(); + + await expect(fs.stat(preview)).rejects.toMatchObject({ code: "ENOENT" }); + await expect(fs.readFile(final, "utf8")).resolves.toBe("final"); + }); + + it("validates and forwards a selected image edit", async () => { + const root = await temporaryDirectory(); + const modelPath = path.join(root, "image-model"); + const sourceImage = path.join(root, "source.png"); + await fs.mkdir(modelPath); + await fs.writeFile(path.join(modelPath, "model_index.json"), "{}"); + await fs.writeFile(sourceImage, "fixture"); + const events: JobEvent[] = []; + const manager = new LocalJobManager({ + runtimeDirectory: () => root, + outputDirectory: () => path.join(root, "outputs"), + imageWorker: () => fixture, + }); + + await manager.startImage( + { + jobId: "image-edit-fixture", + pythonPath: process.execPath, + model: { + id: modelPath, + name: "fixture", + path: modelPath, + format: "diffusers", + architecture: "StableDiffusionXLPipeline", + modifiedAt: new Date().toISOString(), + }, + prompt: "A portrait in a quiet studio", + negativePrompt: "blurred", + width: 512, + height: 512, + steps: 2, + guidance: 5, + seed: 42, + faceFix: false, + faceFixStrength: 0.45, + faceDetectorModelPath: "", + upscale: false, + upscaleFactor: 2, + upscalerModelPath: "", + nsfwSegmentation: false, + nsfwSegmenterModelPath: "", + sourceImage, + editRegion: { x: 0.1, y: 0.2, width: 0.3, height: 0.4 }, + editPrompt: "Move the subject to the left", + editStrength: 0.7, + }, + (event) => events.push(event), + ); + + await waitForDone(events); + const request = JSON.parse( + await fs.readFile( + path.join(root, "outputs", "images", "request.json"), + "utf8", + ), + ) as Record; + expect(request).toMatchObject({ + source_image: sourceImage, + edit_region: { x: 0.1, y: 0.2, width: 0.3, height: 0.4 }, + edit_prompt: "Move the subject to the left", + edit_strength: 0.7, + }); + + await expect( + manager.startImage( + { + jobId: "invalid-image-edit", + pythonPath: process.execPath, + model: { + id: modelPath, + name: "fixture", + path: modelPath, + format: "diffusers", + architecture: "StableDiffusionXLPipeline", + modifiedAt: new Date().toISOString(), + }, + prompt: "A portrait", + negativePrompt: "", + width: 512, + height: 512, + steps: 2, + guidance: 5, + seed: 42, + faceFix: false, + faceFixStrength: 0.45, + faceDetectorModelPath: "", + upscale: false, + upscaleFactor: 2, + upscalerModelPath: "", + nsfwSegmentation: false, + nsfwSegmenterModelPath: "", + sourceImage, + editRegion: { x: 0.8, y: 0.2, width: 0.3, height: 0.4 }, + }, + () => undefined, + ), + ).rejects.toThrow("must stay inside"); + }); + it("runs training only with a local Transformers model and dataset", async () => { const root = await temporaryDirectory(); const modelPath = path.join(root, "language-model"); @@ -228,6 +366,8 @@ describe("LocalJobManager", () => { upscale: false, upscaleFactor: 2 as const, upscalerModelPath: "", + nsfwSegmentation: false, + nsfwSegmenterModelPath: "", }; await expect( @@ -250,6 +390,16 @@ describe("LocalJobManager", () => { () => undefined, ), ).rejects.toThrow("Upscaler model was not found"); + await expect( + manager.startImage( + { + ...request, + nsfwSegmentation: true, + nsfwSegmenterModelPath: path.join(root, "missing-segmenters"), + }, + () => undefined, + ), + ).rejects.toThrow("NSFW segmentation model directory was not found"); }); it("rejects single-file image checkpoints before spawning Python", async () => { @@ -285,6 +435,8 @@ describe("LocalJobManager", () => { upscale: false, upscaleFactor: 2, upscalerModelPath: "", + nsfwSegmentation: false, + nsfwSegmenterModelPath: "", }, () => undefined, ), diff --git a/electron/jobs.ts b/electron/jobs.ts index 25d9f8c..ff8078a 100644 --- a/electron/jobs.ts +++ b/electron/jobs.ts @@ -19,9 +19,16 @@ interface ActiveJob { kind: JobKind; outputDirectory: string; emit: (event: JobEvent) => void; + previewPaths: Set; } const JOB_ID_PATTERN = /^[A-Za-z0-9_-]{1,120}$/; +const PREVIEW_FILE_PATTERN = /^preview-[A-Za-z0-9_-]+\.png$/; +const NSFW_SEGMENTER_FILES = [ + "nsfw-seg-breast-x.pt", + "nsfw-seg-penis-x.pt", + "nsfw-seg-vagina-x.pt", +]; function assertNumber( value: number, @@ -122,6 +129,7 @@ export class LocalJobManager { assertBoolean(request.faceFix, "Face Fix"); assertNumber(request.faceFixStrength, "Face Fix strength", 0.1, 0.8); assertBoolean(request.upscale, "Upscale"); + assertBoolean(request.nsfwSegmentation, "NSFW segmentation"); if (request.upscaleFactor !== 2 && request.upscaleFactor !== 4) { throw new Error("Upscale factor must be 2 or 4."); } @@ -139,6 +147,47 @@ export class LocalJobManager { ".safetensors", ]); } + if (request.nsfwSegmentation) { + await assertDirectory( + request.nsfwSegmenterModelPath, + "NSFW segmentation model directory", + ); + await Promise.all( + NSFW_SEGMENTER_FILES.map((filename) => + assertFile( + path.join(request.nsfwSegmenterModelPath, filename), + `NSFW segmentation model ${filename}`, + ), + ), + ); + } + if (Boolean(request.sourceImage) !== Boolean(request.editRegion)) { + throw new Error( + "Image edits require both a source image and a selection.", + ); + } + if (request.sourceImage && request.editRegion) { + await assertModelFile(request.sourceImage, "Source image", [ + ".png", + ".jpg", + ".jpeg", + ".webp", + ]); + const { x, y, width, height } = request.editRegion; + assertNumber(x, "Selection x", 0, 1); + assertNumber(y, "Selection y", 0, 1); + assertNumber(width, "Selection width", 0.01, 1); + assertNumber(height, "Selection height", 0.01, 1); + if (x + width > 1 || y + height > 1) { + throw new Error( + "The image edit selection must stay inside the source image.", + ); + } + if ((request.editPrompt?.length ?? 0) > 2_000) { + throw new Error("The image edit instruction is too long."); + } + assertNumber(request.editStrength ?? 0.65, "Edit strength", 0.1, 1); + } const outputDirectory = path.join(this.paths.outputDirectory(), "images"); await fs.mkdir(outputDirectory, { recursive: true }); @@ -173,6 +222,16 @@ export class LocalJobManager { upscale: request.upscale, upscale_factor: request.upscaleFactor, upscaler_model: request.upscalerModelPath, + nsfw_segmentation: request.nsfwSegmentation, + nsfw_segmenter_model_dir: request.nsfwSegmenterModelPath, + ...(request.sourceImage && request.editRegion + ? { + source_image: request.sourceImage, + edit_region: request.editRegion, + edit_prompt: request.editPrompt?.trim() ?? "", + edit_strength: request.editStrength ?? 0.65, + } + : {}), }, outputDirectory, emit, @@ -244,6 +303,26 @@ export class LocalJobManager { this.processes.cancelAll(); } + async cleanupPreviews(): Promise { + const outputDirectory = path.join(this.paths.outputDirectory(), "images"); + let entries; + try { + entries = await fs.readdir(outputDirectory, { withFileTypes: true }); + } catch (error) { + if ((error as NodeJS.ErrnoException).code === "ENOENT") return; + throw error; + } + await Promise.all( + entries + .filter( + (entry) => entry.isFile() && PREVIEW_FILE_PATTERN.test(entry.name), + ) + .map((entry) => + fs.rm(path.join(outputDirectory, entry.name), { force: true }), + ), + ); + } + private assertPythonPath(value: string): void { if ( !value.trim() || @@ -266,7 +345,12 @@ export class LocalJobManager { ): void { if (this.active.has(jobId)) throw new Error(`Job ${jobId} is already active.`); - this.active.set(jobId, { kind, outputDirectory, emit }); + this.active.set(jobId, { + kind, + outputDirectory, + emit, + previewPaths: new Set(), + }); emit({ jobId, kind, type: "queued", progress: 0 }); try { this.processes.start( @@ -315,14 +399,23 @@ export class LocalJobManager { return; } if (event.type === "preview") { + let outputPath = ""; try { if (job.kind !== "image" || typeof event.path !== "string") { throw new Error("Worker preview is not a valid image output."); } - const outputPath = path.resolve(event.path); + outputPath = path.resolve(event.path); assertInside(job.outputDirectory, outputPath); + if (!PREVIEW_FILE_PATTERN.test(path.basename(outputPath))) { + throw new Error("Worker preview has an invalid filename."); + } + job.previewPaths.add(outputPath); const stat = await fs.stat(outputPath); if (!stat.isFile()) throw new Error("Worker preview is not a file."); + if (this.active.get(jobId) !== job) { + await fs.rm(outputPath, { force: true }); + return; + } const step = Number(event.step); const total = Number(event.total); const progress = @@ -345,6 +438,8 @@ export class LocalJobManager { : undefined, }); } catch (error) { + if (outputPath) job.previewPaths.delete(outputPath); + if (this.active.get(jobId) !== job) return; job.emit({ jobId, kind: job.kind, @@ -367,6 +462,7 @@ export class LocalJobManager { } if (event.type === "error" || event.type === "cancelled") { this.active.delete(jobId); + await this.cleanupJobPreviews(job); job.emit({ jobId, kind: job.kind, @@ -391,6 +487,7 @@ export class LocalJobManager { throw new Error("Training worker output is not a directory."); } this.active.delete(jobId); + await this.cleanupJobPreviews(job); job.emit({ jobId, kind: job.kind, @@ -406,6 +503,7 @@ export class LocalJobManager { }); } catch (error) { this.active.delete(jobId); + await this.cleanupJobPreviews(job); job.emit({ jobId, kind: job.kind, @@ -415,4 +513,14 @@ export class LocalJobManager { }); } } + + private async cleanupJobPreviews(job: ActiveJob): Promise { + const previewPaths = [...job.previewPaths]; + job.previewPaths.clear(); + await Promise.all( + previewPaths.map((previewPath) => + fs.rm(previewPath, { force: true }).catch(() => undefined), + ), + ); + } } diff --git a/electron/main.ts b/electron/main.ts index 8daf44d..c695943 100644 --- a/electron/main.ts +++ b/electron/main.ts @@ -20,6 +20,9 @@ import { import type { ChatRequest, ChatStreamEvent, + EnhancementInstallResult, + EnhancementModelKind, + EnhancementModelPaths, ImageAttachment, ImageGenerationRequest, ImageModel, @@ -34,8 +37,15 @@ import type { RuntimeHealth, SystemSnapshot, TrainingRequest, + VisionDescribeRequest, + VisionDescribeResult, } from "../src/types"; -import { runMcpChat } from "./chat"; +import { runMcpChat, warmChatModel } from "./chat"; +import { + discoverEnhancementModels, + installEnhancementModel, +} from "./enhancement-models"; +import { describeImageWithOllama } from "./vision"; import { LocalJobManager } from "./jobs"; import { closeMcpConnections, disconnectMcpServer, testMcpServer } from "./mcp"; import { discoverModelCatalog } from "./model-catalog"; @@ -72,8 +82,10 @@ const MODEL_SCAN_IGNORES = new Set([ let mainWindow: BrowserWindow | null = null; let workspaceWrite = Promise.resolve(); -let modelCatalogCache: ModelCatalogResponse | null = null; -let modelCatalogExpiresAt = 0; +const modelCatalogCache = new Map< + boolean, + { value: ModelCatalogResponse; expiresAt: number } +>(); const jobs = new LocalJobManager({ runtimeDirectory: () => app.isPackaged @@ -175,15 +187,50 @@ function emitJob(sender: WebContents, event: JobEvent): void { }); } -async function imageForOllama(value: string): Promise { +function enhancementModelRoots(): string[] { + const appRoot = app.getAppPath(); + return [ + path.join(app.getPath("userData"), "models"), + path.join(appRoot, "models"), + path.join(process.resourcesPath, "models"), + path.resolve(appRoot, "..", "Lavely-LLM", "models"), + path.resolve(process.cwd(), "..", "Lavely-LLM", "models"), + ].filter((root, index, roots) => roots.indexOf(root) === index); +} + +async function localImageFilePath(value: string): Promise { if (value.startsWith(`${ATTACHMENT_SCHEME}:`)) { - return (await fs.readFile(attachmentFilePath(value))).toString("base64"); + return attachmentFilePath(value); + } + if (value.startsWith(`${OUTPUT_SCHEME}:`)) { + return outputImageFilePath(value); + } + if (/^\.?\/demo\/[A-Za-z0-9_-]+\.(?:png|jpe?g|webp)$/i.test(value)) { + const relative = value.replace(/^\.?\//, ""); + const candidates = [ + path.join(app.getAppPath(), "dist", relative), + path.join(app.getAppPath(), "public", relative), + ]; + for (const candidate of candidates) { + try { + if ((await fs.stat(candidate)).isFile()) return candidate; + } catch (error) { + if ((error as NodeJS.ErrnoException).code !== "ENOENT") throw error; + } + } + throw new Error(`Bundled reference image was not found: ${value}`); } + throw new Error("Unsupported Local Forge image URL."); +} + +async function imageForOllama(value: string): Promise { if (value.startsWith("data:")) { const separator = value.indexOf(","); return separator >= 0 ? value.slice(separator + 1) : value; } - return value; + return (await fs.readFile(await localImageFilePath(value))).toString( + "base64", + ); } async function describeDiffusersModel( @@ -552,13 +599,21 @@ async function getSystemSnapshot(): Promise { } } -async function getModelCatalog(refresh = false): Promise { - if (!refresh && modelCatalogCache && modelCatalogExpiresAt > Date.now()) { - return modelCatalogCache; +async function getModelCatalog( + refresh = false, + includeNsfw = false, +): Promise { + const cached = modelCatalogCache.get(includeNsfw); + if (!refresh && cached && cached.expiresAt > Date.now()) { + return cached.value; } - const catalog = await discoverModelCatalog(await getSystemSnapshot()); - modelCatalogCache = catalog; - modelCatalogExpiresAt = Date.now() + MODEL_CATALOG_TTL_MS; + const catalog = await discoverModelCatalog(await getSystemSnapshot(), fetch, { + includeNsfw, + }); + modelCatalogCache.set(includeNsfw, { + value: catalog, + expiresAt: Date.now() + MODEL_CATALOG_TTL_MS, + }); return catalog; } @@ -642,12 +697,29 @@ function registerIpc(): void { ipcMain.handle("ollama:models", (_event, baseUrl: string) => listModels(baseUrl), ); + ipcMain.handle( + "ollama:warm-model", + (_event, baseUrl: string, model: string) => warmChatModel(baseUrl, model), + ); ipcMain.handle("ollama:chat", (event, request: ChatRequest) => streamChat(event.sender, request), ); ipcMain.handle("ollama:cancel-chat", (_event, requestId: string) => { activeChats.get(requestId)?.abort(); }); + ipcMain.handle( + "vision:describe", + async ( + _event, + request: VisionDescribeRequest, + ): Promise => + describeImageWithOllama({ + baseUrl: request.baseUrl, + model: request.model, + imageBase64: await imageForOllama(request.imageUrl), + mode: request.mode, + }), + ); ipcMain.handle("ollama:pull", (event, request: PullRequest) => pullModel(event.sender, request), ); @@ -655,13 +727,30 @@ function registerIpc(): void { activePulls.get(requestId)?.abort(); }); - ipcMain.handle("catalog:list", (_event, refresh: boolean) => - getModelCatalog(Boolean(refresh)), + ipcMain.handle( + "catalog:list", + (_event, refresh: boolean, includeNsfw: boolean) => + getModelCatalog(Boolean(refresh), Boolean(includeNsfw)), ); ipcMain.handle("catalog:open", (_event, url: string) => shell.openExternal(modelCatalogUrl(url)), ); + ipcMain.handle( + "enhancements:discover", + (): Promise => + discoverEnhancementModels(enhancementModelRoots()), + ); + ipcMain.handle( + "enhancements:install", + (_event, kind: EnhancementModelKind): Promise => + installEnhancementModel( + path.join(app.getPath("userData"), "models"), + kind, + (url) => net.fetch(url), + ), + ); + ipcMain.handle("mcp:test-server", (_event, server: McpServerConfig) => testMcpServer(server), ); @@ -672,7 +761,13 @@ function registerIpc(): void { ipcMain.handle( "jobs:start-image", async (event, request: ImageGenerationRequest) => { - await jobs.startImage(request, (jobEvent) => + const resolvedRequest = request.sourceImage + ? { + ...request, + sourceImage: await localImageFilePath(request.sourceImage), + } + : request; + await jobs.startImage(resolvedRequest, (jobEvent) => emitJob(event.sender, jobEvent), ); return { ok: true }; @@ -846,6 +941,18 @@ function registerIpc(): void { return result.canceled ? null : (result.filePaths[0] ?? null); }, ); + + ipcMain.handle( + "dialog:choose-nsfw-segmenter-models", + async (): Promise => { + const result = await dialog.showOpenDialog({ + title: "Select an NSFW segmentation model directory", + buttonLabel: "Select models", + properties: ["openDirectory"], + }); + return result.canceled ? null : (result.filePaths[0] ?? null); + }, + ); } function createWindow(): void { @@ -891,7 +998,7 @@ function createWindow(): void { registerIpc(); -void app.whenReady().then(() => { +void app.whenReady().then(async () => { app.setAppUserModelId("io.localforge.desktop"); void protocol.handle(ATTACHMENT_SCHEME, async (request) => { try { @@ -911,6 +1018,9 @@ void app.whenReady().then(() => { return new Response("Output not found.", { status: 404 }); } }); + await jobs.cleanupPreviews().catch((error) => { + console.warn("Could not remove temporary image previews:", error); + }); createWindow(); app.on("activate", () => { if (BrowserWindow.getAllWindows().length === 0) createWindow(); diff --git a/electron/model-catalog.test.ts b/electron/model-catalog.test.ts index 450e825..be45011 100644 --- a/electron/model-catalog.test.ts +++ b/electron/model-catalog.test.ts @@ -190,5 +190,75 @@ describe("model catalog compatibility", () => { }), ); expect(catalog.items.every((item) => item.verified)).toBe(true); + expect(fetcher).not.toHaveBeenCalledWith( + expect.stringContaining("search=nsfw"), + expect.anything(), + ); + }); + + it("discovers NSFW image pipelines only when explicitly enabled", async () => { + const requestedUrls: string[] = []; + const fetcher = vi.fn(async (input: string | URL | Request) => { + const url = String(input); + requestedUrls.push(url); + if (url.includes("registry.ollama.ai")) { + return jsonResponse({ error: "not found" }, 404); + } + if (url.includes("search=nsfw")) { + return jsonResponse([ + { + id: "fixture/nsfw-flux", + library_name: "diffusers", + pipeline_tag: "text-to-image", + tags: ["diffusers", "nsfw", "diffusers:FluxPipeline"], + }, + { + id: "fixture/nsfw-sdxl", + library_name: "diffusers", + pipeline_tag: "text-to-image", + tags: ["diffusers", "nsfw", "diffusers:StableDiffusionXLPipeline"], + }, + ]); + } + if (url.includes("/api/models/fixture/nsfw-sdxl")) { + return jsonResponse({ + id: "fixture/nsfw-sdxl", + library_name: "diffusers", + pipeline_tag: "text-to-image", + gated: false, + downloads: 500, + tags: ["diffusers", "nsfw", "diffusers:StableDiffusionXLPipeline"], + safetensors: { total: 3.5e9 }, + }); + } + return jsonResponse([]); + }) as unknown as typeof fetch; + + const safeCatalog = await discoverModelCatalog(system, fetcher); + expect(safeCatalog.items.some((item) => item.nsfw)).toBe(false); + expect(requestedUrls.some((url) => url.includes("search=nsfw"))).toBe( + false, + ); + + requestedUrls.length = 0; + const adultCatalog = await discoverModelCatalog(system, fetcher, { + includeNsfw: true, + }); + expect(requestedUrls.some((url) => url.includes("search=nsfw"))).toBe(true); + expect( + requestedUrls.some((url) => url.includes("/fixture/nsfw-flux")), + ).toBe(false); + expect( + adultCatalog.items.find( + (item) => item.id === "huggingface:fixture/nsfw-sdxl", + ), + ).toEqual( + expect.objectContaining({ + category: "image", + runtime: "diffusers", + architecture: "StableDiffusionXLPipeline", + nsfw: true, + }), + ); }); }); diff --git a/electron/model-catalog.ts b/electron/model-catalog.ts index d8dd09a..680ea2c 100644 --- a/electron/model-catalog.ts +++ b/electron/model-catalog.ts @@ -6,6 +6,7 @@ import type { ModelCatalogRuntime, SystemSnapshot, } from "../src/types"; +import { hasNsfwModelTag } from "../src/lib/nsfw"; const GIB = 1024 ** 3; const REQUEST_TIMEOUT_MS = 12_000; @@ -66,6 +67,10 @@ interface CompatibilityInput { unsupportedReason?: string; } +interface ModelCatalogDiscoveryOptions { + includeNsfw?: boolean; +} + const OLLAMA_CANDIDATES: OllamaCandidate[] = [ { tag: "qwen3:8b", @@ -357,6 +362,7 @@ async function discoverOllama( sizeBytes, parameters: candidate.parameters, gated: false, + nsfw: false, verified: true, ...assessModelCompatibility({ runtime: "ollama", @@ -387,7 +393,7 @@ async function discoverOllama( }; } -function huggingFaceSearchUrl(pipeline: string): string { +function huggingFaceSearchUrl(pipeline: string, search?: string): string { const params = new URLSearchParams({ pipeline_tag: pipeline, sort: "downloads", @@ -395,6 +401,7 @@ function huggingFaceSearchUrl(pipeline: string): string { limit: "24", full: "true", }); + if (search) params.set("search", search); return `https://huggingface.co/api/models?${params}`; } @@ -442,6 +449,7 @@ function huggingFaceItem( const parameters = modelParameters(model); const sizeBytes = modelWeightBytes(model); const gated = model.gated !== false && Boolean(model.gated); + const nsfw = hasNsfwModelTag(id, architecture, ...tags); let runtime: ModelCatalogRuntime = "none"; let executorSupported = false; let unsupportedReason: string | undefined; @@ -491,6 +499,7 @@ function huggingFaceItem( likes: numberValue(model.likes), updatedAt: stringValue(model.lastModified) || undefined, gated, + nsfw, verified: true, ...assessModelCompatibility({ runtime, @@ -516,15 +525,23 @@ async function searchHuggingFaceGroup( category: ModelCatalogCategory, system: SystemSnapshot, fetcher: typeof fetch, + options: { search?: string; nsfw?: boolean } = {}, ): Promise { - const search = await fetchJson(huggingFaceSearchUrl(pipeline), fetcher); + const search = await fetchJson( + huggingFaceSearchUrl(pipeline, options.search), + fetcher, + ); if (!Array.isArray(search)) throw new Error("Hugging Face returned an invalid model list."); const candidates = (search as HuggingFaceModel[]) .filter( (model) => stringValue(model.library_name) === library && - !stringArray(model.tags).includes("nsfw") && + (options.nsfw + ? hasNsfwModelTag(modelId(model), ...stringArray(model.tags)) + : !hasNsfwModelTag(modelId(model), ...stringArray(model.tags))) && + (!options.nsfw || + SUPPORTED_IMAGE_PIPELINES.has(modelArchitecture(model))) && Boolean(modelId(model)), ) .slice(0, HUGGING_FACE_RESULTS_PER_GROUP); @@ -566,8 +583,9 @@ async function searchHuggingFaceGroup( async function discoverHuggingFace( system: SystemSnapshot, fetcher: typeof fetch, + options: ModelCatalogDiscoveryOptions, ): Promise { - const groups = await Promise.allSettled([ + const searches = [ searchHuggingFaceGroup( "text-to-image", "diffusers", @@ -596,7 +614,20 @@ async function discoverHuggingFace( system, fetcher, ), - ]); + ]; + if (options.includeNsfw) { + searches.push( + searchHuggingFaceGroup( + "text-to-image", + "diffusers", + "image", + system, + fetcher, + { search: "nsfw", nsfw: true }, + ), + ); + } + const groups = await Promise.allSettled(searches); const items: ModelCatalogItem[] = []; const warnings: string[] = []; for (const group of groups) { @@ -625,10 +656,11 @@ const COMPATIBILITY_ORDER: Record = { export async function discoverModelCatalog( system: SystemSnapshot, fetcher: typeof fetch = fetch, + options: ModelCatalogDiscoveryOptions = {}, ): Promise { const sources = await Promise.allSettled([ discoverOllama(system, fetcher), - discoverHuggingFace(system, fetcher), + discoverHuggingFace(system, fetcher, options), ]); const items: ModelCatalogItem[] = []; const warnings: string[] = []; @@ -644,7 +676,10 @@ export async function discoverModelCatalog( ); } } - items.sort( + const uniqueItems = Array.from( + new Map(items.map((item) => [item.id, item])).values(), + ); + uniqueItems.sort( (left, right) => COMPATIBILITY_ORDER[left.compatibility] - COMPATIBILITY_ORDER[right.compatibility] || @@ -652,7 +687,7 @@ export async function discoverModelCatalog( left.name.localeCompare(right.name), ); return { - items, + items: uniqueItems, fetchedAt: new Date().toISOString(), system, warnings, diff --git a/electron/preload.ts b/electron/preload.ts index ad501f0..840f3e5 100644 --- a/electron/preload.ts +++ b/electron/preload.ts @@ -3,6 +3,9 @@ import type { AppInfo, ChatRequest, ChatStreamEvent, + EnhancementInstallResult, + EnhancementModelKind, + EnhancementModelPaths, ForgeApi, ImageAttachment, ImageGenerationRequest, @@ -18,6 +21,8 @@ import type { RuntimeHealth, SystemSnapshot, TrainingRequest, + VisionDescribeRequest, + VisionDescribeResult, } from "../src/types"; function subscribe( @@ -52,6 +57,8 @@ const api: ForgeApi = { ipcRenderer.invoke("ollama:health", baseUrl) as Promise, models: (baseUrl: string) => ipcRenderer.invoke("ollama:models", baseUrl) as Promise, + warmModel: (baseUrl: string, model: string) => + ipcRenderer.invoke("ollama:warm-model", baseUrl, model) as Promise, chat: (request: ChatRequest) => ipcRenderer.invoke("ollama:chat", request) as Promise<{ ok: boolean }>, cancelChat: (requestId: string) => @@ -65,14 +72,33 @@ const api: ForgeApi = { onPullEvent: (callback: (event: PullStreamEvent) => void) => subscribe("ollama:pull-event", callback), }, + vision: { + describe: (request: VisionDescribeRequest) => + ipcRenderer.invoke( + "vision:describe", + request, + ) as Promise, + }, catalog: { - list: (refresh = false) => + list: (refresh = false, includeNsfw = false) => ipcRenderer.invoke( "catalog:list", refresh, + includeNsfw, ) as Promise, open: (url: string) => ipcRenderer.invoke("catalog:open", url), }, + enhancements: { + discover: () => + ipcRenderer.invoke( + "enhancements:discover", + ) as Promise, + install: (kind: EnhancementModelKind) => + ipcRenderer.invoke( + "enhancements:install", + kind, + ) as Promise, + }, mcp: { testServer: (server: McpServerConfig) => ipcRenderer.invoke("mcp:test-server", server) as Promise, @@ -123,6 +149,10 @@ const api: ForgeApi = { ipcRenderer.invoke("dialog:choose-face-detector-model") as Promise< string | null >, + chooseNsfwSegmenterModels: () => + ipcRenderer.invoke("dialog:choose-nsfw-segmenter-models") as Promise< + string | null + >, }, }; diff --git a/electron/vision.test.ts b/electron/vision.test.ts new file mode 100644 index 0000000..e12ecdf --- /dev/null +++ b/electron/vision.test.ts @@ -0,0 +1,77 @@ +import { createServer } from "node:http"; +import type { AddressInfo } from "node:net"; +import { describe, expect, it } from "vitest"; +import { describeImageWithOllama } from "./vision"; + +describe("describeImageWithOllama", () => { + it("sends a prompt-building request and image to local Ollama", async () => { + let requestBody: Record = {}; + const server = createServer(async (request, response) => { + const chunks: Buffer[] = []; + for await (const chunk of request) chunks.push(Buffer.from(chunk)); + requestBody = JSON.parse(Buffer.concat(chunks).toString("utf8")); + response.writeHead(200, { "content-type": "application/json" }); + response.end( + JSON.stringify({ + message: { + content: "portrait, window light, shallow depth of field", + }, + }), + ); + }); + await new Promise((resolve) => + server.listen(0, "127.0.0.1", resolve), + ); + const address = server.address() as AddressInfo; + + try { + await expect( + describeImageWithOllama({ + baseUrl: `http://127.0.0.1:${address.port}`, + model: "vision-fixture", + imageBase64: "aW1hZ2U=", + mode: "prompt", + }), + ).resolves.toEqual({ + text: "portrait, window light, shallow depth of field", + model: "vision-fixture", + mode: "prompt", + }); + } finally { + await new Promise((resolve, reject) => + server.close((error) => (error ? reject(error) : resolve())), + ); + } + + expect(requestBody).toMatchObject({ + model: "vision-fixture", + stream: false, + messages: [ + { role: "system", content: expect.stringContaining("clearly adult") }, + { role: "user", images: ["aW1hZ2U="] }, + ], + }); + }); + + it("rejects non-local runtime URLs", async () => { + await expect( + describeImageWithOllama({ + baseUrl: "https://example.com", + model: "vision-fixture", + imageBase64: "aW1hZ2U=", + mode: "description", + }), + ).rejects.toThrow("only connects to runtimes on this machine"); + }); + + it("rejects unexpected IPC mode values", async () => { + await expect( + describeImageWithOllama({ + baseUrl: "http://127.0.0.1:11434", + model: "vision-fixture", + imageBase64: "aW1hZ2U=", + mode: "__proto__" as never, + }), + ).rejects.toThrow("Unknown image description mode"); + }); +}); diff --git a/electron/vision.ts b/electron/vision.ts new file mode 100644 index 0000000..f04873c --- /dev/null +++ b/electron/vision.ts @@ -0,0 +1,85 @@ +import type { VisionDescribeMode, VisionDescribeResult } from "../src/types"; + +interface VisionRequest { + baseUrl: string; + model: string; + imageBase64: string; + mode: VisionDescribeMode; +} + +const SYSTEM_PROMPTS: Record = { + description: + "Describe the attached image precisely in two to four sentences. Cover visible subjects, pose or action, clothing, setting, composition, lighting, color, and visual style. If it contains clearly adult nudity or explicit adult sexual content, describe it accurately with neutral anatomical language without moralizing. Do not identify real people or infer protected traits. If any subject may be under 18, omit sexual detail. Return only the description.", + prompt: + "Create one detailed image-generation prompt from the attached image. Preserve the visible subject, pose or action, clothing or nudity for clearly adult subjects, setting, composition, camera perspective, lighting, color, medium, and style. Describe explicit adult details factually when present, but never sexualize a subject who may be under 18. Do not identify real people or invent unsupported details. Return only the prompt as concise comma-separated natural language with no heading or quotation marks.", +}; + +function visionEndpoint(baseUrl: string): string { + const url = new URL(baseUrl); + const localHosts = new Set(["localhost", "127.0.0.1", "::1", "[::1]"]); + if (url.protocol !== "http:" || !localHosts.has(url.hostname)) { + throw new Error("Local Forge only connects to runtimes on this machine."); + } + url.pathname = "/api/chat"; + url.search = ""; + url.hash = ""; + return url.toString(); +} + +export async function describeImageWithOllama( + request: VisionRequest, + fetchRuntime: typeof fetch = fetch, +): Promise { + if (!request.model.trim()) throw new Error("Select a vision model first."); + if (!request.imageBase64) throw new Error("Select an image first."); + if (request.mode !== "description" && request.mode !== "prompt") { + throw new Error("Unknown image description mode."); + } + + const response = await fetchRuntime(visionEndpoint(request.baseUrl), { + method: "POST", + headers: { "content-type": "application/json" }, + body: JSON.stringify({ + model: request.model, + stream: false, + messages: [ + { role: "system", content: SYSTEM_PROMPTS[request.mode] }, + { + role: "user", + content: + request.mode === "prompt" + ? "Build an image-generation prompt from this image." + : "Describe this image.", + images: [request.imageBase64], + }, + ], + options: { + temperature: request.mode === "prompt" ? 0.4 : 0.2, + }, + }), + }); + + if (!response.ok) { + const detail = (await response.text()).trim().slice(0, 500); + throw new Error( + `Vision model request failed with HTTP ${response.status}${ + detail ? `: ${detail}` : "." + }`, + ); + } + + const payload = (await response.json()) as { + error?: unknown; + message?: { content?: unknown }; + }; + if (typeof payload.error === "string" && payload.error) { + throw new Error(payload.error); + } + const text = + typeof payload.message?.content === "string" + ? payload.message.content.trim() + : ""; + if (!text) throw new Error("The vision model returned an empty response."); + + return { text, model: request.model, mode: request.mode }; +} diff --git a/index.html b/index.html index 25b445c..5dd0833 100644 --- a/index.html +++ b/index.html @@ -4,7 +4,7 @@ diff --git a/package-lock.json b/package-lock.json index 7f700d7..fffbd55 100644 --- a/package-lock.json +++ b/package-lock.json @@ -12,6 +12,7 @@ "@fontsource-variable/manrope": "^5.2.5", "@fontsource-variable/space-grotesk": "^5.2.9", "@modelcontextprotocol/sdk": "^1.30.1", + "@tanstack/react-virtual": "^3.14.13", "lucide-react": "^1.47.0", "react": "^19.2.8", "react-dom": "^19.2.8", @@ -2097,6 +2098,33 @@ "node": ">=10" } }, + "node_modules/@tanstack/react-virtual": { + "version": "3.14.13", + "resolved": "https://registry.npmjs.org/@tanstack/react-virtual/-/react-virtual-3.14.13.tgz", + "integrity": "sha512-JbDTAwtzZ99aOeCrAfW5EsE5KSq5RWh6Af2dtFwyLIIk48Ja7vm6n4axu/43T3vfjFnEGJalxJ1wyYjSQD6bSg==", + "license": "MIT", + "dependencies": { + "@tanstack/virtual-core": "3.17.11" + }, + "funding": { + "type": "github", + "url": "https://github.com/sponsors/tannerlinsley" + }, + "peerDependencies": { + "react": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0", + "react-dom": "^16.8.0 || ^17.0.0 || ^18.0.0 || ^19.0.0" + } + }, + "node_modules/@tanstack/virtual-core": { + "version": "3.17.11", + "resolved": "https://registry.npmjs.org/@tanstack/virtual-core/-/virtual-core-3.17.11.tgz", + "integrity": "sha512-+ILjvtHup6Y2hzQ6YzwMgX1Q+oQpxEGOXCEsCNaPoIP0VxMbizIBTmYTDtkerkIQS8/CbP1BRuyt8V/8BCsy1g==", + "license": "MIT", + "funding": { + "type": "github", + "url": "https://github.com/sponsors/tannerlinsley" + } + }, "node_modules/@testing-library/dom": { "version": "10.4.2", "resolved": "https://registry.npmjs.org/@testing-library/dom/-/dom-10.4.2.tgz", diff --git a/package.json b/package.json index 26a18bc..06996b3 100644 --- a/package.json +++ b/package.json @@ -21,6 +21,7 @@ "@fontsource-variable/manrope": "^5.2.5", "@fontsource-variable/space-grotesk": "^5.2.9", "@modelcontextprotocol/sdk": "^1.30.1", + "@tanstack/react-virtual": "^3.14.13", "lucide-react": "^1.47.0", "react": "^19.2.8", "react-dom": "^19.2.8", diff --git a/runtime/image_worker.py b/runtime/image_worker.py index 5d33246..9812028 100644 --- a/runtime/image_worker.py +++ b/runtime/image_worker.py @@ -12,6 +12,12 @@ sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding="utf-8", errors="replace") sys.stderr = io.TextIOWrapper(sys.stderr.buffer, encoding="utf-8", errors="replace") +NSFW_SEGMENTER_MODELS = { + "breast": "nsfw-seg-breast-x.pt", + "penis": "nsfw-seg-penis-x.pt", + "vagina": "nsfw-seg-vagina-x.pt", +} + def emit(payload: dict) -> None: print(json.dumps(payload), flush=True) @@ -165,6 +171,54 @@ def require_model_file( return model_path +def require_segmenter_directory(value: object) -> Path: + model_dir = Path(str(value or "")).expanduser().resolve() + if not model_dir.is_dir(): + raise ValueError(f"NSFW segmentation model directory was not found: {model_dir}") + for filename in NSFW_SEGMENTER_MODELS.values(): + if not (model_dir / filename).is_file(): + raise ValueError(f"NSFW segmentation model was not found: {filename}") + return model_dir + + +def segment_nsfw( + model_dir: Path, + image: Any, + conf: float = 0.3, + model_factory: Any | None = None, +) -> dict[str, Any]: + import numpy as np + from PIL import Image + + if model_factory is None: + try: + from ultralytics import YOLO + except ImportError as error: + raise RuntimeError( + "NSFW segmentation requires ultralytics. " + "Reinstall runtime/requirements-image.txt." + ) from error + model_factory = YOLO + + width, height = image.size + masks: dict[str, Any] = {} + for region, filename in NSFW_SEGMENTER_MODELS.items(): + model = model_factory(str(model_dir / filename)) + results = model.predict(image, conf=conf, verbose=False) + combined = np.zeros((height, width), dtype=np.uint8) + for result in results: + if result.masks is None: + continue + for mask_data in result.masks.data: + mask_array = mask_data.cpu().numpy() + mask = Image.fromarray((mask_array * 255).astype(np.uint8)) + mask = mask.resize((width, height), Image.Resampling.NEAREST) + combined = np.maximum(combined, np.asarray(mask)) + if combined.max() > 0: + masks[region] = Image.fromarray(combined) + return masks + + def expanded_face_box( box: tuple[int, int, int, int], image_width: int, @@ -197,6 +251,50 @@ def feathered_mask(size: tuple[int, int], feather: int = 24): return mask.filter(ImageFilter.GaussianBlur(radius=feather / 2)) +def parse_edit_region(value: object) -> tuple[float, float, float, float]: + if not isinstance(value, dict): + raise ValueError("edit_region must be an object.") + x = bounded_float(value.get("x"), "edit_region.x", 0, 1) + y = bounded_float(value.get("y"), "edit_region.y", 0, 1) + width = bounded_float(value.get("width"), "edit_region.width", 0.01, 1) + height = bounded_float(value.get("height"), "edit_region.height", 0.01, 1) + if x + width > 1 or y + height > 1: + raise ValueError("edit_region must stay inside the source image.") + return x, y, width, height + + +def edit_selection_mask( + size: tuple[int, int], + region: tuple[float, float, float, float], + feather: int = 24, +): + from PIL import Image, ImageChops, ImageDraw, ImageFilter + + image_width, image_height = size + x, y, width, height = region + left = max(0, min(image_width - 1, round(x * image_width))) + top = max(0, min(image_height - 1, round(y * image_height))) + right = max(left + 1, min(image_width, round((x + width) * image_width))) + bottom = max(top + 1, min(image_height, round((y + height) * image_height))) + hard_mask = Image.new("L", size, 0) + ImageDraw.Draw(hard_mask).rectangle( + (left, top, right - 1, bottom - 1), + fill=255, + ) + if feather <= 0: + return hard_mask + softened = hard_mask.filter(ImageFilter.GaussianBlur(radius=feather / 2)) + return ImageChops.multiply(softened, hard_mask) + + +def edit_generation_prompt(prompt: str, instruction: str) -> str: + instruction = instruction.strip() + if not instruction: + return prompt + prompt = prompt.strip().rstrip(" ,.") + return f"{prompt}, {instruction}" if prompt else instruction + + def detect_face_boxes(detector_path: Path, image, maximum: int = 4): try: from ultralytics import YOLO @@ -235,8 +333,6 @@ def refine_faces( prompt: str, negative_prompt: str, strength: float, - steps: int, - guidance: float, generator: Any, ): import torch @@ -283,8 +379,8 @@ def refine_faces( negative_prompt=negative_prompt or None, image=resized, strength=strength, - num_inference_steps=max(20, steps), - guidance_scale=max(5.0, guidance - 1.5), + num_inference_steps=25, + guidance_scale=6.0, generator=generator, ).images[0] refined = refined.resize( @@ -397,6 +493,17 @@ def main() -> None: steps = bounded_int(request.get("steps", 24), "steps", 1, 100) guidance = bounded_float(request.get("guidance", 5.5), "guidance", 0, 30) seed = bounded_int(request.get("seed", 0), "seed", 0, 2**32 - 1) + source_image_value = str(request.get("source_image", "")).strip() + edit_region_value = request.get("edit_region") + if bool(source_image_value) != isinstance(edit_region_value, dict): + raise ValueError("Image edits require both source_image and edit_region.") + edit_region = ( + parse_edit_region(edit_region_value) if source_image_value else None + ) + edit_instruction = str(request.get("edit_prompt", "")).strip() + edit_strength = bounded_float( + request.get("edit_strength", 0.65), "edit_strength", 0.1, 1 + ) face_fix = bool(request.get("face_fix", False)) face_fix_strength = bounded_float( request.get("face_fix_strength", 0.45), @@ -405,6 +512,7 @@ def main() -> None: 0.8, ) upscale = bool(request.get("upscale", False)) + nsfw_segmentation = bool(request.get("nsfw_segmentation", False)) upscale_factor = bounded_int( request.get("upscale_factor", 2), "upscale_factor", 2, 4 ) @@ -417,6 +525,33 @@ def main() -> None: pipeline, generator = load_pipeline(model_path) generator.manual_seed(seed) + generation_pipeline = pipeline + source_image = None + edit_mask = None + if source_image_value and edit_region: + try: + from diffusers import AutoPipelineForInpainting + from PIL import Image + + source_path = Path(source_image_value).expanduser().resolve() + if not source_path.is_file(): + raise ValueError(f"Source image was not found: {source_path}") + with Image.open(source_path) as opened_image: + source_image = opened_image.convert("RGB").resize( + (width, height), Image.Resampling.LANCZOS + ) + edit_mask = edit_selection_mask(source_image.size, edit_region) + generation_pipeline = AutoPipelineForInpainting.from_pipe(pipeline) + prompt = edit_generation_prompt(prompt, edit_instruction) + except ImportError as error: + raise RuntimeError( + "Image editing requires the Diffusers inpainting runtime." + ) from error + except (AttributeError, TypeError, ValueError) as error: + raise RuntimeError( + "The selected Diffusers model cannot run masked image edits." + ) from error + safe_job_id = re.sub(r"[^A-Za-z0-9_-]", "-", arguments.job_id)[:80] run_timestamp = datetime.now().strftime("%Y%m%d_%H%M%S_%f") preview_steps = preview_step_numbers(steps) @@ -474,17 +609,50 @@ def legacy_progress_callback(step_index, _timestep, latents): "height": height, "generator": generator, } - call_parameters = inspect.signature(pipeline.__call__).parameters + if source_image is not None and edit_mask is not None: + call_options.update( + { + "image": source_image, + "mask_image": edit_mask, + "strength": edit_strength, + } + ) + call_parameters = inspect.signature(generation_pipeline.__call__).parameters + if source_image is not None and edit_mask is not None: + required_edit_parameters = {"image", "mask_image", "strength"} + if not required_edit_parameters.issubset(call_parameters): + raise RuntimeError( + "The selected Diffusers model does not expose inpainting inputs." + ) + if "padding_mask_crop" in call_parameters: + call_options["padding_mask_crop"] = 32 if "callback_on_step_end" in call_parameters: call_options["callback_on_step_end"] = progress_callback elif "callback" in call_parameters: call_options["callback"] = legacy_progress_callback call_options["callback_steps"] = 1 - emit({"type": "status", "message": "Generating image..."}) - result = pipeline(**call_options) + emit( + { + "type": "status", + "message": ( + "Reworking selected area..." + if source_image is not None + else "Generating image..." + ), + } + ) + result = generation_pipeline(**call_options) image = result.images[0] processing_suffixes: list[str] = [] + if source_image is not None and edit_mask is not None: + from PIL import Image + + generated = image.convert("RGB") + if generated.size != source_image.size: + generated = generated.resize(source_image.size, Image.Resampling.LANCZOS) + image = Image.composite(generated, source_image, edit_mask) + processing_suffixes.append("edit") if face_fix: try: @@ -501,8 +669,6 @@ def legacy_progress_callback(step_index, _timestep, latents): prompt, negative_prompt, face_fix_strength, - steps, - guidance, generator, ) processing_suffixes.append("face") @@ -544,11 +710,39 @@ def legacy_progress_callback(step_index, _timestep, latents): metadata.add_text("face_fix_strength", str(face_fix_strength)) metadata.add_text("upscale", str(upscale).lower()) metadata.add_text("upscale_factor", str(upscale_factor)) + metadata.add_text("nsfw_segmentation", str(nsfw_segmentation).lower()) + metadata.add_text("edit_prompt", edit_instruction) + metadata.add_text("edit_strength", str(edit_strength)) + if edit_region: + metadata.add_text("edit_region", json.dumps(edit_region)) metadata.add_text("generator", "Local Forge") image.save(output_path, pnginfo=metadata) except ImportError: image.save(output_path) + if nsfw_segmentation: + try: + segmenter_dir = require_segmenter_directory( + request.get("nsfw_segmenter_model_dir") + ) + emit({"type": "status", "message": "Running NSFW segmentation..."}) + masks = segment_nsfw(segmenter_dir, image) + for region, mask in masks.items(): + mask_path = output_dir / f"{output_path.stem}_mask_{region}.png" + mask.save(mask_path) + if masks: + regions = ", ".join(sorted(masks)) + emit({"type": "log", "message": f"Saved NSFW masks: {regions}."}) + else: + emit( + { + "type": "log", + "message": "NSFW segmentation found no matching regions.", + } + ) + except Exception as error: + emit({"type": "log", "message": f"NSFW segmentation skipped: {error}"}) + emit( { "type": "done", diff --git a/runtime/requirements-image.txt b/runtime/requirements-image.txt index 40be050..ae93592 100644 --- a/runtime/requirements-image.txt +++ b/runtime/requirements-image.txt @@ -1,5 +1,6 @@ # Install a CUDA-compatible PyTorch build for your driver first. -torch>=2.5,<3 +torch>=2.5,<2.6 +torchvision>=0.20,<0.21 accelerate>=1.2,<2 diffusers>=0.32,<1 transformers>=4.47,<5 diff --git a/src/App.css b/src/App.css index 0ee01b4..5e66a63 100644 --- a/src/App.css +++ b/src/App.css @@ -210,8 +210,7 @@ flex: 0 0 var(--rail-width); flex-direction: column; align-items: center; - border-right: 1px solid var(--line); - background: var(--bg-elevated); + background: transparent; } .rail-glyph { @@ -1688,17 +1687,23 @@ .studio-view { display: grid; overflow: hidden; + padding: 22px 0 0 24px; grid-template-rows: auto minmax(0, 1fr); } +.studio-view > .tool-header { + padding-right: 24px; +} + .studio-layout { display: grid; min-height: 0; + margin-left: -24px; grid-template-columns: 215px minmax(360px, 1fr) 285px; border: 1px solid var(--line); - border-radius: 8px; + border-radius: 6px 0 0; background: var(--panel); - box-shadow: 0 16px 42px rgb(0 0 0 / 22%); + box-shadow: -4px -3px 14px rgb(0 0 0 / 28%); } .studio-presets, @@ -1775,21 +1780,17 @@ font-size: 8px; } -.studio-history-heading { - display: flex; - align-items: center; - justify-content: space-between; - margin: 20px 1px 8px; - color: var(--faint); - font-size: 9px; - font-weight: 700; - text-transform: uppercase; -} - .studio-model-field { margin-bottom: 16px; } +.studio-nsfw-defaults { + margin-bottom: 12px; + padding-top: 0; + border-top: 0; + border-bottom: 1px solid var(--line); +} + .studio-model-control { display: grid; grid-template-columns: minmax(0, 1fr) 34px; @@ -1849,37 +1850,11 @@ line-height: 1.45; } -.studio-history-grid { - display: grid; - grid-template-columns: repeat(2, minmax(0, 1fr)); - gap: 6px; -} - -.studio-history-grid button { - aspect-ratio: 1; - overflow: hidden; - padding: 2px; - border: 1px solid transparent; - border-radius: 6px; - background: transparent; -} - -.studio-history-grid button.active { - border-color: var(--ember); -} - -.studio-history-grid img { - width: 100%; - height: 100%; - border-radius: 3px; - object-fit: cover; -} - .studio-canvas-wrap { display: grid; min-width: 0; min-height: 0; - grid-template-rows: 43px minmax(0, 1fr) 72px; + grid-template-rows: 43px minmax(0, 1fr); background: var(--bg); } @@ -1932,6 +1907,25 @@ text-transform: uppercase; } +.canvas-toolbar .canvas-toolbar-actions, +.canvas-tool-toggle, +.canvas-zoom-controls { + display: flex; + align-items: center; + gap: 3px; +} + +.canvas-tool-toggle { + margin-right: 4px; + padding-right: 7px; + border-right: 1px solid var(--line); +} + +.canvas-toolbar button:disabled { + cursor: not-allowed; + opacity: 0.35; +} + .studio-canvas { position: relative; display: grid; @@ -1939,6 +1933,7 @@ place-items: center; overflow: hidden; padding: 24px; + container-type: size; } .canvas-grid { @@ -1959,25 +1954,51 @@ } .active-artwork { + --artwork-ratio: 1; + position: relative; z-index: 1; display: flex; - width: min(76%, 620px); - height: min(83%, 610px); - min-height: 260px; + width: min( + 76cqw, + 620px, + calc(83cqh * var(--artwork-ratio)), + calc(610px * var(--artwork-ratio)) + ); + height: auto; + aspect-ratio: var(--artwork-ratio); align-items: center; justify-content: center; + overflow: hidden; margin: 0; border: 1px solid var(--line-strong); background: var(--bg); box-shadow: 0 20px 60px rgb(0 0 0 / 55%); + touch-action: none; transition: transform 180ms ease; + user-select: none; + will-change: transform; +} + +.active-artwork.canvas-pan { + cursor: grab; +} + +.active-artwork.canvas-pan.is-dragging { + cursor: grabbing; + transition: none; +} + +.active-artwork.canvas-select { + cursor: crosshair; } .active-artwork img { + display: block; width: 100%; height: 100%; - object-fit: cover; + object-fit: contain; + pointer-events: none; } .active-artwork.is-denoising { @@ -2006,6 +2027,7 @@ .active-artwork figcaption { position: absolute; + z-index: 3; right: 8px; bottom: 8px; left: 8px; @@ -2019,57 +2041,70 @@ backdrop-filter: blur(8px); color: #e5e9e0; font-size: 8px; + pointer-events: none; +} + +.studio-edit-selection { + position: absolute; + z-index: 2; + border: 2px solid var(--ember); + background: rgb(220 125 71 / 12%); + box-shadow: 0 0 0 9999px rgb(5 7 6 / 48%); + pointer-events: none; +} + +.studio-edit-selection span { + position: absolute; + top: 5px; + left: 5px; + padding: 3px 5px; + border-radius: 3px; + background: rgb(9 11 9 / 82%); + color: var(--text); + font-size: 8px; + font-weight: 700; + text-transform: uppercase; } -.variant-strip { +.studio-edit-panel { + margin-bottom: 16px; + padding: 12px 0 14px; + border-top: 1px solid var(--line); + border-bottom: 1px solid var(--line); +} + +.studio-edit-panel header { display: flex; align-items: center; - justify-content: flex-start; - gap: 7px; - overflow-x: auto; - padding: 8px; - border-top: 1px solid var(--line); - background: var(--bg-elevated); + justify-content: space-between; + gap: 8px; } -.variant-strip button { - position: relative; - width: 50px; - height: 50px; - overflow: hidden; - padding: 2px; - border: 1px solid var(--line); - border-radius: 6px; - background: var(--panel); +.studio-edit-panel header strong { + display: block; + margin-top: 2px; + color: var(--text-soft); + font-size: 11px; } -.variant-strip button.active { - border-color: var(--ember); +.studio-edit-panel .compact-textarea { + min-height: 64px; + resize: vertical; } -.variant-strip img { - width: 100%; - height: 100%; - border-radius: 3px; - object-fit: cover; +.studio-edit-strength { + margin-top: 10px; } -.variant-strip button span { - position: absolute; - right: 3px; - bottom: 3px; - padding: 1px 3px; - border-radius: 3px; - background: rgb(0 0 0 / 75%); - color: #fff; - font-size: 7px; +.studio-edit-action { + width: 100%; + margin-top: 13px; + font-size: 9px; } -.variant-strip .variant-add { - display: grid; - place-items: center; - border-style: dashed; - color: var(--faint); +.studio-edit-action:disabled { + cursor: not-allowed; + opacity: 0.45; } .prompt-field, @@ -2099,6 +2134,72 @@ font-size: 8px; } +.studio-prompt-heading, +.prompt-image-actions, +.studio-vision-output > div { + display: flex; + align-items: center; +} + +.studio-prompt-heading { + justify-content: space-between; + gap: 8px; +} + +.studio-prompt-heading > span { + max-width: 150px; + overflow: hidden; + color: var(--faint); + font-family: var(--font-mono); + font-size: 8px; + text-overflow: ellipsis; + white-space: nowrap; +} + +.prompt-image-actions { + gap: 10px; +} + +.prompt-image-actions button:disabled { + cursor: not-allowed; + opacity: 0.45; +} + +.studio-vision-error { + margin: 7px 0 0; + color: var(--danger); + font-size: 9px; + line-height: 1.4; +} + +.studio-vision-output { + margin-top: 8px; + padding: 8px 0 8px 9px; + border-left: 2px solid var(--ember); +} + +.studio-vision-output > div { + justify-content: space-between; + gap: 8px; +} + +.studio-vision-output strong, +.studio-vision-output button { + font-size: 8px; +} + +.studio-vision-output button { + background: transparent; + color: var(--ember); +} + +.studio-vision-output p { + margin: 5px 0 0; + color: var(--muted); + font-size: 9px; + line-height: 1.5; +} + .seed-field { display: grid; grid-template-columns: 42px minmax(0, 1fr) 32px; @@ -2159,6 +2260,34 @@ font-size: 9px; } +.studio-enhancement-error { + margin: 7px 0 2px; + color: var(--danger); + font-size: 9px; + line-height: 1.4; +} + +.studio-model-install { + display: inline-flex; + min-width: 72px; + min-height: 27px; + align-items: center; + justify-content: center; + gap: 5px; + padding: 5px 8px; + border: 1px solid color-mix(in srgb, var(--ember) 42%, var(--line)); + border-radius: 6px; + background: color-mix(in srgb, var(--ember) 8%, var(--panel-strong)); + color: var(--ember); + font-size: 9px; + font-weight: 700; +} + +.studio-model-install:disabled { + cursor: wait; + opacity: 0.55; +} + .studio-enhancements .toggle-row:first-of-type { margin-top: 4px; } @@ -2342,6 +2471,36 @@ aspect-ratio: auto; } +.asset-grid:not(.virtualized) .asset-tile { + content-visibility: auto; + contain-intrinsic-size: 210px; +} + +.asset-grid.virtualized, +.asset-grid.list.virtualized { + position: relative; + display: block; +} + +.asset-virtualizer { + position: relative; + width: 100%; +} + +.asset-virtual-row { + position: absolute; + top: 0; + left: 0; + display: grid; + width: 100%; + gap: 12px; + padding-bottom: 12px; +} + +.asset-virtual-row.list { + grid-template-columns: 1fr; +} + .asset-details { min-height: 0; overflow-y: auto; @@ -2658,6 +2817,7 @@ .catalog-source, .catalog-verified, +.catalog-adult, .fit-badge { display: inline-flex; align-items: center; @@ -2681,6 +2841,19 @@ color: var(--mint); } +.catalog-source-group { + display: flex; + align-items: center; + gap: 7px; +} + +.catalog-adult { + padding: 2px 4px; + border: 1px solid color-mix(in srgb, var(--danger) 45%, var(--line)); + border-radius: 3px; + color: var(--danger); +} + .catalog-model-heading h3 { min-width: 0; overflow: hidden; @@ -3689,6 +3862,12 @@ object-fit: contain; } +.restricted-output span { + color: var(--faint); + font-size: 10px; + font-weight: 650; +} + .run-error { display: flex; align-items: flex-start; @@ -4104,9 +4283,14 @@ display: block; overflow-y: auto; } + .studio-view > .tool-header { + padding-right: 0; + } .studio-layout { display: flex; + margin-left: 0; flex-direction: column; + border-radius: 6px; } .studio-canvas-wrap { min-height: 500px; diff --git a/src/App.tsx b/src/App.tsx index 28d2917..1787b42 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -28,8 +28,10 @@ import { StudioView } from "./features/studio/StudioView"; import { Workbench } from "./features/workbench/Workbench"; import { useWorkspace } from "./hooks/useWorkspace"; import { forgeApi } from "./lib/forge-api"; -import { appendRunLog } from "./state/workspace"; +import { isNsfwImageModel } from "./lib/nsfw"; +import { appendRunLog, createDefaultStudioState } from "./state/workspace"; import type { + ForgeSettings, ForgeRun, ForgeRunLog, ImageRunRecipe, @@ -38,7 +40,9 @@ import type { } from "./state/workspace"; import type { AppView, + EnhancementModelKind, ImageModel, + StudioImageEdit, JobEvent, OllamaModel, RuntimeHealth, @@ -96,6 +100,75 @@ function App() { const [sidebarCollapsed, setSidebarCollapsed] = useState(false); const [commandOpen, setCommandOpen] = useState(false); const [imagePreview, setImagePreview] = useState(null); + const [installingEnhancement, setInstallingEnhancement] = + useState(null); + const [enhancementInstallError, setEnhancementInstallError] = useState(""); + + const discoverEnhancements = useEffectEvent(async () => { + try { + const discovered = await forgeApi.enhancements.discover(); + setWorkspace((current) => { + const upscalerModelPath = current.settings.upscalerModelPath.trim() + ? current.settings.upscalerModelPath + : discovered.upscalerModelPath; + const faceDetectorModelPath = + current.settings.faceDetectorModelPath.trim() + ? current.settings.faceDetectorModelPath + : discovered.faceDetectorModelPath; + const nsfwSegmenterModelPath = + current.settings.nsfwSegmenterModelPath.trim() + ? current.settings.nsfwSegmenterModelPath + : discovered.nsfwSegmenterModelPath; + if ( + upscalerModelPath === current.settings.upscalerModelPath && + faceDetectorModelPath === current.settings.faceDetectorModelPath && + nsfwSegmenterModelPath === current.settings.nsfwSegmenterModelPath + ) { + return current; + } + return { + ...current, + settings: { + ...current.settings, + upscalerModelPath, + faceDetectorModelPath, + nsfwSegmenterModelPath, + }, + }; + }); + } catch (error) { + console.warn("Could not discover enhancement models:", error); + } + }); + + useEffect(() => { + if (loaded) void discoverEnhancements(); + }, [loaded]); + + async function installEnhancement(kind: EnhancementModelKind) { + setInstallingEnhancement(kind); + setEnhancementInstallError(""); + try { + const installed = await forgeApi.enhancements.install(kind); + setWorkspace((current) => ({ + ...current, + settings: { + ...current.settings, + ...(installed.kind === "upscaler" + ? { upscalerModelPath: installed.path } + : installed.kind === "faceDetector" + ? { faceDetectorModelPath: installed.path } + : { nsfwSegmenterModelPath: installed.path }), + }, + })); + } catch (error) { + setEnhancementInstallError( + error instanceof Error ? error.message : "Model installation failed.", + ); + } finally { + setInstallingEnhancement(null); + } + } async function refreshModels() { try { @@ -151,6 +224,38 @@ function App() { }; }, [workspace.settings.ollamaUrl]); + const warmModelName = + models.find((model) => model.name === activeThread.model)?.name || + models.find((model) => model.name === workspace.settings.selectedModel) + ?.name || + models[0]?.name || + ""; + + useEffect(() => { + if ( + !loaded || + !health.online || + workspace.activeView !== "workbench" || + !warmModelName + ) { + return; + } + const timeout = window.setTimeout(() => { + void forgeApi.ollama + .warmModel(workspace.settings.ollamaUrl, warmModelName) + .catch((error: unknown) => + console.warn("Could not preload the selected model:", error), + ); + }, 400); + return () => window.clearTimeout(timeout); + }, [ + health.online, + loaded, + warmModelName, + workspace.activeView, + workspace.settings.ollamaUrl, + ]); + useEffect(() => { let active = true; async function pollSystem() { @@ -240,6 +345,16 @@ function App() { const imageRecipe = existing.recipe?.kind === "image" ? existing.recipe : null; + const generatedNsfw = + imageRecipe?.nsfwDefaults === true || + current.assets.some( + (asset) => + asset.src === imageRecipe?.sourceImage && asset.nsfw === true, + ) || + current.imageModels.some( + (model) => + model.id === imageRecipe?.modelId && isNsfwImageModel(model), + ); const generated: LibraryAsset = { id: `generated-${event.jobId}`, src: event.outputUrl, @@ -250,6 +365,7 @@ function App() { createdAt: new Date().toISOString(), prompt: imageRecipe?.prompt ?? "", outputPath: event.outputPath, + nsfw: generatedNsfw, }; return { ...current, @@ -257,7 +373,10 @@ function App() { assets: current.assets.some((asset) => asset.id === generated.id) ? current.assets : [generated, ...current.assets], - studio: { ...current.studio, activeAsset: generated.src }, + studio: + generated.nsfw && !current.settings.nsfwConsent + ? current.studio + : { ...current.studio, activeAsset: generated.src }, }; }); }), @@ -278,6 +397,37 @@ function App() { })); } + function updateSettings(patch: Partial) { + setWorkspace((current) => { + const settings = { ...current.settings, ...patch }; + if (settings.nsfwConsent) return { ...current, settings }; + + const defaults = createDefaultStudioState(); + const selectedStudioModel = current.imageModels.find( + (model) => model.id === current.studio.model, + ); + return { + ...current, + settings, + studio: { + ...current.studio, + model: + selectedStudioModel && isNsfwImageModel(selectedStudioModel) + ? "" + : current.studio.model, + prompt: current.studio.nsfwDefaults + ? defaults.prompt + : current.studio.prompt, + negativePrompt: current.studio.nsfwDefaults + ? defaults.negativePrompt + : current.studio.negativePrompt, + nsfwDefaults: false, + nsfwSegmentation: false, + }, + }; + }); + } + function addRun(run: ForgeRun) { setWorkspace((current) => ({ ...current, @@ -312,6 +462,15 @@ function App() { model: ImageModel, name: string, ) { + const sourceIsNsfw = workspace.assets.some( + (asset) => asset.src === recipe.sourceImage && asset.nsfw === true, + ); + const nsfwAllowed = + workspace.settings.nsfwConsent && workspace.studio.nsfwDefaults; + if (sourceIsNsfw && !workspace.settings.nsfwConsent) return; + if ((recipe.nsfwDefaults || isNsfwImageModel(model)) && !nsfwAllowed) { + return; + } setImagePreview(null); const jobId = crypto.randomUUID(); const startedAt = new Date().toISOString(); @@ -350,6 +509,13 @@ function App() { upscale: Boolean(recipe.upscale), upscaleFactor: recipe.upscaleFactor ?? 2, upscalerModelPath: workspace.settings.upscalerModelPath, + nsfwSegmentation: + workspace.settings.nsfwConsent && Boolean(recipe.nsfwSegmentation), + nsfwSegmenterModelPath: workspace.settings.nsfwSegmenterModelPath, + sourceImage: recipe.sourceImage, + editRegion: recipe.editRegion, + editPrompt: recipe.editPrompt, + editStrength: recipe.editStrength, }); } catch (error) { failRun(jobId, error); @@ -396,39 +562,62 @@ function App() { } } - async function startImage(action: "render" | "variant") { + async function startImage( + action: "render" | "variant" | "edit", + edit?: StudioImageEdit, + ) { const model = workspace.imageModels.find( (item) => item.id === workspace.studio.model, ); - if (!model) return; + if (!model || (action === "edit" && !edit)) return; + const editInstruction = edit?.instruction.trim() ?? ""; + const basePrompt = workspace.studio.prompt.trim(); + const prompt = basePrompt || editInstruction; + if (!prompt) return; const seed = - action === "variant" + action === "variant" || action === "edit" ? Math.floor(Math.random() * 2 ** 32) : workspace.studio.seed; const recipe: ImageRunRecipe = { kind: "image", modelId: model.id, modelName: model.name, - prompt: workspace.studio.prompt, + prompt, negativePrompt: workspace.studio.negativePrompt, - width: workspace.studio.width, - height: workspace.studio.height, + width: edit?.width ?? workspace.studio.width, + height: edit?.height ?? workspace.studio.height, steps: workspace.studio.steps, guidance: workspace.studio.guidance, seed, faceFix: + action !== "edit" && workspace.studio.faceFix && Boolean(workspace.settings.faceDetectorModelPath.trim()), faceFixStrength: workspace.studio.faceFixStrength, upscale: + action !== "edit" && workspace.studio.upscale && Boolean(workspace.settings.upscalerModelPath.trim()), upscaleFactor: workspace.studio.upscaleFactor, + nsfwSegmentation: + workspace.studio.nsfwSegmentation && + workspace.settings.nsfwConsent && + Boolean(workspace.settings.nsfwSegmenterModelPath.trim()), + nsfwDefaults: + workspace.settings.nsfwConsent && workspace.studio.nsfwDefaults, + sourceImage: edit?.sourceImage, + editRegion: edit?.region, + editPrompt: basePrompt ? editInstruction : "", + editStrength: edit?.strength, }; await queueImageRun( recipe, model, - action === "variant" ? "Studio variant" : "Studio render", + action === "variant" + ? "Studio variant" + : action === "edit" + ? "Studio selected edit" + : "Studio render", ); } @@ -483,26 +672,43 @@ function App() { const recipe = run.recipe; const outputUrl = run.outputUrl; if (recipe?.kind !== "image" || !outputUrl) return; - setWorkspace((current) => ({ - ...current, - activeView: "studio", - studio: { - ...current.studio, - model: recipe.modelId, - prompt: recipe.prompt, - negativePrompt: recipe.negativePrompt, - width: recipe.width, - height: recipe.height, - steps: recipe.steps, - guidance: recipe.guidance, - seed: recipe.seed, - faceFix: recipe.faceFix ?? false, - faceFixStrength: recipe.faceFixStrength ?? 0.45, - upscale: recipe.upscale ?? false, - upscaleFactor: recipe.upscaleFactor ?? 2, - activeAsset: outputUrl, - }, - })); + setWorkspace((current) => { + const model = current.imageModels.find( + (candidate) => candidate.id === recipe.modelId, + ); + const nsfwRun = + recipe.nsfwDefaults === true || + Boolean(model && isNsfwImageModel(model)) || + current.assets.some( + (asset) => + asset.nsfw === true && + (asset.src === outputUrl || asset.src === recipe.sourceImage), + ); + if (nsfwRun && !current.settings.nsfwConsent) return current; + return { + ...current, + activeView: "studio", + studio: { + ...current.studio, + model: recipe.modelId, + prompt: recipe.prompt, + negativePrompt: recipe.negativePrompt, + width: recipe.width, + height: recipe.height, + steps: recipe.steps, + guidance: recipe.guidance, + seed: recipe.seed, + faceFix: recipe.faceFix ?? false, + faceFixStrength: recipe.faceFixStrength ?? 0.45, + upscale: recipe.upscale ?? false, + upscaleFactor: recipe.upscaleFactor ?? 2, + nsfwSegmentation: + current.settings.nsfwConsent && (recipe.nsfwSegmentation ?? false), + nsfwDefaults: nsfwRun, + activeAsset: outputUrl, + }, + }; + }); } async function scanImageModels(): Promise { @@ -513,6 +719,11 @@ function App() { current.imageModels.map((model) => [model.path, model]), ); for (const model of discovered) modelsByPath.set(model.path, model); + const allowNsfwModels = + current.settings.nsfwConsent && current.studio.nsfwDefaults; + const defaultModel = discovered.find( + (model) => allowNsfwModels || !isNsfwImageModel(model), + ); return { ...current, imageModels: [...modelsByPath.values()].sort((left, right) => @@ -520,7 +731,7 @@ function App() { ), studio: { ...current.studio, - model: current.studio.model || discovered[0].id, + model: current.studio.model || defaultModel?.id || "", }, }; }); @@ -632,6 +843,9 @@ function App() { ""; const selectedModel = models.find((model) => model.name === selectedModelName) ?? null; + const visibleAssets = workspace.settings.nsfwConsent + ? workspace.assets + : workspace.assets.filter((asset) => !asset.nsfw); const activeImageRun = workspace.runs.find( (run) => run.kind === "image" && @@ -751,7 +965,6 @@ function App() { void startImage(action)} + onGenerate={(action, edit) => void startImage(action, edit)} onCancel={(jobId) => void forgeApi.jobs.cancel(jobId)} onScanModels={scanImageModels} onImportAssets={importStudioAssets} onRemoveModel={removeImageModel} onOpenSettings={() => selectView("settings")} + onInstallEnhancement={(kind) => void installEnhancement(kind)} + onDescribeImage={async (imageUrl, mode) => { + const result = await forgeApi.vision.describe({ + baseUrl: workspace.settings.ollamaUrl, + model: selectedModelName, + imageUrl, + mode, + }); + return result.text; + }} /> )} {workspace.activeView === "library" && ( setWorkspace((current) => ({ @@ -808,6 +1039,7 @@ function App() { models={models} baseUrl={workspace.settings.ollamaUrl} selectedModel={selectedModelName} + nsfwConsent={workspace.settings.nsfwConsent} onSelect={selectModel} onRefresh={() => void checkRuntime()} /> @@ -829,6 +1061,7 @@ function App() { {workspace.activeView === "activity" && ( void forgeApi.jobs.cancel(id)} onReveal={(outputPath) => void forgeApi.jobs.revealOutput(outputPath) @@ -849,12 +1082,7 @@ function App() { health={health} system={system} models={models} - onChange={(patch) => - setWorkspace((current) => ({ - ...current, - settings: { ...current.settings, ...patch }, - })) - } + onChange={updateSettings} onCheckRuntime={() => void checkRuntime()} /> )} diff --git a/src/features/library/LibraryView.test.tsx b/src/features/library/LibraryView.test.tsx new file mode 100644 index 0000000..35753bd --- /dev/null +++ b/src/features/library/LibraryView.test.tsx @@ -0,0 +1,46 @@ +import { cleanup, render, screen, waitFor } from "@testing-library/react"; +import { afterEach, describe, expect, it, vi } from "vitest"; +import type { LibraryAsset } from "../../state/workspace"; +import { LibraryView } from "./LibraryView"; + +function createAssets(count: number): LibraryAsset[] { + return Array.from({ length: count }, (_, index) => ({ + id: `asset-${index}`, + src: `./demo/asset-${index}.jpg`, + title: `Asset ${index}`, + source: "imported" as const, + width: 1024, + height: 768, + createdAt: "2026-09-25T00:00:00.000Z", + prompt: `Prompt ${index}`, + })); +} + +afterEach(cleanup); + +describe("LibraryView", () => { + it("virtualizes large galleries and lazily decodes mounted thumbnails", async () => { + const assets = createAssets(100); + const { container } = render( + 0)} + onOpenInStudio={vi.fn()} + onRemove={vi.fn()} + />, + ); + + expect(screen.getByText("100 assets")).toBeVisible(); + await waitFor(() => { + const mountedTiles = container.querySelectorAll(".asset-tile"); + expect(mountedTiles.length).toBeGreaterThan(0); + expect(mountedTiles.length).toBeLessThan(assets.length); + }); + + container.querySelectorAll(".asset-tile img").forEach((thumbnail) => { + expect(thumbnail).toHaveAttribute("loading", "lazy"); + expect(thumbnail).toHaveAttribute("decoding", "async"); + expect(thumbnail).toHaveAttribute("fetchpriority", "low"); + }); + }); +}); diff --git a/src/features/library/LibraryView.tsx b/src/features/library/LibraryView.tsx index a1b7e72..966b818 100644 --- a/src/features/library/LibraryView.tsx +++ b/src/features/library/LibraryView.tsx @@ -7,10 +7,15 @@ import { Trash2, X, } from "lucide-react"; -import { useState } from "react"; +import { useVirtualizer } from "@tanstack/react-virtual"; +import { useDeferredValue, useEffect, useRef, useState } from "react"; import { forgeApi } from "../../lib/forge-api"; import type { LibraryAsset } from "../../state/workspace"; +const ASSET_GAP = 12; +const MIN_GRID_TILE_WIDTH = 150; +const VIRTUALIZE_AFTER = 24; + interface LibraryViewProps { assets: LibraryAsset[]; onImport: () => Promise; @@ -23,6 +28,42 @@ function assetFormat(src: string): string { return extension === "JPG" ? "JPEG" : (extension ?? "Image"); } +function useElementWidth(elementRef: React.RefObject) { + const [width, setWidth] = useState(0); + + useEffect(() => { + const element = elementRef.current; + if (!element) return; + const updateWidth = () => setWidth(element.clientWidth); + updateWidth(); + if (typeof ResizeObserver === "undefined") return; + const observer = new ResizeObserver(updateWidth); + observer.observe(element); + return () => observer.disconnect(); + }, [elementRef]); + + return width; +} + +function useScrollableGalleryLayout() { + const [enabled, setEnabled] = useState( + () => + typeof window.matchMedia !== "function" || + window.matchMedia("(min-width: 821px)").matches, + ); + + useEffect(() => { + if (typeof window.matchMedia !== "function") return; + const media = window.matchMedia("(min-width: 821px)"); + const update = () => setEnabled(media.matches); + update(); + media.addEventListener("change", update); + return () => media.removeEventListener("change", update); + }, []); + + return enabled; +} + export function LibraryView({ assets, onImport, @@ -35,10 +76,97 @@ export function LibraryView({ const [copyStatus, setCopyStatus] = useState<"idle" | "copied" | "failed">( "idle", ); + const assetGridRef = useRef(null); + const deferredQuery = useDeferredValue(query); const selected = assets.find((asset) => asset.id === selectedId) ?? null; const filtered = assets.filter((asset) => - asset.title.toLowerCase().includes(query.toLowerCase()), + asset.title.toLowerCase().includes(deferredQuery.toLowerCase()), ); + const gridWidth = useElementWidth(assetGridRef); + const scrollableGalleryLayout = useScrollableGalleryLayout(); + const maximumColumns = selected ? 3 : 4; + const columnCount = + layout === "list" + ? 1 + : Math.max( + 1, + Math.min( + maximumColumns, + gridWidth + ? Math.floor( + (gridWidth + ASSET_GAP) / (MIN_GRID_TILE_WIDTH + ASSET_GAP), + ) + : maximumColumns, + ), + ); + const shouldVirtualize = + scrollableGalleryLayout && filtered.length > VIRTUALIZE_AFTER; + const rowCount = Math.ceil(filtered.length / columnCount); + const estimatedTileWidth = + ((gridWidth || (selected ? 700 : 900)) - ASSET_GAP * (columnCount - 1)) / + columnCount; + const rowVirtualizer = useVirtualizer({ + count: shouldVirtualize ? rowCount : 0, + getScrollElement: () => assetGridRef.current, + estimateSize: () => + layout === "list" ? 84 : estimatedTileWidth * 0.75 + 58, + getItemKey: (rowIndex) => + `${layout}:${filtered[rowIndex * columnCount]?.id ?? rowIndex}`, + overscan: 3, + initialRect: { width: 900, height: 600 }, + }); + const virtualRows = rowVirtualizer.getVirtualItems(); + const visibleRows = + virtualRows.length > 0 + ? virtualRows + : [{ index: 0, key: "initial", start: 0 }]; + + useEffect(() => { + assetGridRef.current?.scrollTo?.({ top: 0 }); + rowVirtualizer.measure(); + }, [columnCount, deferredQuery, layout, rowVirtualizer]); + + function renderAsset(asset: LibraryAsset) { + return ( + + ); + } async function copyPrompt() { if (!selected) return; @@ -107,40 +235,39 @@ export function LibraryView({ {filtered.length} assets
-
- {filtered.map((asset) => ( - - ))} + {visibleRows.map((virtualRow) => { + const rowStart = virtualRow.index * columnCount; + return ( +
+ {filtered + .slice(rowStart, rowStart + columnCount) + .map(renderAsset)} +
+ ); + })} +
+ ) : ( + filtered.map(renderAsset) + )} {selected && (
@@ -176,39 +416,78 @@ export function StudioView({ {previewLabel} -
- - {zoom}% - - + + +
+
+ + {zoom}% + + +
{artworkSource ? ( { + const { naturalWidth, naturalHeight } = event.currentTarget; + if (naturalWidth > 0 && naturalHeight > 0) { + setImageDimensions({ + width: naturalWidth, + height: naturalHeight, + }); + } + }} /> ) : ( Import an image )} + {selection && !preview && ( +
+ Edit area +
+ )}
{preview ? previewLabel : "Selected reference"} @@ -231,33 +534,6 @@ export function StudioView({
-
- {assets.map((asset, index) => ( - - ))} - -