mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-08-16 18:53:21 -03:00
Compare commits
149 Commits
ea80c2224c
...
v1.2.1
| Author | SHA1 | Date | |
|---|---|---|---|
| 94dd08646d | |||
| 658f88ca48 | |||
| f53352efb2 | |||
| 38809a9d1b | |||
| 395682509c | |||
| ef3e7d7bf4 | |||
| c85b6b64a1 | |||
| 34c87d4934 | |||
| 93472e5d67 | |||
| ae185ee714 | |||
| 795036275a | |||
| d43ab6e32f | |||
| 280181f92e | |||
| f8d98934ad | |||
| 303cca0d85 | |||
| c2f16784b3 | |||
| 5bc6d8286c | |||
| 3f8381ffee | |||
| 1ca99294c9 | |||
| 680f0a57f5 | |||
| 94e3f54571 | |||
| 5c2b2aedcc | |||
| ebc31fb963 | |||
| 9659df6ad9 | |||
| 04d131e9dc | |||
| 78fe6282c7 | |||
| 0c00ee22fc | |||
| 5fd4946b1f | |||
| f1d3ac0cdc | |||
| e2c45905f0 | |||
| b2c68e6a65 | |||
| eb0f6dd3b6 | |||
| 0bf87f9092 | |||
| 1da2433bb2 | |||
| 2d6cf545b9 | |||
| 6a259a14fa | |||
| 41e1fd1e1f | |||
| 95fb3c7fc9 | |||
| 8237e5f9ea | |||
| aa75986178 | |||
| b887922055 | |||
| 68fa0f29c7 | |||
| d9d362c9c9 | |||
| d0bc4be0dc | |||
| 420530f532 | |||
| 3001f0f0ef | |||
| b2a1307d23 | |||
| 64da845a58 | |||
| 27027c4497 | |||
| 86c85c08ec | |||
| 196c8ffc3e | |||
| cfc95ee02a | |||
| 479fa36997 | |||
| 3e1216e9bc | |||
| 007883b7d1 | |||
| dc9200a12c | |||
| d2f955266d | |||
| 8e724538bd | |||
| 6fcdeb799d | |||
| 97b9b1f62b | |||
| 4bf9a4b640 | |||
| c5088772e8 | |||
| 56acefbd6c | |||
| 5ab06c4aae | |||
| c11f4b5c68 | |||
| 86376284f4 | |||
| 2b8a2fc7d8 | |||
| f26e1b41c8 | |||
| c1671af99f | |||
| ac7707d0f6 | |||
| 381cd710a2 | |||
| ad0d18cb79 | |||
| 7980ee77d0 | |||
| 916b8bb327 | |||
| 87e3d4dea9 | |||
| 76a913f5e0 | |||
| d8c192e647 | |||
| c453437620 | |||
| 720fa6d909 | |||
| b4f71089f4 | |||
| 83e6657ead | |||
| 7ea6df4111 | |||
| d9ab92602a | |||
| 5ffadaed31 | |||
| 24f5f7df5d | |||
| daf01fb1d6 | |||
| 0f11b6def9 | |||
| 7df83f44b8 | |||
| 169fa7bed6 | |||
| 027b504fe8 | |||
| 186ef4da78 | |||
| dc674098e7 | |||
| 9087b4b07c | |||
| 8e45c22d7a | |||
| 191c4e03cd | |||
| ab4154c57d | |||
| 28e93d12ff | |||
| 75e63c758b | |||
| 823f71f269 | |||
| 042dd4088d | |||
| eaa791a9eb | |||
| 2228627ff4 | |||
| 4c647ad9c8 | |||
| 8ca3e6c33f | |||
| dd6bdbf297 | |||
| b47dde87e4 | |||
| 99e65cccd8 | |||
| 3bdacb8f46 | |||
| b4f9c224d3 | |||
| 5ec0399c81 | |||
| b464fdc333 | |||
| 53825500db | |||
| f2ac790752 | |||
| 0d8805cdee | |||
| 656e24ac9b | |||
| 6718b37403 | |||
| c9e5e784fc | |||
| f92f958682 | |||
| f63fab0676 | |||
| cfc4903c0c | |||
| a527a847fe | |||
| 91b0bf8933 | |||
| 66d1c96783 | |||
| 986128076e | |||
| 1de0a53241 | |||
| 0ec7eaf606 | |||
| d9fcb0e92b | |||
| f49b4ba4db | |||
| 84e708328b | |||
| 125bed3f09 | |||
| 077e70169d | |||
| e6dc169a05 | |||
| f34c02756d | |||
| 1e4c315481 | |||
| a8283a0d00 | |||
| 55896669fc | |||
| e341e0b9d2 | |||
| e6538c83bb | |||
| 92e1285ea5 | |||
| 2aabd1d90e | |||
| 7b8b778f83 | |||
| 7c8dc57d55 | |||
| fe95fae5f2 | |||
| ce8a95abf7 | |||
| c8e7e543d6 | |||
| a9dbb15ffa | |||
| cf64043f7d | |||
| ccaff92c18 | |||
| 585b5c922a |
@@ -1,47 +1,145 @@
|
|||||||
---
|
---
|
||||||
name: lora-manager-e2e
|
name: lora-manager-e2e
|
||||||
description: End-to-end testing and validation for LoRa Manager features. Use when performing automated E2E validation of LoRa Manager standalone mode, including starting/restarting the server, using Chrome DevTools MCP to interact with the web UI at http://127.0.0.1:8188/loras, and verifying frontend-to-backend functionality. Covers workflow validation, UI interaction testing, and integration testing between the standalone Python backend and the browser frontend.
|
description: End-to-end testing and validation for LoRa Manager features. Use when performing automated E2E validation of LoRa Manager standalone mode in a SANDBOXED, disposable configuration: check the port, start/restart the standalone server on a free port, use Chrome DevTools MCP to interact with the web UI (http://127.0.0.1:{PORT}/loras), and verify frontend-to-backend functionality. Covers workflow validation, UI interaction testing, and integration testing between the standalone Python backend and the browser frontend. Trigger keywords: E2E, standalone, Chrome DevTools MCP, lora-manager-e2e, sandbox.
|
||||||
---
|
---
|
||||||
|
|
||||||
# LoRa Manager E2E Testing
|
# LoRa Manager E2E Testing
|
||||||
|
|
||||||
This skill provides workflows and utilities for end-to-end testing of LoRa Manager using Chrome DevTools MCP.
|
This skill provides workflows and utilities for end-to-end testing of LoRa Manager using Chrome DevTools MCP.
|
||||||
|
|
||||||
|
## Conventions Used in This Document
|
||||||
|
|
||||||
|
- **`{PORT}`**: The server port. The default candidate is `8188`, but **`8188` is commonly occupied by a live ComfyUI process** and MUST NOT be assumed to be free. Always check availability first (see [Port Selection](#port-selection)) and use a free port (e.g. `8199`) for the E2E run. Substitute the actual port for every `{PORT}` in the commands below.
|
||||||
|
- **`<repo-root>`**: The repository/worktree root. Always run commands from the repo or worktree root; never assume a specific absolute path (paths such as `/home/<user>/...` differ per machine). The E2E scripts resolve the project root themselves, but fixture/settings paths are relative to `<repo-root>`.
|
||||||
|
|
||||||
|
## SANDBOX (MANDATORY)
|
||||||
|
|
||||||
|
> **Read this section before running anything.** Every E2E run MUST target a throwaway sandbox, never the real user data. A fresh subagent that skips this section WILL permanently mutate real user recipes.
|
||||||
|
|
||||||
|
1. **Portable settings**: create `<repo-root>/settings.json` (gitignored) with `"use_portable_settings": true` plus sandboxed `folder_paths` (lora/checkpoint roots) and `recipes_path`. This keeps the configuration inside the repo instead of the real user config dir (`~/.config/ComfyUI-LoRA-Manager/settings.json`).
|
||||||
|
2. **Sandboxed paths**: point `folder_paths` / `recipes_path` / `example_images_path` at disposable dirs — e.g. under `/tmp/opencode/<plan-name>-e2e/` (or worktree-local dirs). NEVER point the E2E at the real library (`~/models/...`), real recipe dir, or real settings.
|
||||||
|
3. **Never touch the real config**: the real user config at `~/.config/ComfyUI-LoRA-Manager/settings.json` and the real recipe dir must remain byte-identical before and after the run.
|
||||||
|
4. **Record real-data protection proof** before starting and after finishing:
|
||||||
|
```bash
|
||||||
|
# BEFORE: snapshot real config + recipe library state
|
||||||
|
sha256sum ~/.config/ComfyUI-LoRA-Manager/settings.json > /tmp/opencode/<plan>-e2e/settings.before.sha256
|
||||||
|
ls ~/models/recipes/*.recipe.json 2>/dev/null | wc -l > /tmp/opencode/<plan>-e2e/recipes-count.before.txt
|
||||||
|
find ~/models/recipes -name '*.recipe.json' -newermt "$(date -Iseconds)" | head # expect empty after run
|
||||||
|
# AFTER: record again, then diff the two snapshots. Any change = the run leaked into real data.
|
||||||
|
```
|
||||||
|
Also confirm `<repo-root>/git status` stays clean for `settings.json`/`cache/` (both are gitignored).
|
||||||
|
|
||||||
|
### Portable Settings Example
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"use_portable_settings": true,
|
||||||
|
"folder_paths": {
|
||||||
|
"loras": ["/tmp/opencode/<plan>-e2e/models/loras"],
|
||||||
|
"checkpoints": ["/tmp/opencode/<plan>-e2e/models/checkpoints"],
|
||||||
|
"unet": ["/tmp/opencode/<plan>-e2e/models/checkpoints"],
|
||||||
|
"diffusers": []
|
||||||
|
},
|
||||||
|
"recipes_path": "/tmp/opencode/<plan>-e2e/recipes",
|
||||||
|
"example_images_path": "/tmp/opencode/<plan>-e2e/example_images"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
The scanner computes and persists model hashes during the library scan, so the sandbox model dirs just need the model files + `.metadata.json` sidecars (see [Fixture + Fresh-State Guidance](#fixture--fresh-state-guidance)).
|
||||||
|
|
||||||
|
## Time Budgets & Abort Guidance
|
||||||
|
|
||||||
|
A fresh subagent should complete a sandboxed standalone E2E **in well under 30 minutes**. Budget each phase:
|
||||||
|
|
||||||
|
| Phase | Expected duration | Abort if |
|
||||||
|
| --- | --- | --- |
|
||||||
|
| Port check + sandbox setup | < 2 min | — |
|
||||||
|
| Server start (detached) + readiness | < 30 s | > 60 s (2x) → stop |
|
||||||
|
| Chrome DevTools MCP connect | < 1 min | > 2 min → stop |
|
||||||
|
| Per entry-point run (after fixtures ready) | < 5 min | > 10 min (2x) → stop |
|
||||||
|
| Fixture reset + cache clear between runs | < 1 min | > 2 min → stop |
|
||||||
|
|
||||||
|
**Abort rule**: if a phase exceeds ~2x its budget, OR any single tool call fails/retries 3+ times in a row, **STOP**. Do not loop or retry blindly. Report `BLOCKED` with: the phase, the last observed state (server PID + `ss -tlnp` output, page snapshot, last API response), and the suspected cause. Record the partial state as evidence; a clean BLOCKED report is more valuable than an hour of retries.
|
||||||
|
|
||||||
## Prerequisites
|
## Prerequisites
|
||||||
|
|
||||||
- LoRa Manager project cloned and dependencies installed (`pip install -r requirements.txt`)
|
- LoRa Manager project cloned and dependencies installed (`pip install -r requirements.txt`) — run everything from `<repo-root>`
|
||||||
- Chrome browser available for debugging
|
- Chrome browser available for debugging
|
||||||
- Chrome DevTools MCP connected
|
- Chrome DevTools MCP connected
|
||||||
|
- `ss` (or `lsof`/`netstat`) available for port checks: `ss -tlnp`
|
||||||
|
|
||||||
## Quick Start Workflow
|
## Port Selection
|
||||||
|
|
||||||
### 1. Start LoRa Manager Standalone
|
`8188` is only the *default candidate*. Verify it is actually free before every run:
|
||||||
|
|
||||||
```python
|
|
||||||
# Use the provided script to start the server
|
|
||||||
python .agents/skills/lora-manager-e2e/scripts/start_server.py --port 8188
|
|
||||||
```
|
|
||||||
|
|
||||||
Or manually:
|
|
||||||
```bash
|
|
||||||
cd /home/miao/workspace/ComfyUI/custom_nodes/ComfyUI-Lora-Manager
|
|
||||||
python standalone.py --port 8188
|
|
||||||
```
|
|
||||||
|
|
||||||
Wait for server ready message before proceeding.
|
|
||||||
|
|
||||||
### 2. Open Chrome Debug Mode
|
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
# Chrome with remote debugging on port 9222
|
# Is anything listening on 8188?
|
||||||
google-chrome --remote-debugging-port=9222 --user-data-dir=/tmp/chrome-lora-manager http://127.0.0.1:8188/loras
|
ss -tlnp | grep ':8188' || echo "8188 is free"
|
||||||
```
|
```
|
||||||
|
|
||||||
### 3. Connect Chrome DevTools MCP
|
- If a process holds `8188` (e.g. a live ComfyUI — pid 6575 on this machine), pick a different free port, e.g. `8199`:
|
||||||
|
```bash
|
||||||
|
ss -tlnp | grep ':8199' || echo "8199 is free"
|
||||||
|
```
|
||||||
|
- **Never** kill a process you did not start for this E2E. The live ComfyUI is off-limits. Pick a free port instead.
|
||||||
|
- Use your chosen port for **all** subsequent commands (server, Chrome launch, browser URLs).
|
||||||
|
|
||||||
Ensure the MCP server is connected to Chrome at `http://localhost:9222`.
|
## Quick Start Workflow (sandboxed)
|
||||||
|
|
||||||
### 4. Navigate and Interact
|
### 1. Prepare the sandbox
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cd <repo-root> # ALWAYS run from the repo/worktree root
|
||||||
|
mkdir -p /tmp/opencode/<plan>-e2e/models/{loras,checkpoints}
|
||||||
|
mkdir -p /tmp/opencode/<plan>-e2e/{recipes,example_images,recipes-before}
|
||||||
|
# write <repo-root>/settings.json per the portable-settings example above
|
||||||
|
# record real-data protection proof (see SANDBOX section)
|
||||||
|
```
|
||||||
|
|
||||||
|
### 2. Check port availability
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ss -tlnp | grep ':{PORT}' || echo "port {PORT} is free"
|
||||||
|
```
|
||||||
|
|
||||||
|
If `{PORT}` is occupied by an unrelated process, pick a free one and use it everywhere below. When in doubt use `8199`.
|
||||||
|
|
||||||
|
### 3. Start LoRa Manager Standalone (detached)
|
||||||
|
|
||||||
|
The standalone server **dies with the shell unless launched fully detached** — a plain background `&` from the bash tool is killed when the tool call returns. Launch via the helper script:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python .agents/skills/lora-manager-e2e/scripts/start_server.py --port {PORT} --wait --timeout 30 --detach
|
||||||
|
```
|
||||||
|
|
||||||
|
Or manually (equivalent detached form):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
setsid nohup python standalone.py --port {PORT} --host 127.0.0.1 < /dev/null \
|
||||||
|
>> /tmp/opencode/<plan>-e2e/server.log 2>&1 &
|
||||||
|
echo "started" # record the printed/pidfile PID for cleanup
|
||||||
|
```
|
||||||
|
|
||||||
|
Verify it is listening **before** proceeding (readiness poll is not a substitute for this):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ss -tlnp | grep ':{PORT}'
|
||||||
|
```
|
||||||
|
|
||||||
|
Record the server PID for cleanup: the helper script writes it to `/tmp/lora-manager-e2e-server-{PORT}.pid`; a manual `setsid` launch has no pidfile, so capture it explicitly (e.g. from `ss -tlnp`).
|
||||||
|
|
||||||
|
### 4. Open Chrome Debug Mode
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Chrome with remote debugging on port 9222 (note the {PORT} URL)
|
||||||
|
google-chrome --remote-debugging-port=9222 --user-data-dir=/tmp/chrome-lora-manager http://127.0.0.1:{PORT}/loras
|
||||||
|
```
|
||||||
|
|
||||||
|
### 5. Connect Chrome DevTools MCP
|
||||||
|
|
||||||
|
Ensure the MCP server is connected to Chrome at `http://localhost:9222`. Verify with `list_pages` — if it fails with "browser is already running", see [Chrome DevTools MCP Troubleshooting](#chrome-devtools-mcp-troubleshooting).
|
||||||
|
|
||||||
|
### 6. Navigate and Interact
|
||||||
|
|
||||||
Use Chrome DevTools MCP tools to:
|
Use Chrome DevTools MCP tools to:
|
||||||
- Take snapshots: `take_snapshot`
|
- Take snapshots: `take_snapshot`
|
||||||
@@ -56,7 +154,7 @@ Use Chrome DevTools MCP tools to:
|
|||||||
|
|
||||||
```python
|
```python
|
||||||
# Navigate to LoRA list page
|
# Navigate to LoRA list page
|
||||||
navigate_page(type="url", url="http://127.0.0.1:8188/loras")
|
navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
||||||
|
|
||||||
# Wait for page to load
|
# Wait for page to load
|
||||||
wait_for(text="LoRAs", timeout=10000)
|
wait_for(text="LoRAs", timeout=10000)
|
||||||
@@ -68,9 +166,10 @@ snapshot = take_snapshot()
|
|||||||
### Pattern: Restart Server for Configuration Changes
|
### Pattern: Restart Server for Configuration Changes
|
||||||
|
|
||||||
```python
|
```python
|
||||||
# Stop current server (if running)
|
# Stop current server (if running), start with new configuration.
|
||||||
# Start with new configuration
|
# --restart only kills the E2E server this script started before (via its pidfile);
|
||||||
python .agents/skills/lora-manager-e2e/scripts/start_server.py --port 8188 --restart
|
# it refuses to blindly kill unrelated processes on the port.
|
||||||
|
python .agents/skills/lora-manager-e2e/scripts/start_server.py --port {PORT} --restart --wait --detach
|
||||||
|
|
||||||
# Wait and refresh browser
|
# Wait and refresh browser
|
||||||
navigate_page(type="reload", ignoreCache=True)
|
navigate_page(type="reload", ignoreCache=True)
|
||||||
@@ -130,24 +229,96 @@ click(uid="modal-submit-button")
|
|||||||
wait_for(text="Success", timeout=5000)
|
wait_for(text="Success", timeout=5000)
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Fixture + Fresh-State Guidance
|
||||||
|
|
||||||
|
For rematch/repair E2E runs, seed the **sandboxed** `recipes_path` with hand-written fixture recipes. Rules (validated by the task-8 E2E):
|
||||||
|
|
||||||
|
1. **Filename constraint**: each file MUST be named `f"{id}.recipe.json"` **and** the in-JSON `id` field MUST equal the filename. Discovery accepts any `*.recipe.json`, but persistence resolves the path via `get_recipe_json_path` and `_save_recipe_persistently` returns `False` on a mismatch → the fixture would be counted as an error.
|
||||||
|
- `recipe-a.recipe.json` → in-JSON `"id": "recipe-a"`
|
||||||
|
2. **File format**: mirror an existing recipe JSON — top-level `id`, `file_path`, `title`, `loras`, `fingerprint`, `gen_params`; lora entries per the persistence conventions (`hash`, `file_name`, `modelVersionId`, `isDeleted`, ...).
|
||||||
|
3. **Companion image**: each recipe needs an image (e.g. a `.webp` generated with PIL) referenced by `file_path`, used for EXIF verification (`ExifUtils.append_recipe_metadata` writes a `"Recipe metadata: ..."` marker; a freshly generated `.webp` with no marker is the clean "untouched" control).
|
||||||
|
4. **autov3 three-state contract**: for L3 (autov3-only, renamed-file) fixtures the local model's `.metadata.json` sidecar MUST have the `autov3` key **ABSENT** (the "unchecked" state), NOT `""` — `""` is the TERMINAL "checked but unavailable" state that L3 deliberately skips. The scanner computes + persists `autov3` from the file header during the normal library scan (`model_scanner.py` `_process_model_file`), so the live L3 match resolves through the local autov3/hash cache; the computed-autov3 branch for unchecked items is covered by the unit suite.
|
||||||
|
5. **Fixture design for a rematch run** (mirrors the task-8 E2E):
|
||||||
|
- `recipe-a`: lora entry `isDeleted=True`, `hash` = 12-char autov3 computed from the local model (`calculate_autov3`, `py/utils/file_utils.py`), whose local model file was RENAMED after the recipe was written so `file_name` differs (proves L3 match without filename).
|
||||||
|
- `recipe-b`: parser-convention checkpoint entry (uses `id`, no `modelVersionId`) matching a local checkpoint via L2 — the local checkpoint's `.metadata.json` MUST carry civitai version data with that `id` so `version_index` contains it (L2 cannot match otherwise).
|
||||||
|
- `recipe-c`: healthy recipe (no deleted entries) → must remain untouched.
|
||||||
|
|
||||||
|
### Fresh state between entry-point runs
|
||||||
|
|
||||||
|
Each entry point (global / per-recipe / selection-bulk) must start from the same deleted state. Between runs:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 1. Reset fixtures to the before-state snapshot (copy back from recipes-before/)
|
||||||
|
cp /tmp/opencode/<plan>-e2e/recipes-before/*.recipe.json /tmp/opencode/<plan>-e2e/recipes/
|
||||||
|
# 2. Clear the recipe/FTS caches so the stale in-memory/library state is gone
|
||||||
|
rm -f <repo-root>/cache/recipe/*.sqlite
|
||||||
|
rm -rf <repo-root>/cache/fts/*
|
||||||
|
# 3. Restart the server (fresh process, fresh scan)
|
||||||
|
python .agents/skills/lora-manager-e2e/scripts/start_server.py --port {PORT} --restart --wait --timeout 30 --detach
|
||||||
|
# 4. Re-verify server listening + reload the browser page
|
||||||
|
```
|
||||||
|
|
||||||
|
## Server Lifecycle
|
||||||
|
|
||||||
|
- **Detached launch is mandatory**: the standalone server dies with the shell unless launched via `setsid` (or the helper script's `--detach`). Use `setsid nohup python standalone.py --port {PORT} --host 127.0.0.1 ... < /dev/null &`.
|
||||||
|
- **Verify with `ss -tlnp`** after every (re)start; do not proceed on a blind "server starting" message.
|
||||||
|
- **Never kill pre-existing processes** — only kill the E2E server PID you started (`start_server.py --restart` kills only PIDs it manages via its pidfile). The live ComfyUI or a stale QA Chrome must never be killed as part of cleanup unless explicitly identified as such (see Chrome troubleshooting).
|
||||||
|
- **Record your PID for cleanup**: note the PID printed/pidfile, and stop exactly that PID at the end (`kill <PID>`, then confirm with `ss -tlnp` that `{PORT}` is released).
|
||||||
|
|
||||||
|
## Chrome DevTools MCP Troubleshooting
|
||||||
|
|
||||||
|
### Stale profile lock ("browser is already running" / `list_pages` fails)
|
||||||
|
|
||||||
|
A Chrome profile can be held by a stale Chrome from a prior MCP session, which makes `list_pages` fail with "browser is already running":
|
||||||
|
|
||||||
|
1. Identify the stale Chrome — it owns the profile dir in `--user-data-dir` (e.g. `~/.config/chrome-dev-profile`). Find its process:
|
||||||
|
```bash
|
||||||
|
ps -ef | grep -i '[c]hrome.*user-data-dir'
|
||||||
|
```
|
||||||
|
2. Confirm it is a QA Chrome from a completed task (its parent is an old MCP/browser process, it is NOT the live ComfyUI server, and it is NOT your current MCP instance).
|
||||||
|
3. Kill ONLY that stale Chrome:
|
||||||
|
```bash
|
||||||
|
kill <stale-chrome-pid>
|
||||||
|
```
|
||||||
|
Never kill the live server or unrelated processes.
|
||||||
|
4. Retry `list_pages`. The current MCP will spawn a fresh browser.
|
||||||
|
|
||||||
|
### Screenshot-write restrictions
|
||||||
|
|
||||||
|
The chrome-devtools MCP may refuse to write into paths outside its configured workspace roots (e.g. the worktree `.omo/evidence/...` canonicalizing to an unmapped path). Workaround:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# 1. Save the screenshot to /tmp via the MCP
|
||||||
|
# take_screenshot(filePath="/tmp/<plan>-e2e/recipe-b-after.png", format="png")
|
||||||
|
# 2. Copy it into the evidence dir from the shell
|
||||||
|
mkdir -p <repo-root>/.omo/evidence/screenshots
|
||||||
|
cp /tmp/<plan>-e2e/recipe-b-after.png <repo-root>/.omo/evidence/screenshots/
|
||||||
|
```
|
||||||
|
|
||||||
|
## Cancellation Testing (KNOWN GAP)
|
||||||
|
|
||||||
|
Testing the rematch-cancel path E2E requires a run long enough to cancel mid-flight. A tiny 3-recipe fixture set completes in **seconds** — too fast to reliably cancel. The cancel path is currently **unit-covered only** (`rematch_all_recipes` cancellation tests); do not block an E2E run on cancel-path verification. If you must attempt it, you would need an artificially large/deferred fixture set to create a cancellable window — treat this as a research task, not part of the standard E2E.
|
||||||
|
|
||||||
## Available Scripts
|
## Available Scripts
|
||||||
|
|
||||||
### scripts/start_server.py
|
### scripts/start_server.py
|
||||||
|
|
||||||
Starts or restarts the LoRa Manager standalone server.
|
Starts or restarts the LoRa Manager standalone server for E2E testing.
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/start_server.py [--port PORT] [--restart] [--wait]
|
python scripts/start_server.py [--port PORT] [--restart] [--wait] [--timeout SECONDS] [--detach]
|
||||||
```
|
```
|
||||||
|
|
||||||
Options:
|
Options:
|
||||||
- `--port`: Server port (default: 8188)
|
- `--port`: Server port (default: 8188). The script exits early with a clear message if the port is already in use by an unrelated process.
|
||||||
- `--restart`: Kill existing server before starting
|
- `--restart`: Kill the E2E server this script previously managed (tracked via `/tmp/lora-manager-e2e-server-{PORT}.pid`) before starting. If unrelated processes still hold the port after that, the script reports them and aborts instead of killing them.
|
||||||
- `--wait`: Wait for server to be ready before exiting
|
- `--wait`: Wait for the server to be ready before exiting.
|
||||||
|
- `--timeout`: Readiness wait timeout in seconds (default: 30).
|
||||||
|
- `--detach`: Launch the server fully detached (`setsid`-style, survives shell death — REQUIRED for E2E). Default off: a normal background process that dies with the shell.
|
||||||
|
|
||||||
### scripts/wait_for_server.py
|
### scripts/wait_for_server.py
|
||||||
|
|
||||||
Polls server until ready or timeout.
|
Polls the server until ready or timeout.
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
python scripts/wait_for_server.py [--port PORT] [--timeout SECONDS]
|
python scripts/wait_for_server.py [--port PORT] [--timeout SECONDS]
|
||||||
@@ -196,6 +367,7 @@ results = performance_stop_trace()
|
|||||||
## Cleanup
|
## Cleanup
|
||||||
|
|
||||||
Always ensure proper cleanup after tests:
|
Always ensure proper cleanup after tests:
|
||||||
1. Stop the standalone server
|
1. Stop the standalone server: `kill <recorded-pid>` (only the PID you started), then confirm `ss -tlnp | grep ':{PORT}'` is empty.
|
||||||
2. Close browser pages (keep at least one open)
|
2. Close browser pages (keep at least one open).
|
||||||
3. Clear temporary data if needed
|
3. Remove the sandbox: `rm -rf /tmp/opencode/<plan>-e2e` and `<repo-root>/settings.json` + `<repo-root>/cache` (both gitignored).
|
||||||
|
4. Re-run the real-data protection check from the SANDBOX section and record the result in your evidence.
|
||||||
|
|||||||
@@ -2,11 +2,13 @@
|
|||||||
|
|
||||||
Quick reference for common MCP commands used in LoRa Manager E2E testing.
|
Quick reference for common MCP commands used in LoRa Manager E2E testing.
|
||||||
|
|
||||||
|
> **Port convention**: `{PORT}` is the port chosen for the E2E run (default candidate `8188`, but only if actually free — see the SKILL.md Port Selection section; use e.g. `8199` when `8188` is occupied by a live ComfyUI). Always run against the **sandboxed** standalone server, never a live instance.
|
||||||
|
|
||||||
## Navigation
|
## Navigation
|
||||||
|
|
||||||
```python
|
```python
|
||||||
# Navigate to LoRA list page
|
# Navigate to LoRA list page
|
||||||
navigate_page(type="url", url="http://127.0.0.1:8188/loras")
|
navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
||||||
|
|
||||||
# Reload page with cache clear
|
# Reload page with cache clear
|
||||||
navigate_page(type="reload", ignoreCache=True)
|
navigate_page(type="reload", ignoreCache=True)
|
||||||
@@ -179,7 +181,7 @@ pages = list_pages()
|
|||||||
select_page(pageId=0, bringToFront=True)
|
select_page(pageId=0, bringToFront=True)
|
||||||
|
|
||||||
# Create new page
|
# Create new page
|
||||||
new_page(url="http://127.0.0.1:8188/loras")
|
new_page(url="http://127.0.0.1:{PORT}/loras")
|
||||||
|
|
||||||
# Close page (keep at least one open!)
|
# Close page (keep at least one open!)
|
||||||
close_page(pageId=1)
|
close_page(pageId=1)
|
||||||
@@ -261,7 +263,7 @@ drag(from_uid="draggable-item", to_uid="drop-zone")
|
|||||||
### Verify LoRA Cards Loaded
|
### Verify LoRA Cards Loaded
|
||||||
|
|
||||||
```python
|
```python
|
||||||
navigate_page(type="url", url="http://127.0.0.1:8188/loras")
|
navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
||||||
wait_for(text="LoRAs", timeout=10000)
|
wait_for(text="LoRAs", timeout=10000)
|
||||||
|
|
||||||
# Check if cards loaded
|
# Check if cards loaded
|
||||||
@@ -322,3 +324,37 @@ navigate_page(type="reload")
|
|||||||
errors = list_console_messages(types=["error"])
|
errors = list_console_messages(types=["error"])
|
||||||
assert len(errors) == 0, f"Console errors: {errors}"
|
assert len(errors) == 0, f"Console errors: {errors}"
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Troubleshooting
|
||||||
|
|
||||||
|
### Stale profile lock ("browser is already running" / `list_pages` fails)
|
||||||
|
|
||||||
|
A Chrome profile held by a stale Chrome from a prior MCP session makes `list_pages`
|
||||||
|
fail with "browser is already running". Fix:
|
||||||
|
|
||||||
|
1. Find the stale Chrome that owns the profile dir (e.g. `~/.config/chrome-dev-profile`):
|
||||||
|
```bash
|
||||||
|
ps -ef | grep -i '[c]hrome.*user-data-dir'
|
||||||
|
```
|
||||||
|
2. Confirm it is a QA Chrome from a completed task (NOT the live ComfyUI server, NOT
|
||||||
|
your current MCP instance).
|
||||||
|
3. Kill ONLY that stale Chrome (`kill <stale-pid>`), then retry `list_pages`.
|
||||||
|
|
||||||
|
### Screenshot-write restrictions
|
||||||
|
|
||||||
|
The MCP may refuse to write into paths outside its configured workspace roots
|
||||||
|
(e.g. `.omo/evidence/screenshots/` under a worktree that canonicalizes to an unmapped
|
||||||
|
path). Save the screenshot to `/tmp` via the MCP, then copy it into the evidence dir:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# MCP: take_screenshot(filePath="/tmp/<plan>-e2e/recipe-b-after.png", format="png")
|
||||||
|
# Shell:
|
||||||
|
mkdir -p <repo-root>/.omo/evidence/screenshots
|
||||||
|
cp /tmp/<plan>-e2e/recipe-b-after.png <repo-root>/.omo/evidence/screenshots/
|
||||||
|
```
|
||||||
|
|
||||||
|
### Time budgets & abort rule
|
||||||
|
|
||||||
|
See SKILL.md "Time Budgets & Abort Guidance": if a phase exceeds ~2x its budget or a
|
||||||
|
tool call retries 3+ times in a row, STOP and report BLOCKED with the last observed
|
||||||
|
state (server PID + `ss -tlnp`, page snapshot, last API response). Do not loop.
|
||||||
|
|||||||
@@ -2,6 +2,14 @@
|
|||||||
|
|
||||||
This document provides detailed test scenarios for end-to-end validation of LoRa Manager features.
|
This document provides detailed test scenarios for end-to-end validation of LoRa Manager features.
|
||||||
|
|
||||||
|
> **Run preconditions (from SKILL.md)**: every run uses the **sandboxed** standalone
|
||||||
|
> server on a free port `{PORT}` (default candidate `8188`, only if actually free — pick
|
||||||
|
> e.g. `8199` when `8188` is occupied by a live ComfyUI). Fixtures live in the sandboxed
|
||||||
|
> `recipes_path` as `f"{id}.recipe.json"` files with matching in-JSON `id`; the real user
|
||||||
|
> config and real library are never touched (record protection proof before/after).
|
||||||
|
> Abort if a phase exceeds ~2x its budget or a tool call retries 3+ times (SKILL.md
|
||||||
|
> "Time Budgets & Abort Guidance").
|
||||||
|
|
||||||
## Table of Contents
|
## Table of Contents
|
||||||
|
|
||||||
1. [LoRA List Page](#lora-list-page)
|
1. [LoRA List Page](#lora-list-page)
|
||||||
@@ -19,7 +27,7 @@ This document provides detailed test scenarios for end-to-end validation of LoRa
|
|||||||
**Objective**: Verify the LoRA list page loads correctly and displays models.
|
**Objective**: Verify the LoRA list page loads correctly and displays models.
|
||||||
|
|
||||||
**Steps**:
|
**Steps**:
|
||||||
1. Navigate to `http://127.0.0.1:8188/loras`
|
1. Navigate to `http://127.0.0.1:{PORT}/loras`
|
||||||
2. Wait for page title "LoRAs" to appear
|
2. Wait for page title "LoRAs" to appear
|
||||||
3. Take snapshot to verify:
|
3. Take snapshot to verify:
|
||||||
- Header with "LoRAs" title is visible
|
- Header with "LoRAs" title is visible
|
||||||
@@ -134,7 +142,7 @@ evaluate_script(function="""
|
|||||||
**Objective**: Verify recipes page loads and displays recipes.
|
**Objective**: Verify recipes page loads and displays recipes.
|
||||||
|
|
||||||
**Steps**:
|
**Steps**:
|
||||||
1. Navigate to `http://127.0.0.1:8188/recipes`
|
1. Navigate to `http://127.0.0.1:{PORT}/recipes`
|
||||||
2. Wait for "Recipes" title
|
2. Wait for "Recipes" title
|
||||||
3. Take snapshot
|
3. Take snapshot
|
||||||
|
|
||||||
@@ -176,7 +184,7 @@ evaluate_script(function="""
|
|||||||
**Objective**: Verify settings page displays correctly.
|
**Objective**: Verify settings page displays correctly.
|
||||||
|
|
||||||
**Steps**:
|
**Steps**:
|
||||||
1. Navigate to `http://127.0.0.1:8188/settings`
|
1. Navigate to `http://127.0.0.1:{PORT}/settings`
|
||||||
2. Wait for "Settings" title
|
2. Wait for "Settings" title
|
||||||
3. Take snapshot
|
3. Take snapshot
|
||||||
|
|
||||||
@@ -190,7 +198,7 @@ evaluate_script(function="""
|
|||||||
1. Navigate to settings page
|
1. Navigate to settings page
|
||||||
2. Change a setting (e.g., default view mode)
|
2. Change a setting (e.g., default view mode)
|
||||||
3. Save settings
|
3. Save settings
|
||||||
4. Restart server: `python scripts/start_server.py --restart --wait`
|
4. Restart server: `python scripts/start_server.py --port {PORT} --restart --wait --timeout 30 --detach`
|
||||||
5. Refresh browser page
|
5. Refresh browser page
|
||||||
6. Navigate to settings
|
6. Navigate to settings
|
||||||
|
|
||||||
|
|||||||
@@ -8,186 +8,208 @@ This script shows how to:
|
|||||||
3. Verify functionality end-to-end
|
3. Verify functionality end-to-end
|
||||||
|
|
||||||
Note: This is a template. Actual execution requires Chrome DevTools MCP.
|
Note: This is a template. Actual execution requires Chrome DevTools MCP.
|
||||||
|
|
||||||
|
Port: pick a FREE port for the run — 8188 is commonly occupied by a live
|
||||||
|
ComfyUI (see the skill's Port Selection section). Set PORT below to e.g. 8199
|
||||||
|
when 8188 is taken. Always run against a SANDBOXED standalone server.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import time
|
|
||||||
|
# Choose the E2E port. 8188 is only the default candidate; use 8199 (or any
|
||||||
|
# free port checked with `ss -tlnp`) when 8188 is occupied by a live ComfyUI.
|
||||||
|
PORT = "8188"
|
||||||
|
|
||||||
|
|
||||||
def run_test():
|
def run_test():
|
||||||
"""Run example E2E test flow."""
|
"""Run example E2E test flow."""
|
||||||
|
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
print("LoRa Manager E2E Test Example")
|
print("LoRa Manager E2E Test Example")
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|
||||||
# Step 1: Start server
|
# Step 1: Start server (detached so it survives the shell)
|
||||||
print("\n[1/5] Starting LoRa Manager standalone server...")
|
print("\n[1/5] Starting LoRa Manager standalone server...")
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
[sys.executable, "start_server.py", "--port", "8188", "--wait", "--timeout", "30"],
|
[sys.executable, "start_server.py", "--port", PORT, "--wait", "--timeout", "30", "--detach"],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True
|
text=True,
|
||||||
)
|
)
|
||||||
if result.returncode != 0:
|
if result.returncode != 0:
|
||||||
print(f"Failed to start server: {result.stderr}")
|
print(f"Failed to start server: {result.stderr}")
|
||||||
return 1
|
return 1
|
||||||
print("Server ready!")
|
print("Server ready!")
|
||||||
|
|
||||||
# Step 2: Open Chrome (manual step - show command)
|
# Step 2: Open Chrome (manual step - show command)
|
||||||
print("\n[2/5] Open Chrome with debug mode:")
|
print("\n[2/5] Open Chrome with debug mode:")
|
||||||
print("google-chrome --remote-debugging-port=9222 --user-data-dir=/tmp/chrome-lora-manager http://127.0.0.1:8188/loras")
|
print(
|
||||||
|
f"google-chrome --remote-debugging-port=9222 "
|
||||||
|
f"--user-data-dir=/tmp/chrome-lora-manager http://127.0.0.1:{PORT}/loras"
|
||||||
|
)
|
||||||
print("(In actual test, this would be automated via MCP)")
|
print("(In actual test, this would be automated via MCP)")
|
||||||
|
|
||||||
# Step 3: Navigate and verify page load
|
# Step 3: Navigate and verify page load
|
||||||
print("\n[3/5] Page Load Verification:")
|
print("\n[3/5] Page Load Verification:")
|
||||||
print("""
|
print(
|
||||||
|
f"""
|
||||||
MCP Commands to execute:
|
MCP Commands to execute:
|
||||||
1. navigate_page(type="url", url="http://127.0.0.1:8188/loras")
|
1. navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
||||||
2. wait_for(text="LoRAs", timeout=10000)
|
2. wait_for(text="LoRAs", timeout=10000)
|
||||||
3. snapshot = take_snapshot()
|
3. snapshot = take_snapshot()
|
||||||
""")
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
# Step 4: Test search functionality
|
# Step 4: Test search functionality
|
||||||
print("\n[4/5] Search Functionality Test:")
|
print("\n[4/5] Search Functionality Test:")
|
||||||
print("""
|
print(
|
||||||
|
"""
|
||||||
MCP Commands to execute:
|
MCP Commands to execute:
|
||||||
1. fill(uid="search-input", value="test")
|
1. fill(uid="search-input", value="test")
|
||||||
2. press_key(key="Enter")
|
2. press_key(key="Enter")
|
||||||
3. wait_for(text="Results", timeout=5000)
|
3. wait_for(text="Results", timeout=5000)
|
||||||
4. result = evaluate_script(function="""
|
4. result = evaluate_script(function=`
|
||||||
() => {
|
() => {
|
||||||
const cards = document.querySelectorAll('.lora-card');
|
const cards = document.querySelectorAll('.lora-card');
|
||||||
return { count: cards.length };
|
return { count: cards.length };
|
||||||
}
|
}
|
||||||
""")
|
`)
|
||||||
""")
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
# Step 5: Verify API
|
# Step 5: Verify API
|
||||||
print("\n[5/5] API Verification:")
|
print("\n[5/5] API Verification:")
|
||||||
print("""
|
print(
|
||||||
|
"""
|
||||||
MCP Commands to execute:
|
MCP Commands to execute:
|
||||||
1. api_result = evaluate_script(function="""
|
1. api_result = evaluate_script(function=`
|
||||||
async () => {
|
async () => {
|
||||||
const response = await fetch('/loras/api/list');
|
const response = await fetch('/loras/api/list');
|
||||||
const data = await response.json();
|
const data = await response.json();
|
||||||
return { count: data.length, status: response.status };
|
return { count: data.length, status: response.status };
|
||||||
}
|
}
|
||||||
""")
|
`)
|
||||||
2. Verify api_result['status'] == 200
|
2. Verify api_result['status'] == 200
|
||||||
""")
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
print("\n" + "=" * 60)
|
||||||
print("Test flow completed!")
|
print("Test flow completed!")
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
|
|
||||||
def example_restart_flow():
|
def example_restart_flow():
|
||||||
"""Example: Testing configuration change that requires restart."""
|
"""Example: Testing configuration change that requires restart."""
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
print("\n" + "=" * 60)
|
||||||
print("Example: Server Restart Flow")
|
print("Example: Server Restart Flow")
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|
||||||
print("""
|
print(
|
||||||
|
f"""
|
||||||
Scenario: Change setting and verify after restart
|
Scenario: Change setting and verify after restart
|
||||||
|
|
||||||
Steps:
|
Steps:
|
||||||
1. Navigate to settings page
|
1. Navigate to settings page
|
||||||
- navigate_page(type="url", url="http://127.0.0.1:8188/settings")
|
- navigate_page(type="url", url="http://127.0.0.1:{PORT}/settings")
|
||||||
|
|
||||||
2. Change a setting (e.g., theme)
|
2. Change a setting (e.g., theme)
|
||||||
- fill(uid="theme-select", value="dark")
|
- fill(uid="theme-select", value="dark")
|
||||||
- click(uid="save-settings-button")
|
- click(uid="save-settings-button")
|
||||||
|
|
||||||
3. Restart server
|
3. Restart server
|
||||||
- subprocess.run([python, "start_server.py", "--restart", "--wait"])
|
- subprocess.run([python, "start_server.py", "--port", "{PORT}", "--restart", "--wait", "--detach"])
|
||||||
|
|
||||||
4. Refresh browser
|
4. Refresh browser
|
||||||
- navigate_page(type="reload", ignoreCache=True)
|
- navigate_page(type="reload", ignoreCache=True)
|
||||||
- wait_for(text="LoRAs", timeout=15000)
|
- wait_for(text="LoRAs", timeout=15000)
|
||||||
|
|
||||||
5. Verify setting persisted
|
5. Verify setting persisted
|
||||||
- navigate_page(type="url", url="http://127.0.0.1:8188/settings")
|
- navigate_page(type="url", url="http://127.0.0.1:{PORT}/settings")
|
||||||
- theme = evaluate_script(function="() => document.querySelector('#theme-select').value")
|
- theme = evaluate_script(function="() => document.querySelector('#theme-select').value")
|
||||||
- assert theme == "dark"
|
- assert theme == "dark"
|
||||||
""")
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def example_modal_interaction():
|
def example_modal_interaction():
|
||||||
"""Example: Testing modal dialog interaction."""
|
"""Example: Testing modal dialog interaction."""
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
print("\n" + "=" * 60)
|
||||||
print("Example: Modal Dialog Interaction")
|
print("Example: Modal Dialog Interaction")
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|
||||||
print("""
|
print(
|
||||||
|
"""
|
||||||
Scenario: Add new LoRA via modal
|
Scenario: Add new LoRA via modal
|
||||||
|
|
||||||
Steps:
|
Steps:
|
||||||
1. Open modal
|
1. Open modal
|
||||||
- click(uid="add-lora-button")
|
- click(uid="add-lora-button")
|
||||||
- wait_for(text="Add LoRA", timeout=3000)
|
- wait_for(text="Add LoRA", timeout=3000)
|
||||||
|
|
||||||
2. Fill form
|
2. Fill form
|
||||||
- fill_form(elements=[
|
- fill_form(elements=[
|
||||||
{"uid": "lora-name", "value": "Test Character"},
|
{"uid": "lora-name", "value": "Test Character"},
|
||||||
{"uid": "lora-path", "value": "/models/test.safetensors"},
|
{"uid": "lora-path", "value": "/models/test.safetensors"},
|
||||||
])
|
])
|
||||||
|
|
||||||
3. Submit
|
3. Submit
|
||||||
- click(uid="modal-submit-button")
|
- click(uid="modal-submit-button")
|
||||||
|
|
||||||
4. Verify success
|
4. Verify success
|
||||||
- wait_for(text="Successfully added", timeout=5000)
|
- wait_for(text="Successfully added", timeout=5000)
|
||||||
- snapshot = take_snapshot()
|
- snapshot = take_snapshot()
|
||||||
""")
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def example_network_monitoring():
|
def example_network_monitoring():
|
||||||
"""Example: Network request monitoring."""
|
"""Example: Network request monitoring."""
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
print("\n" + "=" * 60)
|
||||||
print("Example: Network Request Monitoring")
|
print("Example: Network Request Monitoring")
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|
||||||
print("""
|
print(
|
||||||
|
f"""
|
||||||
Scenario: Verify API calls during user interaction
|
Scenario: Verify API calls during user interaction
|
||||||
|
|
||||||
Steps:
|
Steps:
|
||||||
1. Clear network log (implicit on navigation)
|
1. Clear network log (implicit on navigation)
|
||||||
- navigate_page(type="url", url="http://127.0.0.1:8188/loras")
|
- navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
||||||
|
|
||||||
2. Perform action that triggers API call
|
2. Perform action that triggers API call
|
||||||
- fill(uid="search-input", value="character")
|
- fill(uid="search-input", value="character")
|
||||||
- press_key(key="Enter")
|
- press_key(key="Enter")
|
||||||
|
|
||||||
3. List network requests
|
3. List network requests
|
||||||
- requests = list_network_requests(resourceTypes=["xhr", "fetch"])
|
- requests = list_network_requests(resourceTypes=["xhr", "fetch"])
|
||||||
|
|
||||||
4. Find search API call
|
4. Find search API call
|
||||||
- search_requests = [r for r in requests if "/api/search" in r.get("url", "")]
|
- search_requests = [r for r in requests if "/api/search" in r.get("url", "")]
|
||||||
- assert len(search_requests) > 0, "Search API was not called"
|
- assert len(search_requests) > 0, "Search API was not called"
|
||||||
|
|
||||||
5. Get request details
|
5. Get request details
|
||||||
- if search_requests:
|
- if search_requests:
|
||||||
details = get_network_request(reqid=search_requests[0]["reqid"])
|
details = get_network_request(reqid=search_requests[0]["reqid"])
|
||||||
- Verify request method, response status, etc.
|
- Verify request method, response status, etc.
|
||||||
""")
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
print("LoRa Manager E2E Test Examples\n")
|
print("LoRa Manager E2E Test Examples\n")
|
||||||
print("This script demonstrates E2E testing patterns.\n")
|
print("This script demonstrates E2E testing patterns.\n")
|
||||||
print("Note: Actual execution requires Chrome DevTools MCP connection.\n")
|
print("Note: Actual execution requires Chrome DevTools MCP connection.\n")
|
||||||
|
|
||||||
run_test()
|
run_test()
|
||||||
example_restart_flow()
|
example_restart_flow()
|
||||||
example_modal_interaction()
|
example_modal_interaction()
|
||||||
example_network_monitoring()
|
example_network_monitoring()
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
print("\n" + "=" * 60)
|
||||||
print("All examples shown!")
|
print("All examples shown!")
|
||||||
print("=" * 60)
|
print("=" * 60)
|
||||||
|
|||||||
@@ -1,15 +1,78 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""
|
"""
|
||||||
Start or restart LoRa Manager standalone server for E2E testing.
|
Start or restart LoRa Manager standalone server for E2E testing.
|
||||||
|
|
||||||
|
Backward-compatible CLI: --port, --restart, --wait, --timeout all work as before.
|
||||||
|
New options: --detach (setsid-style fully detached launch, survives shell death).
|
||||||
|
|
||||||
|
Safety rules implemented here:
|
||||||
|
- Never kill processes the script did not start. The script tracks the PIDs it
|
||||||
|
manages in a pidfile (/tmp/lora-manager-e2e-server-{PORT}.pid).
|
||||||
|
- If the port is held by an unrelated process (e.g. a live ComfyUI) the script
|
||||||
|
reports the conflict and exits early instead of killing it.
|
||||||
|
- --restart only kills managed PIDs; if unrelated processes still hold the port
|
||||||
|
afterwards, the script reports them and aborts.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
|
import os
|
||||||
|
import signal
|
||||||
|
import socket
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
import socket
|
|
||||||
import signal
|
PIDFILE_PREFIX = "/tmp/lora-manager-e2e-server"
|
||||||
import os
|
|
||||||
|
|
||||||
|
def pidfile_path(port: int) -> str:
|
||||||
|
"""Path of the pidfile that records PIDs this script started for a port."""
|
||||||
|
return f"{PIDFILE_PREFIX}-{port}.pid"
|
||||||
|
|
||||||
|
|
||||||
|
def read_managed_pids(port: int) -> list[int]:
|
||||||
|
"""Read PIDs this script previously managed for the port (may be stale)."""
|
||||||
|
path = pidfile_path(port)
|
||||||
|
if not os.path.exists(path):
|
||||||
|
return []
|
||||||
|
try:
|
||||||
|
with open(path, "r", encoding="utf-8") as fh:
|
||||||
|
return [int(line.strip()) for line in fh if line.strip().isdigit()]
|
||||||
|
except (OSError, ValueError):
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def write_managed_pids(port: int, pids: list[int]) -> None:
|
||||||
|
"""Record PIDs this script manages for the port."""
|
||||||
|
try:
|
||||||
|
with open(pidfile_path(port), "w", encoding="utf-8") as fh:
|
||||||
|
for pid in pids:
|
||||||
|
fh.write(f"{pid}\n")
|
||||||
|
except OSError as exc:
|
||||||
|
print(f"Warning: could not write pidfile for port {port}: {exc}")
|
||||||
|
|
||||||
|
|
||||||
|
def clear_managed_pids(port: int) -> None:
|
||||||
|
"""Remove the pidfile for the port (no longer managed)."""
|
||||||
|
path = pidfile_path(port)
|
||||||
|
try:
|
||||||
|
if os.path.exists(path):
|
||||||
|
os.remove(path)
|
||||||
|
except OSError as exc:
|
||||||
|
print(f"Warning: could not remove pidfile {path}: {exc}")
|
||||||
|
|
||||||
|
|
||||||
|
def process_alive(pid: int) -> bool:
|
||||||
|
"""Return True if a process with the given pid exists."""
|
||||||
|
try:
|
||||||
|
os.kill(pid, 0)
|
||||||
|
return True
|
||||||
|
except ProcessLookupError:
|
||||||
|
return False
|
||||||
|
except PermissionError:
|
||||||
|
return True # exists but owned by someone else
|
||||||
|
|
||||||
|
|
||||||
def find_server_process(port: int) -> list[int]:
|
def find_server_process(port: int) -> list[int]:
|
||||||
@@ -19,7 +82,7 @@ def find_server_process(port: int) -> list[int]:
|
|||||||
["lsof", "-ti", f":{port}"],
|
["lsof", "-ti", f":{port}"],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True,
|
text=True,
|
||||||
check=False
|
check=False,
|
||||||
)
|
)
|
||||||
if result.returncode == 0 and result.stdout.strip():
|
if result.returncode == 0 and result.stdout.strip():
|
||||||
return [int(pid) for pid in result.stdout.strip().split("\n") if pid]
|
return [int(pid) for pid in result.stdout.strip().split("\n") if pid]
|
||||||
@@ -30,7 +93,7 @@ def find_server_process(port: int) -> list[int]:
|
|||||||
["netstat", "-tlnp"],
|
["netstat", "-tlnp"],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True,
|
text=True,
|
||||||
check=False
|
check=False,
|
||||||
)
|
)
|
||||||
pids = []
|
pids = []
|
||||||
for line in result.stdout.split("\n"):
|
for line in result.stdout.split("\n"):
|
||||||
@@ -49,30 +112,48 @@ def find_server_process(port: int) -> list[int]:
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
||||||
def kill_server(port: int) -> None:
|
def describe_processes(pids: list[int]) -> str:
|
||||||
"""Kill processes using the specified port."""
|
"""Human-readable description of a pid list (pid + command line)."""
|
||||||
pids = find_server_process(port)
|
descriptions = []
|
||||||
for pid in pids:
|
for pid in pids:
|
||||||
|
cmdline = ""
|
||||||
|
try:
|
||||||
|
with open(f"/proc/{pid}/cmdline", "rb") as fh:
|
||||||
|
raw = fh.read().replace(b"\x00", b" ").decode("utf-8", "replace")
|
||||||
|
cmdline = raw.strip()
|
||||||
|
except OSError:
|
||||||
|
pass
|
||||||
|
descriptions.append(f"pid {pid}{' (' + cmdline + ')' if cmdline else ''}")
|
||||||
|
return ", ".join(descriptions) if descriptions else "none"
|
||||||
|
|
||||||
|
|
||||||
|
def kill_pids(pids: list[int], what: str) -> None:
|
||||||
|
"""Send SIGTERM (then SIGKILL) to the given PIDs, only after reporting."""
|
||||||
|
for pid in pids:
|
||||||
|
print(f"Sent SIGTERM to {what} pid {pid}")
|
||||||
try:
|
try:
|
||||||
os.kill(pid, signal.SIGTERM)
|
os.kill(pid, signal.SIGTERM)
|
||||||
print(f"Sent SIGTERM to process {pid}")
|
|
||||||
except ProcessLookupError:
|
except ProcessLookupError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# Wait for processes to terminate
|
# Wait for processes to terminate
|
||||||
time.sleep(1)
|
deadline = time.time() + 5
|
||||||
|
while time.time() < deadline:
|
||||||
|
if not any(process_alive(pid) for pid in pids):
|
||||||
|
break
|
||||||
|
time.sleep(0.2)
|
||||||
|
|
||||||
# Force kill if still running
|
# Force kill if still running
|
||||||
pids = find_server_process(port)
|
|
||||||
for pid in pids:
|
for pid in pids:
|
||||||
try:
|
if process_alive(pid):
|
||||||
os.kill(pid, signal.SIGKILL)
|
try:
|
||||||
print(f"Sent SIGKILL to process {pid}")
|
os.kill(pid, signal.SIGKILL)
|
||||||
except ProcessLookupError:
|
print(f"Sent SIGKILL to {what} pid {pid}")
|
||||||
pass
|
except ProcessLookupError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
def is_server_ready(port: int, timeout: float = 0.5) -> bool:
|
def is_server_ready(port: int, timeout: float = 2.0) -> bool:
|
||||||
"""Check if server is accepting connections."""
|
"""Check if server is accepting connections."""
|
||||||
try:
|
try:
|
||||||
with socket.create_connection(("127.0.0.1", port), timeout=timeout):
|
with socket.create_connection(("127.0.0.1", port), timeout=timeout):
|
||||||
@@ -84,9 +165,15 @@ def is_server_ready(port: int, timeout: float = 0.5) -> bool:
|
|||||||
def wait_for_server(port: int, timeout: int = 30) -> bool:
|
def wait_for_server(port: int, timeout: int = 30) -> bool:
|
||||||
"""Wait for server to become ready."""
|
"""Wait for server to become ready."""
|
||||||
start = time.time()
|
start = time.time()
|
||||||
|
last_report = 0.0
|
||||||
while time.time() - start < timeout:
|
while time.time() - start < timeout:
|
||||||
if is_server_ready(port):
|
if is_server_ready(port):
|
||||||
return True
|
return True
|
||||||
|
# Report progress every ~5s so a slow boot is visible, not silent.
|
||||||
|
elapsed = time.time() - start
|
||||||
|
if elapsed - last_report >= 5:
|
||||||
|
print(f" ...still waiting ({int(elapsed)}s/{timeout}s)")
|
||||||
|
last_report = elapsed
|
||||||
time.sleep(0.5)
|
time.sleep(0.5)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -99,68 +186,148 @@ def main() -> int:
|
|||||||
"--port",
|
"--port",
|
||||||
type=int,
|
type=int,
|
||||||
default=8188,
|
default=8188,
|
||||||
help="Server port (default: 8188)"
|
help="Server port (default: 8188)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--restart",
|
"--restart",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
help="Kill existing server before starting"
|
help="Kill the E2E server previously managed by this script for the port "
|
||||||
|
"(tracked via pidfile) before starting; refuse to kill unrelated processes",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--wait",
|
"--wait",
|
||||||
action="store_true",
|
action="store_true",
|
||||||
help="Wait for server to be ready before exiting"
|
help="Wait for server to be ready before exiting",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--timeout",
|
"--timeout",
|
||||||
type=int,
|
type=int,
|
||||||
default=30,
|
default=30,
|
||||||
help="Timeout for waiting (default: 30)"
|
help="Timeout for waiting (default: 30)",
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--detach",
|
||||||
|
action="store_true",
|
||||||
|
help="Launch the server fully detached (setsid-style) so it survives shell "
|
||||||
|
"death. REQUIRED for E2E: a plain background process dies with the shell",
|
||||||
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
# Get project root (parent of .agents directory)
|
# Get project root (parent of .agents directory)
|
||||||
script_dir = os.path.dirname(os.path.abspath(__file__))
|
script_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
skill_dir = os.path.dirname(script_dir)
|
skill_dir = os.path.dirname(script_dir)
|
||||||
project_root = os.path.dirname(os.path.dirname(os.path.dirname(skill_dir)))
|
project_root = os.path.dirname(os.path.dirname(os.path.dirname(skill_dir)))
|
||||||
|
|
||||||
# Restart if requested
|
managed_pids = read_managed_pids(args.port)
|
||||||
|
|
||||||
|
# Restart if requested: kill ONLY managed PIDs.
|
||||||
if args.restart:
|
if args.restart:
|
||||||
print(f"Killing existing server on port {args.port}...")
|
alive_managed = [pid for pid in managed_pids if process_alive(pid)]
|
||||||
kill_server(args.port)
|
if alive_managed:
|
||||||
|
print(
|
||||||
|
f"Killing E2E server previously started by this script on port "
|
||||||
|
f"{args.port} ({describe_processes(alive_managed)})..."
|
||||||
|
)
|
||||||
|
kill_pids(alive_managed, "managed E2E server")
|
||||||
|
else:
|
||||||
|
print(
|
||||||
|
f"No live managed E2E server for port {args.port} "
|
||||||
|
f"(pidfile: {pidfile_path(args.port)})"
|
||||||
|
)
|
||||||
time.sleep(1)
|
time.sleep(1)
|
||||||
|
# Refuse to kill anything the script did not manage.
|
||||||
# Check if already running
|
remaining = find_server_process(args.port)
|
||||||
if is_server_ready(args.port):
|
if remaining:
|
||||||
print(f"Server already running on port {args.port}")
|
print(
|
||||||
return 0
|
f"ERROR: port {args.port} is still held by process(es) this script "
|
||||||
|
f"did not start: {describe_processes(remaining)}."
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
"These may be unrelated (e.g. a live ComfyUI). The script will NOT "
|
||||||
|
"kill them. Pick a different --port, or stop them manually if you "
|
||||||
|
"are certain they are stale E2E servers."
|
||||||
|
)
|
||||||
|
return 2
|
||||||
|
clear_managed_pids(args.port)
|
||||||
|
|
||||||
|
# Port conflict check before starting: never blind-kill.
|
||||||
|
port_pids = find_server_process(args.port)
|
||||||
|
if port_pids:
|
||||||
|
alive_managed = [pid for pid in port_pids if pid in managed_pids]
|
||||||
|
unmanaged = [pid for pid in port_pids if pid not in managed_pids]
|
||||||
|
if alive_managed and not unmanaged:
|
||||||
|
print(
|
||||||
|
f"Server already running on port {args.port} "
|
||||||
|
f"({describe_processes(alive_managed)}, started by this script). "
|
||||||
|
f"Use --restart to recycle it."
|
||||||
|
)
|
||||||
|
return 0
|
||||||
|
print(
|
||||||
|
f"ERROR: port {args.port} is already in use by process(es): "
|
||||||
|
f"{describe_processes(port_pids)}."
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
"This is likely an unrelated process (e.g. a live ComfyUI holding 8188). "
|
||||||
|
"The script will NOT kill it. Pick a free port with --port, e.g. 8199."
|
||||||
|
)
|
||||||
|
return 2
|
||||||
|
|
||||||
# Start server
|
# Start server
|
||||||
print(f"Starting LoRa Manager standalone server on port {args.port}...")
|
print(f"Starting LoRa Manager standalone server on port {args.port}...")
|
||||||
cmd = [sys.executable, "standalone.py", "--port", str(args.port)]
|
cmd = [
|
||||||
|
sys.executable,
|
||||||
# Start in background
|
"standalone.py",
|
||||||
process = subprocess.Popen(
|
"--host",
|
||||||
cmd,
|
"127.0.0.1",
|
||||||
cwd=project_root,
|
"--port",
|
||||||
stdout=subprocess.PIPE,
|
str(args.port),
|
||||||
stderr=subprocess.PIPE,
|
]
|
||||||
start_new_session=True
|
|
||||||
)
|
if args.detach:
|
||||||
|
# Fully detached launch: new session (setsid), no controlling terminal,
|
||||||
print(f"Server process started with PID {process.pid}")
|
# stdin from /dev/null, stdout/stderr to a log file. Survives the shell.
|
||||||
|
log_dir = os.path.join(script_dir, "logs")
|
||||||
|
os.makedirs(log_dir, exist_ok=True)
|
||||||
|
log_path = os.path.join(log_dir, f"server-{args.port}.log")
|
||||||
|
with open(log_path, "ab") as log_fh:
|
||||||
|
process = subprocess.Popen(
|
||||||
|
cmd,
|
||||||
|
cwd=project_root,
|
||||||
|
stdin=subprocess.DEVNULL,
|
||||||
|
stdout=log_fh,
|
||||||
|
stderr=subprocess.STDOUT,
|
||||||
|
start_new_session=True,
|
||||||
|
close_fds=True,
|
||||||
|
)
|
||||||
|
print(f"Detached server process started with PID {process.pid} (setsid)")
|
||||||
|
print(f"Log: {log_path}")
|
||||||
|
else:
|
||||||
|
# Plain background process (legacy behavior): dies with the shell.
|
||||||
|
process = subprocess.Popen(
|
||||||
|
cmd,
|
||||||
|
cwd=project_root,
|
||||||
|
stdout=subprocess.PIPE,
|
||||||
|
stderr=subprocess.PIPE,
|
||||||
|
start_new_session=True,
|
||||||
|
)
|
||||||
|
print(f"Server process started with PID {process.pid}")
|
||||||
|
print(
|
||||||
|
"NOTE: not detached — this process dies when the launching shell exits. "
|
||||||
|
"For E2E use --detach."
|
||||||
|
)
|
||||||
|
|
||||||
|
write_managed_pids(args.port, [process.pid])
|
||||||
|
|
||||||
# Wait for ready if requested
|
# Wait for ready if requested
|
||||||
if args.wait:
|
if args.wait:
|
||||||
print(f"Waiting for server to be ready (timeout: {args.timeout}s)...")
|
print(f"Waiting for server to be ready (timeout: {args.timeout}s)...")
|
||||||
if wait_for_server(args.port, args.timeout):
|
if wait_for_server(args.port, args.timeout):
|
||||||
print(f"Server ready at http://127.0.0.1:{args.port}/loras")
|
print(f"Server ready at http://127.0.0.1:{args.port}/loras")
|
||||||
return 0
|
return 0
|
||||||
else:
|
print(f"Timeout waiting for server on port {args.port}")
|
||||||
print(f"Timeout waiting for server")
|
return 1
|
||||||
return 1
|
|
||||||
|
|
||||||
print(f"Server starting at http://127.0.0.1:{args.port}/loras")
|
print(f"Server starting at http://127.0.0.1:{args.port}/loras")
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
|
|||||||
@@ -1,15 +1,20 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""
|
"""
|
||||||
Wait for LoRa Manager server to become ready.
|
Wait for LoRa Manager server to become ready.
|
||||||
|
|
||||||
|
Timeout is configurable via --timeout (default 30s); the script polls the port
|
||||||
|
until the server accepts connections or the timeout expires.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import socket
|
import socket
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
|
|
||||||
|
|
||||||
def is_server_ready(port: int, timeout: float = 0.5) -> bool:
|
def is_server_ready(port: int, timeout: float = 2.0) -> bool:
|
||||||
"""Check if server is accepting connections."""
|
"""Check if server is accepting connections."""
|
||||||
try:
|
try:
|
||||||
with socket.create_connection(("127.0.0.1", port), timeout=timeout):
|
with socket.create_connection(("127.0.0.1", port), timeout=timeout):
|
||||||
@@ -21,9 +26,15 @@ def is_server_ready(port: int, timeout: float = 0.5) -> bool:
|
|||||||
def wait_for_server(port: int, timeout: int = 30) -> bool:
|
def wait_for_server(port: int, timeout: int = 30) -> bool:
|
||||||
"""Wait for server to become ready."""
|
"""Wait for server to become ready."""
|
||||||
start = time.time()
|
start = time.time()
|
||||||
|
last_report = 0.0
|
||||||
while time.time() - start < timeout:
|
while time.time() - start < timeout:
|
||||||
if is_server_ready(port):
|
if is_server_ready(port):
|
||||||
return True
|
return True
|
||||||
|
# Report progress every ~5s so a slow boot is visible, not silent.
|
||||||
|
elapsed = time.time() - start
|
||||||
|
if elapsed - last_report >= 5:
|
||||||
|
print(f" ...still waiting ({int(elapsed)}s/{timeout}s)")
|
||||||
|
last_report = elapsed
|
||||||
time.sleep(0.5)
|
time.sleep(0.5)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -36,25 +47,24 @@ def main() -> int:
|
|||||||
"--port",
|
"--port",
|
||||||
type=int,
|
type=int,
|
||||||
default=8188,
|
default=8188,
|
||||||
help="Server port (default: 8188)"
|
help="Server port (default: 8188)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--timeout",
|
"--timeout",
|
||||||
type=int,
|
type=int,
|
||||||
default=30,
|
default=30,
|
||||||
help="Timeout in seconds (default: 30)"
|
help="Timeout in seconds (default: 30)",
|
||||||
)
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
print(f"Waiting for server on port {args.port} (timeout: {args.timeout}s)...")
|
print(f"Waiting for server on port {args.port} (timeout: {args.timeout}s)...")
|
||||||
|
|
||||||
if wait_for_server(args.port, args.timeout):
|
if wait_for_server(args.port, args.timeout):
|
||||||
print(f"Server ready at http://127.0.0.1:{args.port}/loras")
|
print(f"Server ready at http://127.0.0.1:{args.port}/loras")
|
||||||
return 0
|
return 0
|
||||||
else:
|
print(f"Timeout: Server not ready after {args.timeout}s")
|
||||||
print(f"Timeout: Server not ready after {args.timeout}s")
|
return 1
|
||||||
return 1
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ model_cache/
|
|||||||
reasonix.toml
|
reasonix.toml
|
||||||
.reasonix/
|
.reasonix/
|
||||||
.codegraph/
|
.codegraph/
|
||||||
|
.playwright-mcp/
|
||||||
|
|
||||||
# Vue widgets development cache (but keep build output)
|
# Vue widgets development cache (but keep build output)
|
||||||
vue-widgets/node_modules/
|
vue-widgets/node_modules/
|
||||||
|
|||||||
@@ -0,0 +1,202 @@
|
|||||||
|
---
|
||||||
|
slug: undo-delete-staging
|
||||||
|
status: drafting
|
||||||
|
intent: clear
|
||||||
|
review_required: false
|
||||||
|
pending-action: write .omo/plans/undo-delete-staging.md
|
||||||
|
approach: "Option B: delayed physical deletion with Undo. Backend: same-volume rename to per-root staging dir (.lm-pending-delete/) [updated 2026-08: model staging moved to a SIBLING dir inside each deleted model's own folder — see 'Symlink fix (2026-08)' under Decisions] + manifest JSON (batch_id, expires_at, staged->original map) + purge (30s TTL timer + startup sweep + opportunistic) + undo-delete endpoint + settings toggle 'skip undo'. Small files (recipes: JSON+preview) copy to global staging under settings dir instead of rename. Frontend: extend toast system with action button + 30s countdown; delete flows (single model / recipe / bulk / duplicates) consume batch_id from delete response and show Undo toast; expired undo -> 'undo expired' toast. Plus confirm-modal friction (C-friction, NO type-to-confirm): delete button delay-activation 1.5s + modal shows file size 'will free X GB' + Cancel gets initial focus. i18n keys + sync_translation_keys.py."
|
||||||
|
---
|
||||||
|
|
||||||
|
# Draft: undo-delete-staging
|
||||||
|
|
||||||
|
## Components (topology ledger)
|
||||||
|
<!-- Lock the SHAPE before depth. One row per top-level component that can succeed or fail independently. -->
|
||||||
|
<!-- id | outcome (one line) | status: active|deferred | evidence path -->
|
||||||
|
- backend staging module (stage/purge/undo + manifest + per-volume dir resolution) | new module, active | pending exploration: model_lifecycle_service.py delete_model / delete_model_artifacts
|
||||||
|
- delete endpoints return batch_id (model/recipe/bulk/duplicates) | active | pending exploration: handlers + response shapes
|
||||||
|
- undo-delete HTTP endpoint + route registration | active | pending exploration: route registrar pattern
|
||||||
|
- purge scheduling (30s timer + startup sweep + opportunistic) | active | pending exploration: app on_startup hooks
|
||||||
|
- settings toggle "skip undo window" | active | pending exploration: settings service read pattern
|
||||||
|
- frontend toast extension (action button + countdown) | active | pending exploration: showToast impl
|
||||||
|
- frontend delete flows consume batch_id + Undo toast | active | pending exploration: call sites
|
||||||
|
- confirm-modal friction (delay-activate + size display + cancel focus) | active | pending exploration: modal focus behavior
|
||||||
|
- i18n keys + sync_translation_keys.py | active | known
|
||||||
|
|
||||||
|
## Open assumptions (announced defaults)
|
||||||
|
<!-- Record any default you adopt instead of asking, so the user can veto it at the gate. -->
|
||||||
|
<!-- assumption | adopted default | rationale | reversible? -->
|
||||||
|
- Undo window TTL = 30s | 30s balances space-freeing intent vs accident recovery | yes (constant)
|
||||||
|
- Staging dir name: `.lm-pending-delete/` under each model root; recipes: `{settings_dir}/.lm-pending-delete/` | hidden, same-volume [updated 2026-08: same-volume is now guaranteed by sibling staging inside the model's own folder, not by the root location], consistent | yes
|
||||||
|
- Staging failure falls back to existing hard delete | user intent is delete; staging is best-effort; hard delete likely fails identically under same conditions | yes
|
||||||
|
- Purge on startup uses expires_at (not purge-all) so a <30s restart with live tab can still undo | robust, matches client-side timer | yes
|
||||||
|
- Settings toggle label: "Delete permanently immediately (skip undo window)" | power users freeing space | yes
|
||||||
|
- C-friction: delete button enabled after 1.5s + modal shows freed size; NO type-to-confirm (user vetoed) | user explicitly rejected type-to-confirm | n/a
|
||||||
|
- Bulk/duplicates delete: one batch id for whole action, one undo restores all | simplest consistent semantics | yes
|
||||||
|
|
||||||
|
## Findings (cited - path:lines)
|
||||||
|
|
||||||
|
### Backend
|
||||||
|
- `delete_model_artifacts` (py/services/model_lifecycle_service.py:19-48) = physical delete via os.remove; patterns: main file + `{name}.metadata.json` + PREVIEW_EXTENSIONS (py/utils/constants.py:22-37). ALSO called by ModelScanner.bulk_delete_models (py/services/model_scanner.py:2221) - single swap point covers bulk models.
|
||||||
|
- `ModelLifecycleService.delete_model` (model_lifecycle_service.py:101-154): fetches `cached_entry` (111-116) - SNAPSHOT available for cache restore; after delete: cache.raw_data removal + resort + bump_cache_version (136-143), `_hash_index.remove_by_path` (145-146), `_sync_update_for_model` (148; update-service only, no recipe JSON rewrites - recipe refs are hash-based, re-resolve on restore), `_persist_current_cache` (150-152), returns `{"success": True, "deleted_files": [...]}` (154).
|
||||||
|
- Handler `delete_model` (py/routes/handlers/model_handlers.py:478-492): POST /api/lm/{prefix}/delete; response passthrough; `_broadcast_models_changed()` (57-74) after success; 400 `{"success":false,"error"}`; 500 plain text.
|
||||||
|
- Recipe delete: handler (recipe_handlers.py:1422-1438) DELETE /api/lm/recipe/{recipe_id} -> persistence_service.delete_recipe (py/services/recipes/persistence_service.py:193-209): os.remove(recipe_json_path) + os.remove(image_path) (204-206), recipe_scanner.remove_recipe (208), returns `{"success": true, "message": ...}`. PersistenceResult dataclass (20-25).
|
||||||
|
- Bulk models: POST /api/lm/{prefix}/bulk-delete (model_route_registrar.py:39) -> handler (model_handlers.py:974-994) -> lifecycle_service.bulk_delete_models (model_lifecycle_service.py:308-318) -> scanner.bulk_delete_models (model_scanner.py:2181-2269) which calls delete_model_artifacts per file (2221) + `_batch_update_cache_for_deleted_models` (2271-2335); response `{"success","status","total_deleted","total_attempted","cache_updated","results"}` (2254-2269).
|
||||||
|
- Bulk recipes: POST /api/lm/recipes/bulk-delete (recipe_route_registrar.py:50) -> handler (recipe_handlers.py:1554-1573) -> persistence_service.bulk_delete (persistence_service.py:439-482): per-id os.remove x2 (464-466), recipe_scanner.bulk_remove (472); response `{"success","deleted","failed","total_deleted","total_failed"}` (474-482).
|
||||||
|
- Duplicates: NO dedicated delete endpoints (find-only: GET /api/lm/{prefix}/find-duplicates model_route_registrar.py:59, GET /api/lm/recipes/find-duplicates recipe_route_registrar.py:49). Duplicate deletion reuses bulk-delete endpoints.
|
||||||
|
- Startup hooks: lora_manager.py:183-187 `app.on_startup.append(lambda app: cls._initialize_services())` (ComfyUI mode, app = PromptServer.instance.app at :78); standalone.py:370-374 same (StandaloneLoraManager.add_routes). Background tasks: `asyncio.create_task(name=...)` (lora_manager.py:224-239; recipe_handlers.py:793). Singleton+asyncio.Lock pattern: model_scanner.py:40-63.
|
||||||
|
- Settings: DEFAULT_SETTINGS (py/services/settings_manager.py:57-119), `get(key, default)` (1390-1392), get_settings_manager() (2215-2228), reset_settings_manager() (2231). Typed-bool getter example: get_skip_previously_downloaded_model_versions (1253-1262). Handlers: base_model_routes.py:70, base_recipe_routes.py:54.
|
||||||
|
- Model roots: ModelScanner.get_model_roots base NotImplementedError (model_scanner.py:1073-1075); impls lora_scanner.py:31-45, checkpoint_scanner.py:428-441, embedding_scanner.py:24-36. `_find_root_for_file(file_path)` (model_scanner.py:1108-1124) returns containing root - for per-root staging dir computation [updated 2026-08: staging no longer uses the containing root; batches are siblings inside the model's own folder]. Business-path rule (AGENTS.md): use os.path.abspath, never realpath, for staging/undo routing.
|
||||||
|
- Cache restore methods: ModelCache has raw_data + resort (conftest mocks: tests/conftest.py:144-154); ModelHashIndex.add_entry(sha256, file_path, autov3) (py/services/model_hash_index.py:16); RecipeScanner.add_recipe(recipe_data) (recipe_scanner.py:2136) -> recipe_cache.add_recipe (recipe_cache.py:64). No single-file incremental model rescan - use snapshot restore instead of rescan.
|
||||||
|
- Route registrar: model_route_registrar.py:177 add_route(method, path, handler), :180 add_prefixed_route - undo endpoint can be a non-prefixed route via add_route.
|
||||||
|
- Tests: tests/services/test_model_lifecycle_service.py (inline tmp_path files, per-test stub scanners ScannerForDelete/VersionAwareScanner etc); conftest MockScanner/MockCache/MockHashIndex (tests/conftest.py:134-212); integration fixtures tests/integration/conftest.py; lifecycle hook tests tests/routes/test_lora_manager_lifecycle.py:177-178, tests/standalone/test_standalone_server.py:83-84.
|
||||||
|
|
||||||
|
### Frontend
|
||||||
|
- 5 delete call sites:
|
||||||
|
a) Single model: static/js/utils/modalUtils.js confirmDelete (27-42) -> getModelApiClient().deleteModel(path); ignores return.
|
||||||
|
b) Recipe single: static/js/components/RecipeCard.js confirmDeleteRecipe (405-449) - RAW fetch DELETE /api/lm/recipe/{id}, checks only response.ok, showToast toast.recipes.deletedSuccessfully, state.virtualScroller.removeItemByFilePath.
|
||||||
|
c) Bulk: static/js/managers/BulkManager.js confirmBulkDelete (633-672) -> getActiveApiClient() (134-142) -> bulkDeleteModels(filePaths); reads result.cancelled/success/deleted_count/error.
|
||||||
|
d) Recipe duplicates: static/js/components/DuplicatesManager.js confirmDeleteDuplicates (457-494) - RAW fetch POST /api/lm/recipes/bulk-delete, reads data.success/data.total_deleted, exitDuplicateMode().
|
||||||
|
e) Model duplicates: static/js/components/ModelDuplicatesManager.js confirmDeleteDuplicates (710-776) - RAW fetch POST /api/lm/{type}/bulk-delete, reads data.total_deleted, then resetAndReload(true) + find-duplicates re-check.
|
||||||
|
Bonus: static/js/components/shared/ModelVersionsTab.js:1136-1144 client.deleteModel (ignores return).
|
||||||
|
- API clients: BaseModelApiClient.deleteModel (static/js/api/baseModelApi.js:184-216) returns true/false, shows its own toasts, does removeItemByFilePath inside; bulkDeleteModels (1591-1642) returns {success, deleted_count, failed_count, errors} or {success:false, cancelled:true}; RecipeSidebarApiClient.bulkDeleteModels (recipeApi.js:623-664) returns {success, deleted_count: total_deleted, ...}. Endpoint map apiConfig.js:56,64.
|
||||||
|
- Toast: showToast(key, params={}, type='info', fallback=null) (static/js/utils/uiHelpers.js:136-193) - textContent only, NO action/button support; durations 2000/5000ms; CSS static/css/components/toast.css (.toast flex gap:12px - button can be added). Closest action pattern: bannerService.registerBanner actions array + onRegister (static/js/managers/BannerService.js; used uiHelpers.js:18-57).
|
||||||
|
- i18n: locales/en.json delete keys (1303-1314 bulkDelete, 1945-1948 recipes, 1987-1991 models, 2124-2130 duplicates, 2166-2170 toast.api); t()/interpolate (static/js/i18n/index.js:193-248); translate wrapper (utils/i18nHelpers.js:13-23); sync script scripts/sync_translation_keys.py (en reference, [TODO: Translate] placeholders).
|
||||||
|
- Refresh after undo: recipes -> window.recipeManager.loadRecipes(true) (recipes.js:359; used by FilterManager.js:752 etc) or refreshRecipes (recipeApi.js:308); models -> resetAndReload(true) from modelApiFactory (used by ModelDuplicatesManager.js:740).
|
||||||
|
- Size for modal: card.dataset.file_size (ModelCard.js:467), formatFileSize (ModelModal.js:615).
|
||||||
|
- Tests: tests/frontend/utils/uiHelpers.dom.test.js (toast), api/recipeApi.bulk.test.js, components/duplicatesManager.test.js, components/modelDuplicatesManager.test.js, pages/*Page.test.js, i18n tests tests/i18n/test_i18n.py.
|
||||||
|
|
||||||
|
## Decisions (with rationale)
|
||||||
|
|
||||||
|
1. Same-volume rename staging for model files (atomic, no copy cost for multi-GB files); cross-volume rename forbidden. [CORRECTED 2026-08: "same-volume because under the containing root" was only true for plain directories — nested symlinked subdirs could cross volumes. Superseded by sibling staging: `.lm-pending-delete/<batch_id>/` inside the deleted model's own folder makes stage/undo same-device by construction; see "Symlink fix (2026-08)" below.]
|
||||||
|
2. Copy-to-global-staging for recipes (small files; avoids recipe JSON vs preview image cross-volume problem).
|
||||||
|
3. Manifest JSON files are the only state - no DB changes. Manifest includes model cached_entry snapshot for exact cache restore (no rescan needed).
|
||||||
|
4. Undo endpoint returns restored paths; expired batch -> 404-style error -> frontend 'undo expired' toast.
|
||||||
|
5. Skip-undo setting honored server-side (no batch_id in response -> no undo toast client-side).
|
||||||
|
6. Staging failure falls back to existing hard delete (best-effort undo, never blocks delete).
|
||||||
|
7. Undo window TTL = 30s constant (PENDING_DELETE_TTL_SECONDS); startup sweep uses expires_at (survives restart; browser-tab timer survives).
|
||||||
|
8. Purge triple-trigger: per-batch asyncio timer task + on_startup sweep + opportunistic purge at each stage/undo.
|
||||||
|
9. Frontend: new showActionToast (keep showToast signature untouched; extract shared createToastElement/appendToast internals); undo click -> shared handleUndoDelete(batchId, refreshFn); full list refresh after undo (recipes: window.recipeManager.loadRecipes(true); models: resetAndReload(true)).
|
||||||
|
10. C-friction wave (NO type-to-confirm - user vetoed): delete buttons delay-activate 1.5s after modal open, initial focus on Cancel, model delete modal gains "permanently deleted from disk" warning + file size display (card.dataset.file_size + formatFileSize).
|
||||||
|
11. Model cache restore on undo: append snapshot to cache.raw_data (dedupe by file_path) + resort + bump_cache_version + _persist_current_cache + _hash_index.add_entry + _broadcast_models_changed. Recipe restore: copy back files + recipe_scanner.add_recipe(recipe_data loaded from restored JSON).
|
||||||
|
|
||||||
|
### Symlink fix (2026-08)
|
||||||
|
|
||||||
|
Post-execution addendum (plan `.omo/plans/undo-delete-symlink-fix.md`, commits 5fd4946b / 0c00ee22):
|
||||||
|
|
||||||
|
12. Model staging moved from `<model_root>/.lm-pending-delete/<batch_id>/` to `<model_dir>/.lm-pending-delete/<batch_id>/` (sibling of the model artifacts, inside the deleted model's own folder). Stage/undo renames are same-device BY CONSTRUCTION — EXDEV is impossible even when the business path traverses nested symlinks to other volumes (the decision-1 "containing root" guarantee covered only plain directories). EXDEV remains possible only for cross-volume merges, which keep the batch_ids-array fallback. Accepted edge: deleting the model's whole FOLDER during the 30s window destroys that batch (undo returns 404). Batch discovery uses an in-memory registry (`_known_batch_dirs`) with a startup reconciliation scan (`purge_expired(scan_roots=True)`) covering restarts and crash leftovers. Recipe batches unchanged (copy-based settings-dir staging with the `_restore_file` EXDEV fallback).
|
||||||
|
|
||||||
|
## Scope IN
|
||||||
|
|
||||||
|
- Model single delete (model_handlers delete_model / model_lifecycle_service)
|
||||||
|
- Recipe delete (recipe_handlers delete_recipe / persistence_service)
|
||||||
|
- Bulk delete (models scanner + recipes persistence) + duplicates (reuse bulk endpoints)
|
||||||
|
- Undo endpoint POST /api/lm/undo-delete (models + recipes, one batch space)
|
||||||
|
- Purge: timer + startup sweep + opportunistic
|
||||||
|
- Settings toggle delete_undo_enabled + settings page checkbox
|
||||||
|
- Frontend: showActionToast + all 5 delete flows + shared undo handler
|
||||||
|
- C-friction modal changes (delay-activate + cancel focus + warning copy + size display)
|
||||||
|
- i18n keys + sync_translation_keys.py
|
||||||
|
- Backend + frontend tests
|
||||||
|
|
||||||
|
## Scope OUT (Must NOT have)
|
||||||
|
|
||||||
|
- NO type-to-confirm / hold-to-confirm friction (user vetoed)
|
||||||
|
- NO OS trash integration (send2trash) in this iteration
|
||||||
|
- NO persistent recycle-bin UI (no trash browsing page)
|
||||||
|
- NO changes to exclude/unexclude flow
|
||||||
|
- NO DB migrations
|
||||||
|
- NO new dependencies (no send2trash)
|
||||||
|
- NO changes to download flows
|
||||||
|
- NO recipe-JSON rewriting on model undo (hash-based refs re-resolve themselves)
|
||||||
|
|
||||||
|
## Open questions
|
||||||
|
|
||||||
|
None - all implementation details resolved by exploration. Design decisions settled in conversation (B+C, no type-to-confirm).
|
||||||
|
|
||||||
|
## Approval gate
|
||||||
|
status: approved
|
||||||
|
<!-- Approach approved -> rerun scaffold without --draft-only, run Metis gap analysis, APPEND todo batches, fill TL;DR last, run structural self-check, then Phase 4 handoff. -->
|
||||||
|
|
||||||
|
## Review round state (ulw-plan-review-round-state-contract)
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"transition": "replace",
|
||||||
|
"phase": "review_round_initialized",
|
||||||
|
"applies_when": ["retry_after_plan_change"],
|
||||||
|
"atomic": true,
|
||||||
|
"review_required": true,
|
||||||
|
"plan_path": ".omo/plans/undo-delete-staging.md",
|
||||||
|
"plan_sha256": "8cf7c9be38a76d8ef1fb832aba043d6d7e82b60465b1bf28c6eafa7045117adc",
|
||||||
|
"review_round_id": "rr-undo-del-20260811-006",
|
||||||
|
"round_status": "active",
|
||||||
|
"pending-action": "review .omo/plans/undo-delete-staging.md",
|
||||||
|
"review": {
|
||||||
|
"momus": { "status": "pending", "workspace_root": "/mnt/data/reinstall-backup-2026-04-12/data/workspace/ComfyUI/custom_nodes/ComfyUI-Lora-Manager", "runtime_home": null, "target": ".omo/plans/undo-delete-staging.md", "round_id": "rr-undo-del-20260811-006", "plan_sha256": "8cf7c9be38a76d8ef1fb832aba043d6d7e82b60465b1bf28c6eafa7045117adc", "launch_id": null, "session": null, "result": null },
|
||||||
|
"independent": { "status": "pending", "workspace_root": "/mnt/data/reinstall-backup-2026-04-12/data/workspace/ComfyUI/custom_nodes/ComfyUI-Lora-Manager", "runtime_home": null, "target": ".omo/plans/undo-delete-staging.md", "round_id": "rr-undo-del-20260811-006", "plan_sha256": "8cf7c9be38a76d8ef1fb832aba043d6d7e82b60465b1bf28c6eafa7045117adc", "launch_id": null, "session": null, "result": null }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Review results + fix/retry ledger
|
||||||
|
|
||||||
|
### Round 1 (rr-undo-del-20260811-001, plan sha256 6c52bf99...)
|
||||||
|
- momus: APPROVE (non-blocking notes: todo1+7 duplicate DEFAULT_SETTINGS key -> fixed todo 7 to verify-only; "batch_ids" plural in todos 8/9 acceptance -> fixed; purge OSError note -> folded into todo 1 purge semantics)
|
||||||
|
- independent (oracle): CHANGES_REQUESTED
|
||||||
|
- BLOCKING S1: scanner walks would index .lm-pending-delete staged files as ghost entries -> fixed: todo 1 now mandates scanner walk exclusion at model_scanner.py:706/:867/:1404/_process_model_file + acceptance (o) scanner-visibility test
|
||||||
|
- BLOCKING S2: manifest lacks model_type, undo could restore into wrong cache/hash index -> fixed: manifest now carries model_type + todo 5 resolves per-type scanner via registrar pattern + acceptance (b) checkpoint-batch test
|
||||||
|
- S3 merged-batch expires_at re-anchor -> fixed: merge_batches re-anchors now+TTL in todo 1 + todo 3/4 assertions
|
||||||
|
- S4 manifest-less dir policy -> fixed: quarantine to <batch_id>.orphaned, never delete (todo 1 + acceptance g)
|
||||||
|
- S5 partial-undo retry semantics -> fixed: per-entry restored flag write-through + retry test (acceptance e)
|
||||||
|
- S6 purge locked-file failure semantics -> fixed: skip file, keep batch, never rmtree past errors (todo 1 + acceptance i)
|
||||||
|
- T8 undo-after-restart test -> fixed: todo 5 acceptance (f)
|
||||||
|
- T9 recipe undo -> re-delete test -> fixed: todo 5 acceptance (h)
|
||||||
|
- T7 rescan-stale-entry test -> fixed: todo 5 acceptance (g)
|
||||||
|
- Route registration pinned to shared routes class per mode (NOT per-model-type registrar which registers 3x) -> fixed: todo 5 now creates py/routes/pending_delete_routes.py registered once in lora_manager.py:170-172 + standalone.py:356-358 + duplicate-route test (e)
|
||||||
|
- Version-index staleness on single-delete undo -> fixed: todo 5 follows bulk cache-update pattern incl. rebuild_version_index (model_scanner.py:2324)
|
||||||
|
- Cancelled-bulk batch_id frontend handling -> fixed: todo 9 shows action toast on cancelled+staged-subset
|
||||||
|
- Single-instance assumption -> added to Scope OUT
|
||||||
|
- Occupied-refusal loss UX -> accepted-intent documented in success criteria + modal copy
|
||||||
|
|
||||||
|
### Round 2 (rr-undo-del-20260811-002, plan sha256 f3d52235...)
|
||||||
|
- momus: APPROVE (all 12 round-1 fixes verified present; zero dead references; non-blocking nits only)
|
||||||
|
- independent (oracle): CHANGES_REQUESTED
|
||||||
|
- BLOCK-1: merge_batches file-movement semantics unspecified (silent data-loss vector) -> fixed: todo 1 now specifies move-into-winner-dir + entry re-point + loser-dirs-removed-only-when-empty + abort-on-move-failure (all batches intact) + merge inside service lock + acceptance (k) file-survival assertions + acceptance (l) merge-failure abort test
|
||||||
|
- BLOCK-2: same-file parallel edits within waves (todo 5 vs 6 on lora_manager.py; todo 8 vs 9 on baseModelApi.js) -> fixed: waves/matrix now serialize 5->6 and 8->9 with explicit reasons; matrix updated
|
||||||
|
- Recommended: checkpoint_scanner.py:331 exclusion -> fixed (todo 1 + acceptance p); S5 pre-check skips restored:true entries -> fixed (todo 1); _tags_count restore on undo -> fixed (todo 5 + acceptance j); undo-blind flows documented (ModelVersionsTab + misc_handlers:2456) -> fixed (todo 8 note + Scope OUT); merge-failure no-merge fallback contract (batch_ids array) -> fixed (todos 3/4/9)
|
||||||
|
|
||||||
|
### Round 3 (rr-undo-del-20260811-003, plan sha256 8f2dfd46...)
|
||||||
|
- momus: APPROVE (all round-2 fixes verified present + spot-checked refs; no new contradictions)
|
||||||
|
- independent (oracle): CHANGES_REQUESTED
|
||||||
|
- BLOCKING A: merged batches never timer-purged after re-anchor (winner's original timer no-ops at old expiry; no fresh timer for re-anchored expiry; idle server -> merged batch lingers, violating "30s purge" success criterion; affects EVERY bulk delete) -> fixed: todo 1 merge_batches now ARMS A FRESH PURGE TIMER for the winner with re-anchored expiry + acceptance (q) fresh-timer test + purge_expired must enumerate ALL scanner types' roots (explicit in todo 1)
|
||||||
|
- BLOCKING B: dependency matrix contradicted same-file policy for todos 8/9<->11 (5 shared files) and 12<->11 -> fixed: todo 11 now "Blocked by: 8, 9 (same files...)"; todo 12 blocked by 11 (sync after 11); wave text updated (11, then 12 AFTER 11); "Can parallelize with" columns corrected
|
||||||
|
- BLOCKING C: frontend batch_ids sequential-undo fallback has NO test + merge->undo loser-restore + merge->purge assertions missing -> fixed: todo 9 acceptance now tests the batch_ids fallback path; todo 1 acceptance now has (k2)/(k3)
|
||||||
|
- Notes folded: sub-second toast-tail expiry race accepted; EXDEV fallback = NORMAL path for cross-volume bulks [annotated 2026-08: after the sibling-staging fix, EXDEV can only arise during cross-volume MERGES, never during single stage/undo renames]
|
||||||
|
|
||||||
|
### Round 4 (rr-undo-del-20260811-004, plan sha256 179e7ff7...)
|
||||||
|
- momus: APPROVE (round-3 fixes verified; one non-blocking nit: todo 11 inline "Blocked by: —" stale -> fixed to "8, 9")
|
||||||
|
- independent (oracle): CHANGES_REQUESTED
|
||||||
|
- BLOCKING GAP-1 (NEW, introduced by round-3 fix): todo 8 handleUndoDelete always-refresh/always-toast contract contradicted todo 9's sequential loop "exactly ONE final refresh" -> fixed: handleUndoDelete(batchId, refreshFn, {showToast, refresh}) suppression options; todo 9 loop uses suppressed calls + one final refresh/toast; acceptance extended (loop failure mid-way -> stop + error toast + no final refresh; 404 body discrimination expired vs occupied)
|
||||||
|
- BLOCKING GAP-2: no cross-type purge enumeration test -> fixed: todo 1 acceptance (r) purges expired batches across lora root + checkpoint root + recipe staging dir in one call
|
||||||
|
- Non-blocking folded: GAP-3 404-copy discrimination -> fixed in todo 8 (d); GAP-4 merge partial-failure rollback direction (move back + restore manifests, extended (l) asserts sequential constituent undo still restores everything) -> fixed in todo 1; GAP-5 post-restart timer-loss residual gap documented -> fixed in todo 6; GAP-6 usage_stats.py:424 walk added to exclusion mandate + todo 5 acceptance (k) embeddings undo test
|
||||||
|
|
||||||
|
### Round 5 (rr-undo-del-20260811-005, plan sha256 dfaa39ea...)
|
||||||
|
- momus: APPROVE (all round-4 fixes verified; no new contradictions)
|
||||||
|
- independent (oracle): CHANGES_REQUESTED
|
||||||
|
- BLOCK-1: lock-ordering deadlock ambiguity (asyncio.Lock not re-entrant: opportunistic purge_expired called while stage/undo hold the lock would deadlock on first use) -> fixed: todo 1 now has explicit LOCK HIERARCHY (lock acquired ONLY by stage/merge/undo/purge_batch; purge_expired is lock-free and must be called BEFORE lock acquisition); todo 6 (c) updated with the same rule + acceptance (u) lock-no-deadlock test
|
||||||
|
- BLOCK-2: purge edge semantics unspecified -> fixed: purge_batch treats missing staged files (partially-restored batches) as already-purged (FileNotFoundError silent no-op); sweep skips `.orphaned`-suffixed dirs (quarantine is terminal); acceptance (s) partially-restored purge + (t) quarantine-terminal tests
|
||||||
|
- Non-blocking folded: todo 2/3 test-file collision -> todo 3's bulk tests moved to tests/services/test_model_scanner.py; todo 9 (d) DuplicatesManager refreshFn stated explicitly (recipes loadRecipes / models resetAndReload); modal-copy + bulk-count trade-offs acknowledged in success criteria; acceptance (r) extended with embeddings root
|
||||||
|
|
||||||
|
### Round 6 (rr-undo-del-20260811-006, plan sha256 8cf7c9be...)
|
||||||
|
- momus: APPROVE (all round-5 fixes verified; no new contradictions; references verified)
|
||||||
|
- independent (oracle): APPROVE — no blocking issues; all round-5 items fixed with working, tested solutions; no new race/data-loss/consistency defects
|
||||||
|
- Deferred optional improvements (non-blocking, recorded for executor awareness; plan file left untouched to preserve the approved digest):
|
||||||
|
1. Tag-count asymmetry: single delete_model never decrements _tags_count (lifecycle 101-154), bulk does (scanner 2297-2303); undo re-increment is exact for bulk, over-counts for single until rescan (cosmetic, self-healing). Optional fix riding in todo 2: decrement tags in the single-delete path to mirror bulk.
|
||||||
|
2. Todo 5 factual nit: ModelCache.resort() already rebuilds the version index — explicit rebuild in undo is belt-and-braces, no action needed.
|
||||||
|
3. Todo 8 premise nit: ModelVersionsTab call ignores deleteModel's return entirely — nothing breaks, no adaptation needed.
|
||||||
|
4. Todo 3's pytest command includes test_model_lifecycle_service.py which todo 2 edits in the same wave — run that file's tests after todo 2 lands.
|
||||||
|
5. merge_batches with a missing/quarantined constituent id: any sane fallback (abort -> batch_ids, or skip missing) acceptable — files stay staged either way.
|
||||||
|
|
||||||
|
## Review lifecycle
|
||||||
|
- rounds: 6 (rr-undo-del-20260811-001..006); final round both lanes APPROVE
|
||||||
|
- final live-plan validation: sha256 = 8cf7c9be38a76d8ef1fb832aba043d6d7e82b60465b1bf28c6eafa7045117adc — MATCHES approved round-6 digest
|
||||||
|
- status: APPROVED — ready for execution handoff ($start-work undo-delete-staging)
|
||||||
File diff suppressed because one or more lines are too long
@@ -31,7 +31,7 @@ COVERAGE_FILE=coverage/backend/.coverage pytest \
|
|||||||
--cov-report=xml:coverage/backend/coverage.xml
|
--cov-report=xml:coverage/backend/coverage.xml
|
||||||
```
|
```
|
||||||
|
|
||||||
### Frontend Development (Standalone Web UI)
|
### Frontend Development (LoRA Manager Web UI)
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
npm install
|
npm install
|
||||||
@@ -137,7 +137,13 @@ npm run test:coverage # Generate coverage report
|
|||||||
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
||||||
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
||||||
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
|
- Run `python scripts/sync_translation_keys.py` after adding UI strings to `locales/en.json`
|
||||||
- Symlinks require normalized paths
|
- Symlinks require normalized paths.
|
||||||
|
**Business paths vs real paths**: All stored paths and operation routing use the
|
||||||
|
original paths as they appear under configured model roots — symlinks are NOT
|
||||||
|
resolved. `os.path.realpath` is only for scanner dedup and the symlink cache.
|
||||||
|
Any path passed to `os.remove`/`os.rename`/`shutil.move` or validated by a
|
||||||
|
containment check MUST use the business path (i.e. `os.path.abspath`, not
|
||||||
|
`realpath`).
|
||||||
|
|
||||||
## Git / Commit Messages
|
## Git / Commit Messages
|
||||||
|
|
||||||
@@ -148,9 +154,9 @@ npm run test:coverage # Generate coverage report
|
|||||||
|
|
||||||
## Frontend UI Architecture
|
## Frontend UI Architecture
|
||||||
|
|
||||||
### 1. Standalone Web UI
|
### 1. LoRA Manager Web UI
|
||||||
- Location: `./static/` and `./templates/`
|
- Location: `./static/` and `./templates/`
|
||||||
- Tech: Vanilla JS + CSS, served by standalone server
|
- Tech: Vanilla JS + CSS, served by the hosting server (ComfyUI app in plugin mode, `standalone.py` in standalone mode)
|
||||||
- Tests via npm in root directory
|
- Tests via npm in root directory
|
||||||
|
|
||||||
### 2. ComfyUI Custom Node Widgets
|
### 2. ComfyUI Custom Node Widgets
|
||||||
|
|||||||
+20
@@ -3,6 +3,8 @@ try: # pragma: no cover - import fallback for pytest collection
|
|||||||
from .py.nodes.lora_loader import LoraLoaderLM, LoraTextLoaderLM
|
from .py.nodes.lora_loader import LoraLoaderLM, LoraTextLoaderLM
|
||||||
from .py.nodes.checkpoint_loader import CheckpointLoaderLM
|
from .py.nodes.checkpoint_loader import CheckpointLoaderLM
|
||||||
from .py.nodes.unet_loader import UNETLoaderLM
|
from .py.nodes.unet_loader import UNETLoaderLM
|
||||||
|
from .py.nodes.random_checkpoint_loader import RandomCheckpointLoaderLM
|
||||||
|
from .py.nodes.random_unet_loader import RandomUNETLoaderLM
|
||||||
from .py.nodes.trigger_word_toggle import TriggerWordToggleLM
|
from .py.nodes.trigger_word_toggle import TriggerWordToggleLM
|
||||||
from .py.nodes.prompt import PromptLM
|
from .py.nodes.prompt import PromptLM
|
||||||
from .py.nodes.text import TextLM
|
from .py.nodes.text import TextLM
|
||||||
@@ -17,6 +19,8 @@ try: # pragma: no cover - import fallback for pytest collection
|
|||||||
from .py.nodes.lora_cycler import LoraCyclerLM
|
from .py.nodes.lora_cycler import LoraCyclerLM
|
||||||
from .py.nodes.lora_info import LoraInfoLM
|
from .py.nodes.lora_info import LoraInfoLM
|
||||||
from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath
|
from .py.nodes.lora_syntax_to_path import LoraSyntaxToPath
|
||||||
|
from .py.nodes.create_hook_lora import CreateHookLoraLM
|
||||||
|
from .py.nodes.metadata_overwrite import MetadataOverwriteLM
|
||||||
from .py.metadata_collector import init as init_metadata_collector
|
from .py.metadata_collector import init as init_metadata_collector
|
||||||
except (
|
except (
|
||||||
ImportError
|
ImportError
|
||||||
@@ -38,6 +42,12 @@ except (
|
|||||||
"py.nodes.checkpoint_loader"
|
"py.nodes.checkpoint_loader"
|
||||||
).CheckpointLoaderLM
|
).CheckpointLoaderLM
|
||||||
UNETLoaderLM = importlib.import_module("py.nodes.unet_loader").UNETLoaderLM
|
UNETLoaderLM = importlib.import_module("py.nodes.unet_loader").UNETLoaderLM
|
||||||
|
RandomCheckpointLoaderLM = importlib.import_module(
|
||||||
|
"py.nodes.random_checkpoint_loader"
|
||||||
|
).RandomCheckpointLoaderLM
|
||||||
|
RandomUNETLoaderLM = importlib.import_module(
|
||||||
|
"py.nodes.random_unet_loader"
|
||||||
|
).RandomUNETLoaderLM
|
||||||
TriggerWordToggleLM = importlib.import_module(
|
TriggerWordToggleLM = importlib.import_module(
|
||||||
"py.nodes.trigger_word_toggle"
|
"py.nodes.trigger_word_toggle"
|
||||||
).TriggerWordToggleLM
|
).TriggerWordToggleLM
|
||||||
@@ -62,6 +72,12 @@ except (
|
|||||||
LoraSyntaxToPath = importlib.import_module(
|
LoraSyntaxToPath = importlib.import_module(
|
||||||
"py.nodes.lora_syntax_to_path"
|
"py.nodes.lora_syntax_to_path"
|
||||||
).LoraSyntaxToPath
|
).LoraSyntaxToPath
|
||||||
|
CreateHookLoraLM = importlib.import_module(
|
||||||
|
"py.nodes.create_hook_lora"
|
||||||
|
).CreateHookLoraLM
|
||||||
|
MetadataOverwriteLM = importlib.import_module(
|
||||||
|
"py.nodes.metadata_overwrite"
|
||||||
|
).MetadataOverwriteLM
|
||||||
init_metadata_collector = importlib.import_module("py.metadata_collector").init
|
init_metadata_collector = importlib.import_module("py.metadata_collector").init
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
@@ -71,6 +87,8 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
LoraTextLoaderLM.NAME: LoraTextLoaderLM,
|
LoraTextLoaderLM.NAME: LoraTextLoaderLM,
|
||||||
CheckpointLoaderLM.NAME: CheckpointLoaderLM,
|
CheckpointLoaderLM.NAME: CheckpointLoaderLM,
|
||||||
UNETLoaderLM.NAME: UNETLoaderLM,
|
UNETLoaderLM.NAME: UNETLoaderLM,
|
||||||
|
RandomCheckpointLoaderLM.NAME: RandomCheckpointLoaderLM,
|
||||||
|
RandomUNETLoaderLM.NAME: RandomUNETLoaderLM,
|
||||||
TriggerWordToggleLM.NAME: TriggerWordToggleLM,
|
TriggerWordToggleLM.NAME: TriggerWordToggleLM,
|
||||||
LoraStackerLM.NAME: LoraStackerLM,
|
LoraStackerLM.NAME: LoraStackerLM,
|
||||||
LoraStackCombinerLM.NAME: LoraStackCombinerLM,
|
LoraStackCombinerLM.NAME: LoraStackCombinerLM,
|
||||||
@@ -83,6 +101,8 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
LoraCyclerLM.NAME: LoraCyclerLM,
|
LoraCyclerLM.NAME: LoraCyclerLM,
|
||||||
LoraInfoLM.NAME: LoraInfoLM,
|
LoraInfoLM.NAME: LoraInfoLM,
|
||||||
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
LoraSyntaxToPath.NAME: LoraSyntaxToPath,
|
||||||
|
CreateHookLoraLM.NAME: CreateHookLoraLM,
|
||||||
|
MetadataOverwriteLM.NAME: MetadataOverwriteLM,
|
||||||
}
|
}
|
||||||
|
|
||||||
WEB_DIRECTORY = "./web/comfyui"
|
WEB_DIRECTORY = "./web/comfyui"
|
||||||
|
|||||||
+361
-307
File diff suppressed because it is too large
Load Diff
@@ -39,6 +39,7 @@ These fields are present in all model metadata files.
|
|||||||
| `metadata_source` | string\|null | ❌ No | ✅ Yes | Last provider that supplied metadata (see below) |
|
| `metadata_source` | string\|null | ❌ No | ✅ Yes | Last provider that supplied metadata (see below) |
|
||||||
| `last_checked_at` | float | ❌ No (default: `0`) | ✅ Yes | Unix timestamp of last metadata check |
|
| `last_checked_at` | float | ❌ No (default: `0`) | ✅ Yes | Unix timestamp of last metadata check |
|
||||||
| `hash_status` | string | ❌ No (default: `"completed"`) | ✅ Yes | Hash calculation status: `"pending"`, `"calculating"`, `"completed"`, `"failed"` |
|
| `hash_status` | string | ❌ No (default: `"completed"`) | ✅ Yes | Hash calculation status: `"pending"`, `"calculating"`, `"completed"`, `"failed"` |
|
||||||
|
| `autov3` | string\|null | ❌ No | ✅ Yes | CivitAI AutoV3 hash (first 12 chars, lowercase hex) sourced from the safetensors embedded metadata (`sshs_model_hash` / `modelspec.hash_sha256`). **Absent** = not yet checked (may be backfilled later); **`null`** = checked but unavailable (header has no recognized hash); **12-char hex string** = value |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -287,6 +288,7 @@ These fields are automatically synchronized with the filesystem:
|
|||||||
- `preview_url` — Updated if preview file is moved/removed
|
- `preview_url` — Updated if preview file is moved/removed
|
||||||
- `sha256` — Updated during hash calculation (when `hash_status="pending"`)
|
- `sha256` — Updated during hash calculation (when `hash_status="pending"`)
|
||||||
- `hash_status` — Updated during hash calculation
|
- `hash_status` — Updated during hash calculation
|
||||||
|
- `autov3` — Set when metadata is first created (from safetensors header); may be backfilled later for entries where it is absent
|
||||||
- `last_checked_at` — Timestamp of scan
|
- `last_checked_at` — Timestamp of scan
|
||||||
- `metadata_source` — Set based on metadata provider
|
- `metadata_source` — Set based on metadata provider
|
||||||
|
|
||||||
@@ -345,6 +347,7 @@ These fields can be edited by users at any time through the Lora Manager UI or b
|
|||||||
| `metadata_source` | `null` |
|
| `metadata_source` | `null` |
|
||||||
| `last_checked_at` | `0` |
|
| `last_checked_at` | `0` |
|
||||||
| `hash_status` | `"completed"` |
|
| `hash_status` | `"completed"` |
|
||||||
|
| `autov3` | absent (not checked) or `null` (checked, no value) |
|
||||||
| `usage_tips` | `"{}"` (LoRA only) |
|
| `usage_tips` | `"{}"` (LoRA only) |
|
||||||
| `model_type` | `"checkpoint"` or `"embedding"` (not present in LoRA models) |
|
| `model_type` | `"checkpoint"` or `"embedding"` (not present in LoRA models) |
|
||||||
|
|
||||||
@@ -354,6 +357,7 @@ These fields can be edited by users at any time through the Lora Manager UI or b
|
|||||||
|
|
||||||
| Version | Date | Changes |
|
| Version | Date | Changes |
|
||||||
|---------|------|---------|
|
|---------|------|---------|
|
||||||
|
| 1.1 | 2026-08 | Added `autov3` field (CivitAI AutoV3 hash with three-state semantics) |
|
||||||
| 1.0 | 2026-03 | Initial schema documentation |
|
| 1.0 | 2026-03 | Initial schema documentation |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|||||||
+2303
-2205
File diff suppressed because it is too large
Load Diff
+103
-5
@@ -186,6 +186,16 @@
|
|||||||
"cancelled": "Repair cancelled. {count} recipes were repaired.",
|
"cancelled": "Repair cancelled. {count} recipes were repaired.",
|
||||||
"error": "Recipe repair failed: {message}"
|
"error": "Recipe repair failed: {message}"
|
||||||
},
|
},
|
||||||
|
"rematchRecipes": {
|
||||||
|
"label": "Rematch recipes to local models",
|
||||||
|
"loading": "Rematching recipes to local models...",
|
||||||
|
"success": "Matched {entries} entries across {recipes} recipes",
|
||||||
|
"successErrors": "Matched {entries} entries across {recipes} recipes, {failures} failed",
|
||||||
|
"allFailed": "Rematch failed for {failures} of {total} recipes",
|
||||||
|
"noMatch": "No local match found for {entries} entries in {recipes} recipes",
|
||||||
|
"cancelled": "Rematch cancelled. {recipes} recipes updated ({entries} entries).",
|
||||||
|
"error": "Recipe rematch failed: {message}"
|
||||||
|
},
|
||||||
"manageExcludedModels": {
|
"manageExcludedModels": {
|
||||||
"label": "Manage Excluded Models"
|
"label": "Manage Excluded Models"
|
||||||
},
|
},
|
||||||
@@ -449,6 +459,12 @@
|
|||||||
"compact": "7 (1080p), 8 (2K), 10 (4K)"
|
"compact": "7 (1080p), 8 (2K), 10 (4K)"
|
||||||
},
|
},
|
||||||
"displayDensityWarning": "Warning: Higher densities may cause performance issues on systems with limited resources.",
|
"displayDensityWarning": "Warning: Higher densities may cause performance issues on systems with limited resources.",
|
||||||
|
"recipesLayout": "Recipes Layout",
|
||||||
|
"recipesLayoutHelp": "Choose how recipe cards are arranged: a uniform grid or a masonry (Pinterest-style) layout that preserves each image's aspect ratio.",
|
||||||
|
"recipesLayoutOptions": {
|
||||||
|
"grid": "Grid",
|
||||||
|
"masonry": "Masonry"
|
||||||
|
},
|
||||||
"showFolderSidebar": "Show Folder Sidebar",
|
"showFolderSidebar": "Show Folder Sidebar",
|
||||||
"showFolderSidebarHelp": "Toggle the folder navigation sidebar on model pages. When disabled, the sidebar and hover area stay hidden.",
|
"showFolderSidebarHelp": "Toggle the folder navigation sidebar on model pages. When disabled, the sidebar and hover area stay hidden.",
|
||||||
"cardInfoDisplay": "Card Info Display",
|
"cardInfoDisplay": "Card Info Display",
|
||||||
@@ -606,6 +622,10 @@
|
|||||||
"label": "Hide Early Access Updates",
|
"label": "Hide Early Access Updates",
|
||||||
"help": "When enabled, models with only early access updates will not show 'Update available' badge"
|
"help": "When enabled, models with only early access updates will not show 'Update available' badge"
|
||||||
},
|
},
|
||||||
|
"hidePaidUpdates": {
|
||||||
|
"label": "Hide Paid Updates",
|
||||||
|
"help": "When enabled, models with only paid updates will not show 'Update available' badge"
|
||||||
|
},
|
||||||
"licenseIcons": {
|
"licenseIcons": {
|
||||||
"useNewStyle": "Use updated license icons",
|
"useNewStyle": "Use updated license icons",
|
||||||
"useNewStyleHelp": "Display license permissions with colored indicators (new style) or restriction-only icons (classic style). Mirroring the current CivitAI design."
|
"useNewStyleHelp": "Display license permissions with colored indicators (new style) or restriction-only icons (classic style). Mirroring the current CivitAI design."
|
||||||
@@ -678,6 +698,7 @@
|
|||||||
"deepseek": "DeepSeek",
|
"deepseek": "DeepSeek",
|
||||||
"groq": "Groq",
|
"groq": "Groq",
|
||||||
"openrouter": "OpenRouter",
|
"openrouter": "OpenRouter",
|
||||||
|
"google": "Gemini",
|
||||||
"opencode-go": "OpenCode Go",
|
"opencode-go": "OpenCode Go",
|
||||||
"custom": "Custom (OpenAI-compatible)"
|
"custom": "Custom (OpenAI-compatible)"
|
||||||
},
|
},
|
||||||
@@ -714,7 +735,9 @@
|
|||||||
"versionsCount": "Local Versions",
|
"versionsCount": "Local Versions",
|
||||||
"versionsCountDesc": "Most versions first",
|
"versionsCountDesc": "Most versions first",
|
||||||
"versionsCountAsc": "Fewest versions first",
|
"versionsCountAsc": "Fewest versions first",
|
||||||
"versionIdDesc": "Newest version first"
|
"versionIdDesc": "Newest version first",
|
||||||
|
"random": "Random",
|
||||||
|
"randomAction": "Randomize (shuffle)"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "Refresh model list",
|
"title": "Refresh model list",
|
||||||
@@ -759,6 +782,7 @@
|
|||||||
"copyAll": "Copy Selected Syntax",
|
"copyAll": "Copy Selected Syntax",
|
||||||
"refreshAll": "Refresh Selected Metadata",
|
"refreshAll": "Refresh Selected Metadata",
|
||||||
"repairMetadata": "Repair Metadata for Selected",
|
"repairMetadata": "Repair Metadata for Selected",
|
||||||
|
"rematchMetadata": "Rematch Selected to Local Models",
|
||||||
"reimportMetadata": "Re-import from Source",
|
"reimportMetadata": "Re-import from Source",
|
||||||
"checkUpdates": "Check Updates for Selected",
|
"checkUpdates": "Check Updates for Selected",
|
||||||
"moveAll": "Move Selected to Folder",
|
"moveAll": "Move Selected to Folder",
|
||||||
@@ -771,6 +795,8 @@
|
|||||||
"deleteAll": "Delete Selected",
|
"deleteAll": "Delete Selected",
|
||||||
"downloadMissingLoras": "Download Missing LoRAs",
|
"downloadMissingLoras": "Download Missing LoRAs",
|
||||||
"downloadExamples": "Download Example Images",
|
"downloadExamples": "Download Example Images",
|
||||||
|
"downloadMissingExamples": "Download Missing",
|
||||||
|
"reprocessExamples": "Re-process All",
|
||||||
"clear": "Clear Selection",
|
"clear": "Clear Selection",
|
||||||
"skipMetadataRefreshCount": "Skip ({count} models)",
|
"skipMetadataRefreshCount": "Skip ({count} models)",
|
||||||
"resumeMetadataRefreshCount": "Resume ({count} models)",
|
"resumeMetadataRefreshCount": "Resume ({count} models)",
|
||||||
@@ -806,10 +832,13 @@
|
|||||||
"sendToWorkflowReplace": "Send to Workflow (Replace)",
|
"sendToWorkflowReplace": "Send to Workflow (Replace)",
|
||||||
"openExamples": "Open Examples Folder",
|
"openExamples": "Open Examples Folder",
|
||||||
"downloadExamples": "Download Example Images",
|
"downloadExamples": "Download Example Images",
|
||||||
|
"downloadMissingExamples": "Download Missing",
|
||||||
|
"reprocessExamples": "Re-process All",
|
||||||
"replacePreview": "Replace Preview",
|
"replacePreview": "Replace Preview",
|
||||||
"setContentRating": "Set Content Rating",
|
"setContentRating": "Set Content Rating",
|
||||||
"moveToFolder": "Move to Folder",
|
"moveToFolder": "Move to Folder",
|
||||||
"repairMetadata": "Repair metadata",
|
"repairMetadata": "Repair metadata",
|
||||||
|
"rematchMetadata": "Rematch to local models",
|
||||||
"reimportMetadata": "Re-import from Source",
|
"reimportMetadata": "Re-import from Source",
|
||||||
"excludeModel": "Exclude Model",
|
"excludeModel": "Exclude Model",
|
||||||
"restoreModel": "Restore Model",
|
"restoreModel": "Restore Model",
|
||||||
@@ -895,7 +924,9 @@
|
|||||||
"dateAsc": "Oldest",
|
"dateAsc": "Oldest",
|
||||||
"lorasCount": "LoRA Count",
|
"lorasCount": "LoRA Count",
|
||||||
"lorasCountDesc": "Most",
|
"lorasCountDesc": "Most",
|
||||||
"lorasCountAsc": "Least"
|
"lorasCountAsc": "Least",
|
||||||
|
"opened": "Recently Opened",
|
||||||
|
"openedDesc": "Recently opened"
|
||||||
},
|
},
|
||||||
"refresh": {
|
"refresh": {
|
||||||
"title": "Refresh recipe list",
|
"title": "Refresh recipe list",
|
||||||
@@ -906,12 +937,25 @@
|
|||||||
"favorites": {
|
"favorites": {
|
||||||
"title": "Show Favorites Only",
|
"title": "Show Favorites Only",
|
||||||
"action": "Favorites"
|
"action": "Favorites"
|
||||||
|
},
|
||||||
|
"layout": {
|
||||||
|
"title": "Recipes Layout",
|
||||||
|
"grid": "Grid layout",
|
||||||
|
"masonry": "Masonry layout (Pinterest-style, preserves image aspect ratio)"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"duplicates": {
|
"duplicates": {
|
||||||
"found": "Found {count} duplicate groups",
|
"found": "Found {count} duplicate groups",
|
||||||
|
"noGroups": "No duplicate groups found with the current matching basis",
|
||||||
"keepLatest": "Keep Latest Versions",
|
"keepLatest": "Keep Latest Versions",
|
||||||
"deleteSelected": "Delete Selected"
|
"deleteSelected": "Delete Selected",
|
||||||
|
"includePromptLabel": "Include prompt in matching",
|
||||||
|
"basis": {
|
||||||
|
"loraCombo": "Matched by: LoRA combination",
|
||||||
|
"loraComboAndPrompt": "Matched by: LoRA combination + prompt",
|
||||||
|
"hintLoraCombo": "Recipes with the same LoRAs at identical strengths are grouped.",
|
||||||
|
"hintPromptIncluded": "Recipes are grouped only when they use the same LoRAs at identical strengths AND have the same prompt."
|
||||||
|
}
|
||||||
},
|
},
|
||||||
"contextMenu": {
|
"contextMenu": {
|
||||||
"copyRecipe": {
|
"copyRecipe": {
|
||||||
@@ -1244,8 +1288,13 @@
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"deleteModel": {
|
"deleteModel": {
|
||||||
|
"freesSpace": "Frees {size}",
|
||||||
"title": "Delete Model",
|
"title": "Delete Model",
|
||||||
"message": "Are you sure you want to delete this model and all associated files?"
|
"message": "Are you sure you want to delete this model and all associated files?",
|
||||||
|
"recoverableWarning": "This will permanently delete the file after 20 seconds unless you undo."
|
||||||
|
},
|
||||||
|
"deleteRecipe": {
|
||||||
|
"recoverableWarning": "This action can be undone for 20 seconds."
|
||||||
},
|
},
|
||||||
"excludeModel": {
|
"excludeModel": {
|
||||||
"title": "Exclude Model",
|
"title": "Exclude Model",
|
||||||
@@ -1510,6 +1559,8 @@
|
|||||||
"newerTooltip": "This version is newer than your latest local version",
|
"newerTooltip": "This version is newer than your latest local version",
|
||||||
"earlyAccess": "Early Access",
|
"earlyAccess": "Early Access",
|
||||||
"earlyAccessTooltip": "This version currently requires Civitai early access",
|
"earlyAccessTooltip": "This version currently requires Civitai early access",
|
||||||
|
"paid": "Paid",
|
||||||
|
"paidTooltip": "This version requires payment to download",
|
||||||
"ignored": "Ignored",
|
"ignored": "Ignored",
|
||||||
"ignoredTooltip": "Update notifications are disabled for this version",
|
"ignoredTooltip": "Update notifications are disabled for this version",
|
||||||
"onSiteOnly": "On-Site Only",
|
"onSiteOnly": "On-Site Only",
|
||||||
@@ -1519,6 +1570,7 @@
|
|||||||
"download": "Download",
|
"download": "Download",
|
||||||
"downloadTooltip": "Download this version",
|
"downloadTooltip": "Download this version",
|
||||||
"downloadEarlyAccessTooltip": "Download this early access version from Civitai",
|
"downloadEarlyAccessTooltip": "Download this early access version from Civitai",
|
||||||
|
"downloadPaidTooltip": "Download this paid version from Civitai",
|
||||||
"downloadNotAllowedTooltip": "This version is only available for on-site generation on Civitai",
|
"downloadNotAllowedTooltip": "This version is only available for on-site generation on Civitai",
|
||||||
"delete": "Delete",
|
"delete": "Delete",
|
||||||
"deleteTooltip": "Delete this local version",
|
"deleteTooltip": "Delete this local version",
|
||||||
@@ -1548,6 +1600,7 @@
|
|||||||
"empty": "No version history available for this model yet.",
|
"empty": "No version history available for this model yet.",
|
||||||
"error": "Failed to load versions.",
|
"error": "Failed to load versions.",
|
||||||
"missingModelId": "This model is missing a Civitai model id.",
|
"missingModelId": "This model is missing a Civitai model id.",
|
||||||
|
"hfGroupInfo": "This is a HuggingFace model group. Open the library to see all versions in the grid.",
|
||||||
"confirm": {
|
"confirm": {
|
||||||
"delete": "Delete this version from your library?"
|
"delete": "Delete this version from your library?"
|
||||||
},
|
},
|
||||||
@@ -1574,6 +1627,21 @@
|
|||||||
"downloadCsv": "Download CSV",
|
"downloadCsv": "Download CSV",
|
||||||
"columnModelName": "Model Name",
|
"columnModelName": "Model Name",
|
||||||
"columnError": "Error"
|
"columnError": "Error"
|
||||||
|
},
|
||||||
|
"downloadBatchSummary": {
|
||||||
|
"title": "Batch Download Summary",
|
||||||
|
"statSuccess": "Success",
|
||||||
|
"statFailed": "Failed",
|
||||||
|
"statTotal": "Total",
|
||||||
|
"successMessage": "All {count} models downloaded successfully",
|
||||||
|
"completedWithErrors": "Completed with errors",
|
||||||
|
"failed": "Download failed",
|
||||||
|
"failedItems": "Failed Items ({count})",
|
||||||
|
"columnName": "Model Name",
|
||||||
|
"columnError": "Error",
|
||||||
|
"close": "Close",
|
||||||
|
"copyReport": "Copy Report",
|
||||||
|
"retryFailed": "Retry Failed ({count})"
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"modelTags": {
|
"modelTags": {
|
||||||
@@ -1672,6 +1740,7 @@
|
|||||||
"recipeReplaced": "Recipe replaced in workflow",
|
"recipeReplaced": "Recipe replaced in workflow",
|
||||||
"recipeFailedToSend": "Failed to send recipe to workflow",
|
"recipeFailedToSend": "Failed to send recipe to workflow",
|
||||||
"noMatchingNodes": "No compatible nodes available in the current workflow",
|
"noMatchingNodes": "No compatible nodes available in the current workflow",
|
||||||
|
"noPromptTargets": "No compatible prompt targets in the workflow.\nRight-click a node in ComfyUI → Mark as → Send Prompt Target",
|
||||||
"noTargetNodeSelected": "No target node selected",
|
"noTargetNodeSelected": "No target node selected",
|
||||||
"modelUpdated": "Model updated in workflow",
|
"modelUpdated": "Model updated in workflow",
|
||||||
"modelFailed": "Failed to update model node",
|
"modelFailed": "Failed to update model node",
|
||||||
@@ -1751,6 +1820,12 @@
|
|||||||
"checkingMessage": "Please wait while we check for the latest version.",
|
"checkingMessage": "Please wait while we check for the latest version.",
|
||||||
"showNotifications": "Show update notifications",
|
"showNotifications": "Show update notifications",
|
||||||
"latestBadge": "Latest",
|
"latestBadge": "Latest",
|
||||||
|
"latestMain": "Latest main",
|
||||||
|
"channel": "Update Channel",
|
||||||
|
"channels": {
|
||||||
|
"release": "Release",
|
||||||
|
"nightly": "Nightly"
|
||||||
|
},
|
||||||
"updateProgress": {
|
"updateProgress": {
|
||||||
"preparing": "Preparing update...",
|
"preparing": "Preparing update...",
|
||||||
"installing": "Installing update...",
|
"installing": "Installing update...",
|
||||||
@@ -1771,6 +1846,15 @@
|
|||||||
"warning": "Warning: Nightly builds may contain experimental features and could be unstable.",
|
"warning": "Warning: Nightly builds may contain experimental features and could be unstable.",
|
||||||
"enable": "Enable Nightly Updates"
|
"enable": "Enable Nightly Updates"
|
||||||
},
|
},
|
||||||
|
"channelSwitch": {
|
||||||
|
"nightlyTitle": "Switch to Nightly Channel",
|
||||||
|
"nightlyMessage": "Switching to Nightly will initialize a Git repository and track the latest main branch commits. Updates will be more frequent but may be unstable. You can switch back to Release at any time.",
|
||||||
|
"releaseTitle": "Switch to Release Channel",
|
||||||
|
"releaseMessage": "Switching to Release will checkout the latest stable release tag. You can switch back to Nightly at any time.",
|
||||||
|
"switching": "Switching to {channel} channel...",
|
||||||
|
"completed": "Successfully switched to {channel} channel",
|
||||||
|
"failed": "Failed to switch channel"
|
||||||
|
},
|
||||||
"banners": {
|
"banners": {
|
||||||
"recent": "Recent messages",
|
"recent": "Recent messages",
|
||||||
"empty": "No recent banners yet.",
|
"empty": "No recent banners yet.",
|
||||||
@@ -1907,6 +1991,12 @@
|
|||||||
"repairBulkComplete": "Repair complete: {repaired} repaired, {skipped} skipped (of {total})",
|
"repairBulkComplete": "Repair complete: {repaired} repaired, {skipped} skipped (of {total})",
|
||||||
"repairBulkSkipped": "No repair needed for any of the {total} selected recipes",
|
"repairBulkSkipped": "No repair needed for any of the {total} selected recipes",
|
||||||
"repairBulkFailed": "Failed to repair selected recipes: {message}",
|
"repairBulkFailed": "Failed to repair selected recipes: {message}",
|
||||||
|
"rematchComplete": "Matched {entries} entries across {recipes} recipes",
|
||||||
|
"rematchCompleteErrors": "Matched {entries} entries across {recipes} recipes, {failures} failed",
|
||||||
|
"rematchAllFailed": "Rematch failed for {failures} of {total} selected recipes",
|
||||||
|
"rematchUnmatched": "No local match found for {entries} entries in {recipes} recipes",
|
||||||
|
"rematchSkipped": "No rematch needed for any of the {total} selected recipes",
|
||||||
|
"rematchFailed": "Failed to rematch selected recipes: {message}",
|
||||||
"reimporting": "Re-importing recipe from source...",
|
"reimporting": "Re-importing recipe from source...",
|
||||||
"reimportSuccess": "Recipe re-imported successfully",
|
"reimportSuccess": "Recipe re-imported successfully",
|
||||||
"reimportBulkComplete": "Re-import complete: {completed} re-imported, {failed} failed (of {total})",
|
"reimportBulkComplete": "Re-import complete: {completed} re-imported, {failed} failed (of {total})",
|
||||||
@@ -2015,7 +2105,6 @@
|
|||||||
"presetNameTooLong": "Preset name must be {max} characters or less",
|
"presetNameTooLong": "Preset name must be {max} characters or less",
|
||||||
"presetNameInvalidChars": "Preset name contains invalid characters",
|
"presetNameInvalidChars": "Preset name contains invalid characters",
|
||||||
"presetNameExists": "A preset with this name already exists",
|
"presetNameExists": "A preset with this name already exists",
|
||||||
"maxPresetsReached": "Maximum {max} presets allowed. Delete one to add more.",
|
|
||||||
"presetNotFound": "Preset not found",
|
"presetNotFound": "Preset not found",
|
||||||
"invalidPreset": "Invalid preset data",
|
"invalidPreset": "Invalid preset data",
|
||||||
"deletePresetFailed": "Failed to delete preset",
|
"deletePresetFailed": "Failed to delete preset",
|
||||||
@@ -2044,6 +2133,14 @@
|
|||||||
"updateFailed": "Failed to update trigger words",
|
"updateFailed": "Failed to update trigger words",
|
||||||
"copyFailed": "Copy failed"
|
"copyFailed": "Copy failed"
|
||||||
},
|
},
|
||||||
|
"undo": {
|
||||||
|
"action": "Undo",
|
||||||
|
"deleted": "Deleted {name}",
|
||||||
|
"deletedBulk": "Deleted {count} item(s)",
|
||||||
|
"expired": "Undo window expired. The item was permanently deleted.",
|
||||||
|
"failed": "Undo failed: {error}",
|
||||||
|
"restored": "Item restored"
|
||||||
|
},
|
||||||
"virtual": {
|
"virtual": {
|
||||||
"loadFailed": "Failed to load items",
|
"loadFailed": "Failed to load items",
|
||||||
"loadMoreFailed": "Failed to load more items",
|
"loadMoreFailed": "Failed to load more items",
|
||||||
@@ -2107,6 +2204,7 @@
|
|||||||
"fileRenameFailed": "Failed to rename file: {error}",
|
"fileRenameFailed": "Failed to rename file: {error}",
|
||||||
"previewUpdated": "Preview updated successfully",
|
"previewUpdated": "Preview updated successfully",
|
||||||
"previewUploadFailed": "Failed to upload preview image",
|
"previewUploadFailed": "Failed to upload preview image",
|
||||||
|
"previewDropInvalid": "Unsupported file type: {name}. Drop an image or MP4 video instead.",
|
||||||
"refreshComplete": "{action} complete",
|
"refreshComplete": "{action} complete",
|
||||||
"refreshFailed": "Failed to {action} {type}s",
|
"refreshFailed": "Failed to {action} {type}s",
|
||||||
"metadataRefreshed": "Metadata refreshed successfully",
|
"metadataRefreshed": "Metadata refreshed successfully",
|
||||||
|
|||||||
+2303
-2205
File diff suppressed because it is too large
Load Diff
+2303
-2205
File diff suppressed because it is too large
Load Diff
+2303
-2205
File diff suppressed because it is too large
Load Diff
+2303
-2205
File diff suppressed because it is too large
Load Diff
+2303
-2205
File diff suppressed because it is too large
Load Diff
+2303
-2205
File diff suppressed because it is too large
Load Diff
+2303
-2205
File diff suppressed because it is too large
Load Diff
+2303
-2205
File diff suppressed because it is too large
Load Diff
+15
-10
@@ -1,9 +1,13 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
import os
|
import os
|
||||||
import platform
|
import platform
|
||||||
import posixpath
|
import posixpath
|
||||||
import threading
|
import threading
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import folder_paths # type: ignore
|
import folder_paths # pyright: ignore[reportMissingImports]
|
||||||
from typing import Any, Dict, Iterable, List, Mapping, Optional, Set, Tuple
|
from typing import Any, Dict, Iterable, List, Mapping, Optional, Set, Tuple
|
||||||
import logging
|
import logging
|
||||||
import json
|
import json
|
||||||
@@ -90,7 +94,7 @@ def _resolve_valid_default_root(
|
|||||||
|
|
||||||
|
|
||||||
def _normalize_folder_paths_for_comparison(
|
def _normalize_folder_paths_for_comparison(
|
||||||
folder_paths: Mapping[str, Iterable[str]],
|
folder_paths: Mapping[str, Any],
|
||||||
) -> Dict[str, Set[str]]:
|
) -> Dict[str, Set[str]]:
|
||||||
"""Normalize folder paths for comparison across libraries."""
|
"""Normalize folder paths for comparison across libraries."""
|
||||||
|
|
||||||
@@ -482,7 +486,7 @@ class Config:
|
|||||||
import ctypes
|
import ctypes
|
||||||
|
|
||||||
FILE_ATTRIBUTE_REPARSE_POINT = 0x400
|
FILE_ATTRIBUTE_REPARSE_POINT = 0x400
|
||||||
attrs = ctypes.windll.kernel32.GetFileAttributesW(str(path)) # type: ignore[attr-defined]
|
attrs = ctypes.windll.kernel32.GetFileAttributesW(str(path)) # pyright: ignore[reportAttributeAccessIssue]
|
||||||
return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT)
|
return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error checking Windows reparse point: {e}")
|
logger.error(f"Error checking Windows reparse point: {e}")
|
||||||
@@ -491,7 +495,7 @@ class Config:
|
|||||||
logger.error(f"Error checking link status for {path}: {e}")
|
logger.error(f"Error checking link status for {path}: {e}")
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def _entry_is_symlink(self, entry: os.DirEntry) -> bool:
|
def _entry_is_symlink(self, entry: os.DirEntry[str]) -> bool:
|
||||||
"""Check if a directory entry is a symlink, including Windows junctions."""
|
"""Check if a directory entry is a symlink, including Windows junctions."""
|
||||||
if entry.is_symlink():
|
if entry.is_symlink():
|
||||||
return True
|
return True
|
||||||
@@ -500,7 +504,7 @@ class Config:
|
|||||||
import ctypes
|
import ctypes
|
||||||
|
|
||||||
FILE_ATTRIBUTE_REPARSE_POINT = 0x400
|
FILE_ATTRIBUTE_REPARSE_POINT = 0x400
|
||||||
attrs = ctypes.windll.kernel32.GetFileAttributesW(entry.path) # type: ignore[attr-defined]
|
attrs = ctypes.windll.kernel32.GetFileAttributesW(entry.path) # pyright: ignore[reportAttributeAccessIssue]
|
||||||
return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT)
|
return attrs != -1 and (attrs & FILE_ATTRIBUTE_REPARSE_POINT)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
@@ -1126,8 +1130,8 @@ class Config:
|
|||||||
|
|
||||||
def _apply_library_paths(
|
def _apply_library_paths(
|
||||||
self,
|
self,
|
||||||
folder_paths: Mapping[str, Iterable[str]],
|
folder_paths: Mapping[str, Any],
|
||||||
extra_folder_paths: Optional[Mapping[str, Iterable[str]]] = None,
|
extra_folder_paths: Optional[Mapping[str, Any]] = None,
|
||||||
recipes_path: str = "",
|
recipes_path: str = "",
|
||||||
) -> None:
|
) -> None:
|
||||||
self._path_mappings.clear()
|
self._path_mappings.clear()
|
||||||
@@ -1432,12 +1436,13 @@ class Config:
|
|||||||
# ('_lm_config_cache') that is NEVER removed from sys.modules (its key does
|
# ('_lm_config_cache') that is NEVER removed from sys.modules (its key does
|
||||||
# NOT start with 'py.'), so it survives re-imports of py.* modules.
|
# NOT start with 'py.'), so it survives re-imports of py.* modules.
|
||||||
_CONFIG_SENTINEL = "_lm_config_cache"
|
_CONFIG_SENTINEL = "_lm_config_cache"
|
||||||
|
config: Config
|
||||||
if _CONFIG_SENTINEL in _sys.modules:
|
if _CONFIG_SENTINEL in _sys.modules:
|
||||||
# Re-import: reuse the existing singleton from the sentinel.
|
# Re-import: reuse the existing singleton from the sentinel.
|
||||||
config: Config = _sys.modules[_CONFIG_SENTINEL].config # type: ignore[valid-type]
|
config = _sys.modules[_CONFIG_SENTINEL].config
|
||||||
else:
|
else:
|
||||||
config: Config = Config()
|
config = Config()
|
||||||
# Register the sentinel so re-imports of py.config find us.
|
# Register the sentinel so re-imports of py.config find us.
|
||||||
_sentinel_mod = _types.ModuleType(_CONFIG_SENTINEL)
|
_sentinel_mod = _types.ModuleType(_CONFIG_SENTINEL)
|
||||||
_sentinel_mod.config = config
|
setattr(_sentinel_mod, "config", config)
|
||||||
_sys.modules[_CONFIG_SENTINEL] = _sentinel_mod
|
_sys.modules[_CONFIG_SENTINEL] = _sentinel_mod
|
||||||
|
|||||||
+18
-1
@@ -14,7 +14,7 @@ standalone_mode = (
|
|||||||
if not standalone_mode:
|
if not standalone_mode:
|
||||||
setup_logging()
|
setup_logging()
|
||||||
|
|
||||||
from server import PromptServer # type: ignore
|
from server import PromptServer # pyright: ignore[reportMissingImports]
|
||||||
|
|
||||||
from .config import config
|
from .config import config
|
||||||
from .services.model_service_factory import (
|
from .services.model_service_factory import (
|
||||||
@@ -25,10 +25,12 @@ from .routes.recipe_routes import RecipeRoutes
|
|||||||
from .routes.stats_routes import StatsRoutes
|
from .routes.stats_routes import StatsRoutes
|
||||||
from .routes.update_routes import UpdateRoutes
|
from .routes.update_routes import UpdateRoutes
|
||||||
from .routes.misc_routes import MiscRoutes
|
from .routes.misc_routes import MiscRoutes
|
||||||
|
from .routes.pending_delete_routes import PendingDeleteRoutes
|
||||||
from .routes.preview_routes import PreviewRoutes
|
from .routes.preview_routes import PreviewRoutes
|
||||||
from .routes.example_images_routes import ExampleImagesRoutes
|
from .routes.example_images_routes import ExampleImagesRoutes
|
||||||
from .services.service_registry import ServiceRegistry
|
from .services.service_registry import ServiceRegistry
|
||||||
from .services.settings_manager import get_settings_manager
|
from .services.settings_manager import get_settings_manager
|
||||||
|
from .services.pending_delete_service import get_pending_delete_service
|
||||||
from .utils.example_images_migration import ExampleImagesMigration
|
from .utils.example_images_migration import ExampleImagesMigration
|
||||||
from .services.websocket_manager import ws_manager
|
from .services.websocket_manager import ws_manager
|
||||||
from .services.example_images_cleanup_service import ExampleImagesCleanupService
|
from .services.example_images_cleanup_service import ExampleImagesCleanupService
|
||||||
@@ -170,6 +172,7 @@ class LoraManager:
|
|||||||
RecipeRoutes.setup_routes(app)
|
RecipeRoutes.setup_routes(app)
|
||||||
UpdateRoutes.setup_routes(app)
|
UpdateRoutes.setup_routes(app)
|
||||||
MiscRoutes.setup_routes(app)
|
MiscRoutes.setup_routes(app)
|
||||||
|
PendingDeleteRoutes.setup_routes(app)
|
||||||
ExampleImagesRoutes.setup_routes(app, ws_manager=ws_manager)
|
ExampleImagesRoutes.setup_routes(app, ws_manager=ws_manager)
|
||||||
PreviewRoutes.setup_routes(app)
|
PreviewRoutes.setup_routes(app)
|
||||||
|
|
||||||
@@ -245,6 +248,20 @@ class LoraManager:
|
|||||||
cls._run_post_initialization_tasks(init_tasks), name="post_init_tasks"
|
cls._run_post_initialization_tasks(init_tasks), name="post_init_tasks"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Startup sweep: purge pending-delete batches that expired during a
|
||||||
|
# previous run. Non-blocking (fire-and-forget); purge_expired only
|
||||||
|
# removes already-expired batches, so a staged undo that survived a
|
||||||
|
# restart stays restorable. scan_roots=True runs the reconciliation
|
||||||
|
# pass first so leftover batches (the in-process registry is empty
|
||||||
|
# after a restart) are re-discovered on disk. Covers both plugin
|
||||||
|
# and standalone modes (StandaloneLoraManager reuses this
|
||||||
|
# classmethod).
|
||||||
|
pending_delete_service = await get_pending_delete_service()
|
||||||
|
asyncio.create_task(
|
||||||
|
pending_delete_service.purge_expired(scan_roots=True),
|
||||||
|
name="pending_delete_startup_sweep",
|
||||||
|
)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"LoRA Manager: All services initialized and background tasks scheduled"
|
"LoRA Manager: All services initialized and background tasks scheduled"
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ if not standalone_mode:
|
|||||||
|
|
||||||
logger.info("ComfyUI Metadata Collector initialized")
|
logger.info("ComfyUI Metadata Collector initialized")
|
||||||
|
|
||||||
def get_metadata(prompt_id=None): # type: ignore[no-redef]
|
def get_metadata(prompt_id=None): # pyright: ignore[reportRedeclaration]
|
||||||
"""Helper function to get metadata from the registry"""
|
"""Helper function to get metadata from the registry"""
|
||||||
registry = MetadataRegistry()
|
registry = MetadataRegistry()
|
||||||
return registry.get_metadata(prompt_id)
|
return registry.get_metadata(prompt_id)
|
||||||
@@ -31,6 +31,6 @@ else:
|
|||||||
def init():
|
def init():
|
||||||
logger.info("ComfyUI Metadata Collector disabled in standalone mode")
|
logger.info("ComfyUI Metadata Collector disabled in standalone mode")
|
||||||
|
|
||||||
def get_metadata(prompt_id=None): # type: ignore[no-redef]
|
def get_metadata(prompt_id=None): # pyright: ignore[reportRedeclaration]
|
||||||
"""Dummy implementation for standalone mode"""
|
"""Dummy implementation for standalone mode"""
|
||||||
return {}
|
return {}
|
||||||
|
|||||||
@@ -1,5 +1,11 @@
|
|||||||
"""Constants used by the metadata collector"""
|
"""Constants used by the metadata collector"""
|
||||||
|
|
||||||
|
# Sentinel value for clip_skip to distinguish "unconnected / widget default"
|
||||||
|
# from "user wired value 0". Both ComfyUI CLIPSetLastLayer (-24..-1) and
|
||||||
|
# A1111 conventions treat 0 as meaningless for clip skipping, but users may
|
||||||
|
# explicitly wire 0 to the overwrite node to express "no clip skip / default".
|
||||||
|
CLIP_SKIP_SENTINEL = -25
|
||||||
|
|
||||||
# Metadata categories
|
# Metadata categories
|
||||||
MODELS = "models"
|
MODELS = "models"
|
||||||
PROMPTS = "prompts"
|
PROMPTS = "prompts"
|
||||||
@@ -9,6 +15,14 @@ EMBEDDINGS = "embeddings"
|
|||||||
SIZE = "size"
|
SIZE = "size"
|
||||||
IMAGES = "images"
|
IMAGES = "images"
|
||||||
IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes
|
IS_SAMPLER = "is_sampler" # New constant to mark sampler nodes
|
||||||
|
OVERWRITE = "overwrite" # Manual metadata overwrite from MetadataOverwriteLM node
|
||||||
|
|
||||||
|
# Field names that the MetadataOverwriteLM node and its extractor share
|
||||||
|
METADATA_OVERWRITE_FIELDS = (
|
||||||
|
"prompt", "negative_prompt", "seed", "steps", "cfg_scale",
|
||||||
|
"sampler", "scheduler", "model", "loras", "size",
|
||||||
|
"clip_skip", "additional_data",
|
||||||
|
)
|
||||||
|
|
||||||
# Complete list of categories to track
|
# Complete list of categories to track
|
||||||
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES]
|
METADATA_CATEGORIES = [MODELS, PROMPTS, SAMPLING, LORAS, EMBEDDINGS, SIZE, IMAGES, OVERWRITE]
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ class MetadataHook:
|
|||||||
execution = None
|
execution = None
|
||||||
try:
|
try:
|
||||||
# Try direct import first
|
# Try direct import first
|
||||||
import execution # type: ignore
|
import execution # pyright: ignore[reportMissingImports]
|
||||||
except ImportError:
|
except ImportError:
|
||||||
# Try to locate from system modules
|
# Try to locate from system modules
|
||||||
for module_name in sys.modules:
|
for module_name in sys.modules:
|
||||||
@@ -83,7 +83,8 @@ class MetadataHook:
|
|||||||
|
|
||||||
# Record inputs before execution
|
# Record inputs before execution
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -114,7 +115,8 @@ class MetadataHook:
|
|||||||
|
|
||||||
# Record outputs after execution
|
# Record outputs after execution
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.update_node_execution(node_id, class_type, results)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -135,10 +137,13 @@ class MetadataHook:
|
|||||||
# Store the dynprompt reference for node lookups
|
# Store the dynprompt reference for node lookups
|
||||||
if hasattr(prompt, 'original_prompt'):
|
if hasattr(prompt, 'original_prompt'):
|
||||||
registry.set_current_prompt(prompt)
|
registry.set_current_prompt(prompt)
|
||||||
|
|
||||||
|
# Store extra_data for accessing full workflow node properties
|
||||||
|
registry.set_extra_data(extra_data)
|
||||||
|
|
||||||
# Execute the original function
|
# Execute the original function
|
||||||
return original_execute(*args, **kwargs)
|
return original_execute(*args, **kwargs)
|
||||||
|
|
||||||
# Replace the functions
|
# Replace the functions
|
||||||
execution._map_node_over_list = map_node_over_list_with_metadata
|
execution._map_node_over_list = map_node_over_list_with_metadata
|
||||||
execution.execute = execute_with_prompt_tracking
|
execution.execute = execute_with_prompt_tracking
|
||||||
@@ -163,7 +168,8 @@ class MetadataHook:
|
|||||||
class_type = obj.__class__.__name__
|
class_type = obj.__class__.__name__
|
||||||
node_id = unique_id
|
node_id = unique_id
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.record_node_execution(node_id, class_type, input_data_all, None)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.record_node_execution(node_id, class_type, input_data_all, None, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (pre-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -180,7 +186,8 @@ class MetadataHook:
|
|||||||
class_type = obj.__class__.__name__
|
class_type = obj.__class__.__name__
|
||||||
node_id = unique_id
|
node_id = unique_id
|
||||||
if node_id is not None:
|
if node_id is not None:
|
||||||
registry.update_node_execution(node_id, class_type, results)
|
return_types = getattr(obj, 'RETURN_TYPES', None)
|
||||||
|
registry.update_node_execution(node_id, class_type, results, return_types=return_types)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
logger.error(f"Error collecting metadata (post-execution): {str(e)}")
|
||||||
|
|
||||||
@@ -202,6 +209,9 @@ class MetadataHook:
|
|||||||
if hasattr(prompt, 'original_prompt'):
|
if hasattr(prompt, 'original_prompt'):
|
||||||
registry.set_current_prompt(prompt)
|
registry.set_current_prompt(prompt)
|
||||||
|
|
||||||
|
# Store extra_data for accessing full workflow node properties
|
||||||
|
registry.set_extra_data(extra_data)
|
||||||
|
|
||||||
# Execute the original function
|
# Execute the original function
|
||||||
return await original_execute(*args, **kwargs)
|
return await original_execute(*args, **kwargs)
|
||||||
|
|
||||||
|
|||||||
@@ -1,15 +1,68 @@
|
|||||||
import json
|
import json
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
from .constants import IMAGES
|
from .constants import IMAGES
|
||||||
|
|
||||||
# Check if running in standalone mode
|
# Check if running in standalone mode
|
||||||
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1" or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
||||||
|
|
||||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER
|
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IS_SAMPLER, OVERWRITE
|
||||||
|
from .node_extractors import NODE_EXTRACTORS
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Keys that identify metadata hint marks stored in node.properties.lm_marker_role
|
||||||
|
_META_MARK_PREFIX = "meta_"
|
||||||
|
_MARK_PRIMARY_MODEL = "primary_model"
|
||||||
|
_MARK_PRIMARY_SAMPLER = "primary_sampler"
|
||||||
|
_MARK_POSITIVE_PROMPT = "positive_prompt"
|
||||||
|
_MARK_NEGATIVE_PROMPT = "negative_prompt"
|
||||||
|
|
||||||
class MetadataProcessor:
|
class MetadataProcessor:
|
||||||
"""Process and format collected metadata"""
|
"""Process and format collected metadata"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_user_marks(metadata):
|
||||||
|
"""Scan workflow nodes (from extra_data.extra_pnginfo.workflow) for user-assigned
|
||||||
|
metadata hint marks stored in node.properties.lm_marker_role.
|
||||||
|
|
||||||
|
Returns a dict mapping mark type keys to node IDs.
|
||||||
|
Example: {'primary_model': '42', 'primary_sampler': '17'}
|
||||||
|
"""
|
||||||
|
marks: dict[str, str] = {}
|
||||||
|
|
||||||
|
# Primary source: extra_data.extra_pnginfo.workflow.nodes (has full properties)
|
||||||
|
extra_data = metadata.get("extra_data")
|
||||||
|
if extra_data and isinstance(extra_data, dict):
|
||||||
|
extra_pnginfo = extra_data.get("extra_pnginfo", {})
|
||||||
|
if isinstance(extra_pnginfo, dict):
|
||||||
|
workflow = extra_pnginfo.get("workflow", {})
|
||||||
|
nodes = workflow.get("nodes", [])
|
||||||
|
for node in nodes:
|
||||||
|
node_id = str(node.get("id", ""))
|
||||||
|
role = node.get("properties", {}).get("lm_marker_role", "")
|
||||||
|
if role.startswith(_META_MARK_PREFIX):
|
||||||
|
mark_type = role[len(_META_MARK_PREFIX):]
|
||||||
|
if mark_type in marks:
|
||||||
|
logger.warning(
|
||||||
|
"Duplicate meta hint '%s': node %s (previous: %s), "
|
||||||
|
"last match wins",
|
||||||
|
mark_type, node_id, marks[mark_type],
|
||||||
|
)
|
||||||
|
marks[mark_type] = node_id
|
||||||
|
|
||||||
|
# Fallback: try prompt.original_prompt (API-only submissions may not have workflow)
|
||||||
|
if not marks:
|
||||||
|
prompt = metadata.get("current_prompt")
|
||||||
|
if prompt and getattr(prompt, "original_prompt", None):
|
||||||
|
for node_id, node_data in prompt.original_prompt.items():
|
||||||
|
role = node_data.get("properties", {}).get("lm_marker_role", "")
|
||||||
|
if role.startswith(_META_MARK_PREFIX):
|
||||||
|
mark_type = role[len(_META_MARK_PREFIX):]
|
||||||
|
marks[mark_type] = node_id
|
||||||
|
|
||||||
|
return marks
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def find_primary_sampler(metadata, downstream_id=None):
|
def find_primary_sampler(metadata, downstream_id=None):
|
||||||
"""
|
"""
|
||||||
@@ -161,6 +214,24 @@ class MetadataProcessor:
|
|||||||
max_denoise = denoise
|
max_denoise = denoise
|
||||||
primary_sampler = sampler_info
|
primary_sampler = sampler_info
|
||||||
primary_sampler_id = node_id
|
primary_sampler_id = node_id
|
||||||
|
|
||||||
|
# Last resort: any registered sampler. Samplers without a denoise or
|
||||||
|
# add_noise parameter (e.g. multi-stage samplers like KreaTwoStageSampler)
|
||||||
|
# are not caught by the criteria above. Prefer execution order so the
|
||||||
|
# first executed sampler wins, matching the downstream_id branch.
|
||||||
|
if primary_sampler is None:
|
||||||
|
sampler_ids = [
|
||||||
|
node_id
|
||||||
|
for node_id, sampler_info in metadata.get(SAMPLING, {}).items()
|
||||||
|
if sampler_info.get(IS_SAMPLER, False)
|
||||||
|
]
|
||||||
|
if sampler_ids:
|
||||||
|
if downstream_id and "execution_order" in metadata:
|
||||||
|
for node_id in metadata["execution_order"]:
|
||||||
|
if node_id in sampler_ids:
|
||||||
|
return node_id, metadata[SAMPLING][node_id]
|
||||||
|
primary_sampler_id = sampler_ids[0]
|
||||||
|
primary_sampler = metadata[SAMPLING][sampler_ids[0]]
|
||||||
|
|
||||||
return primary_sampler_id, primary_sampler
|
return primary_sampler_id, primary_sampler
|
||||||
|
|
||||||
@@ -471,20 +542,57 @@ class MetadataProcessor:
|
|||||||
"checkpoint": None,
|
"checkpoint": None,
|
||||||
"loras": "",
|
"loras": "",
|
||||||
"size": None,
|
"size": None,
|
||||||
"clip_skip": None
|
"clip_skip": None,
|
||||||
|
"additional_data": "",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Get the prompt object for node relationship tracing
|
# Get the prompt object for node relationship tracing
|
||||||
prompt = metadata.get("current_prompt")
|
prompt = metadata.get("current_prompt")
|
||||||
|
|
||||||
# Find the primary KSampler node
|
# ---- User marks: override heuristic inference with user-assigned hints ----
|
||||||
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
user_marks = MetadataProcessor._get_user_marks(metadata)
|
||||||
|
|
||||||
# Directly get checkpoint from metadata instead of tracing
|
# Find the primary KSampler node (user mark takes priority)
|
||||||
# Pass primary_sampler_id to avoid redundant calculation
|
primary_sampler_id = None
|
||||||
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
primary_sampler = None
|
||||||
if checkpoint:
|
if _MARK_PRIMARY_SAMPLER in user_marks:
|
||||||
params["checkpoint"] = checkpoint
|
marked_id = user_marks[_MARK_PRIMARY_SAMPLER]
|
||||||
|
sampler_data = metadata.get(SAMPLING, {}).get(marked_id)
|
||||||
|
if sampler_data and sampler_data.get(IS_SAMPLER):
|
||||||
|
primary_sampler_id = marked_id
|
||||||
|
primary_sampler = sampler_data
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"User-marked primary sampler %s has no runtime metadata, "
|
||||||
|
"falling back to heuristic",
|
||||||
|
marked_id,
|
||||||
|
)
|
||||||
|
if primary_sampler is None:
|
||||||
|
primary_sampler_id, primary_sampler = MetadataProcessor.find_primary_sampler(metadata, id)
|
||||||
|
|
||||||
|
# Resolve checkpoint / model (user mark takes priority)
|
||||||
|
if _MARK_PRIMARY_MODEL in user_marks:
|
||||||
|
marked_id = user_marks[_MARK_PRIMARY_MODEL]
|
||||||
|
if marked_id in metadata.get(MODELS, {}):
|
||||||
|
params["checkpoint"] = metadata[MODELS][marked_id].get("name")
|
||||||
|
else:
|
||||||
|
extra_data = metadata.get("extra_data")
|
||||||
|
extra_pnginfo = extra_data.get("extra_pnginfo", {}) if extra_data and isinstance(extra_data, dict) else {}
|
||||||
|
workflow = extra_pnginfo.get("workflow", {}) if isinstance(extra_pnginfo, dict) else {}
|
||||||
|
node_type = "unknown"
|
||||||
|
for n in workflow.get("nodes", []):
|
||||||
|
if str(n.get("id", "")) == marked_id:
|
||||||
|
node_type = n.get("type", "unknown")
|
||||||
|
break
|
||||||
|
logger.warning(
|
||||||
|
"User-marked primary model %s (type=%s, registered=%s) has no runtime metadata, "
|
||||||
|
"falling back to heuristic",
|
||||||
|
marked_id, node_type, node_type in NODE_EXTRACTORS,
|
||||||
|
)
|
||||||
|
if params["checkpoint"] is None:
|
||||||
|
checkpoint = MetadataProcessor.find_primary_checkpoint(metadata, id, primary_sampler_id)
|
||||||
|
if checkpoint:
|
||||||
|
params["checkpoint"] = checkpoint
|
||||||
|
|
||||||
# Check if guidance parameter exists in any sampling node
|
# Check if guidance parameter exists in any sampling node
|
||||||
for node_id, sampler_info in metadata.get(SAMPLING, {}).items():
|
for node_id, sampler_info in metadata.get(SAMPLING, {}).items():
|
||||||
@@ -539,7 +647,22 @@ class MetadataProcessor:
|
|||||||
|
|
||||||
# For SamplerCustom, handle any additional parameters
|
# For SamplerCustom, handle any additional parameters
|
||||||
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
|
MetadataProcessor.handle_custom_advanced_sampler(metadata, prompt, primary_sampler_id, params)
|
||||||
|
|
||||||
|
# ---- User marks: override prompts with explicitly tagged nodes ----
|
||||||
|
prompts_data = metadata.get(PROMPTS, {})
|
||||||
|
if _MARK_POSITIVE_PROMPT in user_marks:
|
||||||
|
pos_id = user_marks[_MARK_POSITIVE_PROMPT]
|
||||||
|
if pos_id in prompts_data:
|
||||||
|
prompt_text = prompts_data[pos_id].get("text") or prompts_data[pos_id].get("positive_text")
|
||||||
|
if prompt_text:
|
||||||
|
params["prompt"] = prompt_text
|
||||||
|
if _MARK_NEGATIVE_PROMPT in user_marks:
|
||||||
|
neg_id = user_marks[_MARK_NEGATIVE_PROMPT]
|
||||||
|
if neg_id in prompts_data:
|
||||||
|
prompt_text = prompts_data[neg_id].get("text") or prompts_data[neg_id].get("negative_text")
|
||||||
|
if prompt_text:
|
||||||
|
params["negative_prompt"] = prompt_text
|
||||||
|
|
||||||
# Size extraction is same for all sampler types
|
# Size extraction is same for all sampler types
|
||||||
# Check if the sampler itself has size information (from latent_image)
|
# Check if the sampler itself has size information (from latent_image)
|
||||||
if primary_sampler_id in metadata.get(SIZE, {}):
|
if primary_sampler_id in metadata.get(SIZE, {}):
|
||||||
@@ -568,7 +691,26 @@ class MetadataProcessor:
|
|||||||
break
|
break
|
||||||
if params["clip_skip"] is None:
|
if params["clip_skip"] is None:
|
||||||
params["clip_skip"] = "1"
|
params["clip_skip"] = "1"
|
||||||
|
|
||||||
|
# ---- Apply manual metadata overwrites ----
|
||||||
|
for overwrite_info in metadata.get(OVERWRITE, {}).values():
|
||||||
|
overwrite_params = overwrite_info.get("parameters", {})
|
||||||
|
for key, value in overwrite_params.items():
|
||||||
|
if key == "clip_skip":
|
||||||
|
# Accept any value from overwrite node (sentinel -25 already
|
||||||
|
# filtered upstream). Needed because falsy check treats 0
|
||||||
|
# as "not set" even though 0 is a valid wired input here.
|
||||||
|
params[key] = value
|
||||||
|
elif value: # truthy check — only overwrite when user provided a real value
|
||||||
|
params[key] = value
|
||||||
|
|
||||||
|
# Bridge: the overwrite node exposes the field as "model" (more accurate),
|
||||||
|
# but the internal pipeline key remains "checkpoint" for backward compatibility
|
||||||
|
# with A1111 metadata format and downstream consumers.
|
||||||
|
if params.get("model"):
|
||||||
|
params["checkpoint"] = params["model"]
|
||||||
|
del params["model"]
|
||||||
|
|
||||||
return params
|
return params
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
import time
|
import time
|
||||||
from nodes import NODE_CLASS_MAPPINGS # type: ignore
|
from typing import Any
|
||||||
|
from nodes import NODE_CLASS_MAPPINGS # pyright: ignore[reportMissingImports, reportAttributeAccessIssue]
|
||||||
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
|
from .node_extractors import NODE_EXTRACTORS, GenericNodeExtractor
|
||||||
from .constants import METADATA_CATEGORIES, IMAGES
|
from .constants import METADATA_CATEGORIES, IMAGES, OVERWRITE
|
||||||
|
|
||||||
|
|
||||||
class MetadataRegistry:
|
class MetadataRegistry:
|
||||||
@@ -9,6 +10,15 @@ class MetadataRegistry:
|
|||||||
|
|
||||||
_instance = None
|
_instance = None
|
||||||
|
|
||||||
|
current_prompt_id: Any = None
|
||||||
|
current_prompt: Any = None
|
||||||
|
metadata: dict[str, Any] = {}
|
||||||
|
prompt_metadata: dict[str, Any] = {}
|
||||||
|
executed_nodes: set[str] = set()
|
||||||
|
node_cache: dict[str, Any] = {}
|
||||||
|
max_prompt_history: int = 3
|
||||||
|
metadata_categories: list[str] = METADATA_CATEGORIES
|
||||||
|
|
||||||
def __new__(cls):
|
def __new__(cls):
|
||||||
if cls._instance is None:
|
if cls._instance is None:
|
||||||
cls._instance = super().__new__(cls)
|
cls._instance = super().__new__(cls)
|
||||||
@@ -61,6 +71,7 @@ class MetadataRegistry:
|
|||||||
{
|
{
|
||||||
"execution_order": [],
|
"execution_order": [],
|
||||||
"current_prompt": None, # Will store the prompt object
|
"current_prompt": None, # Will store the prompt object
|
||||||
|
"extra_data": None, # Will store the API extra_data for workflow metadata
|
||||||
"timestamp": time.time(),
|
"timestamp": time.time(),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
@@ -75,6 +86,11 @@ class MetadataRegistry:
|
|||||||
# Store the prompt in the metadata for later relationship tracing
|
# Store the prompt in the metadata for later relationship tracing
|
||||||
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
|
self.prompt_metadata[self.current_prompt_id]["current_prompt"] = prompt
|
||||||
|
|
||||||
|
def set_extra_data(self, extra_data):
|
||||||
|
"""Store the API extra_data (contains extra_pnginfo.workflow with node properties)"""
|
||||||
|
if self.current_prompt_id and self.current_prompt_id in self.prompt_metadata:
|
||||||
|
self.prompt_metadata[self.current_prompt_id]["extra_data"] = extra_data
|
||||||
|
|
||||||
def get_metadata(self, prompt_id=None):
|
def get_metadata(self, prompt_id=None):
|
||||||
"""Get collected metadata for a prompt"""
|
"""Get collected metadata for a prompt"""
|
||||||
key = prompt_id if prompt_id is not None else self.current_prompt_id
|
key = prompt_id if prompt_id is not None else self.current_prompt_id
|
||||||
@@ -122,20 +138,28 @@ class MetadataRegistry:
|
|||||||
cache_key = f"{node_id}:{class_type}"
|
cache_key = f"{node_id}:{class_type}"
|
||||||
|
|
||||||
# Check if this node type is relevant for metadata collection
|
# Check if this node type is relevant for metadata collection
|
||||||
if class_type in NODE_EXTRACTORS:
|
if class_type in NODE_EXTRACTORS or cache_key in self.node_cache:
|
||||||
# Check if we have cached metadata for this node
|
# Check if we have cached metadata for this node
|
||||||
if cache_key in self.node_cache:
|
if cache_key in self.node_cache:
|
||||||
cached_data = self.node_cache[cache_key]
|
cached_data = self.node_cache[cache_key]
|
||||||
|
|
||||||
|
# Detect bypass (mode=4) / mute (mode=2) — these nodes
|
||||||
|
# were intentionally disabled and should not contribute
|
||||||
|
# overwrite values from a previous execution's cache.
|
||||||
|
node_mode = node_data.get("mode", 0)
|
||||||
|
node_is_disabled = node_mode in (2, 4)
|
||||||
|
|
||||||
# Apply cached metadata to the current metadata
|
# Apply cached metadata to the current metadata
|
||||||
for category in self.metadata_categories:
|
for category in self.metadata_categories:
|
||||||
|
if category == OVERWRITE and node_is_disabled:
|
||||||
|
continue
|
||||||
if category in cached_data and node_id in cached_data[category]:
|
if category in cached_data and node_id in cached_data[category]:
|
||||||
if node_id not in metadata[category]:
|
if node_id not in metadata[category]:
|
||||||
metadata[category][node_id] = cached_data[category][
|
metadata[category][node_id] = cached_data[category][
|
||||||
node_id
|
node_id
|
||||||
]
|
]
|
||||||
|
|
||||||
def record_node_execution(self, node_id, class_type, inputs, outputs):
|
def record_node_execution(self, node_id, class_type, inputs, outputs, return_types=None):
|
||||||
"""Record information about a node's execution"""
|
"""Record information about a node's execution"""
|
||||||
if not self.current_prompt_id:
|
if not self.current_prompt_id:
|
||||||
return
|
return
|
||||||
@@ -158,17 +182,18 @@ class MetadataRegistry:
|
|||||||
|
|
||||||
# Extract node-specific metadata
|
# Extract node-specific metadata
|
||||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||||
extractor.extract(
|
if extractor is GenericNodeExtractor:
|
||||||
node_id,
|
extractor.extract(node_id, processed_inputs, outputs,
|
||||||
processed_inputs,
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
outputs,
|
return_types=return_types)
|
||||||
self.prompt_metadata[self.current_prompt_id],
|
else:
|
||||||
)
|
extractor.extract(node_id, processed_inputs, outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id])
|
||||||
|
|
||||||
# Cache this node's metadata
|
# Cache this node's metadata
|
||||||
self._cache_node_metadata(node_id, class_type)
|
self._cache_node_metadata(node_id, class_type)
|
||||||
|
|
||||||
def update_node_execution(self, node_id, class_type, outputs):
|
def update_node_execution(self, node_id, class_type, outputs, return_types=None):
|
||||||
"""Update node metadata with output information"""
|
"""Update node metadata with output information"""
|
||||||
if not self.current_prompt_id:
|
if not self.current_prompt_id:
|
||||||
return
|
return
|
||||||
@@ -179,9 +204,17 @@ class MetadataRegistry:
|
|||||||
# Use the same extractor to update with outputs
|
# Use the same extractor to update with outputs
|
||||||
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
extractor = NODE_EXTRACTORS.get(class_type, GenericNodeExtractor)
|
||||||
if hasattr(extractor, "update"):
|
if hasattr(extractor, "update"):
|
||||||
extractor.update(
|
if extractor is GenericNodeExtractor:
|
||||||
node_id, processed_outputs, self.prompt_metadata[self.current_prompt_id]
|
extractor.update(
|
||||||
)
|
node_id, processed_outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
|
return_types=return_types,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
extractor.update(
|
||||||
|
node_id, processed_outputs,
|
||||||
|
self.prompt_metadata[self.current_prompt_id],
|
||||||
|
)
|
||||||
|
|
||||||
# Update the cached metadata for this node
|
# Update the cached metadata for this node
|
||||||
self._cache_node_metadata(node_id, class_type)
|
self._cache_node_metadata(node_id, class_type)
|
||||||
|
|||||||
@@ -2,7 +2,8 @@ import json
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
|
||||||
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER
|
from .constants import MODELS, PROMPTS, SAMPLING, LORAS, SIZE, IMAGES, IS_SAMPLER, OVERWRITE
|
||||||
|
from .overwrite_utils import collect_overwrite_params
|
||||||
|
|
||||||
|
|
||||||
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
def _store_checkpoint_metadata(metadata, node_id, model_name):
|
||||||
@@ -31,11 +32,95 @@ class NodeMetadataExtractor:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
class GenericNodeExtractor(NodeMetadataExtractor):
|
class GenericNodeExtractor(NodeMetadataExtractor):
|
||||||
"""Default extractor for nodes without specific handling"""
|
"""Fallback extractor with type-signature-based detection.
|
||||||
|
|
||||||
|
When a node is not in the NODE_EXTRACTORS registry, the hook layer
|
||||||
|
passes ``return_types`` from ``obj.RETURN_TYPES``:
|
||||||
|
|
||||||
|
* ``MODEL`` output: common input fields (ckpt_name, unet_name, etc.)
|
||||||
|
are checked for a model file name and stored as checkpoint metadata.
|
||||||
|
* ``CONDITIONING`` output: common text input fields are checked for
|
||||||
|
prompt text, and conditioning inputs are tracked through transforms.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Input field names that carry a model path in loader-style nodes.
|
||||||
|
_MODEL_NAME_FIELDS = (
|
||||||
|
"ckpt_name", "unet_name", "model_path", "model_name", "gguf_name",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Extensions used by checkpoint_scanner.py — only record values that look
|
||||||
|
# like real model filenames to avoid capturing unrelated string fields.
|
||||||
|
_MODEL_EXTENSIONS = {
|
||||||
|
".ckpt", ".pt", ".pt2", ".bin", ".pth", ".safetensors", ".pkl", ".sft", ".gguf",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Input field names that may carry prompt text in encoder-style nodes.
|
||||||
|
_TEXT_FIELDS = ("text", "clip_l", "t5xxl", "prompt", "positive", "negative")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def extract(node_id, inputs, outputs, metadata):
|
def extract(node_id, inputs, outputs, metadata, return_types=None):
|
||||||
pass
|
if return_types is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
# — MODEL loader detection (checkpoint / UNET / GGUF) —
|
||||||
|
if "MODEL" in return_types or any("MODEL" in str(t) for t in return_types):
|
||||||
|
for field in GenericNodeExtractor._MODEL_NAME_FIELDS:
|
||||||
|
val = inputs.get(field)
|
||||||
|
if val and isinstance(val, str) and val.strip():
|
||||||
|
name = val.strip()
|
||||||
|
if not any(name.lower().endswith(ext) for ext in GenericNodeExtractor._MODEL_EXTENSIONS):
|
||||||
|
continue
|
||||||
|
_store_checkpoint_metadata(metadata, node_id, name)
|
||||||
|
return
|
||||||
|
|
||||||
|
# — CONDITIONING encoder / transform detection —
|
||||||
|
if "CONDITIONING" in return_types or any("CONDITIONING" in str(t) for t in return_types):
|
||||||
|
text = None
|
||||||
|
for field in GenericNodeExtractor._TEXT_FIELDS:
|
||||||
|
val = inputs.get(field)
|
||||||
|
if val and isinstance(val, str) and val.strip():
|
||||||
|
text = val.strip()
|
||||||
|
break
|
||||||
|
|
||||||
|
input_conditionings = _collect_conditioning_inputs(inputs)
|
||||||
|
if text or input_conditionings:
|
||||||
|
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
||||||
|
if text:
|
||||||
|
prompt_metadata["text"] = text
|
||||||
|
if input_conditionings:
|
||||||
|
prompt_metadata["orig_conditionings"] = input_conditionings
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def update(node_id, outputs, metadata, return_types=None):
|
||||||
|
if return_types is None:
|
||||||
|
return
|
||||||
|
if "CONDITIONING" not in return_types and not any(
|
||||||
|
"CONDITIONING" in str(t) for t in return_types
|
||||||
|
):
|
||||||
|
return
|
||||||
|
if node_id not in metadata.get(PROMPTS, {}):
|
||||||
|
return
|
||||||
|
output_tuple = _first_output_tuple(outputs)
|
||||||
|
if not output_tuple or len(output_tuple) < 1:
|
||||||
|
return
|
||||||
|
|
||||||
|
conditioning_index = _first_conditioning_index(return_types)
|
||||||
|
if conditioning_index is None or len(output_tuple) <= conditioning_index:
|
||||||
|
return
|
||||||
|
|
||||||
|
output_conditioning = output_tuple[conditioning_index]
|
||||||
|
if output_conditioning is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
prompt_metadata = metadata[PROMPTS][node_id]
|
||||||
|
prompt_metadata["conditioning"] = output_conditioning
|
||||||
|
_record_conditioning_source(
|
||||||
|
metadata,
|
||||||
|
node_id,
|
||||||
|
output_conditioning,
|
||||||
|
prompt_metadata.get("orig_conditionings", []),
|
||||||
|
)
|
||||||
|
|
||||||
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
class CheckpointLoaderExtractor(NodeMetadataExtractor):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def extract(node_id, inputs, outputs, metadata):
|
def extract(node_id, inputs, outputs, metadata):
|
||||||
@@ -349,6 +434,34 @@ def _first_output_tuple(outputs):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _first_conditioning_index(return_types):
|
||||||
|
"""Return the index of the first CONDITIONING output slot, or None."""
|
||||||
|
if not return_types:
|
||||||
|
return None
|
||||||
|
for index, return_type in enumerate(return_types):
|
||||||
|
if "CONDITIONING" in str(return_type):
|
||||||
|
return index
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _collect_conditioning_inputs(inputs):
|
||||||
|
"""Collect conditioning object inputs (``conditioning*`` keys).
|
||||||
|
|
||||||
|
Primitive values (None, str, int, float, bool) are excluded so scalar
|
||||||
|
fields like ``conditioning_strength`` are not mistaken for conditioning
|
||||||
|
objects during provenance tracking.
|
||||||
|
"""
|
||||||
|
if not inputs:
|
||||||
|
return []
|
||||||
|
return [
|
||||||
|
value
|
||||||
|
for input_name, value in inputs.items()
|
||||||
|
if input_name.startswith("conditioning")
|
||||||
|
and value is not None
|
||||||
|
and not isinstance(value, (str, int, float, bool))
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def _record_conditioning_source(
|
def _record_conditioning_source(
|
||||||
metadata, node_id, output_conditioning, input_conditionings
|
metadata, node_id, output_conditioning, input_conditionings
|
||||||
):
|
):
|
||||||
@@ -361,6 +474,14 @@ def _record_conditioning_source(
|
|||||||
if not sources:
|
if not sources:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# Identity-preserving selectors return one of their inputs unchanged:
|
||||||
|
# only that input contributed to the output, so record it alone instead
|
||||||
|
# of treating every input as a combination source.
|
||||||
|
for conditioning in sources:
|
||||||
|
if id(conditioning) == id(output_conditioning):
|
||||||
|
sources = [conditioning]
|
||||||
|
break
|
||||||
|
|
||||||
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
||||||
prompt_metadata.setdefault("conditioning_sources", []).append(
|
prompt_metadata.setdefault("conditioning_sources", []).append(
|
||||||
{
|
{
|
||||||
@@ -440,13 +561,7 @@ class ConditioningCombineExtractor(NodeMetadataExtractor):
|
|||||||
if not inputs:
|
if not inputs:
|
||||||
return
|
return
|
||||||
|
|
||||||
input_conditionings = []
|
input_conditionings = _collect_conditioning_inputs(inputs)
|
||||||
for input_name in inputs:
|
|
||||||
if (
|
|
||||||
input_name.startswith("conditioning")
|
|
||||||
and inputs[input_name] is not None
|
|
||||||
):
|
|
||||||
input_conditionings.append(inputs[input_name])
|
|
||||||
|
|
||||||
if input_conditionings:
|
if input_conditionings:
|
||||||
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
prompt_metadata = _ensure_prompt_metadata(metadata, node_id)
|
||||||
@@ -746,6 +861,65 @@ class TSCKSamplerAdvancedExtractor(KSamplerAdvancedExtractor, TSCSamplerBaseExtr
|
|||||||
|
|
||||||
# Update method is inherited from TSCSamplerBaseExtractor
|
# Update method is inherited from TSCSamplerBaseExtractor
|
||||||
|
|
||||||
|
class KreaTwoStageSamplerExtractor(BaseSamplerExtractor):
|
||||||
|
"""Extractor for Krea Two/Three Stage Samplers (Auryg/Krea-2-Two-Stage-Sampler).
|
||||||
|
|
||||||
|
The node samples in two (or three) stages with per-stage settings
|
||||||
|
(stage1_steps/stage2_steps, stage1_cfg/stage2_cfg, ...). The canonical
|
||||||
|
metadata fields consumed by ``extract_generation_params`` (steps, cfg,
|
||||||
|
sampler_name, scheduler) are derived from the base stage (stage 1; the
|
||||||
|
three-stage variant reuses stage 1 settings for stage 3), while the full
|
||||||
|
per-stage breakdown is preserved in the raw parameters.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# All per-stage parameter keys present on both node variants.
|
||||||
|
_STAGE_PARAM_KEYS = (
|
||||||
|
"stage1_steps", "stage1_cfg", "stage1_sampler_name", "stage1_scheduler",
|
||||||
|
"stage2_steps", "stage2_cfg", "stage2_sampler_name", "stage2_scheduler",
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def extract(node_id, inputs, outputs, metadata):
|
||||||
|
if not inputs:
|
||||||
|
return
|
||||||
|
|
||||||
|
BaseSamplerExtractor.extract_sampling_params(
|
||||||
|
node_id,
|
||||||
|
inputs,
|
||||||
|
metadata,
|
||||||
|
("seed", "handoff_percent", "stage3_handoff_percent")
|
||||||
|
+ KreaTwoStageSamplerExtractor._STAGE_PARAM_KEYS,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Derive the canonical fields expected by extract_generation_params.
|
||||||
|
sampling_params = metadata[SAMPLING][node_id]["parameters"]
|
||||||
|
if "stage1_steps" in sampling_params or "stage2_steps" in sampling_params:
|
||||||
|
sampling_params["steps"] = (
|
||||||
|
(sampling_params.get("stage1_steps") or 0)
|
||||||
|
+ (sampling_params.get("stage2_steps") or 0)
|
||||||
|
)
|
||||||
|
if "stage1_cfg" in sampling_params:
|
||||||
|
sampling_params["cfg"] = sampling_params["stage1_cfg"]
|
||||||
|
if "stage1_sampler_name" in sampling_params:
|
||||||
|
sampling_params["sampler_name"] = sampling_params["stage1_sampler_name"]
|
||||||
|
if "stage1_scheduler" in sampling_params:
|
||||||
|
sampling_params["scheduler"] = sampling_params["stage1_scheduler"]
|
||||||
|
|
||||||
|
BaseSamplerExtractor.extract_conditioning(node_id, inputs, metadata)
|
||||||
|
|
||||||
|
# Prefer the final generation resolution; latent dims are the fallback.
|
||||||
|
BaseSamplerExtractor.extract_latent_dimensions(node_id, inputs, metadata)
|
||||||
|
final_width = inputs.get("final_width")
|
||||||
|
final_height = inputs.get("final_height")
|
||||||
|
if final_width and final_height:
|
||||||
|
if SIZE not in metadata:
|
||||||
|
metadata[SIZE] = {}
|
||||||
|
metadata[SIZE][node_id] = {
|
||||||
|
"width": final_width,
|
||||||
|
"height": final_height,
|
||||||
|
"node_id": node_id,
|
||||||
|
}
|
||||||
|
|
||||||
class LoraLoaderExtractor(NodeMetadataExtractor):
|
class LoraLoaderExtractor(NodeMetadataExtractor):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def extract(node_id, inputs, outputs, metadata):
|
def extract(node_id, inputs, outputs, metadata):
|
||||||
@@ -786,6 +960,37 @@ class ImageSizeExtractor(NodeMetadataExtractor):
|
|||||||
"node_id": node_id
|
"node_id": node_id
|
||||||
}
|
}
|
||||||
|
|
||||||
|
class KreaDualResolutionSelectorExtractor(NodeMetadataExtractor):
|
||||||
|
"""Extract base resolution from Krea Dual Resolution Selector outputs
|
||||||
|
(Auryg/Krea-2-Two-Stage-Sampler).
|
||||||
|
|
||||||
|
The node computes base/final dimensions at runtime from aspect ratio and
|
||||||
|
megapixel settings, so the values are only available in the update phase
|
||||||
|
(outputs: base_width, base_height, final_width, final_height, seed).
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def extract(node_id, inputs, outputs, metadata):
|
||||||
|
# Dimensions are computed at runtime; nothing to do here.
|
||||||
|
pass
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def update(node_id, outputs, metadata):
|
||||||
|
output_tuple = _first_output_tuple(outputs)
|
||||||
|
if not output_tuple or len(output_tuple) < 2:
|
||||||
|
return
|
||||||
|
width, height = output_tuple[0], output_tuple[1]
|
||||||
|
if not isinstance(width, int) or not isinstance(height, int):
|
||||||
|
return
|
||||||
|
|
||||||
|
if SIZE not in metadata:
|
||||||
|
metadata[SIZE] = {}
|
||||||
|
metadata[SIZE][node_id] = {
|
||||||
|
"width": width,
|
||||||
|
"height": height,
|
||||||
|
"node_id": node_id,
|
||||||
|
}
|
||||||
|
|
||||||
class RgthreePowerLoraLoaderExtractor(NodeMetadataExtractor):
|
class RgthreePowerLoraLoaderExtractor(NodeMetadataExtractor):
|
||||||
"""Extract LoRA metadata from rgthree Power Lora Loader.
|
"""Extract LoRA metadata from rgthree Power Lora Loader.
|
||||||
|
|
||||||
@@ -1154,6 +1359,28 @@ class CR_ApplyControlNetStackExtractor(NodeMetadataExtractor):
|
|||||||
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
|
metadata[PROMPTS][node_id]["positive_encoded"] = transformed_positive
|
||||||
metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative
|
metadata[PROMPTS][node_id]["negative_encoded"] = transformed_negative
|
||||||
|
|
||||||
|
class MetadataOverwriteExtractor(NodeMetadataExtractor):
|
||||||
|
"""Extract manually specified metadata from MetadataOverwriteLM node.
|
||||||
|
|
||||||
|
Stores truthy input values under the OVERWRITE category so that
|
||||||
|
extract_generation_params can merge them over the inferred params.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def extract(node_id, inputs, outputs, metadata):
|
||||||
|
if not inputs:
|
||||||
|
return
|
||||||
|
|
||||||
|
overwrite_params = collect_overwrite_params(inputs)
|
||||||
|
|
||||||
|
if overwrite_params:
|
||||||
|
metadata.setdefault(OVERWRITE, {})
|
||||||
|
metadata[OVERWRITE][node_id] = {
|
||||||
|
"parameters": overwrite_params,
|
||||||
|
"node_id": node_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
# Registry of node-specific extractors
|
# Registry of node-specific extractors
|
||||||
# Keys are node class names
|
# Keys are node class names
|
||||||
NODE_EXTRACTORS = {
|
NODE_EXTRACTORS = {
|
||||||
@@ -1165,6 +1392,8 @@ NODE_EXTRACTORS = {
|
|||||||
"ClownsharKSampler_Beta": SamplerExtractor,
|
"ClownsharKSampler_Beta": SamplerExtractor,
|
||||||
"TSC_KSampler": TSCKSamplerExtractor, # Efficient Nodes
|
"TSC_KSampler": TSCKSamplerExtractor, # Efficient Nodes
|
||||||
"TSC_KSamplerAdvanced": TSCKSamplerAdvancedExtractor, # Efficient Nodes
|
"TSC_KSamplerAdvanced": TSCKSamplerAdvancedExtractor, # Efficient Nodes
|
||||||
|
"KreaTwoStageSampler": KreaTwoStageSamplerExtractor, # Auryg/Krea-2-Two-Stage-Sampler
|
||||||
|
"KreaThreeStageSampler": KreaTwoStageSamplerExtractor, # Auryg/Krea-2-Two-Stage-Sampler
|
||||||
"KSamplerBasicPipe": KSamplerBasicPipeExtractor, # comfyui-impact-pack
|
"KSamplerBasicPipe": KSamplerBasicPipeExtractor, # comfyui-impact-pack
|
||||||
"KSamplerAdvancedBasicPipe": KSamplerAdvancedBasicPipeExtractor, # comfyui-impact-pack
|
"KSamplerAdvancedBasicPipe": KSamplerAdvancedBasicPipeExtractor, # comfyui-impact-pack
|
||||||
"KSampler_inspire_pipe": KSamplerBasicPipeExtractor, # comfyui-inspire-pack
|
"KSampler_inspire_pipe": KSamplerBasicPipeExtractor, # comfyui-inspire-pack
|
||||||
@@ -1216,10 +1445,13 @@ NODE_EXTRACTORS = {
|
|||||||
"GetNode": GetNodeExtractor,
|
"GetNode": GetNodeExtractor,
|
||||||
# Latent
|
# Latent
|
||||||
"EmptyLatentImage": ImageSizeExtractor,
|
"EmptyLatentImage": ImageSizeExtractor,
|
||||||
|
"KreaDualResolutionSelector": KreaDualResolutionSelectorExtractor, # Auryg/Krea-2-Two-Stage-Sampler
|
||||||
# Flux
|
# Flux
|
||||||
"FluxGuidance": FluxGuidanceExtractor, # Add FluxGuidance
|
"FluxGuidance": FluxGuidanceExtractor, # Add FluxGuidance
|
||||||
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
"CFGGuider": CFGGuiderExtractor, # Add CFGGuider
|
||||||
# Image
|
# Image
|
||||||
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
"VAEDecode": VAEDecodeExtractor, # Added VAEDecode extractor
|
||||||
|
# Metadata overwrite
|
||||||
|
"MetadataOverwriteLM": MetadataOverwriteExtractor,
|
||||||
# Add other nodes as needed
|
# Add other nodes as needed
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,51 @@
|
|||||||
|
"""Shared helpers for Metadata Overwrite node metadata collection.
|
||||||
|
|
||||||
|
Used by both the MetadataOverwriteLM node (execution time) and the
|
||||||
|
MetadataOverwriteExtractor (hook time) so the conversion/filtering logic
|
||||||
|
cannot drift between the two paths.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict
|
||||||
|
|
||||||
|
from ..utils.utils import model_patcher_to_name, sampler_object_to_name
|
||||||
|
from .constants import CLIP_SKIP_SENTINEL, METADATA_OVERWRITE_FIELDS
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def collect_overwrite_params(values: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Convert node input values into non-default overwrite parameters.
|
||||||
|
|
||||||
|
For most fields, a falsy value (empty string, 0) means "not set" and is
|
||||||
|
skipped. clip_skip uses a dedicated sentinel (-25) so that a wired value
|
||||||
|
of 0 is preserved. The ``model`` field accepts either a manual string or
|
||||||
|
a wired MODEL (ModelPatcher) connection; in the latter case the source
|
||||||
|
model name is extracted from the patcher's ``cached_patcher_init`` and
|
||||||
|
stored as a ComfyUI-style relative path. The ``sampler`` field likewise
|
||||||
|
accepts a manual string or a wired SAMPLER (KSAMPLER) connection, from
|
||||||
|
which the sampler name is extracted via the sampler function's name.
|
||||||
|
"""
|
||||||
|
result: Dict[str, Any] = {}
|
||||||
|
for key in METADATA_OVERWRITE_FIELDS:
|
||||||
|
value = values.get(key)
|
||||||
|
if key == "model" and not isinstance(value, str):
|
||||||
|
value = model_patcher_to_name(value)
|
||||||
|
if value is None:
|
||||||
|
logger.warning(
|
||||||
|
"Could not extract model name from wired MODEL input "
|
||||||
|
"(no cached_patcher_init); model metadata overwrite skipped"
|
||||||
|
)
|
||||||
|
elif key == "sampler" and not isinstance(value, str):
|
||||||
|
value = sampler_object_to_name(value)
|
||||||
|
if value is None:
|
||||||
|
logger.warning(
|
||||||
|
"Could not extract sampler name from wired SAMPLER input "
|
||||||
|
"(unrecognized sampler function); sampler metadata overwrite skipped"
|
||||||
|
)
|
||||||
|
if key == "clip_skip":
|
||||||
|
if value != CLIP_SKIP_SENTINEL:
|
||||||
|
result[key] = value
|
||||||
|
elif value:
|
||||||
|
result[key] = value
|
||||||
|
return result
|
||||||
@@ -43,7 +43,7 @@ SCANNER_GETTER_NAMES = tuple(SCANNER_TYPE_MAP.keys())
|
|||||||
|
|
||||||
async def _find_model_entry(
|
async def _find_model_entry(
|
||||||
model_path: str,
|
model_path: str,
|
||||||
) -> tuple[object, object, str | None] | tuple[None, None, None]:
|
) -> tuple[Any, object, str | None] | tuple[None, None, None]:
|
||||||
"""Iterate all scanners and return the first (scanner, entry, getter_name)
|
"""Iterate all scanners and return the first (scanner, entry, getter_name)
|
||||||
that owns *model_path*. Returns ``(None, None, None)`` when no scanner
|
that owns *model_path*. Returns ``(None, None, None)`` when no scanner
|
||||||
claims it.
|
claims it.
|
||||||
@@ -73,7 +73,7 @@ async def _find_model_entry(
|
|||||||
|
|
||||||
async def _find_scanner_for_model(
|
async def _find_scanner_for_model(
|
||||||
model_path: str,
|
model_path: str,
|
||||||
) -> tuple[object, object] | tuple[None, None]:
|
) -> tuple[Any, object] | tuple[None, None]:
|
||||||
"""Find the (scanner, cache_entry) responsible for *model_path*."""
|
"""Find the (scanner, cache_entry) responsible for *model_path*."""
|
||||||
scanner, entry, _ = await _find_model_entry(model_path)
|
scanner, entry, _ = await _find_model_entry(model_path)
|
||||||
return scanner, entry
|
return scanner, entry
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import List, Tuple
|
import os
|
||||||
import comfy.sd # type: ignore
|
from typing import Any, List, Tuple
|
||||||
import folder_paths # type: ignore
|
import comfy.sd # pyright: ignore[reportMissingImports]
|
||||||
|
import folder_paths # pyright: ignore[reportMissingImports]
|
||||||
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
|
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -18,9 +19,9 @@ class CheckpointLoaderLM:
|
|||||||
CATEGORY = "Lora Manager/loaders"
|
CATEGORY = "Lora Manager/loaders"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(cls):
|
||||||
# Get list of checkpoint names from scanner (includes extra folder paths)
|
# Get list of checkpoint names from scanner (includes extra folder paths)
|
||||||
checkpoint_names = s._get_checkpoint_names()
|
checkpoint_names = cls._get_checkpoint_names()
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"ckpt_name": (
|
"ckpt_name": (
|
||||||
@@ -58,7 +59,10 @@ class CheckpointLoaderLM:
|
|||||||
for item in cache.raw_data:
|
for item in cache.raw_data:
|
||||||
if item.get("sub_type") == "checkpoint":
|
if item.get("sub_type") == "checkpoint":
|
||||||
file_path = item.get("file_path", "")
|
file_path = item.get("file_path", "")
|
||||||
if file_path:
|
# Only offer models that still exist on disk so ComfyUI
|
||||||
|
# flags missing checkpoints at queue time via
|
||||||
|
# "value not in list" (the scanner cache can be stale).
|
||||||
|
if file_path and os.path.exists(file_path):
|
||||||
# Format using relative path with OS-native separator
|
# Format using relative path with OS-native separator
|
||||||
formatted_name = _format_model_name_for_comfyui(
|
formatted_name = _format_model_name_for_comfyui(
|
||||||
file_path, model_roots
|
file_path, model_roots
|
||||||
@@ -89,7 +93,7 @@ class CheckpointLoaderLM:
|
|||||||
logger.error(f"Error getting checkpoint names: {e}")
|
logger.error(f"Error getting checkpoint names: {e}")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
def load_checkpoint(self, ckpt_name: str) -> Tuple:
|
def load_checkpoint(self, ckpt_name: str) -> Tuple[Any, Any, Any]:
|
||||||
"""Load a checkpoint by name, supporting extra folder paths
|
"""Load a checkpoint by name, supporting extra folder paths
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
|
|||||||
@@ -0,0 +1,123 @@
|
|||||||
|
"""Create Hook LoRA (LoraManager) — multi-LoRA hook node compatible with ComfyUI's built-in hook pipeline.
|
||||||
|
|
||||||
|
Produces ``("HOOKS",)`` output that chains seamlessly with downstream hook consumers
|
||||||
|
(ConditioningSetProperties, SetHookKeyframes, CombineHooks, SetClipHooks, etc.).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
|
||||||
|
from ..utils.utils import get_lora_info_absolute
|
||||||
|
from .utils import (
|
||||||
|
FlexibleOptionalInputType,
|
||||||
|
any_type,
|
||||||
|
apply_lora_syntax_format,
|
||||||
|
get_loras_list,
|
||||||
|
validate_lora_entries,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class CreateHookLoraLM:
|
||||||
|
NAME = "Create Hook LoRA (LoraManager)"
|
||||||
|
CATEGORY = "Lora Manager/hooks"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"text": (
|
||||||
|
"AUTOCOMPLETE_TEXT_LORAS",
|
||||||
|
{
|
||||||
|
"placeholder": "Search LoRAs to add...",
|
||||||
|
"tooltip": (
|
||||||
|
"Search and select LoRAs. Each LoRA gets its own "
|
||||||
|
"model/clip strength. Hooks chain with prev_hooks."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
},
|
||||||
|
"optional": FlexibleOptionalInputType(any_type),
|
||||||
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def VALIDATE_INPUTS(cls, loras=None):
|
||||||
|
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||||
|
return validate_lora_entries({"loras": loras}) or True
|
||||||
|
|
||||||
|
RETURN_TYPES = ("HOOKS", "STRING", "STRING")
|
||||||
|
RETURN_NAMES = ("HOOKS", "trigger_words", "active_loras")
|
||||||
|
FUNCTION = "create_hook"
|
||||||
|
|
||||||
|
def create_hook(self, text: str, **kwargs):
|
||||||
|
"""Create a HookGroup from the selected LoRAs, chained with prev_hooks.
|
||||||
|
|
||||||
|
Each active LoRA from the widget is loaded and wrapped in a WeightHook
|
||||||
|
via :func:`comfy.hooks.create_hook_lora`. All hooks are combined into a
|
||||||
|
single group and returned alongside trigger words and a human-readable
|
||||||
|
summary of the active LoRAs.
|
||||||
|
"""
|
||||||
|
del text # used by the frontend widget only
|
||||||
|
|
||||||
|
# Lazy imports: comfy is not available in CI/test environment at module level
|
||||||
|
import comfy.hooks # pyright: ignore[reportMissingImports] # noqa: C0415
|
||||||
|
import comfy.utils # pyright: ignore[reportMissingImports] # noqa: C0415
|
||||||
|
|
||||||
|
prev_hooks: comfy.hooks.HookGroup | None = kwargs.get("prev_hooks")
|
||||||
|
|
||||||
|
hook_group = prev_hooks.clone() if prev_hooks is not None else comfy.hooks.HookGroup()
|
||||||
|
|
||||||
|
all_trigger_words: list[str] = []
|
||||||
|
active_loras: list[tuple[str, float, float]] = []
|
||||||
|
|
||||||
|
for lora in get_loras_list(kwargs):
|
||||||
|
if not lora.get("active", False):
|
||||||
|
continue
|
||||||
|
|
||||||
|
lora_name = apply_lora_syntax_format(lora["name"])
|
||||||
|
model_strength = float(lora["strength"])
|
||||||
|
clip_strength = float(lora.get("clipStrength", model_strength))
|
||||||
|
|
||||||
|
# Skip useless no-op entries (both strengths are zero)
|
||||||
|
if model_strength == 0.0 and clip_strength == 0.0:
|
||||||
|
continue
|
||||||
|
|
||||||
|
lora_path, trigger_words = get_lora_info_absolute(lora_name)
|
||||||
|
if not lora_path or not os.path.isfile(lora_path):
|
||||||
|
logger.warning("LoRA '%s' not found — skipping", lora_name)
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
lora_weights = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||||
|
|
||||||
|
lora_hooks = comfy.hooks.create_hook_lora(
|
||||||
|
lora=lora_weights,
|
||||||
|
strength_model=model_strength,
|
||||||
|
strength_clip=clip_strength,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Failed to load LoRA '%s' — skipping", lora_name)
|
||||||
|
continue
|
||||||
|
hook_group = hook_group.clone_and_combine(lora_hooks)
|
||||||
|
|
||||||
|
active_loras.append((lora_name, model_strength, clip_strength))
|
||||||
|
all_trigger_words.extend(trigger_words)
|
||||||
|
|
||||||
|
# Format trigger words (group mode separator)
|
||||||
|
trigger_words_text = ",, ".join(all_trigger_words) if all_trigger_words else ""
|
||||||
|
|
||||||
|
# Format active LoRAs summary
|
||||||
|
formatted_loras = []
|
||||||
|
for name, model_s, clip_s in active_loras:
|
||||||
|
if abs(model_s - clip_s) > 0.001:
|
||||||
|
formatted_loras.append(
|
||||||
|
f"<lora:{name}:{model_s}:{clip_s}>"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
formatted_loras.append(f"<lora:{name}:{model_s}>")
|
||||||
|
active_loras_text = " ".join(formatted_loras)
|
||||||
|
|
||||||
|
return (hook_group, trigger_words_text, active_loras_text)
|
||||||
@@ -1,8 +1,8 @@
|
|||||||
import importlib
|
import importlib
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
import comfy.sd # type: ignore
|
import comfy.sd # pyright: ignore[reportMissingImports]
|
||||||
import comfy.utils # type: ignore
|
import comfy.utils # pyright: ignore[reportMissingImports]
|
||||||
|
|
||||||
from ..utils.utils import get_lora_info_absolute
|
from ..utils.utils import get_lora_info_absolute
|
||||||
from .utils import (
|
from .utils import (
|
||||||
@@ -14,6 +14,7 @@ from .utils import (
|
|||||||
get_loras_list,
|
get_loras_list,
|
||||||
nunchaku_load_lora,
|
nunchaku_load_lora,
|
||||||
parse_lora_syntax,
|
parse_lora_syntax,
|
||||||
|
validate_lora_entries,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -142,6 +143,11 @@ class LoraLoaderLM:
|
|||||||
"optional": FlexibleOptionalInputType(any_type),
|
"optional": FlexibleOptionalInputType(any_type),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def VALIDATE_INPUTS(cls, loras=None):
|
||||||
|
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||||
|
return validate_lora_entries({"loras": loras}) or True
|
||||||
|
|
||||||
RETURN_TYPES = ("MODEL", "CLIP", "STRING", "STRING")
|
RETURN_TYPES = ("MODEL", "CLIP", "STRING", "STRING")
|
||||||
RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras")
|
RETURN_NAMES = ("MODEL", "CLIP", "trigger_words", "loaded_loras")
|
||||||
FUNCTION = "load_loras"
|
FUNCTION = "load_loras"
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ and tracks the last used combination for reuse.
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from ..utils.utils import get_lora_info
|
from ..utils.utils import get_lora_info
|
||||||
|
from .utils import validate_lora_entries
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -31,6 +32,11 @@ class LoraRandomizerLM:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def VALIDATE_INPUTS(cls, loras=None):
|
||||||
|
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||||
|
return validate_lora_entries({"loras": loras}) or True
|
||||||
|
|
||||||
RETURN_TYPES = ("LORA_STACK",)
|
RETURN_TYPES = ("LORA_STACK",)
|
||||||
RETURN_NAMES = ("LORA_STACK",)
|
RETURN_NAMES = ("LORA_STACK",)
|
||||||
|
|
||||||
|
|||||||
@@ -1,26 +1,102 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
import re
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
_STACK_INPUT_PATTERN = re.compile(r"^lora_stack(?:_([ab])|(\d+))$")
|
||||||
|
|
||||||
|
|
||||||
|
def _is_stack_input(name: str) -> bool:
|
||||||
|
return bool(_STACK_INPUT_PATTERN.match(name))
|
||||||
|
|
||||||
|
|
||||||
|
def _stack_slot_number(name: str) -> int:
|
||||||
|
"""Numeric slot used to order stack inputs; legacy a/b map to 1/2."""
|
||||||
|
match = _STACK_INPUT_PATTERN.match(name)
|
||||||
|
if not match:
|
||||||
|
return -1
|
||||||
|
letter, digits = match.group(1), match.group(2)
|
||||||
|
if digits is not None:
|
||||||
|
return int(digits)
|
||||||
|
return 1 if letter == "a" else 2
|
||||||
|
|
||||||
|
|
||||||
|
class _LoraStackOptionalInputs:
|
||||||
|
"""Lookup that preserves explicit optional inputs and dynamic lora_stack slots."""
|
||||||
|
|
||||||
|
def __init__(self, explicit_inputs: dict[str, tuple[str, dict[str, Any]]]) -> None:
|
||||||
|
self._explicit_inputs = explicit_inputs
|
||||||
|
|
||||||
|
def __contains__(self, item: object) -> bool:
|
||||||
|
if not isinstance(item, str):
|
||||||
|
return False
|
||||||
|
return item in self._explicit_inputs or _is_stack_input(item)
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> tuple[str, dict[str, Any]]:
|
||||||
|
if key in self._explicit_inputs:
|
||||||
|
return self._explicit_inputs[key]
|
||||||
|
if _is_stack_input(key):
|
||||||
|
return (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
raise KeyError(key)
|
||||||
|
|
||||||
|
|
||||||
class LoraStackCombinerLM:
|
class LoraStackCombinerLM:
|
||||||
NAME = "Lora Stack Combiner (LoraManager)"
|
NAME = "Lora Stack Combiner (LoraManager)"
|
||||||
CATEGORY = "Lora Manager/stackers"
|
CATEGORY = "Lora Manager/stackers"
|
||||||
|
DESCRIPTION = (
|
||||||
|
"Combines multiple LoRA stacks into a single stack. "
|
||||||
|
"Supports dynamic inputs: connect a stack to add more inputs."
|
||||||
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
|
optional_inputs: dict[str, tuple[str, dict[str, Any]]] = {
|
||||||
|
"lora_stack1": (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"lora_stack2": (
|
||||||
|
"LORA_STACK",
|
||||||
|
{
|
||||||
|
"tooltip": "A LoRA stack to combine. Connect to add more inputs.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
stack = inspect.stack()
|
||||||
|
if len(stack) > 2 and stack[2].function == "get_input_info":
|
||||||
|
optional_inputs = _LoraStackOptionalInputs(optional_inputs) # pyright: ignore[reportAssignmentType]
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {},
|
||||||
"lora_stack_a": ("LORA_STACK",),
|
"optional": optional_inputs,
|
||||||
"lora_stack_b": ("LORA_STACK",),
|
|
||||||
},
|
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("LORA_STACK",)
|
RETURN_TYPES = ("LORA_STACK",)
|
||||||
RETURN_NAMES = ("LORA_STACK",)
|
RETURN_NAMES = ("LORA_STACK",)
|
||||||
FUNCTION = "combine_stacks"
|
FUNCTION = "combine_stacks"
|
||||||
|
|
||||||
def combine_stacks(self, lora_stack_a, lora_stack_b):
|
def combine_stacks(self, lora_stack1=None, lora_stack2=None, **kwargs):
|
||||||
combined_stack = []
|
stacks = {
|
||||||
|
"lora_stack1": lora_stack1,
|
||||||
|
"lora_stack2": lora_stack2,
|
||||||
|
}
|
||||||
|
for key, value in kwargs.items():
|
||||||
|
if _is_stack_input(key) and value is not None:
|
||||||
|
stacks[key] = value
|
||||||
|
|
||||||
if lora_stack_a:
|
combined_stack = []
|
||||||
combined_stack.extend(lora_stack_a)
|
for key in sorted(stacks, key=_stack_slot_number):
|
||||||
if lora_stack_b:
|
stack = stacks[key]
|
||||||
combined_stack.extend(lora_stack_b)
|
if stack:
|
||||||
|
combined_stack.extend(stack)
|
||||||
|
|
||||||
return (combined_stack,)
|
return (combined_stack,)
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
import os
|
import os
|
||||||
from ..utils.utils import get_lora_info
|
from ..utils.utils import get_lora_info
|
||||||
from .utils import FlexibleOptionalInputType, any_type, apply_lora_syntax_format, extract_lora_name, get_loras_list
|
from .utils import FlexibleOptionalInputType, any_type, apply_lora_syntax_format, extract_lora_name, get_loras_list, validate_lora_entries
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
@@ -22,6 +22,11 @@ class LoraStackerLM:
|
|||||||
"optional": FlexibleOptionalInputType(any_type),
|
"optional": FlexibleOptionalInputType(any_type),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def VALIDATE_INPUTS(cls, loras=None):
|
||||||
|
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||||
|
return validate_lora_entries({"loras": loras}) or True
|
||||||
|
|
||||||
RETURN_TYPES = ("LORA_STACK", "STRING", "STRING")
|
RETURN_TYPES = ("LORA_STACK", "STRING", "STRING")
|
||||||
RETURN_NAMES = ("LORA_STACK", "trigger_words", "active_loras")
|
RETURN_NAMES = ("LORA_STACK", "trigger_words", "active_loras")
|
||||||
FUNCTION = "stack_loras"
|
FUNCTION = "stack_loras"
|
||||||
|
|||||||
@@ -0,0 +1,179 @@
|
|||||||
|
"""Metadata Overwrite node — allows users to manually specify generation parameters
|
||||||
|
that override the automatically collected/inferred metadata.
|
||||||
|
|
||||||
|
Most inputs have falsy defaults (empty string / 0) which are skipped.
|
||||||
|
clip_skip uses a sentinel default (-25) so that a wired value of 0 is
|
||||||
|
preserved — both ComfyUI and A1111 conventions have no meaningful 0 value,
|
||||||
|
but users may wire 0 to express "no clip skip / default".
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from ..metadata_collector.constants import CLIP_SKIP_SENTINEL as _CLIP_SKIP_SENTINEL
|
||||||
|
from ..metadata_collector.overwrite_utils import collect_overwrite_params
|
||||||
|
|
||||||
|
|
||||||
|
class MetadataOverwriteLM:
|
||||||
|
NAME = "Metadata Overwrite (LoraManager)"
|
||||||
|
CATEGORY = "Lora Manager/utils"
|
||||||
|
DESCRIPTION = (
|
||||||
|
"Manually specify generation parameters to override automatically collected "
|
||||||
|
"metadata. Only filled/connected inputs will take effect — empty defaults "
|
||||||
|
"are ignored."
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls) -> dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"optional": {
|
||||||
|
"prompt": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"multiline": True,
|
||||||
|
"tooltip": "Positive prompt. Only overwrites when non-empty.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"negative_prompt": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"multiline": True,
|
||||||
|
"tooltip": "Negative prompt. Only overwrites when non-empty.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"seed": (
|
||||||
|
"INT",
|
||||||
|
{
|
||||||
|
"default": 0,
|
||||||
|
"min": 0,
|
||||||
|
"max": 0xFFFFFFFFFFFFFFFF,
|
||||||
|
"control_after_generate": False,
|
||||||
|
"tooltip": "Seed value. Only overwrites when > 0.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"steps": (
|
||||||
|
"INT",
|
||||||
|
{
|
||||||
|
"default": 0,
|
||||||
|
"min": 0,
|
||||||
|
"max": 10000,
|
||||||
|
"tooltip": "Number of steps. Only overwrites when > 0.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"cfg_scale": (
|
||||||
|
"FLOAT",
|
||||||
|
{
|
||||||
|
"default": 0.0,
|
||||||
|
"min": 0.0,
|
||||||
|
"max": 100.0,
|
||||||
|
"tooltip": "CFG scale. Only overwrites when > 0.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"sampler": (
|
||||||
|
"STRING,SAMPLER",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"widgetType": "STRING",
|
||||||
|
"tooltip": (
|
||||||
|
"Sampler name. Fill in the name manually or "
|
||||||
|
"connect a SAMPLER output (e.g. KSamplerSelect) "
|
||||||
|
"— the sampler name is then extracted "
|
||||||
|
"automatically. Note: ddim is recorded as "
|
||||||
|
"euler (ComfyUI internal representation). "
|
||||||
|
"Only overwrites when non-empty."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"scheduler": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"tooltip": "Scheduler name. Only overwrites when non-empty.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"model": (
|
||||||
|
"STRING,MODEL",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"widgetType": "STRING",
|
||||||
|
"tooltip": (
|
||||||
|
"The checkpoint or diffusion model (UNet) used "
|
||||||
|
"for generation. Fill in the name manually or "
|
||||||
|
"connect a MODEL output — the model name is then "
|
||||||
|
"extracted automatically. Only overwrites when "
|
||||||
|
"non-empty."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"loras": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"multiline": True,
|
||||||
|
"tooltip": (
|
||||||
|
"LoRA syntax, e.g. <lora:name:strength> "
|
||||||
|
"or <lora:name:model_strength:clip_strength>, "
|
||||||
|
"separated by spaces. Only overwrites when non-empty."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"size": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"tooltip": (
|
||||||
|
"Image size in WIDTHxHEIGHT format (e.g. 512x768). "
|
||||||
|
"Only overwrites when non-empty."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"clip_skip": (
|
||||||
|
"INT",
|
||||||
|
{
|
||||||
|
"default": _CLIP_SKIP_SENTINEL,
|
||||||
|
"min": -25,
|
||||||
|
"max": 24,
|
||||||
|
"tooltip": (
|
||||||
|
"Clip skip (ComfyUI: -24..-1, A1111: 1+). "
|
||||||
|
"Default -25 means not set — any other value "
|
||||||
|
"overwrites."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"additional_data": (
|
||||||
|
"STRING",
|
||||||
|
{
|
||||||
|
"default": "",
|
||||||
|
"multiline": True,
|
||||||
|
"tooltip": (
|
||||||
|
"Additional data to embed in the image metadata. "
|
||||||
|
"Inserted between Clip skip and Model hash in the "
|
||||||
|
"A1111-compatible parameters string. "
|
||||||
|
'Example: "Copyright": "Some license info"'
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("METADATA",)
|
||||||
|
RETURN_NAMES = ("metadata",)
|
||||||
|
FUNCTION = "collect_metadata"
|
||||||
|
OUTPUT_NODE = True
|
||||||
|
|
||||||
|
def collect_metadata(self, **kwargs: Any) -> tuple[dict[str, Any]]:
|
||||||
|
"""Collect non-default input values into a metadata dict.
|
||||||
|
|
||||||
|
For most fields, a falsy value (empty string, 0) means "not set"
|
||||||
|
and is skipped. clip_skip uses a dedicated sentinel (-25) so that
|
||||||
|
a wired value of 0 is preserved and reaches the metadata pipeline.
|
||||||
|
|
||||||
|
The ``model`` field accepts either a manual string or a wired MODEL
|
||||||
|
(ModelPatcher) connection; in the latter case the underlying model
|
||||||
|
name is extracted from the patcher's ``cached_patcher_init`` and
|
||||||
|
stored as a ComfyUI-style relative path. The ``sampler`` field
|
||||||
|
likewise accepts a manual string or a wired SAMPLER (KSAMPLER)
|
||||||
|
connection, from which the sampler name is extracted automatically.
|
||||||
|
"""
|
||||||
|
return (collect_overwrite_params(kwargs),)
|
||||||
+12
-13
@@ -15,15 +15,15 @@ import os
|
|||||||
import re
|
import re
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Optional, Tuple, Union
|
from typing import Any, Dict, List, Optional, Tuple, Union, cast
|
||||||
|
|
||||||
import comfy.utils # type: ignore
|
import comfy.utils # pyright: ignore[reportMissingImports]
|
||||||
import folder_paths # type: ignore
|
import folder_paths # pyright: ignore[reportMissingImports]
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from safetensors import safe_open
|
from safetensors import safe_open
|
||||||
|
|
||||||
from nunchaku.lora.flux.nunchaku_converter import (
|
from nunchaku.lora.flux.nunchaku_converter import ( # pyright: ignore[reportMissingTypeStubs]
|
||||||
pack_lowrank_weight,
|
pack_lowrank_weight,
|
||||||
unpack_lowrank_weight,
|
unpack_lowrank_weight,
|
||||||
)
|
)
|
||||||
@@ -87,10 +87,6 @@ def _rename_layer_underscore_layer_name(old_name: str) -> str:
|
|||||||
return new_name
|
return new_name
|
||||||
|
|
||||||
|
|
||||||
def _is_indexable_module(module):
|
|
||||||
return isinstance(module, (nn.ModuleList, nn.Sequential, list, tuple))
|
|
||||||
|
|
||||||
|
|
||||||
def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]:
|
def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]:
|
||||||
if not name:
|
if not name:
|
||||||
return model
|
return model
|
||||||
@@ -100,7 +96,7 @@ def _get_module_by_name(model: nn.Module, name: str) -> Optional[nn.Module]:
|
|||||||
continue
|
continue
|
||||||
if hasattr(module, part):
|
if hasattr(module, part):
|
||||||
module = getattr(module, part)
|
module = getattr(module, part)
|
||||||
elif part.isdigit() and _is_indexable_module(module):
|
elif part.isdigit() and isinstance(module, (nn.ModuleList, nn.Sequential, list, tuple)):
|
||||||
try:
|
try:
|
||||||
module = module[int(part)]
|
module = module[int(part)]
|
||||||
except (IndexError, TypeError):
|
except (IndexError, TypeError):
|
||||||
@@ -267,7 +263,9 @@ def _handle_proj_out_split(lora_dict: Dict[str, Dict[str, torch.Tensor]], base_k
|
|||||||
return result, consumed
|
return result, consumed
|
||||||
|
|
||||||
|
|
||||||
def _apply_lora_to_module(module: nn.Module, a_tensor: torch.Tensor, b_tensor: torch.Tensor, module_name: str, model: nn.Module) -> None:
|
def _apply_lora_to_module(module: Any, a_tensor: torch.Tensor, b_tensor: torch.Tensor, module_name: str, model: Any) -> None:
|
||||||
|
# These modules are dynamic torch containers; monkey-patched attributes
|
||||||
|
# below are set at runtime, so the module/model types are deliberately Any.
|
||||||
if not hasattr(module, "in_features") or not hasattr(module, "out_features"):
|
if not hasattr(module, "in_features") or not hasattr(module, "out_features"):
|
||||||
raise ValueError(f"{module_name}: unsupported module without in/out features")
|
raise ValueError(f"{module_name}: unsupported module without in/out features")
|
||||||
if a_tensor.shape[1] != module.in_features or b_tensor.shape[0] != module.out_features:
|
if a_tensor.shape[1] != module.in_features or b_tensor.shape[0] != module.out_features:
|
||||||
@@ -336,7 +334,7 @@ def _apply_lora_to_module(module: nn.Module, a_tensor: torch.Tensor, b_tensor: t
|
|||||||
raise ValueError(f"{module_name}: unsupported module type {type(module)}")
|
raise ValueError(f"{module_name}: unsupported module type {type(module)}")
|
||||||
|
|
||||||
|
|
||||||
def reset_lora_v2(model: nn.Module) -> None:
|
def reset_lora_v2(model: Any) -> None:
|
||||||
slots = getattr(model, "_lora_slots", None)
|
slots = getattr(model, "_lora_slots", None)
|
||||||
if not slots:
|
if not slots:
|
||||||
return
|
return
|
||||||
@@ -344,6 +342,7 @@ def reset_lora_v2(model: nn.Module) -> None:
|
|||||||
module = _get_module_by_name(model, name)
|
module = _get_module_by_name(model, name)
|
||||||
if module is None:
|
if module is None:
|
||||||
continue
|
continue
|
||||||
|
module = cast(Any, module)
|
||||||
module_type = info.get("type", "nunchaku")
|
module_type = info.get("type", "nunchaku")
|
||||||
if module_type == "nunchaku":
|
if module_type == "nunchaku":
|
||||||
base_rank = info["base_rank"]
|
base_rank = info["base_rank"]
|
||||||
@@ -371,7 +370,7 @@ def reset_lora_v2(model: nn.Module) -> None:
|
|||||||
def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]], apply_awq_mod: bool = True) -> bool:
|
def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]], apply_awq_mod: bool = True) -> bool:
|
||||||
del apply_awq_mod # retained for interface compatibility
|
del apply_awq_mod # retained for interface compatibility
|
||||||
reset_lora_v2(model)
|
reset_lora_v2(model)
|
||||||
aggregated_weights: Dict[str, List[Dict[str, object]]] = defaultdict(list)
|
aggregated_weights: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
|
||||||
saw_supported_format = False
|
saw_supported_format = False
|
||||||
unresolved_targets = 0
|
unresolved_targets = 0
|
||||||
|
|
||||||
@@ -471,7 +470,7 @@ def compose_loras_v2(model: nn.Module, lora_configs: List[Tuple[Union[str, Path,
|
|||||||
class ComfyQwenImageWrapperLM(nn.Module):
|
class ComfyQwenImageWrapperLM(nn.Module):
|
||||||
def __init__(self, model: nn.Module, config=None, apply_awq_mod: bool = True):
|
def __init__(self, model: nn.Module, config=None, apply_awq_mod: bool = True):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.model = model
|
self.model: Any = model
|
||||||
self.config = {} if config is None else config
|
self.config = {} if config is None else config
|
||||||
self.dtype = next(model.parameters()).dtype
|
self.dtype = next(model.parameters()).dtype
|
||||||
self.loras: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]] = []
|
self.loras: List[Tuple[Union[str, Path, Dict[str, torch.Tensor]], float]] = []
|
||||||
|
|||||||
+2
-2
@@ -67,7 +67,7 @@ class PromptLM:
|
|||||||
|
|
||||||
stack = inspect.stack()
|
stack = inspect.stack()
|
||||||
if len(stack) > 2 and stack[2].function == "get_input_info":
|
if len(stack) > 2 and stack[2].function == "get_input_info":
|
||||||
optional_inputs = _PromptOptionalInputs(optional_inputs) # type: ignore[assignment]
|
optional_inputs = _PromptOptionalInputs(optional_inputs) # pyright: ignore[reportAssignmentType]
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
@@ -126,7 +126,7 @@ class PromptLM:
|
|||||||
else:
|
else:
|
||||||
prompt = expanded_text
|
prompt = expanded_text
|
||||||
|
|
||||||
from nodes import CLIPTextEncode # type: ignore
|
from nodes import CLIPTextEncode # pyright: ignore[reportMissingImports, reportAttributeAccessIssue]
|
||||||
|
|
||||||
conditioning = CLIPTextEncode().encode(clip, prompt)[0]
|
conditioning = CLIPTextEncode().encode(clip, prompt)[0]
|
||||||
return (conditioning, prompt)
|
return (conditioning, prompt)
|
||||||
|
|||||||
@@ -0,0 +1,214 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import random
|
||||||
|
from typing import Any, List, Optional, Tuple
|
||||||
|
import comfy.sd # pyright: ignore[reportMissingImports]
|
||||||
|
import folder_paths # pyright: ignore[reportMissingImports]
|
||||||
|
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class RandomCheckpointLoaderLM:
|
||||||
|
"""Checkpoint Loader that can randomly pick a checkpoint from the pool
|
||||||
|
|
||||||
|
Loads checkpoints from both standard ComfyUI folders and LoRA Manager's
|
||||||
|
extra folder paths. When select_at_random is enabled, ignores ckpt_name
|
||||||
|
and picks a random checkpoint (optionally filtered by base_model) on
|
||||||
|
every run.
|
||||||
|
"""
|
||||||
|
|
||||||
|
NAME = "Random Checkpoint Loader (LoraManager)"
|
||||||
|
CATEGORY = "Lora Manager/loaders"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
# Get list of checkpoint names from scanner (includes extra folder paths)
|
||||||
|
checkpoint_names = cls._get_checkpoint_names()
|
||||||
|
base_models = cls._get_available_base_models()
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"ckpt_name": (
|
||||||
|
checkpoint_names,
|
||||||
|
{"tooltip": "The name of the checkpoint (model) to load."},
|
||||||
|
),
|
||||||
|
"select_at_random": (
|
||||||
|
"BOOLEAN",
|
||||||
|
{
|
||||||
|
"default": False,
|
||||||
|
"tooltip": (
|
||||||
|
"Ignore ckpt_name and pick a random checkpoint from the "
|
||||||
|
"pool (optionally filtered by base_model) on every run."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"base_model": (
|
||||||
|
base_models,
|
||||||
|
{
|
||||||
|
"default": "Any",
|
||||||
|
"tooltip": "Restrict random selection to this base model. 'Any' uses the full pool.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("MODEL", "CLIP", "VAE", "STRING")
|
||||||
|
RETURN_NAMES = ("MODEL", "CLIP", "VAE", "model_name")
|
||||||
|
OUTPUT_TOOLTIPS = (
|
||||||
|
"The model used for denoising latents.",
|
||||||
|
"The CLIP model used for encoding text prompts.",
|
||||||
|
"The VAE model used for encoding and decoding images to and from latent space.",
|
||||||
|
"The name of the checkpoint that was loaded (useful when select_at_random is enabled).",
|
||||||
|
)
|
||||||
|
FUNCTION = "load_checkpoint"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def IS_CHANGED(cls, ckpt_name, select_at_random=False, base_model="Any"):
|
||||||
|
# Force re-execution on every run while randomizing, since the widget
|
||||||
|
# values themselves don't change between queue runs.
|
||||||
|
if select_at_random:
|
||||||
|
return float("nan")
|
||||||
|
return ckpt_name
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _run_async(coro_fn):
|
||||||
|
"""Run an async fetcher, handling the case where an event loop is already running."""
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
try:
|
||||||
|
asyncio.get_running_loop()
|
||||||
|
import concurrent.futures
|
||||||
|
|
||||||
|
def run_in_thread():
|
||||||
|
new_loop = asyncio.new_event_loop()
|
||||||
|
asyncio.set_event_loop(new_loop)
|
||||||
|
try:
|
||||||
|
return new_loop.run_until_complete(coro_fn())
|
||||||
|
finally:
|
||||||
|
new_loop.close()
|
||||||
|
|
||||||
|
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||||
|
future = executor.submit(run_in_thread)
|
||||||
|
return future.result()
|
||||||
|
except RuntimeError:
|
||||||
|
return asyncio.run(coro_fn())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_checkpoint_names(cls, base_model: Optional[str] = None) -> List[str]:
|
||||||
|
"""Get list of checkpoint names from scanner cache in ComfyUI format (relative path with extension)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
base_model: If given (and not "Any"), only include checkpoints matching this base model.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from ..services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
async def _get_names():
|
||||||
|
scanner = await ServiceRegistry.get_checkpoint_scanner()
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
|
||||||
|
# Get all model roots for calculating relative paths
|
||||||
|
model_roots = scanner.get_model_roots()
|
||||||
|
|
||||||
|
# Filter only checkpoint type (not diffusion_model) and format names
|
||||||
|
names = []
|
||||||
|
for item in cache.raw_data:
|
||||||
|
if item.get("sub_type") != "checkpoint":
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
base_model
|
||||||
|
and base_model != "Any"
|
||||||
|
and item.get("base_model") != base_model
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
file_path = item.get("file_path", "")
|
||||||
|
# Only offer models that still exist on disk so ComfyUI
|
||||||
|
# flags missing checkpoints at queue time via
|
||||||
|
# "value not in list" (the scanner cache can be stale).
|
||||||
|
if file_path and os.path.exists(file_path):
|
||||||
|
# Format using relative path with OS-native separator
|
||||||
|
formatted_name = _format_model_name_for_comfyui(
|
||||||
|
file_path, model_roots
|
||||||
|
)
|
||||||
|
if formatted_name:
|
||||||
|
names.append(formatted_name)
|
||||||
|
|
||||||
|
return sorted(names)
|
||||||
|
|
||||||
|
return cls._run_async(_get_names)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting checkpoint names: {e}")
|
||||||
|
return []
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_available_base_models(cls) -> List[str]:
|
||||||
|
"""Get distinct base_model values present among indexed checkpoints, for the random-selection filter."""
|
||||||
|
try:
|
||||||
|
from ..services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
async def _get_base_models():
|
||||||
|
scanner = await ServiceRegistry.get_checkpoint_scanner()
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
|
||||||
|
base_models = set()
|
||||||
|
for item in cache.raw_data:
|
||||||
|
if item.get("sub_type") != "checkpoint":
|
||||||
|
continue
|
||||||
|
base_model = item.get("base_model")
|
||||||
|
file_path = item.get("file_path", "")
|
||||||
|
if base_model and file_path and os.path.exists(file_path):
|
||||||
|
base_models.add(base_model)
|
||||||
|
|
||||||
|
return sorted(base_models)
|
||||||
|
|
||||||
|
return ["Any"] + cls._run_async(_get_base_models)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting available base models: {e}")
|
||||||
|
return ["Any"]
|
||||||
|
|
||||||
|
def load_checkpoint(
|
||||||
|
self,
|
||||||
|
ckpt_name: str,
|
||||||
|
select_at_random: bool = False,
|
||||||
|
base_model: str = "Any",
|
||||||
|
) -> Tuple[Any, Any, Any, str]:
|
||||||
|
"""Load a checkpoint by name, supporting extra folder paths
|
||||||
|
|
||||||
|
Args:
|
||||||
|
ckpt_name: The name of the checkpoint to load (relative path with extension)
|
||||||
|
select_at_random: If True, ignore ckpt_name and pick randomly from the pool
|
||||||
|
base_model: Restricts random selection to this base model ("Any" = no filter)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (MODEL, CLIP, VAE, model_name)
|
||||||
|
"""
|
||||||
|
if select_at_random:
|
||||||
|
pool = self._get_checkpoint_names(base_model)
|
||||||
|
if not pool:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"No checkpoints found for base model '{base_model}'. "
|
||||||
|
"Pick a different base model or disable 'select_at_random'."
|
||||||
|
)
|
||||||
|
ckpt_name = random.choice(pool)
|
||||||
|
logger.info(
|
||||||
|
f"[RandomCheckpointLoaderLM] Randomly selected checkpoint: {ckpt_name}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Get absolute path from cache using ComfyUI-style name
|
||||||
|
ckpt_path, metadata = get_checkpoint_info_absolute(ckpt_name)
|
||||||
|
|
||||||
|
if metadata is None:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Checkpoint '{ckpt_name}' not found in LoRA Manager cache. "
|
||||||
|
"Make sure the checkpoint is indexed and try again."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Load regular checkpoint using ComfyUI's API
|
||||||
|
logger.info(f"Loading checkpoint from: {ckpt_path}")
|
||||||
|
out = comfy.sd.load_checkpoint_guess_config(
|
||||||
|
ckpt_path,
|
||||||
|
output_vae=True,
|
||||||
|
output_clip=True,
|
||||||
|
embedding_directory=folder_paths.get_folder_paths("embeddings"),
|
||||||
|
)
|
||||||
|
return out[:3] + (ckpt_name,)
|
||||||
@@ -0,0 +1,326 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import random
|
||||||
|
from typing import Any, List, Optional, Tuple
|
||||||
|
import comfy.sd # pyright: ignore[reportMissingImports]
|
||||||
|
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _reload_gguf_unet(
|
||||||
|
unet_path: str, weight_dtype: str, disable_dynamic: bool = False
|
||||||
|
) -> object:
|
||||||
|
"""Reload a GGUF diffusion model from disk (cached_patcher_init factory).
|
||||||
|
|
||||||
|
Mirrors the GGUF branch of RandomUNETLoaderLM.load_unet so ModelPatcher
|
||||||
|
deepclone/dynamic machinery can rebuild GGUF models with the correct
|
||||||
|
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
|
||||||
|
with core ComfyUI loaders.
|
||||||
|
"""
|
||||||
|
loader = RandomUNETLoaderLM()
|
||||||
|
model, _unet_name = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
class RandomUNETLoaderLM:
|
||||||
|
"""UNET Loader that can randomly pick a diffusion model from the pool
|
||||||
|
|
||||||
|
Loads diffusion models/UNets from both standard ComfyUI folders and LoRA
|
||||||
|
Manager's extra folder paths. Supports both regular diffusion models and
|
||||||
|
GGUF format models. When select_at_random is enabled, ignores unet_name
|
||||||
|
and picks a random diffusion model (optionally filtered by base_model)
|
||||||
|
on every run.
|
||||||
|
"""
|
||||||
|
|
||||||
|
NAME = "Random Unet Loader (LoraManager)"
|
||||||
|
CATEGORY = "Lora Manager/loaders"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
# Get list of unet names from scanner (includes extra folder paths)
|
||||||
|
unet_names = cls._get_unet_names()
|
||||||
|
base_models = cls._get_available_base_models()
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"unet_name": (
|
||||||
|
unet_names,
|
||||||
|
{"tooltip": "The name of the diffusion model to load."},
|
||||||
|
),
|
||||||
|
"weight_dtype": (
|
||||||
|
["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2"],
|
||||||
|
{"tooltip": "The dtype to use for the model weights."},
|
||||||
|
),
|
||||||
|
"select_at_random": (
|
||||||
|
"BOOLEAN",
|
||||||
|
{
|
||||||
|
"default": False,
|
||||||
|
"tooltip": (
|
||||||
|
"Ignore unet_name and pick a random diffusion model from "
|
||||||
|
"the pool (optionally filtered by base_model) on every run."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"base_model": (
|
||||||
|
base_models,
|
||||||
|
{
|
||||||
|
"default": "Any",
|
||||||
|
"tooltip": "Restrict random selection to this base model. 'Any' uses the full pool.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("MODEL", "STRING")
|
||||||
|
RETURN_NAMES = ("MODEL", "model_name")
|
||||||
|
OUTPUT_TOOLTIPS = (
|
||||||
|
"The model used for denoising latents.",
|
||||||
|
"The name of the diffusion model that was loaded (useful when select_at_random is enabled).",
|
||||||
|
)
|
||||||
|
FUNCTION = "load_unet"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def IS_CHANGED(
|
||||||
|
cls, unet_name, weight_dtype, select_at_random=False, base_model="Any"
|
||||||
|
):
|
||||||
|
# Force re-execution on every run while randomizing, since the widget
|
||||||
|
# values themselves don't change between queue runs.
|
||||||
|
if select_at_random:
|
||||||
|
return float("nan")
|
||||||
|
return unet_name
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _run_async(coro_fn):
|
||||||
|
"""Run an async fetcher, handling the case where an event loop is already running."""
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
try:
|
||||||
|
asyncio.get_running_loop()
|
||||||
|
import concurrent.futures
|
||||||
|
|
||||||
|
def run_in_thread():
|
||||||
|
new_loop = asyncio.new_event_loop()
|
||||||
|
asyncio.set_event_loop(new_loop)
|
||||||
|
try:
|
||||||
|
return new_loop.run_until_complete(coro_fn())
|
||||||
|
finally:
|
||||||
|
new_loop.close()
|
||||||
|
|
||||||
|
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||||
|
future = executor.submit(run_in_thread)
|
||||||
|
return future.result()
|
||||||
|
except RuntimeError:
|
||||||
|
return asyncio.run(coro_fn())
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_unet_names(cls, base_model: Optional[str] = None) -> List[str]:
|
||||||
|
"""Get list of diffusion model names from scanner cache in ComfyUI format (relative path with extension)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
base_model: If given (and not "Any"), only include models matching this base model.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from ..services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
async def _get_names():
|
||||||
|
scanner = await ServiceRegistry.get_checkpoint_scanner()
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
|
||||||
|
# Get all model roots for calculating relative paths
|
||||||
|
model_roots = scanner.get_model_roots()
|
||||||
|
|
||||||
|
# Filter only diffusion_model type and format names
|
||||||
|
names = []
|
||||||
|
for item in cache.raw_data:
|
||||||
|
if item.get("sub_type") != "diffusion_model":
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
base_model
|
||||||
|
and base_model != "Any"
|
||||||
|
and item.get("base_model") != base_model
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
file_path = item.get("file_path", "")
|
||||||
|
# Only offer models that still exist on disk so ComfyUI
|
||||||
|
# flags missing diffusion models at queue time via
|
||||||
|
# "value not in list" (the scanner cache can be stale).
|
||||||
|
if file_path and os.path.exists(file_path):
|
||||||
|
# Format using relative path with OS-native separator
|
||||||
|
formatted_name = _format_model_name_for_comfyui(
|
||||||
|
file_path, model_roots
|
||||||
|
)
|
||||||
|
if formatted_name:
|
||||||
|
names.append(formatted_name)
|
||||||
|
|
||||||
|
return sorted(names)
|
||||||
|
|
||||||
|
return cls._run_async(_get_names)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting unet names: {e}")
|
||||||
|
return []
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_available_base_models(cls) -> List[str]:
|
||||||
|
"""Get distinct base_model values present among indexed diffusion models, for the random-selection filter."""
|
||||||
|
try:
|
||||||
|
from ..services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
async def _get_base_models():
|
||||||
|
scanner = await ServiceRegistry.get_checkpoint_scanner()
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
|
||||||
|
base_models = set()
|
||||||
|
for item in cache.raw_data:
|
||||||
|
if item.get("sub_type") != "diffusion_model":
|
||||||
|
continue
|
||||||
|
base_model = item.get("base_model")
|
||||||
|
file_path = item.get("file_path", "")
|
||||||
|
if base_model and file_path and os.path.exists(file_path):
|
||||||
|
base_models.add(base_model)
|
||||||
|
|
||||||
|
return sorted(base_models)
|
||||||
|
|
||||||
|
return ["Any"] + cls._run_async(_get_base_models)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting available base models: {e}")
|
||||||
|
return ["Any"]
|
||||||
|
|
||||||
|
def load_unet(
|
||||||
|
self,
|
||||||
|
unet_name: str,
|
||||||
|
weight_dtype: str,
|
||||||
|
select_at_random: bool = False,
|
||||||
|
base_model: str = "Any",
|
||||||
|
) -> Tuple[Any, ...]:
|
||||||
|
"""Load a diffusion model by name, supporting extra folder paths
|
||||||
|
|
||||||
|
Args:
|
||||||
|
unet_name: The name of the diffusion model to load (relative path with extension)
|
||||||
|
weight_dtype: The dtype to use for model weights
|
||||||
|
select_at_random: If True, ignore unet_name and pick randomly from the pool
|
||||||
|
base_model: Restricts random selection to this base model ("Any" = no filter)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (MODEL, model_name)
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if select_at_random:
|
||||||
|
pool = self._get_unet_names(base_model)
|
||||||
|
if not pool:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"No diffusion models found for base model '{base_model}'. "
|
||||||
|
"Pick a different base model or disable 'select_at_random'."
|
||||||
|
)
|
||||||
|
unet_name = random.choice(pool)
|
||||||
|
logger.info(
|
||||||
|
f"[RandomUNETLoaderLM] Randomly selected diffusion model: {unet_name}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Get absolute path from cache using ComfyUI-style name
|
||||||
|
unet_path, metadata = get_checkpoint_info_absolute(unet_name)
|
||||||
|
|
||||||
|
if metadata is None:
|
||||||
|
raise FileNotFoundError(
|
||||||
|
f"Diffusion model '{unet_name}' not found in LoRA Manager cache. "
|
||||||
|
"Make sure the model is indexed and try again."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check if it's a GGUF model
|
||||||
|
if unet_path.endswith(".gguf"):
|
||||||
|
return self._load_gguf_unet(unet_path, unet_name, weight_dtype)
|
||||||
|
|
||||||
|
# Load regular diffusion model using ComfyUI's API
|
||||||
|
logger.info(f"Loading diffusion model from: {unet_path}")
|
||||||
|
|
||||||
|
# Build model options based on weight_dtype
|
||||||
|
model_options = {}
|
||||||
|
if weight_dtype == "fp8_e4m3fn":
|
||||||
|
model_options["dtype"] = torch.float8_e4m3fn
|
||||||
|
elif weight_dtype == "fp8_e4m3fn_fast":
|
||||||
|
model_options["dtype"] = torch.float8_e4m3fn
|
||||||
|
model_options["fp8_optimizations"] = True
|
||||||
|
elif weight_dtype == "fp8_e5m2":
|
||||||
|
model_options["dtype"] = torch.float8_e5m2
|
||||||
|
|
||||||
|
model = comfy.sd.load_diffusion_model(unet_path, model_options=model_options)
|
||||||
|
return (model, unet_name)
|
||||||
|
|
||||||
|
def _load_gguf_unet(
|
||||||
|
self, unet_path: str, unet_name: str, weight_dtype: str
|
||||||
|
) -> Tuple[Any, ...]:
|
||||||
|
"""Load a GGUF format diffusion model
|
||||||
|
|
||||||
|
Args:
|
||||||
|
unet_path: Absolute path to the GGUF file
|
||||||
|
unet_name: Name of the model for error messages
|
||||||
|
weight_dtype: The dtype to use for model weights
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (MODEL, model_name)
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
from .gguf_import_helper import get_gguf_modules
|
||||||
|
|
||||||
|
# Get ComfyUI-GGUF modules using helper (handles various import scenarios)
|
||||||
|
try:
|
||||||
|
loader_module, ops_module, nodes_module = get_gguf_modules()
|
||||||
|
gguf_sd_loader = getattr(loader_module, "gguf_sd_loader")
|
||||||
|
GGMLOps = getattr(ops_module, "GGMLOps")
|
||||||
|
GGUFModelPatcher = getattr(nodes_module, "GGUFModelPatcher")
|
||||||
|
except RuntimeError as e:
|
||||||
|
raise RuntimeError(f"Cannot load GGUF model '{unet_name}'. {str(e)}")
|
||||||
|
|
||||||
|
logger.info(f"Loading GGUF diffusion model from: {unet_path}")
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Load GGUF state dict
|
||||||
|
sd, extra = gguf_sd_loader(unet_path)
|
||||||
|
|
||||||
|
# Prepare kwargs for metadata if supported
|
||||||
|
kwargs = {}
|
||||||
|
import inspect
|
||||||
|
|
||||||
|
valid_params = inspect.signature(
|
||||||
|
comfy.sd.load_diffusion_model_state_dict
|
||||||
|
).parameters
|
||||||
|
if "metadata" in valid_params:
|
||||||
|
kwargs["metadata"] = extra.get("metadata", {})
|
||||||
|
|
||||||
|
# Setup custom operations with GGUF support
|
||||||
|
ops = GGMLOps()
|
||||||
|
|
||||||
|
# Handle weight_dtype for GGUF models
|
||||||
|
if weight_dtype in ("default", None):
|
||||||
|
ops.Linear.dequant_dtype = None
|
||||||
|
elif weight_dtype in ["target"]:
|
||||||
|
ops.Linear.dequant_dtype = weight_dtype
|
||||||
|
else:
|
||||||
|
ops.Linear.dequant_dtype = getattr(torch, weight_dtype, None)
|
||||||
|
|
||||||
|
# Load the model
|
||||||
|
model = comfy.sd.load_diffusion_model_state_dict(
|
||||||
|
sd, model_options={"custom_operations": ops}, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
if model is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Could not detect model type for GGUF diffusion model: {unet_path}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Wrap with GGUFModelPatcher
|
||||||
|
model = GGUFModelPatcher.clone(model)
|
||||||
|
|
||||||
|
# Register a reload factory so the MODEL carries its source path
|
||||||
|
# (cached_patcher_init) like core ComfyUI loaders do — required
|
||||||
|
# for model-name extraction downstream and for ModelPatcher
|
||||||
|
# deepclone/dynamic machinery.
|
||||||
|
model.cached_patcher_init = (_reload_gguf_unet, (unet_path, weight_dtype))
|
||||||
|
|
||||||
|
return (model, unet_name)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error loading GGUF diffusion model '{unet_name}': {e}")
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Failed to load GGUF diffusion model '{unet_name}': {str(e)}"
|
||||||
|
)
|
||||||
+362
-130
@@ -5,7 +5,7 @@ import time
|
|||||||
import uuid
|
import uuid
|
||||||
from typing import Any, Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import folder_paths # type: ignore
|
import folder_paths # pyright: ignore[reportMissingImports]
|
||||||
from ..services.service_registry import ServiceRegistry
|
from ..services.service_registry import ServiceRegistry
|
||||||
from ..metadata_collector.metadata_processor import MetadataProcessor
|
from ..metadata_collector.metadata_processor import MetadataProcessor
|
||||||
from ..metadata_collector import get_metadata
|
from ..metadata_collector import get_metadata
|
||||||
@@ -13,9 +13,159 @@ from ..utils.constants import CARD_PREVIEW_WIDTH
|
|||||||
from ..utils.exif_utils import ExifUtils
|
from ..utils.exif_utils import ExifUtils
|
||||||
from ..utils.utils import calculate_recipe_fingerprint, sanitize_folder_name
|
from ..utils.utils import calculate_recipe_fingerprint, sanitize_folder_name
|
||||||
from PIL import Image, PngImagePlugin
|
from PIL import Image, PngImagePlugin
|
||||||
import piexif
|
import piexif # pyright: ignore[reportMissingTypeStubs]
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
|
# Civitai-compatible sampler name mapping: ComfyUI internal → A1111 display name
|
||||||
|
CIVITAI_SAMPLER_MAP = {
|
||||||
|
"euler": "Euler",
|
||||||
|
"euler_ancestral": "Euler a",
|
||||||
|
"lms": "LMS",
|
||||||
|
"heun": "Heun",
|
||||||
|
"dpm_2": "DPM2",
|
||||||
|
"dpm_2_ancestral": "DPM2 a",
|
||||||
|
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
||||||
|
"dpmpp_2m": "DPM++ 2M",
|
||||||
|
"dpmpp_sde": "DPM++ SDE",
|
||||||
|
"dpmpp_sde_gpu": "DPM++ SDE",
|
||||||
|
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
||||||
|
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
||||||
|
"dpmpp_3m_sde": "DPM++ 3M SDE",
|
||||||
|
"dpm_fast": "DPM fast",
|
||||||
|
"dpm_adaptive": "DPM adaptive",
|
||||||
|
"ddim": "DDIM",
|
||||||
|
"plms": "PLMS",
|
||||||
|
"uni_pc_bh2": "UniPC",
|
||||||
|
"uni_pc": "UniPC",
|
||||||
|
"lcm": "LCM",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Base model display name → AIR URN slug
|
||||||
|
# Sourced from civitai source: src/shared/constants/basemodel.constants.ts
|
||||||
|
BASE_MODEL_AIR_SLUG = {
|
||||||
|
# Stable Diffusion family
|
||||||
|
"SD 1.4": "sd1",
|
||||||
|
"SD 1.5": "sd1",
|
||||||
|
"SD 1.5 LCM": "sd1",
|
||||||
|
"SD 1.5 Hyper": "sd1",
|
||||||
|
"SD 2.0": "sd2",
|
||||||
|
"SD 2.0 768": "sd2",
|
||||||
|
"SD 2.1": "sd2",
|
||||||
|
"SD 2.1 768": "sd2",
|
||||||
|
"SD 2.1 Unclip": "sd2",
|
||||||
|
"SD 3.0": "sd3",
|
||||||
|
"SD 3.5": "sd35",
|
||||||
|
"SD 3.5 Large": "sd35",
|
||||||
|
"SD 3.5 Large Turbo": "sd35",
|
||||||
|
"SD 3.5 Medium": "sd35",
|
||||||
|
"SDXL 0.9": "sdxl",
|
||||||
|
"SDXL 1.0": "sdxl",
|
||||||
|
"SDXL 1.0 LCM": "sdxl",
|
||||||
|
"SDXL Lightning": "sdxl",
|
||||||
|
"SDXL Hyper": "sdxl",
|
||||||
|
"SDXL Turbo": "sdxl",
|
||||||
|
"SDXL Distilled": "sdxldistilled",
|
||||||
|
"Stable Cascade": "scascade",
|
||||||
|
"Stable Video Diffusion": "svd",
|
||||||
|
"SVD": "svd",
|
||||||
|
"SVD XT": "svdxt",
|
||||||
|
|
||||||
|
# SDXL community fine-tunes
|
||||||
|
"Pony": "pony",
|
||||||
|
"Pony Diffusion": "pony",
|
||||||
|
"Illustrious": "illustrious",
|
||||||
|
"NoobAI": "noobai",
|
||||||
|
"Animagine": "illustrious",
|
||||||
|
|
||||||
|
# Flux family
|
||||||
|
"Flux.1": "flux1",
|
||||||
|
"Flux.1 D": "flux1",
|
||||||
|
"Flux.1 S": "flux1",
|
||||||
|
"Flux.1 Krea": "fluxkrea",
|
||||||
|
"Flux.1 Kontext": "flux1kontext",
|
||||||
|
"Flux.2": "flux2",
|
||||||
|
"Flux.2 D": "flux2",
|
||||||
|
"Flux.2 Klein 9B": "flux2klein_9b",
|
||||||
|
"Flux.2 Klein 9B Base": "flux2klein_9b_base",
|
||||||
|
"Flux.2 Klein 4B": "flux2klein_4b",
|
||||||
|
"Flux.2 Klein 4B Base": "flux2klein_4b_base",
|
||||||
|
|
||||||
|
# Other image models (sorted alphabetically)
|
||||||
|
"AuraFlow": "auraflow",
|
||||||
|
"Chroma": "chroma",
|
||||||
|
"HiDream": "hidream",
|
||||||
|
"HiDream-O1": "hidream-o1",
|
||||||
|
"Hunyuan DiT": "hydit1",
|
||||||
|
"Hunyuan Video": "hyv1",
|
||||||
|
"Kolors": "kolors",
|
||||||
|
"Lumina": "lumina",
|
||||||
|
"Mochi": "mochi",
|
||||||
|
"ODOR": "odor",
|
||||||
|
"PixArt Alpha": "pixarta",
|
||||||
|
"PixArt Sigma": "pixarte",
|
||||||
|
"Playground v2": "playgroundv2",
|
||||||
|
"Playground v2.5": "playgroundv2",
|
||||||
|
"Pony Diffusion V7": "ponyv7",
|
||||||
|
|
||||||
|
# Video models
|
||||||
|
"CogVideoX": "cogvideox",
|
||||||
|
"LTX Video": "ltxv",
|
||||||
|
"LTX Video 2": "ltxv2",
|
||||||
|
"LTX Video 2.3": "ltxv23",
|
||||||
|
"Wan Video": "wanvideo",
|
||||||
|
"Wan Video 1.3B T2V": "wanvideo_13b_t2v",
|
||||||
|
"Wan Video 14B T2V": "wanvideo_14b_t2v",
|
||||||
|
"Wan Video 14B I2V 480p": "wanvideo_14b_i2v_480p",
|
||||||
|
"Wan Video 14B I2V 720p": "wanvideo_14b_i2v_720p",
|
||||||
|
|
||||||
|
# Third-party / proprietary image models
|
||||||
|
"Boogu": "boogu",
|
||||||
|
"Ernie": "ernie",
|
||||||
|
"Grok": "grok",
|
||||||
|
"HappyHorse": "happyhorse",
|
||||||
|
"Ideogram": "ideogram",
|
||||||
|
"Ideogram 4.0": "ideogram",
|
||||||
|
"Imagen": "imagen4",
|
||||||
|
"Imagen 4": "imagen4",
|
||||||
|
"Krea": "krea2",
|
||||||
|
"Krea 2": "krea2",
|
||||||
|
"Lens": "lens",
|
||||||
|
"MAI": "mai",
|
||||||
|
"Nano Banana": "nanobanana",
|
||||||
|
"OpenAI": "openai",
|
||||||
|
"Reve": "reve",
|
||||||
|
"Reve 2": "reve",
|
||||||
|
"Reve 2.1": "reve",
|
||||||
|
"Seedream": "seedream",
|
||||||
|
"Sora": "sora2",
|
||||||
|
"Sora 2": "sora2",
|
||||||
|
"Veo": "veo3",
|
||||||
|
"Veo 2": "veo3",
|
||||||
|
"Veo 3": "veo3",
|
||||||
|
"ZImageTurbo": "zimageturbo",
|
||||||
|
"ZImageBase": "zimagebase",
|
||||||
|
"ZImage": "zimagebase",
|
||||||
|
|
||||||
|
# Third-party video models
|
||||||
|
"Hailuo by MiniMax": "minimax",
|
||||||
|
"Haiper": "haiper",
|
||||||
|
"Kling": "kling",
|
||||||
|
"Lightricks": "lightricks",
|
||||||
|
"Seedance": "seedance",
|
||||||
|
"Vidu": "vidu",
|
||||||
|
|
||||||
|
# Qwen family
|
||||||
|
"Qwen": "qwen",
|
||||||
|
"Qwen 2": "qwen2",
|
||||||
|
|
||||||
|
# Anima
|
||||||
|
"Anima": "anima",
|
||||||
|
|
||||||
|
# Special
|
||||||
|
"Upscaler": "upscaler",
|
||||||
|
"Other": "other",
|
||||||
|
}
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -70,11 +220,29 @@ class SaveImageLM:
|
|||||||
"tooltip": "Compression quality for JPEG and lossy WebP formats (1-100). Higher values mean better quality but larger files.",
|
"tooltip": "Compression quality for JPEG and lossy WebP formats (1-100). Higher values mean better quality but larger files.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
|
"webp_method": (
|
||||||
|
"INT",
|
||||||
|
{
|
||||||
|
"default": 6,
|
||||||
|
"min": 0,
|
||||||
|
"max": 6,
|
||||||
|
"tooltip": "WebP compression method (0-6). 0=fastest/largest, 6=slowest/smallest. Only applies when file_format is 'webp'.",
|
||||||
|
},
|
||||||
|
),
|
||||||
|
"jpeg_subsampling": (
|
||||||
|
"INT",
|
||||||
|
{
|
||||||
|
"default": 0,
|
||||||
|
"min": 0,
|
||||||
|
"max": 2,
|
||||||
|
"tooltip": "JPEG chroma subsampling level. 0=4:4:4 (best quality), 1=4:2:2, 2=4:2:0 (smallest files). Only applies when file_format is 'jpeg'.",
|
||||||
|
},
|
||||||
|
),
|
||||||
"embed_workflow": (
|
"embed_workflow": (
|
||||||
"BOOLEAN",
|
"BOOLEAN",
|
||||||
{
|
{
|
||||||
"default": False,
|
"default": False,
|
||||||
"tooltip": "Embeds the complete workflow data into the image metadata. Only works with PNG and WebP formats.",
|
"tooltip": "When enabled, saved images store the complete workflow. Drag the image back into ComfyUI to restore the original node graph. PNG and WebP only.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
"save_with_metadata": (
|
"save_with_metadata": (
|
||||||
@@ -84,6 +252,13 @@ class SaveImageLM:
|
|||||||
"tooltip": "When enabled, embeds generation parameters into the saved image metadata. Disable to skip writing generation metadata.",
|
"tooltip": "When enabled, embeds generation parameters into the saved image metadata. Disable to skip writing generation metadata.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
|
"add_loras_to_prompt": (
|
||||||
|
"BOOLEAN",
|
||||||
|
{
|
||||||
|
"default": False,
|
||||||
|
"tooltip": "When enabled, appends the LoRA syntax line (e.g. <lora:name:strength>) after the positive prompt in the saved metadata.",
|
||||||
|
},
|
||||||
|
),
|
||||||
"add_counter_to_filename": (
|
"add_counter_to_filename": (
|
||||||
"BOOLEAN",
|
"BOOLEAN",
|
||||||
{
|
{
|
||||||
@@ -142,148 +317,197 @@ class SaveImageLM:
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def format_metadata(self, metadata_dict):
|
def _resolve_model_cache_entry(self, scanner_type: str, name: str):
|
||||||
"""Format metadata in the requested format similar to userComment example"""
|
"""Resolve model hash, civitai metadata, and base_model from scanner cache.
|
||||||
if not metadata_dict:
|
Returns (hash_str, civitai_dict, base_model_str). All values are empty defaults when not found."""
|
||||||
return ""
|
scanner = ServiceRegistry.get_service_sync(scanner_type)
|
||||||
|
if scanner is None or not name:
|
||||||
|
return "", {}, ""
|
||||||
|
|
||||||
# Helper function to only add parameter if value is not None
|
entry = self._get_cached_model_by_name(scanner, name)
|
||||||
def add_param_if_not_none(param_list, label, value):
|
if entry is None:
|
||||||
if value is not None:
|
basename = os.path.splitext(os.path.basename(name))[0]
|
||||||
param_list.append(f"{label}: {value}")
|
hash_val = scanner.get_hash_by_filename(basename)
|
||||||
|
return (hash_val or "").lower(), {}, ""
|
||||||
|
|
||||||
|
hash_val = (entry.get("sha256") or "").lower()
|
||||||
|
civitai = entry.get("civitai") or {}
|
||||||
|
base_model = entry.get("base_model") or ""
|
||||||
|
return hash_val, civitai, base_model
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _get_civitai_sampler_name(sampler_name: str, scheduler: str) -> str:
|
||||||
|
if sampler_name in CIVITAI_SAMPLER_MAP:
|
||||||
|
civitai_name = CIVITAI_SAMPLER_MAP[sampler_name]
|
||||||
|
if scheduler == "karras":
|
||||||
|
civitai_name += " Karras"
|
||||||
|
elif scheduler == "exponential":
|
||||||
|
civitai_name += " Exponential"
|
||||||
|
return civitai_name
|
||||||
|
else:
|
||||||
|
if scheduler and scheduler != "normal":
|
||||||
|
return f"{sampler_name}_{scheduler}"
|
||||||
|
return sampler_name
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _build_air_string(base_model: str, model_type: str, model_id: int, version_id: int) -> str:
|
||||||
|
slug = BASE_MODEL_AIR_SLUG.get(base_model, "other")
|
||||||
|
type_lower = model_type.lower() if model_type else "other"
|
||||||
|
return f"urn:air:{slug}:{type_lower}:civitai:{model_id}@{version_id}"
|
||||||
|
|
||||||
|
def format_metadata(self, metadata_dict: dict[str, Any], add_loras_to_prompt: bool = False) -> str:
|
||||||
|
"""Format metadata as A1111-compatible parameters string with Hashes JSON and Civitai resources."""
|
||||||
|
if not metadata_dict: return ""
|
||||||
|
|
||||||
# Extract the prompt and negative prompt
|
|
||||||
prompt = metadata_dict.get("prompt", "")
|
prompt = metadata_dict.get("prompt", "")
|
||||||
negative_prompt = metadata_dict.get("negative_prompt", "")
|
negative_prompt = metadata_dict.get("negative_prompt", "")
|
||||||
|
steps = metadata_dict.get("steps")
|
||||||
# Extract loras from the prompt if present
|
cfg = metadata_dict.get("guidance")
|
||||||
|
if cfg is None:
|
||||||
|
cfg = metadata_dict.get("cfg_scale")
|
||||||
|
if cfg is None:
|
||||||
|
cfg = metadata_dict.get("cfg")
|
||||||
|
seed = metadata_dict.get("seed")
|
||||||
|
size = metadata_dict.get("size")
|
||||||
|
sampler = metadata_dict.get("sampler") or ""
|
||||||
|
scheduler = metadata_dict.get("scheduler") or "normal"
|
||||||
|
checkpoint = metadata_dict.get("checkpoint") or ""
|
||||||
loras_text = metadata_dict.get("loras", "")
|
loras_text = metadata_dict.get("loras", "")
|
||||||
lora_hashes = {}
|
clip_skip = metadata_dict.get("clip_skip")
|
||||||
|
|
||||||
# If loras are found, add them on a new line after the prompt
|
# Parse LoRA entries from <lora:name:strength> format
|
||||||
|
lora_entries: list[tuple[str, float]] = []
|
||||||
if loras_text:
|
if loras_text:
|
||||||
prompt_with_loras = f"{prompt}\n{loras_text}"
|
for match in re.findall(r"<lora:([^:]+):([^>]+)>", loras_text):
|
||||||
|
lora_name, strength_str = match
|
||||||
|
try:
|
||||||
|
strength = float(strength_str)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
strength = 1.0
|
||||||
|
lora_entries.append((lora_name, strength))
|
||||||
|
|
||||||
# Extract lora names from the format <lora:name:strength>
|
# Resolve checkpoint hash and Civitai data from local cache
|
||||||
lora_matches = re.findall(r"<lora:([^:]+):([^>]+)>", loras_text)
|
ckpt_hash, ckpt_civitai, ckpt_base_model = "", {}, ""
|
||||||
|
ckpt_display_name = ""
|
||||||
|
if checkpoint:
|
||||||
|
ckpt_hash, ckpt_civitai, ckpt_base_model = self._resolve_model_cache_entry(
|
||||||
|
"checkpoint_scanner", checkpoint
|
||||||
|
)
|
||||||
|
ckpt_display_name = os.path.splitext(os.path.basename(checkpoint))[0]
|
||||||
|
|
||||||
# Get hash for each lora
|
# Resolve LoRA hash and Civitai data from local cache
|
||||||
for lora_name, strength in lora_matches:
|
loras_data: list[dict[str, Any]] = []
|
||||||
hash_value = self.get_lora_hash(lora_name)
|
for lora_name, strength in lora_entries:
|
||||||
if hash_value:
|
lora_hash, lora_civitai, lora_base_model = self._resolve_model_cache_entry(
|
||||||
lora_hashes[lora_name] = hash_value
|
"lora_scanner", lora_name
|
||||||
else:
|
)
|
||||||
prompt_with_loras = prompt
|
loras_data.append({
|
||||||
|
"name": lora_name,
|
||||||
|
"strength": strength,
|
||||||
|
"hash": lora_hash,
|
||||||
|
"civitai": lora_civitai,
|
||||||
|
"base_model": lora_base_model,
|
||||||
|
})
|
||||||
|
|
||||||
# Format the first part (prompt and loras)
|
# Build Hashes JSON (A1111 / Civitai standard format)
|
||||||
metadata_parts = [prompt_with_loras]
|
hashes: dict[str, str] = {}
|
||||||
|
if ckpt_hash:
|
||||||
|
hashes["model"] = ckpt_hash[:10].upper()
|
||||||
|
for lora in loras_data:
|
||||||
|
if lora["hash"]:
|
||||||
|
hashes[f"LORA:{lora['name']}"] = lora["hash"][:10].upper()
|
||||||
|
|
||||||
# Add negative prompt
|
# Build Civitai resources JSON array
|
||||||
|
civitai_resources: list[dict[str, Any]] = []
|
||||||
|
if ckpt_civitai.get("id", 0) > 0:
|
||||||
|
ckpt_resource: dict[str, Any] = {}
|
||||||
|
ckpt_type = (ckpt_civitai.get("model") or {}).get("type", "Checkpoint")
|
||||||
|
model_id = ckpt_civitai.get("modelId", 0)
|
||||||
|
version_id = ckpt_civitai.get("id", 0)
|
||||||
|
if model_id and version_id:
|
||||||
|
ckpt_resource["air"] = self._build_air_string(
|
||||||
|
ckpt_base_model, ckpt_type, int(model_id), int(version_id)
|
||||||
|
)
|
||||||
|
elif version_id:
|
||||||
|
ckpt_resource["modelVersionId"] = int(version_id)
|
||||||
|
if ckpt_civitai.get("name"):
|
||||||
|
ckpt_resource["versionName"] = ckpt_civitai["name"]
|
||||||
|
if ckpt_resource:
|
||||||
|
civitai_resources.append(ckpt_resource)
|
||||||
|
|
||||||
|
for lora in loras_data:
|
||||||
|
lora_civitai = lora["civitai"]
|
||||||
|
if not lora_civitai or lora_civitai.get("id", 0) <= 0:
|
||||||
|
continue
|
||||||
|
lora_resource: dict[str, Any] = {"weight": lora["strength"]}
|
||||||
|
lora_type = (lora_civitai.get("model") or {}).get("type", "LORA")
|
||||||
|
model_id = lora_civitai.get("modelId", 0)
|
||||||
|
version_id = lora_civitai.get("id", 0)
|
||||||
|
if model_id and version_id:
|
||||||
|
lora_resource["air"] = self._build_air_string(
|
||||||
|
lora["base_model"], lora_type, int(model_id), int(version_id)
|
||||||
|
)
|
||||||
|
elif version_id:
|
||||||
|
lora_resource["modelVersionId"] = int(version_id)
|
||||||
|
if lora_civitai.get("name"):
|
||||||
|
lora_resource["versionName"] = lora_civitai["name"]
|
||||||
|
civitai_resources.append(lora_resource)
|
||||||
|
|
||||||
|
sampler_name = CIVITAI_SAMPLER_MAP.get(sampler, sampler) if sampler else None
|
||||||
|
|
||||||
|
scheduler_mapping = {
|
||||||
|
"normal": "Normal",
|
||||||
|
"karras": "Karras",
|
||||||
|
"exponential": "Exponential",
|
||||||
|
"sgm_uniform": "SGM Uniform",
|
||||||
|
"sgm_quadratic": "SGM Quadratic",
|
||||||
|
}
|
||||||
|
scheduler_name = scheduler_mapping.get(scheduler, scheduler) if scheduler else None
|
||||||
|
|
||||||
|
# Build output lines
|
||||||
|
prompt_line = prompt if prompt else ""
|
||||||
|
if add_loras_to_prompt and loras_text:
|
||||||
|
prompt_line = f"{prompt_line}\n{loras_text}" if prompt_line else loras_text
|
||||||
|
lines = [prompt_line] if prompt_line else [""]
|
||||||
if negative_prompt:
|
if negative_prompt:
|
||||||
metadata_parts.append(f"Negative prompt: {negative_prompt}")
|
lines.append(f"Negative prompt: {negative_prompt}")
|
||||||
|
|
||||||
# Format the second part (generation parameters)
|
params: list[str] = []
|
||||||
params = []
|
if steps is not None:
|
||||||
|
params.append(f"Steps: {steps}")
|
||||||
# Add standard parameters in the correct order
|
|
||||||
if "steps" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "Steps", metadata_dict.get("steps"))
|
|
||||||
|
|
||||||
# Combine sampler and scheduler information
|
|
||||||
sampler_name = None
|
|
||||||
scheduler_name = None
|
|
||||||
|
|
||||||
if "sampler" in metadata_dict:
|
|
||||||
sampler = metadata_dict.get("sampler")
|
|
||||||
# Convert ComfyUI sampler names to user-friendly names
|
|
||||||
sampler_mapping = {
|
|
||||||
"euler": "Euler",
|
|
||||||
"euler_ancestral": "Euler a",
|
|
||||||
"dpm_2": "DPM2",
|
|
||||||
"dpm_2_ancestral": "DPM2 a",
|
|
||||||
"heun": "Heun",
|
|
||||||
"dpm_fast": "DPM fast",
|
|
||||||
"dpm_adaptive": "DPM adaptive",
|
|
||||||
"lms": "LMS",
|
|
||||||
"dpmpp_2s_ancestral": "DPM++ 2S a",
|
|
||||||
"dpmpp_sde": "DPM++ SDE",
|
|
||||||
"dpmpp_sde_gpu": "DPM++ SDE",
|
|
||||||
"dpmpp_2m": "DPM++ 2M",
|
|
||||||
"dpmpp_2m_sde": "DPM++ 2M SDE",
|
|
||||||
"dpmpp_2m_sde_gpu": "DPM++ 2M SDE",
|
|
||||||
"ddim": "DDIM",
|
|
||||||
}
|
|
||||||
sampler_name = sampler_mapping.get(sampler, sampler)
|
|
||||||
|
|
||||||
if "scheduler" in metadata_dict:
|
|
||||||
scheduler = metadata_dict.get("scheduler")
|
|
||||||
scheduler_mapping = {
|
|
||||||
"normal": "Simple",
|
|
||||||
"karras": "Karras",
|
|
||||||
"exponential": "Exponential",
|
|
||||||
"sgm_uniform": "SGM Uniform",
|
|
||||||
"sgm_quadratic": "SGM Quadratic",
|
|
||||||
}
|
|
||||||
scheduler_name = scheduler_mapping.get(scheduler, scheduler)
|
|
||||||
|
|
||||||
# Add combined sampler and scheduler information
|
|
||||||
if sampler_name:
|
if sampler_name:
|
||||||
if scheduler_name:
|
if scheduler_name:
|
||||||
params.append(f"Sampler: {sampler_name} {scheduler_name}")
|
params.append(f"Sampler: {sampler_name} {scheduler_name}")
|
||||||
else:
|
else:
|
||||||
params.append(f"Sampler: {sampler_name}")
|
params.append(f"Sampler: {sampler_name}")
|
||||||
|
if cfg is not None:
|
||||||
|
params.append(f"CFG scale: {cfg}")
|
||||||
|
if seed is not None:
|
||||||
|
params.append(f"Seed: {seed}")
|
||||||
|
if size:
|
||||||
|
params.append(f"Size: {size}")
|
||||||
|
if clip_skip is not None:
|
||||||
|
try:
|
||||||
|
params.append(f"Clip skip: {abs(int(clip_skip))}")
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
pass
|
||||||
|
additional_data = metadata_dict.get("additional_data", "")
|
||||||
|
if additional_data:
|
||||||
|
params.append(additional_data)
|
||||||
|
if ckpt_hash:
|
||||||
|
params.append(f"Model hash: {ckpt_hash[:10].upper()}")
|
||||||
|
if ckpt_display_name:
|
||||||
|
params.append(f"Model: {ckpt_display_name}")
|
||||||
|
if hashes:
|
||||||
|
params.append(f"Hashes: {json.dumps(hashes, separators=(',', ':'))}")
|
||||||
|
params.append("Version: ComfyUI")
|
||||||
|
if civitai_resources:
|
||||||
|
params.append(
|
||||||
|
f"Civitai resources: {json.dumps(civitai_resources, separators=(',', ':'))}"
|
||||||
|
)
|
||||||
|
|
||||||
# CFG scale (Use guidance if available, otherwise fall back to cfg_scale or cfg)
|
lines.append(", ".join(params))
|
||||||
if "guidance" in metadata_dict:
|
return "\n".join(lines)
|
||||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("guidance"))
|
|
||||||
elif "cfg_scale" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg_scale"))
|
|
||||||
elif "cfg" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "CFG scale", metadata_dict.get("cfg"))
|
|
||||||
|
|
||||||
# Seed
|
|
||||||
if "seed" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "Seed", metadata_dict.get("seed"))
|
|
||||||
|
|
||||||
# Size
|
|
||||||
if "size" in metadata_dict:
|
|
||||||
add_param_if_not_none(params, "Size", metadata_dict.get("size"))
|
|
||||||
|
|
||||||
# Model info
|
|
||||||
if "checkpoint" in metadata_dict:
|
|
||||||
# Ensure checkpoint is a string before processing
|
|
||||||
checkpoint = metadata_dict.get("checkpoint")
|
|
||||||
if checkpoint is not None:
|
|
||||||
# Get model hash
|
|
||||||
model_hash = self.get_checkpoint_hash(checkpoint)
|
|
||||||
|
|
||||||
# Extract basename without path
|
|
||||||
checkpoint_name = os.path.basename(checkpoint)
|
|
||||||
# Remove extension if present
|
|
||||||
checkpoint_name = os.path.splitext(checkpoint_name)[0]
|
|
||||||
|
|
||||||
# Add model hash if available
|
|
||||||
if model_hash:
|
|
||||||
params.append(
|
|
||||||
f"Model hash: {model_hash[:10]}, Model: {checkpoint_name}"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
params.append(f"Model: {checkpoint_name}")
|
|
||||||
|
|
||||||
# Add LoRA hashes if available
|
|
||||||
if lora_hashes:
|
|
||||||
lora_hash_parts = []
|
|
||||||
for lora_name, hash_value in lora_hashes.items():
|
|
||||||
lora_hash_parts.append(f"{lora_name}: {hash_value[:10]}")
|
|
||||||
|
|
||||||
if lora_hash_parts:
|
|
||||||
params.append(f'Lora hashes: "{", ".join(lora_hash_parts)}"')
|
|
||||||
|
|
||||||
# Combine all parameters with commas
|
|
||||||
metadata_parts.append(", ".join(params))
|
|
||||||
|
|
||||||
# Join all parts with a new line
|
|
||||||
return "\n".join(metadata_parts)
|
|
||||||
|
|
||||||
# credit to nkchocoai
|
# credit to nkchocoai
|
||||||
# Add format_filename method to handle pattern substitution
|
# Add format_filename method to handle pattern substitution
|
||||||
@@ -573,10 +797,13 @@ class SaveImageLM:
|
|||||||
extra_pnginfo=None,
|
extra_pnginfo=None,
|
||||||
lossless_webp=True,
|
lossless_webp=True,
|
||||||
quality=100,
|
quality=100,
|
||||||
|
webp_method=6,
|
||||||
|
jpeg_subsampling=0,
|
||||||
embed_workflow=False,
|
embed_workflow=False,
|
||||||
save_with_metadata=True,
|
save_with_metadata=True,
|
||||||
add_counter_to_filename=True,
|
add_counter_to_filename=True,
|
||||||
save_as_recipe=False,
|
save_as_recipe=False,
|
||||||
|
add_loras_to_prompt=False,
|
||||||
):
|
):
|
||||||
"""Save images with metadata"""
|
"""Save images with metadata"""
|
||||||
results = []
|
results = []
|
||||||
@@ -585,7 +812,7 @@ class SaveImageLM:
|
|||||||
raw_metadata = get_metadata()
|
raw_metadata = get_metadata()
|
||||||
metadata_dict = MetadataProcessor.to_dict(raw_metadata, id)
|
metadata_dict = MetadataProcessor.to_dict(raw_metadata, id)
|
||||||
|
|
||||||
metadata = self.format_metadata(metadata_dict)
|
metadata = self.format_metadata(metadata_dict, add_loras_to_prompt)
|
||||||
|
|
||||||
# Process filename_prefix with pattern substitution
|
# Process filename_prefix with pattern substitution
|
||||||
filename_prefix = self.format_filename(filename_prefix, metadata_dict)
|
filename_prefix = self.format_filename(filename_prefix, metadata_dict)
|
||||||
@@ -627,15 +854,14 @@ class SaveImageLM:
|
|||||||
elif file_format == "jpeg":
|
elif file_format == "jpeg":
|
||||||
file = base_filename + ".jpg"
|
file = base_filename + ".jpg"
|
||||||
file_extension = ".jpg"
|
file_extension = ".jpg"
|
||||||
save_kwargs = {"quality": quality, "optimize": True}
|
save_kwargs = {"quality": quality, "optimize": True, "subsampling": jpeg_subsampling}
|
||||||
elif file_format == "webp":
|
elif file_format == "webp":
|
||||||
file = base_filename + ".webp"
|
file = base_filename + ".webp"
|
||||||
file_extension = ".webp"
|
file_extension = ".webp"
|
||||||
# Add optimization param to control performance
|
|
||||||
save_kwargs = {
|
save_kwargs = {
|
||||||
"quality": quality,
|
"quality": quality,
|
||||||
"lossless": lossless_webp,
|
"lossless": lossless_webp,
|
||||||
"method": 0,
|
"method": webp_method,
|
||||||
}
|
}
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Unsupported file format: {file_format}")
|
raise ValueError(f"Unsupported file format: {file_format}")
|
||||||
@@ -722,10 +948,13 @@ class SaveImageLM:
|
|||||||
extra_pnginfo=None,
|
extra_pnginfo=None,
|
||||||
lossless_webp=True,
|
lossless_webp=True,
|
||||||
quality=100,
|
quality=100,
|
||||||
|
webp_method=6,
|
||||||
|
jpeg_subsampling=0,
|
||||||
embed_workflow=False,
|
embed_workflow=False,
|
||||||
save_with_metadata=True,
|
save_with_metadata=True,
|
||||||
add_counter_to_filename=True,
|
add_counter_to_filename=True,
|
||||||
save_as_recipe=False,
|
save_as_recipe=False,
|
||||||
|
add_loras_to_prompt=False,
|
||||||
):
|
):
|
||||||
"""Process and save image with metadata"""
|
"""Process and save image with metadata"""
|
||||||
# Make sure the output directory exists
|
# Make sure the output directory exists
|
||||||
@@ -751,10 +980,13 @@ class SaveImageLM:
|
|||||||
extra_pnginfo,
|
extra_pnginfo,
|
||||||
lossless_webp,
|
lossless_webp,
|
||||||
quality,
|
quality,
|
||||||
|
webp_method,
|
||||||
|
jpeg_subsampling,
|
||||||
embed_workflow,
|
embed_workflow,
|
||||||
save_with_metadata,
|
save_with_metadata,
|
||||||
add_counter_to_filename,
|
add_counter_to_filename,
|
||||||
save_as_recipe,
|
save_as_recipe,
|
||||||
|
add_loras_to_prompt,
|
||||||
)
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
|
|||||||
+31
-7
@@ -1,12 +1,27 @@
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from typing import List, Tuple
|
from typing import Any, List, Tuple
|
||||||
import comfy.sd # type: ignore
|
import comfy.sd # pyright: ignore[reportMissingImports]
|
||||||
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
|
from ..utils.utils import get_checkpoint_info_absolute, _format_model_name_for_comfyui
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _reload_gguf_unet(
|
||||||
|
unet_path: str, weight_dtype: str, disable_dynamic: bool = False
|
||||||
|
) -> object:
|
||||||
|
"""Reload a GGUF diffusion model from disk (cached_patcher_init factory).
|
||||||
|
|
||||||
|
Mirrors the GGUF branch of UNETLoaderLM.load_unet so ModelPatcher
|
||||||
|
deepclone/dynamic machinery can rebuild GGUF models with the correct
|
||||||
|
GGMLOps. ``disable_dynamic`` is accepted for signature compatibility
|
||||||
|
with core ComfyUI loaders.
|
||||||
|
"""
|
||||||
|
loader = UNETLoaderLM()
|
||||||
|
model, = loader._load_gguf_unet(unet_path, unet_path, weight_dtype)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
class UNETLoaderLM:
|
class UNETLoaderLM:
|
||||||
"""UNET Loader with support for extra folder paths
|
"""UNET Loader with support for extra folder paths
|
||||||
|
|
||||||
@@ -19,9 +34,9 @@ class UNETLoaderLM:
|
|||||||
CATEGORY = "Lora Manager/loaders"
|
CATEGORY = "Lora Manager/loaders"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(cls):
|
||||||
# Get list of unet names from scanner (includes extra folder paths)
|
# Get list of unet names from scanner (includes extra folder paths)
|
||||||
unet_names = s._get_unet_names()
|
unet_names = cls._get_unet_names()
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"unet_name": (
|
"unet_name": (
|
||||||
@@ -59,7 +74,10 @@ class UNETLoaderLM:
|
|||||||
for item in cache.raw_data:
|
for item in cache.raw_data:
|
||||||
if item.get("sub_type") == "diffusion_model":
|
if item.get("sub_type") == "diffusion_model":
|
||||||
file_path = item.get("file_path", "")
|
file_path = item.get("file_path", "")
|
||||||
if file_path:
|
# Only offer models that still exist on disk so ComfyUI
|
||||||
|
# flags missing diffusion models at queue time via
|
||||||
|
# "value not in list" (the scanner cache can be stale).
|
||||||
|
if file_path and os.path.exists(file_path):
|
||||||
# Format using relative path with OS-native separator
|
# Format using relative path with OS-native separator
|
||||||
formatted_name = _format_model_name_for_comfyui(
|
formatted_name = _format_model_name_for_comfyui(
|
||||||
file_path, model_roots
|
file_path, model_roots
|
||||||
@@ -90,7 +108,7 @@ class UNETLoaderLM:
|
|||||||
logger.error(f"Error getting unet names: {e}")
|
logger.error(f"Error getting unet names: {e}")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
def load_unet(self, unet_name: str, weight_dtype: str) -> Tuple:
|
def load_unet(self, unet_name: str, weight_dtype: str) -> Tuple[Any, ...]:
|
||||||
"""Load a diffusion model by name, supporting extra folder paths
|
"""Load a diffusion model by name, supporting extra folder paths
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -133,7 +151,7 @@ class UNETLoaderLM:
|
|||||||
|
|
||||||
def _load_gguf_unet(
|
def _load_gguf_unet(
|
||||||
self, unet_path: str, unet_name: str, weight_dtype: str
|
self, unet_path: str, unet_name: str, weight_dtype: str
|
||||||
) -> Tuple:
|
) -> Tuple[Any, ...]:
|
||||||
"""Load a GGUF format diffusion model
|
"""Load a GGUF format diffusion model
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -196,6 +214,12 @@ class UNETLoaderLM:
|
|||||||
# Wrap with GGUFModelPatcher
|
# Wrap with GGUFModelPatcher
|
||||||
model = GGUFModelPatcher.clone(model)
|
model = GGUFModelPatcher.clone(model)
|
||||||
|
|
||||||
|
# Register a reload factory so the MODEL carries its source path
|
||||||
|
# (cached_patcher_init) like core ComfyUI loaders do — required
|
||||||
|
# for model-name extraction downstream and for ModelPatcher
|
||||||
|
# deepclone/dynamic machinery.
|
||||||
|
model.cached_patcher_init = (_reload_gguf_unet, (unet_path, weight_dtype))
|
||||||
|
|
||||||
return (model,)
|
return (model,)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
+159
-3
@@ -1,3 +1,6 @@
|
|||||||
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
class AnyType(str):
|
class AnyType(str):
|
||||||
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
"""A special class that is always equal in not equal comparisons. Credit to pythongosssss"""
|
||||||
|
|
||||||
@@ -6,7 +9,7 @@ class AnyType(str):
|
|||||||
|
|
||||||
|
|
||||||
# Credit to Regis Gaughan, III (rgthree)
|
# Credit to Regis Gaughan, III (rgthree)
|
||||||
class FlexibleOptionalInputType(dict):
|
class FlexibleOptionalInputType(dict[str, Any]):
|
||||||
"""A special class to make flexible nodes that pass data to our python handlers.
|
"""A special class to make flexible nodes that pass data to our python handlers.
|
||||||
|
|
||||||
Enables both flexible/dynamic input types (like for Any Switch) or a dynamic number of inputs
|
Enables both flexible/dynamic input types (like for Any Switch) or a dynamic number of inputs
|
||||||
@@ -23,6 +26,7 @@ class FlexibleOptionalInputType(dict):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, type):
|
def __init__(self, type):
|
||||||
|
super().__init__()
|
||||||
self.type = type
|
self.type = type
|
||||||
|
|
||||||
def __getitem__(self, key):
|
def __getitem__(self, key):
|
||||||
@@ -40,7 +44,8 @@ import re
|
|||||||
import logging
|
import logging
|
||||||
import copy
|
import copy
|
||||||
import sys
|
import sys
|
||||||
import folder_paths # type: ignore
|
import asyncio
|
||||||
|
import folder_paths # pyright: ignore[reportMissingImports]
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -70,7 +75,7 @@ def extract_lora_name(lora_path):
|
|||||||
return apply_lora_syntax_format(name_no_ext)
|
return apply_lora_syntax_format(name_no_ext)
|
||||||
|
|
||||||
|
|
||||||
def parse_lora_syntax(text: str) -> list[dict]:
|
def parse_lora_syntax(text: str) -> list[dict[str, Any]]:
|
||||||
"""Parse <lora:name:strength> syntax from text input into a list of dicts.
|
"""Parse <lora:name:strength> syntax from text input into a list of dicts.
|
||||||
|
|
||||||
Each entry contains: name, model_strength, clip_strength.
|
Each entry contains: name, model_strength, clip_strength.
|
||||||
@@ -107,6 +112,157 @@ def get_loras_list(kwargs):
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
_LORA_EXTENSIONS = (".safetensors", ".ckpt", ".pt", ".bin")
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_lora_extension(name: str) -> str:
|
||||||
|
"""Strip a known LoRA model extension from a name (case-insensitive)."""
|
||||||
|
lowered = name.lower()
|
||||||
|
for ext in _LORA_EXTENSIONS:
|
||||||
|
if lowered.endswith(ext):
|
||||||
|
return name[: -len(ext)]
|
||||||
|
return name
|
||||||
|
|
||||||
|
|
||||||
|
def _find_missing_loras(names: list[str]) -> list[str]:
|
||||||
|
"""Return the names that cannot be resolved to an existing local LoRA file.
|
||||||
|
|
||||||
|
Mirrors the matching semantics of ``get_lora_info_absolute``
|
||||||
|
(py/utils/utils.py): after stripping the extension, a name matches a cached
|
||||||
|
LoRA when it equals the cached file name or the ``folder/file`` path. As a
|
||||||
|
fallback, a name containing a folder that only matches by basename resolves
|
||||||
|
to the first basename match (same behavior as the runtime resolver). Raw
|
||||||
|
absolute paths that exist on disk are always considered available.
|
||||||
|
|
||||||
|
The scanner cache is fetched once for all names; the cache may be stale, so
|
||||||
|
resolved paths are additionally verified with ``os.path.isfile``.
|
||||||
|
"""
|
||||||
|
if not names:
|
||||||
|
return []
|
||||||
|
|
||||||
|
async def _check() -> list[str]:
|
||||||
|
from ..services.service_registry import ServiceRegistry
|
||||||
|
|
||||||
|
scanner = await ServiceRegistry.get_lora_scanner()
|
||||||
|
# The scanner cache may not be hydrated yet (startup, library path
|
||||||
|
# change). An empty cache is not authoritative — treat it as "cannot
|
||||||
|
# verify" and skip validation instead of flagging every active LoRA
|
||||||
|
# as missing.
|
||||||
|
if getattr(scanner, "_cache", None) is None or getattr(
|
||||||
|
scanner, "_is_initializing", False
|
||||||
|
):
|
||||||
|
return []
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
|
||||||
|
lookup = {}
|
||||||
|
basename_candidates = {}
|
||||||
|
for item in cache.raw_data:
|
||||||
|
file_path = item.get("file_path")
|
||||||
|
if not file_path:
|
||||||
|
continue
|
||||||
|
file_name = item.get("file_name", "")
|
||||||
|
folder = item.get("folder", "")
|
||||||
|
file_name_no_ext = _strip_lora_extension(file_name)
|
||||||
|
path_name_no_ext = (
|
||||||
|
f"{folder}/{file_name_no_ext}".replace("\\", "/")
|
||||||
|
if folder
|
||||||
|
else file_name_no_ext
|
||||||
|
)
|
||||||
|
lookup.setdefault(file_name_no_ext, file_path)
|
||||||
|
lookup.setdefault(path_name_no_ext, file_path)
|
||||||
|
basename_candidates.setdefault(file_name_no_ext, []).append(
|
||||||
|
(folder, file_path)
|
||||||
|
)
|
||||||
|
|
||||||
|
missing = []
|
||||||
|
for name in names:
|
||||||
|
if not name:
|
||||||
|
continue
|
||||||
|
normalized = name.replace("\\", "/")
|
||||||
|
# Raw absolute paths (outside the library) are usable as-is.
|
||||||
|
if os.path.isfile(normalized):
|
||||||
|
continue
|
||||||
|
no_ext = _strip_lora_extension(normalized)
|
||||||
|
file_path = lookup.get(no_ext)
|
||||||
|
if file_path is None and "/" in no_ext:
|
||||||
|
# A name with a folder that matches only by basename resolves
|
||||||
|
# at runtime like get_lora_info_absolute's fallback does:
|
||||||
|
# prefer a candidate whose folder prefixes the name, else the
|
||||||
|
# first basename match.
|
||||||
|
folder, basename = no_ext.rsplit("/", 1)
|
||||||
|
candidates = basename_candidates.get(basename, [])
|
||||||
|
file_path = next(
|
||||||
|
(
|
||||||
|
fp
|
||||||
|
for fld, fp in candidates
|
||||||
|
if fld and no_ext.startswith(fld + "/")
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if file_path is None and candidates:
|
||||||
|
file_path = candidates[0][1]
|
||||||
|
if file_path is None or not os.path.isfile(file_path):
|
||||||
|
missing.append(name)
|
||||||
|
return missing
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Check if we're already in an event loop
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
# If we're in a running loop, run the async check in a separate thread
|
||||||
|
import concurrent.futures
|
||||||
|
|
||||||
|
def run_in_thread():
|
||||||
|
new_loop = asyncio.new_event_loop()
|
||||||
|
asyncio.set_event_loop(new_loop)
|
||||||
|
try:
|
||||||
|
return new_loop.run_until_complete(_check())
|
||||||
|
finally:
|
||||||
|
new_loop.close()
|
||||||
|
|
||||||
|
with concurrent.futures.ThreadPoolExecutor() as executor:
|
||||||
|
future = executor.submit(run_in_thread)
|
||||||
|
return future.result()
|
||||||
|
except RuntimeError:
|
||||||
|
# No event loop is running, we can use asyncio.run()
|
||||||
|
return asyncio.run(_check())
|
||||||
|
|
||||||
|
|
||||||
|
def validate_lora_entries(kwargs):
|
||||||
|
"""Validate active LoRA widget entries against the local library.
|
||||||
|
|
||||||
|
Used by node ``VALIDATE_INPUTS`` implementations so ComfyUI rejects the
|
||||||
|
prompt at queue time (``custom_validation_failed``) when an active entry
|
||||||
|
references a LoRA that is not available locally — mirroring how built-in
|
||||||
|
loader nodes flag missing models before execution starts.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
None when every active entry resolves to an existing local file,
|
||||||
|
otherwise a descriptive error string listing the missing LoRAs.
|
||||||
|
Verification failures (e.g. scanner not ready) are treated as valid
|
||||||
|
so queueing is never blocked by validation machinery itself.
|
||||||
|
"""
|
||||||
|
# Missing/empty loras input is always valid; skip get_loras_list so it
|
||||||
|
# does not log a warning for the None case on every queue.
|
||||||
|
if not kwargs.get("loras"):
|
||||||
|
return None
|
||||||
|
loras = get_loras_list(kwargs)
|
||||||
|
active_names = []
|
||||||
|
for lora in loras:
|
||||||
|
if not isinstance(lora, dict):
|
||||||
|
continue
|
||||||
|
if not lora.get("active", False):
|
||||||
|
continue
|
||||||
|
active_names.append(apply_lora_syntax_format(str(lora.get("name") or "")))
|
||||||
|
try:
|
||||||
|
missing = _find_missing_loras(active_names)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Failed to validate LoRA entries against the local library")
|
||||||
|
return None
|
||||||
|
if not missing:
|
||||||
|
return None
|
||||||
|
return "Missing LoRA(s) in local library: " + ", ".join(missing)
|
||||||
|
|
||||||
|
|
||||||
def load_state_dict_in_safetensors(path, device="cpu", filter_prefix=""):
|
def load_state_dict_in_safetensors(path, device="cpu", filter_prefix=""):
|
||||||
"""Simplified version of load_state_dict_in_safetensors that just loads from a local path"""
|
"""Simplified version of load_state_dict_in_safetensors that just loads from a local path"""
|
||||||
import safetensors.torch
|
import safetensors.torch
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import os
|
import os
|
||||||
from ..utils.utils import get_lora_info_absolute
|
from ..utils.utils import get_lora_info_absolute
|
||||||
from ..config import config
|
from ..config import config
|
||||||
from .utils import FlexibleOptionalInputType, any_type, get_loras_list
|
from .utils import FlexibleOptionalInputType, any_type, get_loras_list, validate_lora_entries
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -35,6 +35,11 @@ class WanVideoLoraSelectLM:
|
|||||||
"optional": FlexibleOptionalInputType(any_type),
|
"optional": FlexibleOptionalInputType(any_type),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def VALIDATE_INPUTS(cls, loras=None):
|
||||||
|
"""Queue-time validation: reject missing local LoRAs before execution."""
|
||||||
|
return validate_lora_entries({"loras": loras}) or True
|
||||||
|
|
||||||
RETURN_TYPES = ("WANVIDLORA", "STRING", "STRING")
|
RETURN_TYPES = ("WANVIDLORA", "STRING", "STRING")
|
||||||
RETURN_NAMES = ("lora", "trigger_words", "active_loras")
|
RETURN_NAMES = ("lora", "trigger_words", "active_loras")
|
||||||
FUNCTION = "process_loras"
|
FUNCTION = "process_loras"
|
||||||
|
|||||||
+31
-9
@@ -1,3 +1,7 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
"""Base classes for recipe parsers."""
|
"""Base classes for recipe parsers."""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
@@ -7,7 +11,7 @@ import re
|
|||||||
from typing import Dict, List, Any, Optional, Tuple
|
from typing import Dict, List, Any, Optional, Tuple
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from ..config import config
|
from ..config import config
|
||||||
from ..utils.constants import VALID_LORA_TYPES, VALID_CHECKPOINT_SUB_TYPES
|
from ..utils.constants import MODEL_WEIGHT_FILE_TYPES, VALID_LORA_TYPES, VALID_CHECKPOINT_SUB_TYPES
|
||||||
from ..utils.civitai_utils import rewrite_preview_url
|
from ..utils.civitai_utils import rewrite_preview_url
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -38,7 +42,7 @@ class RecipeMetadataParser(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def populate_lora_from_civitai(lora_entry: Dict[str, Any], civitai_info_tuple: Tuple[Dict[str, Any], Optional[str]],
|
async def populate_lora_from_civitai(lora_entry: Dict[str, Any], civitai_info_tuple: Tuple[Dict[str, Any] | None, str | None] | Dict[str, Any],
|
||||||
recipe_scanner=None, base_model_counts=None, hash_value=None) -> Optional[Dict[str, Any]]:
|
recipe_scanner=None, base_model_counts=None, hash_value=None) -> Optional[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
Populate a lora entry with information from Civitai API response
|
Populate a lora entry with information from Civitai API response
|
||||||
@@ -151,9 +155,9 @@ class RecipeMetadataParser(ABC):
|
|||||||
|
|
||||||
# Process file information if available
|
# Process file information if available
|
||||||
if 'files' in civitai_info:
|
if 'files' in civitai_info:
|
||||||
# Find the primary model file (type="Model" and primary=true) in the files list
|
# Find the primary model file (weights-type and primary=true) in the files list
|
||||||
model_file = next((file for file in civitai_info.get('files', [])
|
model_file = next((file for file in civitai_info.get('files', [])
|
||||||
if file.get('type') == 'Model' and file.get('primary') == True), None)
|
if file.get('type') in MODEL_WEIGHT_FILE_TYPES and file.get('primary') == True), None)
|
||||||
|
|
||||||
if model_file:
|
if model_file:
|
||||||
# Get size
|
# Get size
|
||||||
@@ -175,10 +179,18 @@ class RecipeMetadataParser(ABC):
|
|||||||
lora_entry['localPath'] = local_path
|
lora_entry['localPath'] = local_path
|
||||||
lora_entry['file_name'] = os.path.splitext(os.path.basename(local_path))[0]
|
lora_entry['file_name'] = os.path.splitext(os.path.basename(local_path))[0]
|
||||||
|
|
||||||
# Get thumbnail from local preview if available
|
# Get thumbnail from local preview if available.
|
||||||
|
# Match the cache item by local path first (get_path_by_hash
|
||||||
|
# cascade: 10-char autov2 / 12-char autov3), then by hash.
|
||||||
lora_cache = await lora_scanner.get_cached_data()
|
lora_cache = await lora_scanner.get_cached_data()
|
||||||
lora_item = next((item for item in lora_cache.raw_data
|
h = (lora_entry.get("hash") or "").lower()
|
||||||
if item['sha256'].lower() == lora_entry['hash'].lower()), None)
|
lora_item = next((item for item in lora_cache.raw_data
|
||||||
|
if (item.get("file_path") or "") == local_path), None)
|
||||||
|
if lora_item is None:
|
||||||
|
lora_item = next((item for item in lora_cache.raw_data
|
||||||
|
if (item.get("sha256") or "").lower() == h
|
||||||
|
or (item.get("autov3") or "").lower() == h
|
||||||
|
or (item.get("sha256") or "")[:10].lower() == h), None)
|
||||||
if lora_item and 'preview_url' in lora_item:
|
if lora_item and 'preview_url' in lora_item:
|
||||||
lora_entry['thumbnailUrl'] = config.get_preview_static_url(lora_item['preview_url'])
|
lora_entry['thumbnailUrl'] = config.get_preview_static_url(lora_item['preview_url'])
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -194,7 +206,7 @@ class RecipeMetadataParser(ABC):
|
|||||||
return lora_entry
|
return lora_entry
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def populate_checkpoint_from_civitai(checkpoint: Dict[str, Any], civitai_info: Dict[str, Any]) -> Dict[str, Any]:
|
async def populate_checkpoint_from_civitai(checkpoint: Dict[str, Any], civitai_info: Dict[str, Any] | Tuple[Dict[str, Any] | None, str | None] | None) -> Dict[str, Any]:
|
||||||
"""
|
"""
|
||||||
Populate checkpoint information from Civitai API response
|
Populate checkpoint information from Civitai API response
|
||||||
|
|
||||||
@@ -249,11 +261,21 @@ class RecipeMetadataParser(ABC):
|
|||||||
checkpoint['id'] = civitai_data.get('id', 0)
|
checkpoint['id'] = civitai_data.get('id', 0)
|
||||||
|
|
||||||
if 'files' in civitai_data:
|
if 'files' in civitai_data:
|
||||||
|
# Prefer the file CivitAI marked primary; fall back to any
|
||||||
|
# weights-type file (providers without primary flags).
|
||||||
model_file = next(
|
model_file = next(
|
||||||
(
|
(
|
||||||
file
|
file
|
||||||
for file in civitai_data.get('files', [])
|
for file in civitai_data.get('files', [])
|
||||||
if file.get('type') == 'Model'
|
if file.get('type') in MODEL_WEIGHT_FILE_TYPES
|
||||||
|
and file.get('primary') is True
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
) or next(
|
||||||
|
(
|
||||||
|
file
|
||||||
|
for file in civitai_data.get('files', [])
|
||||||
|
if file.get('type') in MODEL_WEIGHT_FILE_TYPES
|
||||||
),
|
),
|
||||||
None,
|
None,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,3 +1,7 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
import logging
|
import logging
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Factory for creating recipe metadata parsers."""
|
"""Factory for creating recipe metadata parsers."""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
from typing import Any
|
||||||
from .parsers import (
|
from .parsers import (
|
||||||
RecipeFormatParser,
|
RecipeFormatParser,
|
||||||
ComfyMetadataParser,
|
ComfyMetadataParser,
|
||||||
@@ -31,7 +32,8 @@ class RecipeParserFactory:
|
|||||||
# First, try CivitaiApiMetadataParser for dict input
|
# First, try CivitaiApiMetadataParser for dict input
|
||||||
if isinstance(metadata, dict):
|
if isinstance(metadata, dict):
|
||||||
try:
|
try:
|
||||||
if CivitaiApiMetadataParser().is_metadata_matching(metadata):
|
user_comment: Any = metadata
|
||||||
|
if CivitaiApiMetadataParser().is_metadata_matching(user_comment):
|
||||||
return CivitaiApiMetadataParser()
|
return CivitaiApiMetadataParser()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug(f"CivitaiApiMetadataParser check failed: {e}")
|
logger.debug(f"CivitaiApiMetadataParser check failed: {e}")
|
||||||
|
|||||||
@@ -52,7 +52,7 @@ class AutomaticMetadataParser(RecipeMetadataParser):
|
|||||||
negative_and_params = ""
|
negative_and_params = ""
|
||||||
|
|
||||||
# Initialize metadata
|
# Initialize metadata
|
||||||
metadata = {
|
metadata: Dict[str, Any] = {
|
||||||
"prompt": prompt,
|
"prompt": prompt,
|
||||||
"loras": []
|
"loras": []
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ import json
|
|||||||
import logging
|
import logging
|
||||||
from typing import Dict, Any, Union
|
from typing import Dict, Any, Union
|
||||||
from ..base import RecipeMetadataParser
|
from ..base import RecipeMetadataParser
|
||||||
from ..constants import GEN_PARAM_KEYS
|
from ..constants import GEN_PARAM_KEYS, VALID_LORA_TYPES
|
||||||
from ...services.metadata_service import get_default_metadata_provider
|
from ...services.metadata_service import get_default_metadata_provider
|
||||||
from ...config import config
|
from ...config import config
|
||||||
|
|
||||||
@@ -14,15 +14,16 @@ logger = logging.getLogger(__name__)
|
|||||||
class CivitaiApiMetadataParser(RecipeMetadataParser):
|
class CivitaiApiMetadataParser(RecipeMetadataParser):
|
||||||
"""Parser for Civitai image metadata format"""
|
"""Parser for Civitai image metadata format"""
|
||||||
|
|
||||||
def is_metadata_matching(self, metadata) -> bool:
|
def is_metadata_matching(self, user_comment) -> bool:
|
||||||
"""Check if the metadata matches the Civitai image metadata format
|
"""Check if the metadata matches the Civitai image metadata format
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
metadata: The metadata from the image (dict)
|
user_comment: The metadata from the image (dict)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
bool: True if this parser can handle the metadata
|
bool: True if this parser can handle the metadata
|
||||||
"""
|
"""
|
||||||
|
metadata = user_comment
|
||||||
if not metadata or not isinstance(metadata, dict):
|
if not metadata or not isinstance(metadata, dict):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -73,7 +74,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
async def parse_metadata( # type: ignore[override]
|
async def parse_metadata( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||||
self, user_comment, recipe_scanner=None, civitai_client=None,
|
self, user_comment, recipe_scanner=None, civitai_client=None,
|
||||||
local_cache: dict[str, Any] | None = None,
|
local_cache: dict[str, Any] | None = None,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
@@ -89,8 +90,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
Returns:
|
Returns:
|
||||||
Dict containing parsed recipe data
|
Dict containing parsed recipe data
|
||||||
"""
|
"""
|
||||||
metadata: Dict[str, Any] = user_comment # type: ignore[assignment]
|
metadata: Dict[str, Any] = user_comment
|
||||||
metadata = user_comment
|
|
||||||
try:
|
try:
|
||||||
# Get metadata provider instead of using civitai_client directly
|
# Get metadata provider instead of using civitai_client directly
|
||||||
metadata_provider = await get_default_metadata_provider()
|
metadata_provider = await get_default_metadata_provider()
|
||||||
@@ -116,7 +116,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
metadata = inner_meta
|
metadata = inner_meta
|
||||||
|
|
||||||
# Initialize result structure
|
# Initialize result structure
|
||||||
result = {
|
result: Dict[str, Any] = {
|
||||||
"base_model": None,
|
"base_model": None,
|
||||||
"loras": [],
|
"loras": [],
|
||||||
"model": None,
|
"model": None,
|
||||||
@@ -125,10 +125,10 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
}
|
}
|
||||||
|
|
||||||
# Track already added LoRAs to prevent duplicates
|
# Track already added LoRAs to prevent duplicates
|
||||||
added_loras = {} # key: model_version_id or hash, value: index in result["loras"]
|
added_loras: Dict[str, Any] = {} # key: model_version_id or hash, value: index in result["loras"]
|
||||||
|
|
||||||
# Extract hash information from hashes field for LoRA matching
|
# Extract hash information from hashes field for LoRA matching
|
||||||
lora_hashes = {}
|
lora_hashes: Dict[str, Any] = {}
|
||||||
if "hashes" in metadata and isinstance(metadata["hashes"], dict):
|
if "hashes" in metadata and isinstance(metadata["hashes"], dict):
|
||||||
for key, hash_value in metadata["hashes"].items():
|
for key, hash_value in metadata["hashes"].items():
|
||||||
key_str = str(key)
|
key_str = str(key)
|
||||||
@@ -184,7 +184,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
if model_info:
|
if model_info:
|
||||||
result["base_model"] = model_info.get("baseModel", "")
|
result["base_model"] = model_info.get("baseModel", "")
|
||||||
|
|
||||||
base_model_counts = {}
|
base_model_counts: Dict[str, int] = {}
|
||||||
|
|
||||||
# Process standard resources array
|
# Process standard resources array
|
||||||
if "resources" in metadata and isinstance(metadata["resources"], list):
|
if "resources" in metadata and isinstance(metadata["resources"], list):
|
||||||
@@ -196,7 +196,7 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
# identification because it has an explicit type field and hash,
|
# identification because it has an explicit type field and hash,
|
||||||
# unlike modelVersionIds which is a flat list with no type info.
|
# unlike modelVersionIds which is a flat list with no type info.
|
||||||
if resource_type == "model":
|
if resource_type == "model":
|
||||||
checkpoint_entry = {
|
checkpoint_entry: Dict[str, Any] = {
|
||||||
"id": 0,
|
"id": 0,
|
||||||
"modelId": 0,
|
"modelId": 0,
|
||||||
"name": resource.get("name", "Unknown Model"),
|
"name": resource.get("name", "Unknown Model"),
|
||||||
@@ -216,7 +216,8 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
# Try to look up base model from the checkpoint hash
|
# Try to look up base model from the checkpoint hash
|
||||||
cp_hash = checkpoint_entry.get("hash")
|
cp_hash = checkpoint_entry.get("hash")
|
||||||
if cp_hash and metadata_provider:
|
if cp_hash and metadata_provider:
|
||||||
local_cached = local_cache.get(cp_hash) if local_cache else None
|
# local_cache keys are stored lowercase
|
||||||
|
local_cached = local_cache.get(cp_hash.lower()) if local_cache else None
|
||||||
if local_cached:
|
if local_cached:
|
||||||
self._populate_entry_from_cache(
|
self._populate_entry_from_cache(
|
||||||
checkpoint_entry, local_cached
|
checkpoint_entry, local_cached
|
||||||
@@ -294,8 +295,15 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
|
|
||||||
# Try to get info from Civitai if hash is available
|
# Try to get info from Civitai if hash is available
|
||||||
if lora_hash and metadata_provider:
|
if lora_hash and metadata_provider:
|
||||||
local_cached = local_cache.get(lora_hash) if local_cache else None
|
# local_cache keys are stored lowercase
|
||||||
|
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
|
||||||
if local_cached:
|
if local_cached:
|
||||||
|
cached_type = self._cache_item_model_type(local_cached)
|
||||||
|
if cached_type and cached_type not in VALID_LORA_TYPES:
|
||||||
|
logger.debug(
|
||||||
|
f"Skipping non-LoRA cache item for hash {lora_hash}"
|
||||||
|
)
|
||||||
|
continue
|
||||||
self._populate_entry_from_cache(
|
self._populate_entry_from_cache(
|
||||||
lora_entry, local_cached
|
lora_entry, local_cached
|
||||||
)
|
)
|
||||||
@@ -304,6 +312,12 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
added_loras[str(lora_entry["id"])] = len(
|
added_loras[str(lora_entry["id"])] = len(
|
||||||
result["loras"]
|
result["loras"]
|
||||||
)
|
)
|
||||||
|
# Mirror base.py:150-151 counts for API-path loras
|
||||||
|
bm = local_cached.get("base_model") or ""
|
||||||
|
if bm:
|
||||||
|
base_model_counts[bm] = base_model_counts.get(
|
||||||
|
bm, 0
|
||||||
|
) + 1
|
||||||
else:
|
else:
|
||||||
try:
|
try:
|
||||||
civitai_info = (
|
civitai_info = (
|
||||||
@@ -649,30 +663,47 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
}
|
}
|
||||||
|
|
||||||
if metadata_provider:
|
if metadata_provider:
|
||||||
try:
|
# local_cache keys are stored lowercase
|
||||||
civitai_info = await metadata_provider.get_model_by_hash(
|
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
|
||||||
lora_hash
|
if local_cached:
|
||||||
)
|
cached_type = self._cache_item_model_type(local_cached)
|
||||||
|
if cached_type and cached_type not in VALID_LORA_TYPES:
|
||||||
populated_entry = await self.populate_lora_from_civitai(
|
logger.debug(
|
||||||
lora_entry,
|
f"Skipping non-LoRA cache item for hash {lora_hash}"
|
||||||
civitai_info,
|
)
|
||||||
recipe_scanner,
|
|
||||||
base_model_counts,
|
|
||||||
lora_hash,
|
|
||||||
)
|
|
||||||
|
|
||||||
if populated_entry is None:
|
|
||||||
continue
|
continue
|
||||||
|
self._populate_entry_from_cache(lora_entry, local_cached)
|
||||||
lora_entry = populated_entry
|
# Mirror base.py:150-151 counts for API-path loras
|
||||||
|
bm = local_cached.get("base_model") or ""
|
||||||
|
if bm:
|
||||||
|
base_model_counts[bm] = base_model_counts.get(bm, 0) + 1
|
||||||
if "id" in lora_entry and lora_entry["id"]:
|
if "id" in lora_entry and lora_entry["id"]:
|
||||||
added_loras[str(lora_entry["id"])] = len(result["loras"])
|
added_loras[str(lora_entry["id"])] = len(result["loras"])
|
||||||
except Exception as e:
|
else:
|
||||||
logger.error(
|
try:
|
||||||
f"Error fetching Civitai info for LoRA hash {lora_hash}: {e}"
|
civitai_info = await metadata_provider.get_model_by_hash(
|
||||||
)
|
lora_hash
|
||||||
|
)
|
||||||
|
|
||||||
|
populated_entry = await self.populate_lora_from_civitai(
|
||||||
|
lora_entry,
|
||||||
|
civitai_info,
|
||||||
|
recipe_scanner,
|
||||||
|
base_model_counts,
|
||||||
|
lora_hash,
|
||||||
|
)
|
||||||
|
|
||||||
|
if populated_entry is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
lora_entry = populated_entry
|
||||||
|
|
||||||
|
if "id" in lora_entry and lora_entry["id"]:
|
||||||
|
added_loras[str(lora_entry["id"])] = len(result["loras"])
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
f"Error fetching Civitai info for LoRA hash {lora_hash}: {e}"
|
||||||
|
)
|
||||||
|
|
||||||
added_loras[lora_hash] = len(result["loras"])
|
added_loras[lora_hash] = len(result["loras"])
|
||||||
result["loras"].append(lora_entry)
|
result["loras"].append(lora_entry)
|
||||||
@@ -711,32 +742,51 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
|
|
||||||
# Try to get info from Civitai if hash is available
|
# Try to get info from Civitai if hash is available
|
||||||
if lora_entry["hash"] and metadata_provider:
|
if lora_entry["hash"] and metadata_provider:
|
||||||
try:
|
# local_cache keys are stored lowercase
|
||||||
civitai_info = await metadata_provider.get_model_by_hash(
|
local_cached = local_cache.get(lora_hash.lower()) if local_cache else None
|
||||||
lora_hash
|
if local_cached:
|
||||||
)
|
cached_type = self._cache_item_model_type(local_cached)
|
||||||
|
if cached_type and cached_type not in VALID_LORA_TYPES:
|
||||||
populated_entry = await self.populate_lora_from_civitai(
|
logger.debug(
|
||||||
lora_entry,
|
f"Skipping non-LoRA cache item for hash {lora_hash}"
|
||||||
civitai_info,
|
)
|
||||||
recipe_scanner,
|
|
||||||
base_model_counts,
|
|
||||||
lora_hash,
|
|
||||||
)
|
|
||||||
|
|
||||||
if populated_entry is None:
|
|
||||||
lora_index += 1
|
lora_index += 1
|
||||||
continue # Skip invalid LoRA types
|
continue # Skip non-LoRA cache items
|
||||||
|
self._populate_entry_from_cache(lora_entry, local_cached)
|
||||||
lora_entry = populated_entry
|
# Mirror base.py:150-151 counts for API-path loras
|
||||||
|
bm = local_cached.get("base_model") or ""
|
||||||
|
if bm:
|
||||||
|
base_model_counts[bm] = base_model_counts.get(bm, 0) + 1
|
||||||
# If we have a version ID from Civitai, track it for deduplication
|
# If we have a version ID from Civitai, track it for deduplication
|
||||||
if "id" in lora_entry and lora_entry["id"]:
|
if "id" in lora_entry and lora_entry["id"]:
|
||||||
added_loras[str(lora_entry["id"])] = len(result["loras"])
|
added_loras[str(lora_entry["id"])] = len(result["loras"])
|
||||||
except Exception as e:
|
else:
|
||||||
logger.error(
|
try:
|
||||||
f"Error fetching Civitai info for LoRA hash {lora_entry['hash']}: {e}"
|
civitai_info = await metadata_provider.get_model_by_hash(
|
||||||
)
|
lora_hash
|
||||||
|
)
|
||||||
|
|
||||||
|
populated_entry = await self.populate_lora_from_civitai(
|
||||||
|
lora_entry,
|
||||||
|
civitai_info,
|
||||||
|
recipe_scanner,
|
||||||
|
base_model_counts,
|
||||||
|
lora_hash,
|
||||||
|
)
|
||||||
|
|
||||||
|
if populated_entry is None:
|
||||||
|
lora_index += 1
|
||||||
|
continue # Skip invalid LoRA types
|
||||||
|
|
||||||
|
lora_entry = populated_entry
|
||||||
|
|
||||||
|
# If we have a version ID from Civitai, track it for deduplication
|
||||||
|
if "id" in lora_entry and lora_entry["id"]:
|
||||||
|
added_loras[str(lora_entry["id"])] = len(result["loras"])
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
f"Error fetching Civitai info for LoRA hash {lora_entry['hash']}: {e}"
|
||||||
|
)
|
||||||
|
|
||||||
# Track by hash if we have it
|
# Track by hash if we have it
|
||||||
if lora_hash:
|
if lora_hash:
|
||||||
@@ -795,3 +845,14 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
base_model = cache_item.get("base_model", "")
|
base_model = cache_item.get("base_model", "")
|
||||||
if base_model:
|
if base_model:
|
||||||
entry["baseModel"] = base_model
|
entry["baseModel"] = base_model
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _cache_item_model_type(cache_item: dict[str, Any]) -> str:
|
||||||
|
"""Lowercased civitai.model.type of a cache item, or '' when unknown."""
|
||||||
|
civ = cache_item.get("civitai")
|
||||||
|
if not isinstance(civ, dict):
|
||||||
|
return ""
|
||||||
|
model_info = civ.get("model")
|
||||||
|
if not isinstance(model_info, dict):
|
||||||
|
return ""
|
||||||
|
return (model_info.get("type") or "").lower()
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ class MetaFormatParser(RecipeMetadataParser):
|
|||||||
prompt = parts[0].strip()
|
prompt = parts[0].strip()
|
||||||
|
|
||||||
# Initialize metadata
|
# Initialize metadata
|
||||||
metadata = {"prompt": prompt, "loras": []}
|
metadata: Dict[str, Any] = {"prompt": prompt, "loras": []}
|
||||||
|
|
||||||
# Extract negative prompt and parameters if available
|
# Extract negative prompt and parameters if available
|
||||||
if len(parts) > 1:
|
if len(parts) > 1:
|
||||||
|
|||||||
@@ -91,7 +91,15 @@ class RecipeFormatParser(RecipeMetadataParser):
|
|||||||
exists_locally = lora_scanner.has_hash(lora['hash'])
|
exists_locally = lora_scanner.has_hash(lora['hash'])
|
||||||
if exists_locally:
|
if exists_locally:
|
||||||
lora_cache = await lora_scanner.get_cached_data()
|
lora_cache = await lora_scanner.get_cached_data()
|
||||||
lora_item = next((item for item in lora_cache.raw_data if item['sha256'].lower() == lora['hash'].lower()), None)
|
# Cascade match: full sha256, stored autov3, or autov2 (sha256[:10]).
|
||||||
|
h = (lora.get('hash') or '').lower()
|
||||||
|
lora_item = next(
|
||||||
|
(item for item in lora_cache.raw_data
|
||||||
|
if (item.get("sha256") or "").lower() == h
|
||||||
|
or (item.get("autov3") or "").lower() == h
|
||||||
|
or (item.get("sha256") or "")[:10].lower() == h),
|
||||||
|
None
|
||||||
|
)
|
||||||
if lora_item:
|
if lora_item:
|
||||||
lora_entry['existsLocally'] = True
|
lora_entry['existsLocally'] = True
|
||||||
lora_entry['inLibrary'] = True
|
lora_entry['inLibrary'] = True
|
||||||
@@ -148,7 +156,7 @@ class RecipeFormatParser(RecipeMetadataParser):
|
|||||||
checkpoint_data = recipe_metadata.get('checkpoint') or {}
|
checkpoint_data = recipe_metadata.get('checkpoint') or {}
|
||||||
if isinstance(checkpoint_data, dict) and checkpoint_data:
|
if isinstance(checkpoint_data, dict) and checkpoint_data:
|
||||||
version_id = checkpoint_data.get('modelVersionId') or checkpoint_data.get('id')
|
version_id = checkpoint_data.get('modelVersionId') or checkpoint_data.get('id')
|
||||||
checkpoint_entry = {
|
checkpoint_entry: Dict[str, Any] = {
|
||||||
'id': version_id or 0,
|
'id': version_id or 0,
|
||||||
'modelId': checkpoint_data.get('modelId', 0),
|
'modelId': checkpoint_data.get('modelId', 0),
|
||||||
'name': checkpoint_data.get('name', 'Unknown Checkpoint'),
|
'name': checkpoint_data.get('name', 'Unknown Checkpoint'),
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import TYPE_CHECKING, Callable, Dict, Mapping
|
from typing import TYPE_CHECKING, Awaitable, Callable, Dict, Mapping
|
||||||
|
|
||||||
import jinja2
|
import jinja2
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
@@ -30,6 +30,7 @@ from ..services.websocket_progress_callback import (
|
|||||||
WebSocketProgressCallback,
|
WebSocketProgressCallback,
|
||||||
)
|
)
|
||||||
from ..utils.exif_utils import ExifUtils
|
from ..utils.exif_utils import ExifUtils
|
||||||
|
from ..utils.constants import MODEL_WEIGHT_FILE_TYPES
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from .model_route_registrar import COMMON_ROUTE_DEFINITIONS, ModelRouteRegistrar
|
from .model_route_registrar import COMMON_ROUTE_DEFINITIONS, ModelRouteRegistrar
|
||||||
from .handlers.model_handlers import (
|
from .handlers.model_handlers import (
|
||||||
@@ -84,7 +85,7 @@ class BaseModelRoutes(ABC):
|
|||||||
self.metadata_progress_callback = WebSocketBroadcastCallback()
|
self.metadata_progress_callback = WebSocketBroadcastCallback()
|
||||||
|
|
||||||
self._handler_set: ModelHandlerSet | None = None
|
self._handler_set: ModelHandlerSet | None = None
|
||||||
self._handler_mapping: Dict[str, Callable[[web.Request], web.StreamResponse]] | None = None
|
self._handler_mapping: Dict[str, Callable[[web.Request], Awaitable[web.Response]]] | None = None
|
||||||
|
|
||||||
self._preview_service = PreviewAssetService(
|
self._preview_service = PreviewAssetService(
|
||||||
metadata_manager=MetadataManager,
|
metadata_manager=MetadataManager,
|
||||||
@@ -131,7 +132,7 @@ class BaseModelRoutes(ABC):
|
|||||||
self._handler_set = None
|
self._handler_set = None
|
||||||
self._handler_mapping = None
|
self._handler_mapping = None
|
||||||
|
|
||||||
def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
|
def _ensure_handler_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||||
if self._handler_mapping is None:
|
if self._handler_mapping is None:
|
||||||
handler_set = self._create_handler_set()
|
handler_set = self._create_handler_set()
|
||||||
self._handler_set = handler_set
|
self._handler_set = handler_set
|
||||||
@@ -220,7 +221,7 @@ class BaseModelRoutes(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def route_handlers(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
|
def route_handlers(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||||
return self._ensure_handler_mapping()
|
return self._ensure_handler_mapping()
|
||||||
|
|
||||||
def setup_routes(self, app: web.Application, prefix: str) -> None:
|
def setup_routes(self, app: web.Application, prefix: str) -> None:
|
||||||
@@ -237,7 +238,7 @@ class BaseModelRoutes(ABC):
|
|||||||
"""Setup model-specific routes."""
|
"""Setup model-specific routes."""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
def _parse_specific_params(self, request: web.Request) -> Dict:
|
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
|
||||||
"""Parse model-specific parameters - to be overridden by subclasses."""
|
"""Parse model-specific parameters - to be overridden by subclasses."""
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
@@ -251,9 +252,9 @@ class BaseModelRoutes(ABC):
|
|||||||
|
|
||||||
def _find_model_file(self, files):
|
def _find_model_file(self, files):
|
||||||
"""Find the appropriate model file from the files list - can be overridden by subclasses."""
|
"""Find the appropriate model file from the files list - can be overridden by subclasses."""
|
||||||
return next((file for file in files if file.get("type") in ("Model", "Diffusion Model") and file.get("primary") is True), None)
|
return next((file for file in files if file.get("type") in MODEL_WEIGHT_FILE_TYPES and file.get("primary") is True), None)
|
||||||
|
|
||||||
def get_handler(self, name: str) -> Callable[[web.Request], web.StreamResponse]:
|
def get_handler(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]:
|
||||||
"""Expose handlers for subclasses or tests."""
|
"""Expose handlers for subclasses or tests."""
|
||||||
return self._ensure_handler_mapping()[name]
|
return self._ensure_handler_mapping()[name]
|
||||||
|
|
||||||
@@ -285,7 +286,7 @@ class BaseModelRoutes(ABC):
|
|||||||
)
|
)
|
||||||
return self.model_lifecycle_service
|
return self.model_lifecycle_service
|
||||||
|
|
||||||
def _make_handler_proxy(self, name: str) -> Callable[[web.Request], web.StreamResponse]:
|
def _make_handler_proxy(self, name: str) -> Callable[[web.Request], Awaitable[web.StreamResponse]]:
|
||||||
async def proxy(request: web.Request) -> web.StreamResponse:
|
async def proxy(request: web.Request) -> web.StreamResponse:
|
||||||
try:
|
try:
|
||||||
handler = self.get_handler(name)
|
handler = self.get_handler(name)
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from typing import Callable, Mapping
|
from typing import Awaitable, Callable, Mapping
|
||||||
|
|
||||||
import jinja2
|
import jinja2
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
@@ -61,7 +61,9 @@ class BaseRecipeRoutes:
|
|||||||
self._i18n_registered = False
|
self._i18n_registered = False
|
||||||
self._startup_hooks_registered = False
|
self._startup_hooks_registered = False
|
||||||
self._handler_set: RecipeHandlerSet | None = None
|
self._handler_set: RecipeHandlerSet | None = None
|
||||||
self._handler_mapping: dict[str, Callable] | None = None
|
self._handler_mapping: Mapping[
|
||||||
|
str, Callable[[web.Request], Awaitable[web.StreamResponse]]
|
||||||
|
] | None = None
|
||||||
|
|
||||||
async def attach_dependencies(self, app: web.Application | None = None) -> None:
|
async def attach_dependencies(self, app: web.Application | None = None) -> None:
|
||||||
"""Resolve shared services from the registry."""
|
"""Resolve shared services from the registry."""
|
||||||
@@ -84,7 +86,9 @@ class BaseRecipeRoutes:
|
|||||||
app.on_startup.append(self.attach_dependencies)
|
app.on_startup.append(self.attach_dependencies)
|
||||||
self._startup_hooks_registered = True
|
self._startup_hooks_registered = True
|
||||||
|
|
||||||
def to_route_mapping(self) -> Mapping[str, Callable]:
|
def to_route_mapping(
|
||||||
|
self,
|
||||||
|
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||||
"""Return a mapping of handler name to coroutine for registrar binding."""
|
"""Return a mapping of handler name to coroutine for registrar binding."""
|
||||||
|
|
||||||
if self._handler_mapping is None:
|
if self._handler_mapping is None:
|
||||||
@@ -124,17 +128,17 @@ class BaseRecipeRoutes:
|
|||||||
or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
or os.environ.get("HF_HUB_DISABLE_TELEMETRY", "0") == "0"
|
||||||
)
|
)
|
||||||
if not standalone_mode:
|
if not standalone_mode:
|
||||||
from ..metadata_collector import get_metadata # type: ignore[import-not-found]
|
from ..metadata_collector import get_metadata # pyright: ignore[reportMissingImports]
|
||||||
from ..metadata_collector.metadata_processor import ( # type: ignore[import-not-found]
|
from ..metadata_collector.metadata_processor import ( # pyright: ignore[reportMissingImports]
|
||||||
MetadataProcessor,
|
MetadataProcessor,
|
||||||
)
|
)
|
||||||
from ..metadata_collector.metadata_registry import ( # type: ignore[import-not-found]
|
from ..metadata_collector.metadata_registry import ( # pyright: ignore[reportMissingImports]
|
||||||
MetadataRegistry,
|
MetadataRegistry,
|
||||||
)
|
)
|
||||||
else: # pragma: no cover - optional dependency path
|
else: # pragma: no cover - optional dependency path
|
||||||
get_metadata = None # type: ignore[assignment]
|
get_metadata = None # pyright: ignore[reportAssignmentType]
|
||||||
MetadataProcessor = None # type: ignore[assignment]
|
MetadataProcessor = None # pyright: ignore[reportAssignmentType]
|
||||||
MetadataRegistry = None # type: ignore[assignment]
|
MetadataRegistry = None # pyright: ignore[reportAssignmentType]
|
||||||
|
|
||||||
analysis_service = RecipeAnalysisService(
|
analysis_service = RecipeAnalysisService(
|
||||||
exif_utils=ExifUtils,
|
exif_utils=ExifUtils,
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import Dict, List, Set
|
from typing import Any, Dict, List, Set
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
from .base_model_routes import BaseModelRoutes
|
from .base_model_routes import BaseModelRoutes
|
||||||
@@ -28,13 +28,13 @@ class CheckpointRoutes(BaseModelRoutes):
|
|||||||
# Attach service dependencies
|
# Attach service dependencies
|
||||||
self.attach_service(self.service)
|
self.attach_service(self.service)
|
||||||
|
|
||||||
def setup_routes(self, app: web.Application):
|
def setup_routes(self, app: web.Application, prefix: str = "checkpoints"):
|
||||||
"""Setup Checkpoint routes"""
|
"""Setup Checkpoint routes"""
|
||||||
# Schedule service initialization on app startup
|
# Schedule service initialization on app startup
|
||||||
app.on_startup.append(lambda _: self.initialize_services())
|
app.on_startup.append(lambda _: self.initialize_services())
|
||||||
|
|
||||||
# Setup common routes with 'checkpoints' prefix (includes page route)
|
# Setup common routes with 'checkpoints' prefix (includes page route)
|
||||||
super().setup_routes(app, 'checkpoints')
|
super().setup_routes(app, prefix)
|
||||||
|
|
||||||
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
||||||
"""Setup Checkpoint-specific routes"""
|
"""Setup Checkpoint-specific routes"""
|
||||||
@@ -53,9 +53,9 @@ class CheckpointRoutes(BaseModelRoutes):
|
|||||||
"""Get expected model types string for error messages"""
|
"""Get expected model types string for error messages"""
|
||||||
return "Checkpoint"
|
return "Checkpoint"
|
||||||
|
|
||||||
def _parse_specific_params(self, request: web.Request) -> Dict:
|
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
|
||||||
"""Parse Checkpoint-specific parameters"""
|
"""Parse Checkpoint-specific parameters"""
|
||||||
params: Dict = {}
|
params: Dict[str, Any] = {}
|
||||||
|
|
||||||
if 'checkpoint_hash' in request.query:
|
if 'checkpoint_hash' in request.query:
|
||||||
params['hash_filters'] = {'single_hash': request.query['checkpoint_hash'].lower()}
|
params['hash_filters'] = {'single_hash': request.query['checkpoint_hash'].lower()}
|
||||||
@@ -70,7 +70,7 @@ class CheckpointRoutes(BaseModelRoutes):
|
|||||||
"""Get detailed information for a specific checkpoint by name"""
|
"""Get detailed information for a specific checkpoint by name"""
|
||||||
try:
|
try:
|
||||||
name = request.match_info.get('name', '')
|
name = request.match_info.get('name', '')
|
||||||
checkpoint_info = await self.service.get_model_info_by_name(name)
|
checkpoint_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
|
||||||
|
|
||||||
if checkpoint_info:
|
if checkpoint_info:
|
||||||
return web.json_response(checkpoint_info)
|
return web.json_response(checkpoint_info)
|
||||||
@@ -89,7 +89,7 @@ class CheckpointRoutes(BaseModelRoutes):
|
|||||||
roots.extend(config.checkpoints_roots or [])
|
roots.extend(config.checkpoints_roots or [])
|
||||||
roots.extend(config.extra_checkpoints_roots or [])
|
roots.extend(config.extra_checkpoints_roots or [])
|
||||||
# Remove duplicates while preserving order
|
# Remove duplicates while preserving order
|
||||||
seen: set = set()
|
seen: set[str] = set()
|
||||||
unique_roots: List[str] = []
|
unique_roots: List[str] = []
|
||||||
for root in roots:
|
for root in roots:
|
||||||
if root and root not in seen:
|
if root and root not in seen:
|
||||||
@@ -114,7 +114,7 @@ class CheckpointRoutes(BaseModelRoutes):
|
|||||||
roots.extend(config.unet_roots or [])
|
roots.extend(config.unet_roots or [])
|
||||||
roots.extend(config.extra_unet_roots or [])
|
roots.extend(config.extra_unet_roots or [])
|
||||||
# Remove duplicates while preserving order
|
# Remove duplicates while preserving order
|
||||||
seen: set = set()
|
seen: set[str] = set()
|
||||||
unique_roots: List[str] = []
|
unique_roots: List[str] = []
|
||||||
for root in roots:
|
for root in roots:
|
||||||
if root and root not in seen:
|
if root and root not in seen:
|
||||||
|
|||||||
@@ -26,13 +26,13 @@ class EmbeddingRoutes(BaseModelRoutes):
|
|||||||
# Attach service dependencies
|
# Attach service dependencies
|
||||||
self.attach_service(self.service)
|
self.attach_service(self.service)
|
||||||
|
|
||||||
def setup_routes(self, app: web.Application):
|
def setup_routes(self, app: web.Application, prefix: str = "embeddings"):
|
||||||
"""Setup Embedding routes"""
|
"""Setup Embedding routes"""
|
||||||
# Schedule service initialization on app startup
|
# Schedule service initialization on app startup
|
||||||
app.on_startup.append(lambda _: self.initialize_services())
|
app.on_startup.append(lambda _: self.initialize_services())
|
||||||
|
|
||||||
# Setup common routes with 'embeddings' prefix (includes page route)
|
# Setup common routes with 'embeddings' prefix (includes page route)
|
||||||
super().setup_routes(app, 'embeddings')
|
super().setup_routes(app, prefix)
|
||||||
|
|
||||||
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
||||||
"""Setup Embedding-specific routes"""
|
"""Setup Embedding-specific routes"""
|
||||||
@@ -51,7 +51,7 @@ class EmbeddingRoutes(BaseModelRoutes):
|
|||||||
"""Get detailed information for a specific embedding by name"""
|
"""Get detailed information for a specific embedding by name"""
|
||||||
try:
|
try:
|
||||||
name = request.match_info.get('name', '')
|
name = request.match_info.get('name', '')
|
||||||
embedding_info = await self.service.get_model_info_by_name(name)
|
embedding_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
|
||||||
|
|
||||||
if embedding_info:
|
if embedding_info:
|
||||||
return web.json_response(embedding_info)
|
return web.json_response(embedding_info)
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import Callable, Mapping
|
from typing import Any, Awaitable, Callable, Mapping
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
@@ -35,7 +35,7 @@ class ExampleImagesRoutes:
|
|||||||
*,
|
*,
|
||||||
ws_manager,
|
ws_manager,
|
||||||
download_manager: DownloadManager | None = None,
|
download_manager: DownloadManager | None = None,
|
||||||
processor=ExampleImagesProcessor,
|
processor: Any = ExampleImagesProcessor,
|
||||||
file_manager=ExampleImagesFileManager,
|
file_manager=ExampleImagesFileManager,
|
||||||
cleanup_service: ExampleImagesCleanupService | None = None,
|
cleanup_service: ExampleImagesCleanupService | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -46,7 +46,9 @@ class ExampleImagesRoutes:
|
|||||||
self._file_manager = file_manager
|
self._file_manager = file_manager
|
||||||
self._cleanup_service = cleanup_service or ExampleImagesCleanupService()
|
self._cleanup_service = cleanup_service or ExampleImagesCleanupService()
|
||||||
self._handler_set: ExampleImagesHandlerSet | None = None
|
self._handler_set: ExampleImagesHandlerSet | None = None
|
||||||
self._handler_mapping: Mapping[str, Callable[[web.Request], web.StreamResponse]] | None = None
|
self._handler_mapping: Mapping[
|
||||||
|
str, Callable[[web.Request], Awaitable[web.StreamResponse]]
|
||||||
|
] | None = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setup_routes(cls, app: web.Application, *, ws_manager) -> None:
|
def setup_routes(cls, app: web.Application, *, ws_manager) -> None:
|
||||||
@@ -61,7 +63,9 @@ class ExampleImagesRoutes:
|
|||||||
registrar = ExampleImagesRouteRegistrar(app)
|
registrar = ExampleImagesRouteRegistrar(app)
|
||||||
registrar.register_routes(self.to_route_mapping())
|
registrar.register_routes(self.to_route_mapping())
|
||||||
|
|
||||||
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
|
def to_route_mapping(
|
||||||
|
self,
|
||||||
|
) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||||
"""Return the registrar-compatible mapping of handler names to callables."""
|
"""Return the registrar-compatible mapping of handler names to callables."""
|
||||||
|
|
||||||
if self._handler_mapping is None:
|
if self._handler_mapping is None:
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Callable, Mapping
|
from typing import Awaitable, Callable, Mapping
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
@@ -170,7 +170,7 @@ class ExampleImagesHandlerSet:
|
|||||||
management: ExampleImagesManagementHandler
|
management: ExampleImagesManagementHandler
|
||||||
files: ExampleImagesFileHandler
|
files: ExampleImagesFileHandler
|
||||||
|
|
||||||
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], web.StreamResponse]]:
|
def to_route_mapping(self) -> Mapping[str, Callable[[web.Request], Awaitable[web.StreamResponse]]]:
|
||||||
"""Flatten handler methods into the registrar mapping."""
|
"""Flatten handler methods into the registrar mapping."""
|
||||||
|
|
||||||
return {
|
return {
|
||||||
|
|||||||
@@ -276,7 +276,7 @@ def _collect_comfyui_session_logs(
|
|||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
if log_entries is None:
|
if log_entries is None:
|
||||||
try:
|
try:
|
||||||
import app.logger as comfy_logger
|
import app.logger as comfy_logger # pyright: ignore[reportMissingImports]
|
||||||
|
|
||||||
log_entries = list(comfy_logger.get_logs() or [])
|
log_entries = list(comfy_logger.get_logs() or [])
|
||||||
except Exception as exc: # pragma: no cover - environment dependent
|
except Exception as exc: # pragma: no cover - environment dependent
|
||||||
@@ -422,10 +422,10 @@ class PromptServerProtocol(Protocol):
|
|||||||
"""Subset of PromptServer used by the handlers."""
|
"""Subset of PromptServer used by the handlers."""
|
||||||
|
|
||||||
instance: "PromptServerProtocol"
|
instance: "PromptServerProtocol"
|
||||||
sockets: dict # maps clientId (sid) → WebSocketResponse
|
sockets: dict[str, Any] # maps clientId (sid) → WebSocketResponse
|
||||||
|
|
||||||
def send_sync(
|
def send_sync(
|
||||||
self, event: str, payload: dict | None = None, sid: str | None = None
|
self, event: str, payload: dict[str, Any] | None = None, sid: str | None = None
|
||||||
) -> None: # pragma: no cover - protocol
|
) -> None: # pragma: no cover - protocol
|
||||||
...
|
...
|
||||||
|
|
||||||
@@ -443,7 +443,12 @@ class UsageStatsFactory(Protocol):
|
|||||||
class MetadataProviderProtocol(Protocol):
|
class MetadataProviderProtocol(Protocol):
|
||||||
async def get_model_versions(
|
async def get_model_versions(
|
||||||
self, model_id: int
|
self, model_id: int
|
||||||
) -> dict | None: # pragma: no cover - protocol
|
) -> dict[str, Any] | None: # pragma: no cover - protocol
|
||||||
|
...
|
||||||
|
|
||||||
|
async def get_user_models(
|
||||||
|
self, username: str, cursor: str | None = None
|
||||||
|
) -> Any: # pragma: no cover - protocol
|
||||||
...
|
...
|
||||||
|
|
||||||
|
|
||||||
@@ -466,16 +471,16 @@ class MetadataArchiveManagerProtocol(Protocol):
|
|||||||
class BackupServiceProtocol(Protocol):
|
class BackupServiceProtocol(Protocol):
|
||||||
async def create_snapshot(
|
async def create_snapshot(
|
||||||
self, *, snapshot_type: str = "manual", persist: bool = False
|
self, *, snapshot_type: str = "manual", persist: bool = False
|
||||||
) -> dict: # pragma: no cover - protocol
|
) -> dict[str, Any]: # pragma: no cover - protocol
|
||||||
...
|
...
|
||||||
|
|
||||||
async def restore_snapshot(self, archive_path: str) -> dict: # pragma: no cover - protocol
|
async def restore_snapshot(self, archive_path: str) -> dict[str, Any]: # pragma: no cover - protocol
|
||||||
...
|
...
|
||||||
|
|
||||||
def get_status(self) -> dict: # pragma: no cover - protocol
|
def get_status(self) -> dict[str, Any]: # pragma: no cover - protocol
|
||||||
...
|
...
|
||||||
|
|
||||||
def get_available_snapshots(self) -> list[dict]: # pragma: no cover - protocol
|
def get_available_snapshots(self) -> list[dict[str, Any]]: # pragma: no cover - protocol
|
||||||
...
|
...
|
||||||
|
|
||||||
|
|
||||||
@@ -491,7 +496,7 @@ class NodeRegistry:
|
|||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self._lock = asyncio.Lock()
|
self._lock = asyncio.Lock()
|
||||||
# sid → {unique_id → node_info}
|
# sid → {unique_id → node_info}
|
||||||
self._tab_nodes: Dict[str, Dict[str, dict]] = {}
|
self._tab_nodes: Dict[str, Dict[str, dict[str, Any]]] = {}
|
||||||
self._ready = asyncio.Event()
|
self._ready = asyncio.Event()
|
||||||
self._waiting_clients: set[str] = set()
|
self._waiting_clients: set[str] = set()
|
||||||
|
|
||||||
@@ -504,7 +509,7 @@ class NodeRegistry:
|
|||||||
# Helpers to build one node dict (extracted so it's reused for each tab)
|
# Helpers to build one node dict (extracted so it's reused for each tab)
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _build_node_dict(node: dict) -> dict:
|
def _build_node_dict(node: dict[str, Any]) -> dict[str, Any]:
|
||||||
node_id = node["node_id"]
|
node_id = node["node_id"]
|
||||||
graph_id = str(node["graph_id"])
|
graph_id = str(node["graph_id"])
|
||||||
unique_id = f"{graph_id}:{node_id}"
|
unique_id = f"{graph_id}:{node_id}"
|
||||||
@@ -513,11 +518,11 @@ class NodeRegistry:
|
|||||||
bgcolor = node.get("bgcolor") or DEFAULT_NODE_COLOR
|
bgcolor = node.get("bgcolor") or DEFAULT_NODE_COLOR
|
||||||
|
|
||||||
raw_capabilities = node.get("capabilities")
|
raw_capabilities = node.get("capabilities")
|
||||||
capabilities: dict = {}
|
capabilities: dict[str, Any] = {}
|
||||||
if isinstance(raw_capabilities, dict):
|
if isinstance(raw_capabilities, dict):
|
||||||
capabilities = dict(raw_capabilities)
|
capabilities = dict(raw_capabilities)
|
||||||
|
|
||||||
raw_widget_names: list | None = node.get("widget_names")
|
raw_widget_names: list[Any] | None = node.get("widget_names")
|
||||||
if not isinstance(raw_widget_names, list):
|
if not isinstance(raw_widget_names, list):
|
||||||
capability_widget_names = capabilities.get("widget_names")
|
capability_widget_names = capabilities.get("widget_names")
|
||||||
raw_widget_names = (
|
raw_widget_names = (
|
||||||
@@ -565,9 +570,9 @@ class NodeRegistry:
|
|||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Public API
|
# Public API
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
async def register_nodes(self, sid: str, nodes: list[dict]) -> None:
|
async def register_nodes(self, sid: str, nodes: list[dict[str, Any]]) -> None:
|
||||||
"""Register/replace the node list for a single ComfyUI tab (identified by *sid*)."""
|
"""Register/replace the node list for a single ComfyUI tab (identified by *sid*)."""
|
||||||
tab_nodes: dict[str, dict] = {}
|
tab_nodes: dict[str, dict[str, Any]] = {}
|
||||||
for node in nodes:
|
for node in nodes:
|
||||||
nd = self._build_node_dict(node)
|
nd = self._build_node_dict(node)
|
||||||
tab_nodes[nd["unique_id"]] = nd
|
tab_nodes[nd["unique_id"]] = nd
|
||||||
@@ -602,7 +607,7 @@ class NodeRegistry:
|
|||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict:
|
async def get_merged_registry(self, active_sids: set[str] | None = None) -> dict[str, Any]:
|
||||||
"""Return the union of all known tab nodes, pruning any tab that is no
|
"""Return the union of all known tab nodes, pruning any tab that is no
|
||||||
longer connected."""
|
longer connected."""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
@@ -619,8 +624,8 @@ class NodeRegistry:
|
|||||||
len(stale_sids), stale_sids,
|
len(stale_sids), stale_sids,
|
||||||
)
|
)
|
||||||
|
|
||||||
merged: dict[str, dict] = {}
|
merged: dict[str, dict[str, Any]] = {}
|
||||||
tab_info: dict[str, dict] = {}
|
tab_info: dict[str, dict[str, Any]] = {}
|
||||||
for sid, nodes in self._tab_nodes.items():
|
for sid, nodes in self._tab_nodes.items():
|
||||||
tab_info[sid] = {
|
tab_info[sid] = {
|
||||||
"node_count": len(nodes),
|
"node_count": len(nodes),
|
||||||
@@ -653,7 +658,7 @@ class SupportersHandler:
|
|||||||
def __init__(self, logger: logging.Logger | None = None) -> None:
|
def __init__(self, logger: logging.Logger | None = None) -> None:
|
||||||
self._logger = logger or logging.getLogger(__name__)
|
self._logger = logger or logging.getLogger(__name__)
|
||||||
|
|
||||||
def _load_supporters(self) -> dict:
|
def _load_supporters(self) -> dict[str, Any]:
|
||||||
"""Load supporters data from JSON file."""
|
"""Load supporters data from JSON file."""
|
||||||
try:
|
try:
|
||||||
current_file = os.path.abspath(__file__)
|
current_file = os.path.abspath(__file__)
|
||||||
@@ -1229,10 +1234,8 @@ class DoctorHandler:
|
|||||||
settings_snapshot = _sanitize_sensitive_data(
|
settings_snapshot = _sanitize_sensitive_data(
|
||||||
getattr(self._settings, "settings", {}) or {}
|
getattr(self._settings, "settings", {}) or {}
|
||||||
)
|
)
|
||||||
startup_messages_getter = getattr(self._settings, "get_startup_messages", None)
|
startup_messages_getter: Any = getattr(self._settings, "get_startup_messages", None)
|
||||||
startup_messages = (
|
startup_messages = list(startup_messages_getter()) if startup_messages_getter else []
|
||||||
list(startup_messages_getter()) if callable(startup_messages_getter) else []
|
|
||||||
)
|
|
||||||
|
|
||||||
environment = {
|
environment = {
|
||||||
"app_version": app_version,
|
"app_version": app_version,
|
||||||
@@ -1439,7 +1442,7 @@ class SettingsHandler:
|
|||||||
*,
|
*,
|
||||||
settings_service=None,
|
settings_service=None,
|
||||||
metadata_provider_updater: Callable[
|
metadata_provider_updater: Callable[
|
||||||
[], Awaitable[None]
|
[], Awaitable[Any]
|
||||||
] = update_metadata_providers,
|
] = update_metadata_providers,
|
||||||
downloader_factory: Callable[
|
downloader_factory: Callable[
|
||||||
[], Awaitable[DownloaderProtocol]
|
[], Awaitable[DownloaderProtocol]
|
||||||
@@ -1484,8 +1487,8 @@ class SettingsHandler:
|
|||||||
settings_file = getattr(self._settings, "settings_file", None)
|
settings_file = getattr(self._settings, "settings_file", None)
|
||||||
if settings_file:
|
if settings_file:
|
||||||
response_data["settings_file"] = settings_file
|
response_data["settings_file"] = settings_file
|
||||||
messages_getter = getattr(self._settings, "get_startup_messages", None)
|
messages_getter: Any = getattr(self._settings, "get_startup_messages", None)
|
||||||
messages = list(messages_getter()) if callable(messages_getter) else []
|
messages = list(messages_getter()) if messages_getter else []
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
"success": True,
|
"success": True,
|
||||||
@@ -1562,6 +1565,11 @@ class SettingsHandler:
|
|||||||
{"success": False, "error": validation_error}
|
{"success": False, "error": validation_error}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if key == "update_channel" and value not in ("release", "nightly"):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "update_channel must be 'release' or 'nightly'"}
|
||||||
|
)
|
||||||
|
|
||||||
if value == "__DELETE__" and key in (
|
if value == "__DELETE__" and key in (
|
||||||
"proxy_username",
|
"proxy_username",
|
||||||
"proxy_password",
|
"proxy_password",
|
||||||
@@ -2000,11 +2008,11 @@ async def _noop_backup_service() -> None:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ServiceRegistryAdapter:
|
class ServiceRegistryAdapter:
|
||||||
get_lora_scanner: Callable[[], Awaitable]
|
get_lora_scanner: Callable[[], Awaitable[Any]]
|
||||||
get_checkpoint_scanner: Callable[[], Awaitable]
|
get_checkpoint_scanner: Callable[[], Awaitable[Any]]
|
||||||
get_embedding_scanner: Callable[[], Awaitable]
|
get_embedding_scanner: Callable[[], Awaitable[Any]]
|
||||||
get_downloaded_version_history_service: Callable[[], Awaitable]
|
get_downloaded_version_history_service: Callable[[], Awaitable[Any]]
|
||||||
get_backup_service: Callable[[], Awaitable] = _noop_backup_service
|
get_backup_service: Callable[[], Awaitable[Any]] = _noop_backup_service
|
||||||
|
|
||||||
|
|
||||||
class ModelLibraryHandler:
|
class ModelLibraryHandler:
|
||||||
@@ -2045,8 +2053,8 @@ class ModelLibraryHandler:
|
|||||||
return await self._service_registry.get_downloaded_version_history_service()
|
return await self._service_registry.get_downloaded_version_history_service()
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _with_downloaded_flag(versions: list[dict]) -> list[dict]:
|
def _with_downloaded_flag(versions: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||||
enriched: list[dict] = []
|
enriched: list[dict[str, Any]] = []
|
||||||
for version in versions:
|
for version in versions:
|
||||||
entry = dict(version)
|
entry = dict(version)
|
||||||
entry.setdefault("hasBeenDownloaded", True)
|
entry.setdefault("hasBeenDownloaded", True)
|
||||||
@@ -2239,7 +2247,7 @@ class ModelLibraryHandler:
|
|||||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||||
embedding_scanner = await self._service_registry.get_embedding_scanner()
|
embedding_scanner = await self._service_registry.get_embedding_scanner()
|
||||||
|
|
||||||
results: list[dict] = []
|
results: list[dict[str, Any]] = []
|
||||||
for model_id in model_ids:
|
for model_id in model_ids:
|
||||||
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
|
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
|
||||||
if lora_versions:
|
if lora_versions:
|
||||||
@@ -2348,7 +2356,7 @@ class ModelLibraryHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
model_version_id = int(data.get("modelVersionId"))
|
model_version_id = int(data.get("modelVersionId")) # pyright: ignore[reportArgumentType]
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "Parameter modelVersionId must be an integer"},
|
{"success": False, "error": "Parameter modelVersionId must be an integer"},
|
||||||
@@ -2460,10 +2468,11 @@ class ModelLibraryHandler:
|
|||||||
"checkpoint": checkpoint_scanner,
|
"checkpoint": checkpoint_scanner,
|
||||||
"embedding": embedding_scanner,
|
"embedding": embedding_scanner,
|
||||||
}
|
}
|
||||||
scanner = scanner_map.get(found_type)
|
scanner = scanner_map.get(found_type or "")
|
||||||
if scanner:
|
if scanner:
|
||||||
persist = getattr(scanner, "_persist_current_cache", None)
|
scanner.bump_cache_version()
|
||||||
if callable(persist):
|
persist: Any = getattr(scanner, "_persist_current_cache", None)
|
||||||
|
if persist:
|
||||||
await persist()
|
await persist()
|
||||||
|
|
||||||
history_service = await self._get_download_history_service()
|
history_service = await self._get_download_history_service()
|
||||||
@@ -2585,6 +2594,8 @@ class ModelLibraryHandler:
|
|||||||
status=400,
|
status=400,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
cursor = request.query.get("cursor")
|
||||||
|
|
||||||
metadata_provider = await self._metadata_provider_factory()
|
metadata_provider = await self._metadata_provider_factory()
|
||||||
if not metadata_provider:
|
if not metadata_provider:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2593,7 +2604,7 @@ class ModelLibraryHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
models = await metadata_provider.get_user_models(username)
|
result = await metadata_provider.get_user_models(username, cursor)
|
||||||
except NotImplementedError:
|
except NotImplementedError:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
@@ -2603,14 +2614,35 @@ class ModelLibraryHandler:
|
|||||||
status=501,
|
status=501,
|
||||||
)
|
)
|
||||||
|
|
||||||
if models is None:
|
if result is None:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "Failed to fetch user models"},
|
{"success": False, "error": "Failed to fetch user models"},
|
||||||
status=502,
|
status=502,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if isinstance(result, dict):
|
||||||
|
models = result.get("items")
|
||||||
|
next_cursor = result.get("nextCursor")
|
||||||
|
else:
|
||||||
|
# Defensive: tolerate providers that still return a raw list
|
||||||
|
models = result
|
||||||
|
next_cursor = None
|
||||||
|
|
||||||
if not isinstance(models, list):
|
if not isinstance(models, list):
|
||||||
models = []
|
models = []
|
||||||
|
if next_cursor is not None and not isinstance(next_cursor, str):
|
||||||
|
next_cursor = str(next_cursor)
|
||||||
|
|
||||||
|
estimated_total = None
|
||||||
|
if cursor is None:
|
||||||
|
get_count = getattr(metadata_provider, "get_creator_model_count", None)
|
||||||
|
if get_count is not None:
|
||||||
|
try:
|
||||||
|
estimated_total = await get_count(username)
|
||||||
|
except Exception: # best-effort only
|
||||||
|
estimated_total = None
|
||||||
|
if not isinstance(estimated_total, int):
|
||||||
|
estimated_total = None
|
||||||
|
|
||||||
lora_scanner = await self._service_registry.get_lora_scanner()
|
lora_scanner = await self._service_registry.get_lora_scanner()
|
||||||
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
checkpoint_scanner = await self._service_registry.get_checkpoint_scanner()
|
||||||
@@ -2621,15 +2653,16 @@ class ModelLibraryHandler:
|
|||||||
}
|
}
|
||||||
lora_type_aliases = {model_type.lower() for model_type in VALID_LORA_TYPES}
|
lora_type_aliases = {model_type.lower() for model_type in VALID_LORA_TYPES}
|
||||||
|
|
||||||
type_scanner_map: Dict[str, object | None] = {
|
type_scanner_map: Dict[str, Any] = {
|
||||||
**{alias: lora_scanner for alias in lora_type_aliases},
|
**{alias: lora_scanner for alias in lora_type_aliases},
|
||||||
"checkpoint": checkpoint_scanner,
|
"checkpoint": checkpoint_scanner,
|
||||||
"textualinversion": embedding_scanner,
|
"textualinversion": embedding_scanner,
|
||||||
}
|
}
|
||||||
|
|
||||||
versions: list[dict] = []
|
versions: list[dict[str, Any]] = []
|
||||||
history_service = await self._get_download_history_service()
|
history_service = await self._get_download_history_service()
|
||||||
model_ids: list[int] = []
|
model_ids: list[int] = []
|
||||||
|
model_count = 0
|
||||||
for model in models:
|
for model in models:
|
||||||
try:
|
try:
|
||||||
model_ids.append(int(model.get("id")))
|
model_ids.append(int(model.get("id")))
|
||||||
@@ -2663,6 +2696,8 @@ class ModelLibraryHandler:
|
|||||||
if model_type not in normalized_allowed_types:
|
if model_type not in normalized_allowed_types:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
model_count += 1
|
||||||
|
|
||||||
scanner = type_scanner_map.get(model_type)
|
scanner = type_scanner_map.get(model_type)
|
||||||
if scanner is None:
|
if scanner is None:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2676,6 +2711,8 @@ class ModelLibraryHandler:
|
|||||||
tags_value = model.get("tags")
|
tags_value = model.get("tags")
|
||||||
tags = tags_value if isinstance(tags_value, list) else []
|
tags = tags_value if isinstance(tags_value, list) else []
|
||||||
model_id = model.get("id")
|
model_id = model.get("id")
|
||||||
|
if model_id is None:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
model_id_int = int(model_id)
|
model_id_int = int(model_id)
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
@@ -2691,6 +2728,8 @@ class ModelLibraryHandler:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
version_id = version.get("id")
|
version_id = version.get("id")
|
||||||
|
if version_id is None:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
version_id_int = int(version_id)
|
version_id_int = int(version_id)
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
@@ -2728,7 +2767,15 @@ class ModelLibraryHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": True, "username": username, "versions": versions}
|
{
|
||||||
|
"success": True,
|
||||||
|
"username": username,
|
||||||
|
"versions": versions,
|
||||||
|
"modelCount": model_count,
|
||||||
|
"nextCursor": next_cursor,
|
||||||
|
"hasMore": next_cursor is not None,
|
||||||
|
"estimatedTotal": estimated_total,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
except Exception as exc: # pragma: no cover - defensive logging
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
logger.error("Failed to get Civitai user models: %s", exc, exc_info=True)
|
||||||
@@ -2744,7 +2791,7 @@ class MetadataArchiveHandler:
|
|||||||
] = get_metadata_archive_manager,
|
] = get_metadata_archive_manager,
|
||||||
settings_service=None,
|
settings_service=None,
|
||||||
metadata_provider_updater: Callable[
|
metadata_provider_updater: Callable[
|
||||||
[], Awaitable[None]
|
[], Awaitable[Any]
|
||||||
] = update_metadata_providers,
|
] = update_metadata_providers,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._metadata_archive_manager_factory = metadata_archive_manager_factory
|
self._metadata_archive_manager_factory = metadata_archive_manager_factory
|
||||||
@@ -2891,7 +2938,7 @@ class BackupHandler:
|
|||||||
|
|
||||||
if request.content_type.startswith("multipart/"):
|
if request.content_type.startswith("multipart/"):
|
||||||
reader = await request.multipart()
|
reader = await request.multipart()
|
||||||
field = await reader.next()
|
field: Any = await reader.next()
|
||||||
uploaded = False
|
uploaded = False
|
||||||
while field is not None:
|
while field is not None:
|
||||||
if getattr(field, "filename", None):
|
if getattr(field, "filename", None):
|
||||||
@@ -3510,7 +3557,7 @@ class NodeRegistryHandler:
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
parsed_node_id = node_identifier
|
parsed_node_id = node_identifier
|
||||||
|
|
||||||
payload: dict = {
|
payload: dict[str, Any] = {
|
||||||
"id": parsed_node_id,
|
"id": parsed_node_id,
|
||||||
"value": value,
|
"value": value,
|
||||||
"mode": mode,
|
"mode": mode,
|
||||||
@@ -3634,7 +3681,7 @@ class NodeRegistryHandler:
|
|||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
parsed_node_id = node_identifier
|
parsed_node_id = node_identifier
|
||||||
|
|
||||||
payload: dict = {
|
payload: dict[str, Any] = {
|
||||||
"id": parsed_node_id,
|
"id": parsed_node_id,
|
||||||
"value": value,
|
"value": value,
|
||||||
"mode": mode,
|
"mode": mode,
|
||||||
@@ -3701,8 +3748,8 @@ class MiscHandlerSet:
|
|||||||
doctor: DoctorHandler,
|
doctor: DoctorHandler,
|
||||||
example_workflows: ExampleWorkflowsHandler,
|
example_workflows: ExampleWorkflowsHandler,
|
||||||
base_model: BaseModelHandlerSet,
|
base_model: BaseModelHandlerSet,
|
||||||
hf_handler: HfHandler | None = None,
|
hf_handler: Any = None,
|
||||||
agent_handler: AgentHandler | None = None,
|
agent_handler: Any = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.health = health
|
self.health = health
|
||||||
self.settings = settings
|
self.settings = settings
|
||||||
|
|||||||
@@ -51,6 +51,29 @@ LICENSE_FIELDS = (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_broadcast_models_changed_tasks: set = set()
|
||||||
|
|
||||||
|
|
||||||
|
def _broadcast_models_changed() -> None:
|
||||||
|
"""Notify connected clients that the local model library changed.
|
||||||
|
|
||||||
|
The ComfyUI graph page listens for this event to invalidate its cached
|
||||||
|
model availability data (loras widget missing-model cues / error flags)
|
||||||
|
without waiting for the cache TTL to expire.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from ...services.websocket_manager import ws_manager
|
||||||
|
|
||||||
|
task = asyncio.create_task(ws_manager.broadcast({"type": "models_changed"}))
|
||||||
|
# Keep a reference so the task is not garbage-collected mid-await.
|
||||||
|
_broadcast_models_changed_tasks.add(task)
|
||||||
|
task.add_done_callback(_broadcast_models_changed_tasks.discard)
|
||||||
|
except Exception:
|
||||||
|
logging.getLogger(__name__).debug(
|
||||||
|
"Failed to broadcast models_changed", exc_info=True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ModelPageView:
|
class ModelPageView:
|
||||||
"""Render the HTML view for model listings."""
|
"""Render the HTML view for model listings."""
|
||||||
|
|
||||||
@@ -71,7 +94,7 @@ class ModelPageView:
|
|||||||
self._server_i18n = server_i18n
|
self._server_i18n = server_i18n
|
||||||
self._logger = logger
|
self._logger = logger
|
||||||
|
|
||||||
def _load_supporters(self) -> dict:
|
def _load_supporters(self) -> dict[str, Any]:
|
||||||
"""Load supporters data from JSON file."""
|
"""Load supporters data from JSON file."""
|
||||||
try:
|
try:
|
||||||
current_file = os.path.abspath(__file__)
|
current_file = os.path.abspath(__file__)
|
||||||
@@ -152,7 +175,7 @@ class ModelPageView:
|
|||||||
self._template_env.filters["t"] = (
|
self._template_env.filters["t"] = (
|
||||||
self._server_i18n.create_template_filter()
|
self._server_i18n.create_template_filter()
|
||||||
)
|
)
|
||||||
self._template_env._i18n_filter_added = True # type: ignore[attr-defined]
|
self._template_env._i18n_filter_added = True # pyright: ignore[reportAttributeAccessIssue]
|
||||||
|
|
||||||
from ...services.llm_service import PROVIDER_PRESETS
|
from ...services.llm_service import PROVIDER_PRESETS
|
||||||
|
|
||||||
@@ -199,7 +222,7 @@ class ModelListingHandler:
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
service,
|
service,
|
||||||
parse_specific_params: Callable[[web.Request], Dict],
|
parse_specific_params: Callable[[web.Request], Dict[str, Any]],
|
||||||
logger: logging.Logger,
|
logger: logging.Logger,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._service = service
|
self._service = service
|
||||||
@@ -287,7 +310,7 @@ class ModelListingHandler:
|
|||||||
)
|
)
|
||||||
return web.json_response({"error": str(exc)}, status=500)
|
return web.json_response({"error": str(exc)}, status=500)
|
||||||
|
|
||||||
def _parse_common_params(self, request: web.Request) -> Dict:
|
def _parse_common_params(self, request: web.Request) -> Dict[str, Any]:
|
||||||
page = int(request.query.get("page", "1"))
|
page = int(request.query.get("page", "1"))
|
||||||
page_size = min(int(request.query.get("page_size", "20")), 100)
|
page_size = min(int(request.query.get("page_size", "20")), 100)
|
||||||
sort_by = request.query.get("sort_by", "name")
|
sort_by = request.query.get("sort_by", "name")
|
||||||
@@ -394,12 +417,14 @@ class ModelListingHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# View-local-versions filter: show all local versions of a specific model
|
# View-local-versions filter: show all local versions of a specific model
|
||||||
|
# Accepts either a CivitAI modelId (int) or a HF group key like "hf:user/repo"
|
||||||
civitai_model_id = request.query.get("civitai_model_id")
|
civitai_model_id = request.query.get("civitai_model_id")
|
||||||
if civitai_model_id is not None:
|
if civitai_model_id is not None:
|
||||||
try:
|
try:
|
||||||
civitai_model_id = int(civitai_model_id)
|
civitai_model_id = int(civitai_model_id)
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
civitai_model_id = None
|
# Keep as string — could be an HF group key (e.g. "hf:user/repo")
|
||||||
|
pass
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"page": page,
|
"page": page,
|
||||||
@@ -458,6 +483,7 @@ class ModelManagementHandler:
|
|||||||
return web.Response(text="Model path is required", status=400)
|
return web.Response(text="Model path is required", status=400)
|
||||||
|
|
||||||
result = await self._lifecycle_service.delete_model(file_path)
|
result = await self._lifecycle_service.delete_model(file_path)
|
||||||
|
_broadcast_models_changed()
|
||||||
return web.json_response(result)
|
return web.json_response(result)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=400)
|
return web.json_response({"success": False, "error": str(exc)}, status=400)
|
||||||
@@ -537,6 +563,7 @@ class ModelManagementHandler:
|
|||||||
# Update model_data with new hash
|
# Update model_data with new hash
|
||||||
model_data["sha256"] = sha256
|
model_data["sha256"] = sha256
|
||||||
model_data["hash_status"] = "completed"
|
model_data["hash_status"] = "completed"
|
||||||
|
hash_status = "completed"
|
||||||
else:
|
else:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "No SHA256 hash found"}, status=400
|
{"success": False, "error": "No SHA256 hash found"}, status=400
|
||||||
@@ -544,6 +571,32 @@ class ModelManagementHandler:
|
|||||||
|
|
||||||
await MetadataManager.hydrate_model_data(model_data)
|
await MetadataManager.hydrate_model_data(model_data)
|
||||||
|
|
||||||
|
# hydrate_model_data replaces model_data with .metadata.json content,
|
||||||
|
# which may lack sha256. Restore from cache and persist the fix.
|
||||||
|
if not model_data.get("sha256"):
|
||||||
|
if sha256:
|
||||||
|
model_data["sha256"] = sha256
|
||||||
|
model_data["hash_status"] = model_data.get("hash_status", hash_status)
|
||||||
|
data_to_save = model_data.copy()
|
||||||
|
data_to_save.pop("folder", None)
|
||||||
|
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||||
|
else:
|
||||||
|
sha256 = await calculate_sha256(file_path)
|
||||||
|
if sha256:
|
||||||
|
model_data["sha256"] = sha256.lower()
|
||||||
|
model_data["hash_status"] = "completed"
|
||||||
|
data_to_save = model_data.copy()
|
||||||
|
data_to_save.pop("folder", None)
|
||||||
|
await MetadataManager.save_metadata(file_path, data_to_save)
|
||||||
|
else:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "Failed to compute SHA256 hash for model",
|
||||||
|
},
|
||||||
|
status=500,
|
||||||
|
)
|
||||||
|
|
||||||
success, error = await self._metadata_sync.fetch_and_update_model(
|
success, error = await self._metadata_sync.fetch_and_update_model(
|
||||||
sha256=model_data["sha256"],
|
sha256=model_data["sha256"],
|
||||||
file_path=file_path,
|
file_path=file_path,
|
||||||
@@ -566,7 +619,12 @@ class ModelManagementHandler:
|
|||||||
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
|
{"success": False, "error": OFFLINE_FRIENDLY_MESSAGE},
|
||||||
status=503,
|
status=503,
|
||||||
)
|
)
|
||||||
self._logger.error("Error fetching from CivitAI: %s", exc, exc_info=True)
|
self._logger.error(
|
||||||
|
"Error fetching from CivitAI for %s: %s",
|
||||||
|
locals().get("file_path", "unknown"),
|
||||||
|
exc,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
async def relink_civitai(self, request: web.Request) -> web.Response:
|
async def relink_civitai(self, request: web.Request) -> web.Response:
|
||||||
@@ -624,7 +682,7 @@ class ModelManagementHandler:
|
|||||||
try:
|
try:
|
||||||
reader = await request.multipart()
|
reader = await request.multipart()
|
||||||
|
|
||||||
field = await reader.next()
|
field: Any = await reader.next()
|
||||||
if field is None or field.name != "preview_file":
|
if field is None or field.name != "preview_file":
|
||||||
raise ValueError("Expected 'preview_file' field")
|
raise ValueError("Expected 'preview_file' field")
|
||||||
content_type = field.headers.get("Content-Type", "image/png")
|
content_type = field.headers.get("Content-Type", "image/png")
|
||||||
@@ -666,7 +724,7 @@ class ModelManagementHandler:
|
|||||||
{
|
{
|
||||||
"success": True,
|
"success": True,
|
||||||
"preview_url": config.get_preview_static_url(
|
"preview_url": config.get_preview_static_url(
|
||||||
result["preview_path"]
|
str(result["preview_path"])
|
||||||
),
|
),
|
||||||
"preview_nsfw_level": result["preview_nsfw_level"],
|
"preview_nsfw_level": result["preview_nsfw_level"],
|
||||||
}
|
}
|
||||||
@@ -747,7 +805,7 @@ class ModelManagementHandler:
|
|||||||
|
|
||||||
result = await self._preview_service.replace_preview(
|
result = await self._preview_service.replace_preview(
|
||||||
model_path=model_path,
|
model_path=model_path,
|
||||||
preview_data=preview_data,
|
preview_data=preview_bytes,
|
||||||
content_type=content_type,
|
content_type=content_type,
|
||||||
original_filename=original_filename,
|
original_filename=original_filename,
|
||||||
nsfw_level=nsfw_level,
|
nsfw_level=nsfw_level,
|
||||||
@@ -759,7 +817,7 @@ class ModelManagementHandler:
|
|||||||
{
|
{
|
||||||
"success": True,
|
"success": True,
|
||||||
"preview_url": config.get_preview_static_url(
|
"preview_url": config.get_preview_static_url(
|
||||||
result["preview_path"]
|
str(result["preview_path"])
|
||||||
),
|
),
|
||||||
"preview_nsfw_level": result["preview_nsfw_level"],
|
"preview_nsfw_level": result["preview_nsfw_level"],
|
||||||
}
|
}
|
||||||
@@ -897,6 +955,8 @@ class ModelManagementHandler:
|
|||||||
file_path=file_path, new_file_name=new_file_name
|
file_path=file_path, new_file_name=new_file_name
|
||||||
)
|
)
|
||||||
|
|
||||||
|
_broadcast_models_changed()
|
||||||
|
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
**result,
|
**result,
|
||||||
@@ -925,6 +985,7 @@ class ModelManagementHandler:
|
|||||||
)
|
)
|
||||||
|
|
||||||
result = await self._lifecycle_service.bulk_delete_models(file_paths)
|
result = await self._lifecycle_service.bulk_delete_models(file_paths)
|
||||||
|
_broadcast_models_changed()
|
||||||
return web.json_response(result)
|
return web.json_response(result)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=400)
|
return web.json_response({"success": False, "error": str(exc)}, status=400)
|
||||||
@@ -1027,6 +1088,7 @@ class ModelQueryHandler:
|
|||||||
await self._service.scan_models(
|
await self._service.scan_models(
|
||||||
force_refresh=True, rebuild_cache=full_rebuild
|
force_refresh=True, rebuild_cache=full_rebuild
|
||||||
)
|
)
|
||||||
|
_broadcast_models_changed()
|
||||||
if self._service.scanner.is_cancelled():
|
if self._service.scanner.is_cancelled():
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
@@ -1454,8 +1516,73 @@ class ModelQueryHandler:
|
|||||||
search = request.query.get("search", "").strip()
|
search = request.query.get("search", "").strip()
|
||||||
limit = min(int(request.query.get("limit", "15")), 100)
|
limit = min(int(request.query.get("limit", "15")), 100)
|
||||||
offset = max(0, int(request.query.get("offset", "0")))
|
offset = max(0, int(request.query.get("offset", "0")))
|
||||||
|
|
||||||
|
folder = request.query.get("folder")
|
||||||
|
recursive = request.query.get("recursive", "true").lower() == "true"
|
||||||
|
base_models = list(request.query.getall("base_model", []))
|
||||||
|
model_types = list(request.query.getall("model_type", []))
|
||||||
|
|
||||||
|
tag_filters: Dict[str, str] = {}
|
||||||
|
for tag in request.query.getall("tag_include", []):
|
||||||
|
if tag:
|
||||||
|
tag_filters[tag] = "include"
|
||||||
|
for tag in request.query.getall("tag_exclude", []):
|
||||||
|
if tag:
|
||||||
|
tag_filters[tag] = "exclude"
|
||||||
|
|
||||||
|
auto_tag_filters: Dict[str, str] = {}
|
||||||
|
for tag in request.query.getall("auto_tag_include", []):
|
||||||
|
if tag:
|
||||||
|
auto_tag_filters[tag] = "include"
|
||||||
|
for tag in request.query.getall("auto_tag_exclude", []):
|
||||||
|
if tag:
|
||||||
|
auto_tag_filters[tag] = "exclude"
|
||||||
|
|
||||||
|
tag_logic = request.query.get("tag_logic", "any").lower()
|
||||||
|
if tag_logic not in ("any", "all"):
|
||||||
|
tag_logic = "any"
|
||||||
|
|
||||||
|
credit_required = request.query.get("credit_required")
|
||||||
|
if credit_required is not None:
|
||||||
|
credit_required = credit_required.lower() not in ("false", "0", "")
|
||||||
|
|
||||||
|
allow_selling_generated_content = request.query.get(
|
||||||
|
"allow_selling_generated_content"
|
||||||
|
)
|
||||||
|
if allow_selling_generated_content is not None:
|
||||||
|
allow_selling_generated_content = (
|
||||||
|
allow_selling_generated_content.lower() not in ("false", "0", "")
|
||||||
|
)
|
||||||
|
|
||||||
|
# The presence of the recursive param (always sent by the loras
|
||||||
|
# widget when filter mode is on) signals that the filter pipeline
|
||||||
|
# must run even when no concrete filter is set, so global settings
|
||||||
|
# like show_only_sfw stay consistent with the list endpoint.
|
||||||
|
apply_filters = (
|
||||||
|
"recursive" in request.query
|
||||||
|
or folder is not None
|
||||||
|
or bool(base_models)
|
||||||
|
or bool(model_types)
|
||||||
|
or bool(tag_filters)
|
||||||
|
or bool(auto_tag_filters)
|
||||||
|
or credit_required is not None
|
||||||
|
or allow_selling_generated_content is not None
|
||||||
|
)
|
||||||
|
|
||||||
matching_paths = await self._service.search_relative_paths(
|
matching_paths = await self._service.search_relative_paths(
|
||||||
search, limit, offset
|
search,
|
||||||
|
limit,
|
||||||
|
offset,
|
||||||
|
folder=folder,
|
||||||
|
recursive=recursive,
|
||||||
|
base_models=base_models,
|
||||||
|
model_types=model_types,
|
||||||
|
tags=tag_filters,
|
||||||
|
auto_tags=auto_tag_filters,
|
||||||
|
tag_logic=tag_logic,
|
||||||
|
credit_required=credit_required,
|
||||||
|
allow_selling_generated_content=allow_selling_generated_content,
|
||||||
|
apply_filters=apply_filters,
|
||||||
)
|
)
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": True, "relative_paths": matching_paths}
|
{"success": True, "relative_paths": matching_paths}
|
||||||
@@ -1961,7 +2088,7 @@ class ModelCivitaiHandler:
|
|||||||
settings_service: SettingsManager,
|
settings_service: SettingsManager,
|
||||||
ws_manager: WebSocketManager,
|
ws_manager: WebSocketManager,
|
||||||
logger: logging.Logger,
|
logger: logging.Logger,
|
||||||
metadata_provider_factory: Callable[[], Awaitable],
|
metadata_provider_factory: Callable[[], Awaitable[Any]],
|
||||||
validate_model_type: Callable[[str], bool],
|
validate_model_type: Callable[[str], bool],
|
||||||
expected_model_types: Callable[[], str],
|
expected_model_types: Callable[[], str],
|
||||||
find_model_file: Callable[
|
find_model_file: Callable[
|
||||||
@@ -2026,7 +2153,7 @@ class ModelCivitaiHandler:
|
|||||||
downloaded_version_ids = set(
|
downloaded_version_ids = set(
|
||||||
await history_service.get_downloaded_version_ids(
|
await history_service.get_downloaded_version_ids(
|
||||||
self._service.model_type,
|
self._service.model_type,
|
||||||
model_id,
|
int(model_id),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
except Exception as exc: # pragma: no cover - defensive logging
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
@@ -2136,6 +2263,8 @@ class ModelMoveHandler:
|
|||||||
result = await self._move_service.move_model(
|
result = await self._move_service.move_model(
|
||||||
file_path, target_path, use_default_paths=use_default_paths
|
file_path, target_path, use_default_paths=use_default_paths
|
||||||
)
|
)
|
||||||
|
if result.get("success"):
|
||||||
|
_broadcast_models_changed()
|
||||||
status = 200 if result.get("success") else 500
|
status = 200 if result.get("success") else 500
|
||||||
return web.json_response(result, status=status)
|
return web.json_response(result, status=status)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -2155,6 +2284,8 @@ class ModelMoveHandler:
|
|||||||
result = await self._move_service.move_models_bulk(
|
result = await self._move_service.move_models_bulk(
|
||||||
file_paths, target_path, use_default_paths=use_default_paths
|
file_paths, target_path, use_default_paths=use_default_paths
|
||||||
)
|
)
|
||||||
|
if result.get("success"):
|
||||||
|
_broadcast_models_changed()
|
||||||
return web.json_response(result)
|
return web.json_response(result)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self._logger.error("Error moving models in bulk: %s", exc, exc_info=True)
|
self._logger.error("Error moving models in bulk: %s", exc, exc_info=True)
|
||||||
@@ -2200,6 +2331,7 @@ class ModelAutoOrganizeHandler:
|
|||||||
progress_callback=self._progress_callback,
|
progress_callback=self._progress_callback,
|
||||||
exclusion_patterns=exclusion_patterns,
|
exclusion_patterns=exclusion_patterns,
|
||||||
)
|
)
|
||||||
|
_broadcast_models_changed()
|
||||||
return web.json_response(result.to_dict())
|
return web.json_response(result.to_dict())
|
||||||
except AutoOrganizeInProgressError:
|
except AutoOrganizeInProgressError:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2303,8 +2435,8 @@ class ModelUpdateHandler:
|
|||||||
self._logger.error("Failed to fetch license info: %s", exc, exc_info=True)
|
self._logger.error("Failed to fetch license info: %s", exc, exc_info=True)
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
updated: List[Dict[str, str]] = []
|
updated: List[Dict[str, Any]] = []
|
||||||
errors: List[Dict[str, str]] = []
|
errors: List[Dict[str, Any]] = []
|
||||||
for model_id in model_ids:
|
for model_id in model_ids:
|
||||||
license_payload = license_map.get(model_id)
|
license_payload = license_map.get(model_id)
|
||||||
if not license_payload:
|
if not license_payload:
|
||||||
@@ -2317,6 +2449,7 @@ class ModelUpdateHandler:
|
|||||||
model_section = civitai_section.get("model")
|
model_section = civitai_section.get("model")
|
||||||
if not isinstance(model_section, Mapping):
|
if not isinstance(model_section, Mapping):
|
||||||
model_section = {}
|
model_section = {}
|
||||||
|
model_section = dict(model_section)
|
||||||
model_section.update(resolved_payload)
|
model_section.update(resolved_payload)
|
||||||
civitai_section["model"] = model_section
|
civitai_section["model"] = model_section
|
||||||
metadata_payload["civitai"] = civitai_section
|
metadata_payload["civitai"] = civitai_section
|
||||||
@@ -2332,7 +2465,7 @@ class ModelUpdateHandler:
|
|||||||
)
|
)
|
||||||
errors.append({"filePath": metadata_path, "error": str(exc)})
|
errors.append({"filePath": metadata_path, "error": str(exc)})
|
||||||
|
|
||||||
response_payload = {"success": True, "updated": updated}
|
response_payload: Dict[str, Any] = {"success": True, "updated": updated}
|
||||||
missing_model_ids = [mid for mid in model_ids if mid not in license_map]
|
missing_model_ids = [mid for mid in model_ids if mid not in license_map]
|
||||||
if missing_model_ids:
|
if missing_model_ids:
|
||||||
response_payload["missingModelIds"] = missing_model_ids
|
response_payload["missingModelIds"] = missing_model_ids
|
||||||
@@ -2402,6 +2535,7 @@ class ModelUpdateHandler:
|
|||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
hide_early_access = False
|
hide_early_access = False
|
||||||
|
hide_paid = False
|
||||||
if self._settings is not None:
|
if self._settings is not None:
|
||||||
try:
|
try:
|
||||||
hide_early_access = bool(
|
hide_early_access = bool(
|
||||||
@@ -2409,12 +2543,17 @@ class ModelUpdateHandler:
|
|||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
try:
|
||||||
|
hide_paid = bool(self._settings.get("hide_paid_updates", False))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
serialized_records = []
|
serialized_records = []
|
||||||
for record in records.values():
|
for record in records.values():
|
||||||
has_update_fn = getattr(record, "has_update", None)
|
has_update_fn = getattr(record, "has_update", None)
|
||||||
if callable(has_update_fn) and has_update_fn(
|
if callable(has_update_fn) and has_update_fn(
|
||||||
hide_early_access=hide_early_access
|
hide_early_access=hide_early_access,
|
||||||
|
hide_paid=hide_paid,
|
||||||
):
|
):
|
||||||
serialized_records.append(self._serialize_record(record))
|
serialized_records.append(self._serialize_record(record))
|
||||||
|
|
||||||
@@ -2568,10 +2707,16 @@ class ModelUpdateHandler:
|
|||||||
if not record or not record.versions:
|
if not record or not record.versions:
|
||||||
return record
|
return record
|
||||||
|
|
||||||
# Find versions that need enrichment
|
# Find versions that need enrichment. Permanent paid versions are not
|
||||||
|
# early access (mirror _is_early_access_active) and never carry an end
|
||||||
|
# time, so skip them to avoid pointless per-version API calls.
|
||||||
versions_needing_update = []
|
versions_needing_update = []
|
||||||
for version in record.versions:
|
for version in record.versions:
|
||||||
if version.is_early_access and not version.early_access_ends_at:
|
if (
|
||||||
|
version.is_early_access
|
||||||
|
and not version.early_access_ends_at
|
||||||
|
and not getattr(version, "is_paid", False)
|
||||||
|
):
|
||||||
versions_needing_update.append(version)
|
versions_needing_update.append(version)
|
||||||
|
|
||||||
if not versions_needing_update:
|
if not versions_needing_update:
|
||||||
@@ -2681,6 +2826,7 @@ class ModelUpdateHandler:
|
|||||||
civitai_payload = metadata_payload.get("civitai")
|
civitai_payload = metadata_payload.get("civitai")
|
||||||
if not isinstance(civitai_payload, Mapping):
|
if not isinstance(civitai_payload, Mapping):
|
||||||
civitai_payload = {}
|
civitai_payload = {}
|
||||||
|
civitai_payload = dict(civitai_payload)
|
||||||
|
|
||||||
model_payload = civitai_payload.get("model")
|
model_payload = civitai_payload.get("model")
|
||||||
if not isinstance(model_payload, Mapping):
|
if not isinstance(model_payload, Mapping):
|
||||||
@@ -2725,7 +2871,7 @@ class ModelUpdateHandler:
|
|||||||
|
|
||||||
return aggregated
|
return aggregated
|
||||||
|
|
||||||
def _extract_target_model_ids(self, payload: Dict) -> Optional[List[int]]:
|
def _extract_target_model_ids(self, payload: Dict[str, Any]) -> Optional[List[int]]:
|
||||||
if not isinstance(payload, Mapping):
|
if not isinstance(payload, Mapping):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -2753,7 +2899,7 @@ class ModelUpdateHandler:
|
|||||||
return {}
|
return {}
|
||||||
|
|
||||||
to_dict = getattr(metadata, "to_dict", None)
|
to_dict = getattr(metadata, "to_dict", None)
|
||||||
if callable(to_dict):
|
if to_dict:
|
||||||
try:
|
try:
|
||||||
return to_dict()
|
return to_dict()
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -2764,7 +2910,7 @@ class ModelUpdateHandler:
|
|||||||
|
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
async def _read_json(self, request: web.Request) -> Dict:
|
async def _read_json(self, request: web.Request) -> Dict[str, Any]:
|
||||||
if not request.can_read_body:
|
if not request.can_read_body:
|
||||||
return {}
|
return {}
|
||||||
try:
|
try:
|
||||||
@@ -2796,10 +2942,11 @@ class ModelUpdateHandler:
|
|||||||
record,
|
record,
|
||||||
*,
|
*,
|
||||||
version_context: Optional[Dict[int, Dict[str, Any]]] = None,
|
version_context: Optional[Dict[int, Dict[str, Any]]] = None,
|
||||||
) -> Dict:
|
) -> Dict[str, Any]:
|
||||||
context = version_context or {}
|
context = version_context or {}
|
||||||
# Check user setting for hiding early access versions
|
# Check user setting for hiding early access versions
|
||||||
hide_early_access = False
|
hide_early_access = False
|
||||||
|
hide_paid = False
|
||||||
if self._settings is not None:
|
if self._settings is not None:
|
||||||
try:
|
try:
|
||||||
hide_early_access = bool(
|
hide_early_access = bool(
|
||||||
@@ -2807,6 +2954,10 @@ class ModelUpdateHandler:
|
|||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
try:
|
||||||
|
hide_paid = bool(self._settings.get("hide_paid_updates", False))
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
return {
|
return {
|
||||||
"modelType": record.model_type,
|
"modelType": record.model_type,
|
||||||
"modelId": record.model_id,
|
"modelId": record.model_id,
|
||||||
@@ -2815,7 +2966,10 @@ class ModelUpdateHandler:
|
|||||||
"inLibraryVersionIds": record.in_library_version_ids,
|
"inLibraryVersionIds": record.in_library_version_ids,
|
||||||
"lastCheckedAt": record.last_checked_at,
|
"lastCheckedAt": record.last_checked_at,
|
||||||
"shouldIgnore": record.should_ignore_model,
|
"shouldIgnore": record.should_ignore_model,
|
||||||
"hasUpdate": record.has_update(hide_early_access=hide_early_access),
|
"hasUpdate": record.has_update(
|
||||||
|
hide_early_access=hide_early_access,
|
||||||
|
hide_paid=hide_paid,
|
||||||
|
),
|
||||||
"versions": [
|
"versions": [
|
||||||
self._serialize_version(version, context.get(version.version_id))
|
self._serialize_version(version, context.get(version.version_id))
|
||||||
for version in record.versions
|
for version in record.versions
|
||||||
@@ -2825,7 +2979,7 @@ class ModelUpdateHandler:
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def _serialize_version(
|
def _serialize_version(
|
||||||
version, context: Optional[Dict[str, Any]]
|
version, context: Optional[Dict[str, Any]]
|
||||||
) -> Dict:
|
) -> Dict[str, Any]:
|
||||||
context = context or {}
|
context = context or {}
|
||||||
preview_override = context.get("preview_override")
|
preview_override = context.get("preview_override")
|
||||||
preview_url = (
|
preview_url = (
|
||||||
@@ -2834,8 +2988,11 @@ class ModelUpdateHandler:
|
|||||||
|
|
||||||
# Determine if version is currently in early access
|
# Determine if version is currently in early access
|
||||||
# Two-phase detection: use exact end time if available, otherwise fallback to basic flag
|
# Two-phase detection: use exact end time if available, otherwise fallback to basic flag
|
||||||
|
# Mirror _is_early_access_active: permanent paid versions (no end time) are NOT early access
|
||||||
is_early_access = False
|
is_early_access = False
|
||||||
if version.early_access_ends_at:
|
if getattr(version, "is_paid", False) and not version.early_access_ends_at:
|
||||||
|
is_early_access = False
|
||||||
|
elif version.early_access_ends_at:
|
||||||
try:
|
try:
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
@@ -2850,6 +3007,13 @@ class ModelUpdateHandler:
|
|||||||
# Fallback to basic EA flag from bulk API
|
# Fallback to basic EA flag from bulk API
|
||||||
is_early_access = True
|
is_early_access = True
|
||||||
|
|
||||||
|
paid_access_payload = None
|
||||||
|
if getattr(version, "paid_access", None):
|
||||||
|
try:
|
||||||
|
paid_access_payload = json.loads(version.paid_access)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
paid_access_payload = None
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"versionId": version.version_id,
|
"versionId": version.version_id,
|
||||||
"name": version.name,
|
"name": version.name,
|
||||||
@@ -2863,6 +3027,8 @@ class ModelUpdateHandler:
|
|||||||
"earlyAccessEndsAt": version.early_access_ends_at,
|
"earlyAccessEndsAt": version.early_access_ends_at,
|
||||||
"isEarlyAccess": is_early_access,
|
"isEarlyAccess": is_early_access,
|
||||||
"usageControl": version.usage_control,
|
"usageControl": version.usage_control,
|
||||||
|
"isPaid": bool(getattr(version, "is_paid", False)),
|
||||||
|
"paidAccess": paid_access_payload,
|
||||||
"filePath": context.get("file_path"),
|
"filePath": context.get("file_path"),
|
||||||
"fileName": context.get("file_name"),
|
"fileName": context.get("file_name"),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,323 @@
|
|||||||
|
"""Handler for the pending-delete undo endpoint.
|
||||||
|
|
||||||
|
Restores a staged delete batch (models or recipes) via
|
||||||
|
``PendingDeleteService.undo`` and then repairs the affected library caches:
|
||||||
|
the model cache entry is restored from the manifest's ``model_snapshot``
|
||||||
|
(including the version index and hash index), tag counts are re-incremented,
|
||||||
|
and the recipe cache is re-populated via ``RecipeScanner.add_recipe``.
|
||||||
|
|
||||||
|
The per-type scanner is resolved from the manifest's ``model_type`` page value
|
||||||
|
through the SAME ServiceRegistry getters the model route registrars use
|
||||||
|
(lora/checkpoint/embedding) - never a hardcoded lora scanner.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
from typing import Any, Awaitable, Callable, Dict, List, Optional, Set, cast
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
from ...services.pending_delete_service import get_pending_delete_service
|
||||||
|
from .model_handlers import _broadcast_models_changed
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Manifest ``model_type`` page values -> ServiceRegistry scanner getter names.
|
||||||
|
# The model route registrars resolve per-type scanners via these getters
|
||||||
|
# (lora_routes / checkpoint_routes / embedding_routes); undo must do the same
|
||||||
|
# so the CORRECT cache is restored for the deleted model's type.
|
||||||
|
_MODEL_TYPE_GETTER_NAMES: Dict[str, str] = {
|
||||||
|
"loras": "get_lora_scanner",
|
||||||
|
"checkpoints": "get_checkpoint_scanner",
|
||||||
|
"embeddings": "get_embedding_scanner",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Staged batch ids are ``uuid.uuid4().hex`` (32 lowercase hex chars). The id is
|
||||||
|
# joined into filesystem paths by ``_find_batch_dir``, so reject anything that
|
||||||
|
# does not match this exact shape (blocks path-traversal via batch_id).
|
||||||
|
_BATCH_ID_RE = re.compile(r"^[0-9a-f]{32}$")
|
||||||
|
|
||||||
|
|
||||||
|
class PendingDeleteHandler:
|
||||||
|
"""Handle undo requests for staged model/recipe deletions."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
service_factory: Callable[[], Awaitable[Any]] = get_pending_delete_service,
|
||||||
|
scanner_getter: Optional[Callable[[str], Awaitable[Any]]] = None,
|
||||||
|
recipe_scanner_getter: Optional[Callable[[], Awaitable[Any]]] = None,
|
||||||
|
) -> None:
|
||||||
|
self._service_factory: Callable[[], Awaitable[Any]] = service_factory
|
||||||
|
self._scanner_getter: Callable[[str], Awaitable[Any]] = (
|
||||||
|
scanner_getter or self._resolve_scanner
|
||||||
|
)
|
||||||
|
self._recipe_scanner_getter: Callable[[], Awaitable[Any]] = (
|
||||||
|
recipe_scanner_getter or self._resolve_recipe_scanner
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _resolve_scanner(model_type: str) -> Any:
|
||||||
|
"""Resolve the per-type scanner for a manifest ``model_type``.
|
||||||
|
|
||||||
|
The getter is looked up on the ServiceRegistry module namespace at call
|
||||||
|
time so tests (and the registry stubs) can patch it.
|
||||||
|
"""
|
||||||
|
from ...services import service_registry
|
||||||
|
|
||||||
|
getter_name = _MODEL_TYPE_GETTER_NAMES.get(model_type)
|
||||||
|
if getter_name is None:
|
||||||
|
raise ValueError(f"Unknown model type: {model_type}")
|
||||||
|
getter = getattr(service_registry.ServiceRegistry, getter_name, None)
|
||||||
|
if not callable(getter):
|
||||||
|
raise ValueError(f"No scanner getter for model type: {model_type}")
|
||||||
|
scanner = await cast(Callable[[], Awaitable[Any]], getter)()
|
||||||
|
if scanner is None:
|
||||||
|
raise ValueError(f"No scanner registered for model type: {model_type}")
|
||||||
|
return scanner
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _resolve_recipe_scanner() -> Any:
|
||||||
|
"""Resolve the recipe scanner via the ServiceRegistry module namespace."""
|
||||||
|
from ...services import service_registry
|
||||||
|
|
||||||
|
getter = getattr(service_registry.ServiceRegistry, "get_recipe_scanner", None)
|
||||||
|
if not callable(getter):
|
||||||
|
raise ValueError("Recipe scanner getter unavailable")
|
||||||
|
scanner = await cast(Callable[[], Awaitable[Any]], getter)()
|
||||||
|
if scanner is None:
|
||||||
|
raise ValueError("No recipe scanner registered")
|
||||||
|
return scanner
|
||||||
|
|
||||||
|
async def undo_delete(self, request: web.Request) -> web.Response:
|
||||||
|
"""Restore a staged batch and its library cache entry.
|
||||||
|
|
||||||
|
Body: ``{"batch_id": str}``. On success returns
|
||||||
|
``{"success": True, "restored": [<original paths>], "kind": kind}``.
|
||||||
|
Expired/unknown batches and occupied target paths -> 404.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
except Exception:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||||
|
)
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||||
|
)
|
||||||
|
batch_id = data.get("batch_id")
|
||||||
|
if not batch_id or not isinstance(batch_id, str):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "batch_id is required"}, status=400
|
||||||
|
)
|
||||||
|
if not _BATCH_ID_RE.fullmatch(batch_id):
|
||||||
|
# batch_id is joined into a path by _find_batch_dir - restrict to
|
||||||
|
# the exact staged-id shape so traversal payloads get 400.
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid batch_id"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
service = await self._service_factory()
|
||||||
|
try:
|
||||||
|
# Read the manifest BEFORE undo: undo() removes the batch dir.
|
||||||
|
manifest = await self._read_staged_manifest(service, batch_id)
|
||||||
|
result = await service.undo(batch_id)
|
||||||
|
except ValueError as exc:
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=404)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Unexpected error undoing batch %s: %s", batch_id, exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
kind = result.get("kind")
|
||||||
|
try:
|
||||||
|
if kind == "model":
|
||||||
|
if manifest is not None:
|
||||||
|
await self._restore_model_cache(manifest)
|
||||||
|
else:
|
||||||
|
# undo() raises when the manifest is missing, so this only
|
||||||
|
# happens defensively - files are restored regardless.
|
||||||
|
logger.warning(
|
||||||
|
"Manifest missing after undo of %s; skipping cache restore",
|
||||||
|
batch_id,
|
||||||
|
)
|
||||||
|
_broadcast_models_changed()
|
||||||
|
elif kind == "recipe":
|
||||||
|
# Recipe undo is client-refresh only: re-add to the scanner
|
||||||
|
# cache, no models_changed broadcast.
|
||||||
|
if manifest is not None:
|
||||||
|
await self._restore_recipe_cache(result, manifest)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Manifest missing after undo of %s; skipping cache restore",
|
||||||
|
batch_id,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
# Files are already restored; only the cache restoration failed.
|
||||||
|
logger.error(
|
||||||
|
"Cache restoration failed after undo of %s: %s",
|
||||||
|
batch_id,
|
||||||
|
exc,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": True,
|
||||||
|
"restored": result.get("restored", []),
|
||||||
|
"kind": kind,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def _read_staged_manifest(
|
||||||
|
service: Any, batch_id: str
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Locate and read the batch manifest while it still exists on disk."""
|
||||||
|
batch_dir = await service._find_batch_dir(batch_id)
|
||||||
|
if not batch_dir:
|
||||||
|
return None
|
||||||
|
manifest_path = os.path.join(batch_dir, "manifest.json")
|
||||||
|
try:
|
||||||
|
with open(manifest_path, "r", encoding="utf-8") as handle:
|
||||||
|
payload = json.load(handle)
|
||||||
|
except (OSError, json.JSONDecodeError) as exc:
|
||||||
|
logger.debug("Failed to read manifest for batch %s: %s", batch_id, exc)
|
||||||
|
return None
|
||||||
|
return payload if isinstance(payload, dict) else None
|
||||||
|
|
||||||
|
async def _restore_model_cache(self, manifest: Dict[str, Any]) -> None:
|
||||||
|
"""Re-add every deleted model's cache entry from the manifest.
|
||||||
|
|
||||||
|
Each main-file entry carries the deleted model's ``snapshot`` (added at
|
||||||
|
stage time), so a merged bulk manifest holds ALL snapshots - undo must
|
||||||
|
restore every one, not just the top-level winner's. Old-format
|
||||||
|
manifests without entry snapshots fall back to the top-level
|
||||||
|
``model_snapshot`` (backward compat / single-delete path).
|
||||||
|
"""
|
||||||
|
model_type = manifest.get("model_type")
|
||||||
|
if not model_type or not isinstance(model_type, str):
|
||||||
|
raise ValueError(f"Manifest carries no model_type: {manifest.get('batch_id')}")
|
||||||
|
scanner = await self._scanner_getter(model_type)
|
||||||
|
|
||||||
|
# Collect one snapshot per distinct file_path from the entry snapshots.
|
||||||
|
snapshots: List[Dict[str, Any]] = []
|
||||||
|
seen: Set[str] = set()
|
||||||
|
for entry in manifest.get("entries") or []:
|
||||||
|
snapshot = entry.get("snapshot")
|
||||||
|
if not isinstance(snapshot, dict):
|
||||||
|
continue
|
||||||
|
file_path = snapshot.get("file_path")
|
||||||
|
if not file_path or not isinstance(file_path, str):
|
||||||
|
continue
|
||||||
|
if file_path in seen:
|
||||||
|
continue
|
||||||
|
seen.add(file_path)
|
||||||
|
snapshots.append(snapshot)
|
||||||
|
|
||||||
|
if not snapshots:
|
||||||
|
# Backward compat: pre-F3 manifests carry only the top-level
|
||||||
|
# model_snapshot (single-delete path, unchanged behavior).
|
||||||
|
top = manifest.get("model_snapshot")
|
||||||
|
if isinstance(top, dict) and top.get("file_path"):
|
||||||
|
snapshots = [top]
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Manifest %s has no restorable model snapshot; skipping cache restore",
|
||||||
|
manifest.get("batch_id"),
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
if cache is None:
|
||||||
|
logger.warning(
|
||||||
|
"Scanner cache unavailable for %s; skipping cache restore", model_type
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
for snapshot in snapshots:
|
||||||
|
file_path = str(snapshot["file_path"])
|
||||||
|
# A rescan between delete and undo may have re-added a stale entry
|
||||||
|
# for this path - drop it so exactly one (the snapshot) remains.
|
||||||
|
cache.raw_data = [
|
||||||
|
item for item in cache.raw_data if item.get("file_path") != file_path
|
||||||
|
]
|
||||||
|
|
||||||
|
# Restore tag counts (mirror of the bulk-delete decrement in
|
||||||
|
# _batch_update_cache_for_deleted_models: undo re-increments).
|
||||||
|
tags = snapshot.get("tags")
|
||||||
|
if isinstance(tags, list):
|
||||||
|
for tag in tags:
|
||||||
|
if not isinstance(tag, str) or not tag:
|
||||||
|
continue
|
||||||
|
scanner._tags_count[tag] = scanner._tags_count.get(tag, 0) + 1
|
||||||
|
|
||||||
|
cache.raw_data.append(dict(snapshot))
|
||||||
|
|
||||||
|
# Re-register the path in the hash index (add_entry guards a
|
||||||
|
# missing sha256 internally; still guard defensively here).
|
||||||
|
sha256 = snapshot.get("sha256") or ""
|
||||||
|
autov3 = snapshot.get("autov3")
|
||||||
|
hash_index = getattr(scanner, "_hash_index", None)
|
||||||
|
if hash_index is not None and sha256 and file_path:
|
||||||
|
hash_index.add_entry(sha256, file_path, autov3)
|
||||||
|
|
||||||
|
# Follow the bulk-delete cache-update pattern ONCE after all entries,
|
||||||
|
# including the explicit version-index rebuild so the version index
|
||||||
|
# does not go stale.
|
||||||
|
cache.rebuild_version_index()
|
||||||
|
await cache.resort()
|
||||||
|
|
||||||
|
scanner.bump_cache_version()
|
||||||
|
|
||||||
|
persist = getattr(scanner, "_persist_current_cache", None)
|
||||||
|
if callable(persist):
|
||||||
|
result = persist()
|
||||||
|
if inspect.isawaitable(result):
|
||||||
|
await result
|
||||||
|
|
||||||
|
async def _restore_recipe_cache(
|
||||||
|
self, result: Dict[str, Any], manifest: Dict[str, Any]
|
||||||
|
) -> None:
|
||||||
|
"""Re-add a restored recipe via ``RecipeScanner.add_recipe``.
|
||||||
|
|
||||||
|
The recipe JSON embeds the full recipe_data (incl. id/file_path);
|
||||||
|
``add_recipe`` only READS the ``_json_path_map`` so the forced frontend
|
||||||
|
refresh self-heals any transient path-map gap.
|
||||||
|
"""
|
||||||
|
restored = result.get("restored") or []
|
||||||
|
json_path = next(
|
||||||
|
(p for p in restored if isinstance(p, str) and p.endswith(".json")),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if not json_path or not os.path.exists(json_path):
|
||||||
|
# Defensive fallback to the manifest's recipe_snapshot file_path.
|
||||||
|
snapshot = manifest.get("recipe_snapshot") or {}
|
||||||
|
fallback = snapshot.get("file_path")
|
||||||
|
if fallback and os.path.exists(fallback):
|
||||||
|
json_path = fallback
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Restored recipe JSON not found in %s; skipping cache restore",
|
||||||
|
restored,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
with open(json_path, "r", encoding="utf-8") as handle:
|
||||||
|
recipe_data = json.load(handle)
|
||||||
|
except (OSError, json.JSONDecodeError) as exc:
|
||||||
|
logger.warning("Failed to load restored recipe JSON %s: %s", json_path, exc)
|
||||||
|
return
|
||||||
|
if not isinstance(recipe_data, dict):
|
||||||
|
return
|
||||||
|
recipe_scanner = await self._recipe_scanner_getter()
|
||||||
|
await recipe_scanner.add_recipe(recipe_data)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["PendingDeleteHandler"]
|
||||||
@@ -10,7 +10,7 @@ import asyncio
|
|||||||
import tempfile
|
import tempfile
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional
|
from typing import Any, Awaitable, Callable, Dict, List, Mapping, Optional, Tuple
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
@@ -34,6 +34,7 @@ from ...utils.civitai_utils import (
|
|||||||
)
|
)
|
||||||
from ...utils.constants import NSFW_LEVELS
|
from ...utils.constants import NSFW_LEVELS
|
||||||
from ...utils.exif_utils import ExifUtils
|
from ...utils.exif_utils import ExifUtils
|
||||||
|
from ...utils.recipe_open_stats import RecipeOpenStats
|
||||||
from ...recipes.merger import GenParamsMerger
|
from ...recipes.merger import GenParamsMerger
|
||||||
from ...recipes.enrichment import RecipeEnricher
|
from ...recipes.enrichment import RecipeEnricher
|
||||||
from ...services.websocket_manager import ws_manager as default_ws_manager
|
from ...services.websocket_manager import ws_manager as default_ws_manager
|
||||||
@@ -44,6 +45,22 @@ EnsureDependenciesCallable = Callable[[], Awaitable[None]]
|
|||||||
RecipeScannerGetter = Callable[[], Any]
|
RecipeScannerGetter = Callable[[], Any]
|
||||||
CivitaiClientGetter = Callable[[], Any]
|
CivitaiClientGetter = Callable[[], Any]
|
||||||
|
|
||||||
|
# Cap concurrent preview-dimension reads across requests. With a cold LRU
|
||||||
|
# cache one page can touch up to page_size image files; 16 balances SSD and
|
||||||
|
# HDD throughput without starving the event loop.
|
||||||
|
_DIMS_READ_SEMAPHORE = asyncio.Semaphore(16)
|
||||||
|
|
||||||
|
|
||||||
|
async def _read_preview_dims(path: str) -> Optional[Tuple[int, int]]:
|
||||||
|
"""Read preview dimensions off the event loop under the concurrency cap.
|
||||||
|
|
||||||
|
PIL I/O runs in a worker thread so it never blocks the event loop, and the
|
||||||
|
semaphore bounds how many files are opened at once even when many list
|
||||||
|
requests land together.
|
||||||
|
"""
|
||||||
|
async with _DIMS_READ_SEMAPHORE:
|
||||||
|
return await asyncio.to_thread(ExifUtils.get_image_dimensions, path)
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class RecipeHandlerSet:
|
class RecipeHandlerSet:
|
||||||
@@ -82,6 +99,7 @@ class RecipeHandlerSet:
|
|||||||
"download_shared_recipe": self.sharing.download_shared_recipe,
|
"download_shared_recipe": self.sharing.download_shared_recipe,
|
||||||
"get_recipe_syntax": self.query.get_recipe_syntax,
|
"get_recipe_syntax": self.query.get_recipe_syntax,
|
||||||
"update_recipe": self.management.update_recipe,
|
"update_recipe": self.management.update_recipe,
|
||||||
|
"record_recipe_open": self.management.record_recipe_open,
|
||||||
"reconnect_lora": self.management.reconnect_lora,
|
"reconnect_lora": self.management.reconnect_lora,
|
||||||
"find_duplicates": self.query.find_duplicates,
|
"find_duplicates": self.query.find_duplicates,
|
||||||
"move_recipes_bulk": self.management.move_recipes_bulk,
|
"move_recipes_bulk": self.management.move_recipes_bulk,
|
||||||
@@ -96,6 +114,11 @@ class RecipeHandlerSet:
|
|||||||
"repair_recipe": self.management.repair_recipe,
|
"repair_recipe": self.management.repair_recipe,
|
||||||
"repair_recipes_bulk": self.management.repair_recipes_bulk,
|
"repair_recipes_bulk": self.management.repair_recipes_bulk,
|
||||||
"get_repair_progress": self.management.get_repair_progress,
|
"get_repair_progress": self.management.get_repair_progress,
|
||||||
|
"rematch_recipes": self.management.rematch_recipes,
|
||||||
|
"cancel_rematch": self.management.cancel_rematch,
|
||||||
|
"rematch_recipe": self.management.rematch_recipe,
|
||||||
|
"rematch_recipes_bulk": self.management.rematch_recipes_bulk,
|
||||||
|
"get_rematch_progress": self.management.get_rematch_progress,
|
||||||
"start_batch_import": self.batch_import.start_batch_import,
|
"start_batch_import": self.batch_import.start_batch_import,
|
||||||
"get_batch_import_progress": self.batch_import.get_batch_import_progress,
|
"get_batch_import_progress": self.batch_import.get_batch_import_progress,
|
||||||
"cancel_batch_import": self.batch_import.cancel_batch_import,
|
"cancel_batch_import": self.batch_import.cancel_batch_import,
|
||||||
@@ -246,7 +269,8 @@ class RecipeListingHandler:
|
|||||||
recursive=recursive,
|
recursive=recursive,
|
||||||
)
|
)
|
||||||
|
|
||||||
for item in result.get("items", []):
|
items = result.get("items", [])
|
||||||
|
for item in items:
|
||||||
file_path = item.get("file_path")
|
file_path = item.get("file_path")
|
||||||
if file_path:
|
if file_path:
|
||||||
item["file_url"] = self.format_recipe_file_url(file_path)
|
item["file_url"] = self.format_recipe_file_url(file_path)
|
||||||
@@ -255,6 +279,26 @@ class RecipeListingHandler:
|
|||||||
item.setdefault("loras", [])
|
item.setdefault("loras", [])
|
||||||
item.setdefault("base_model", "")
|
item.setdefault("base_model", "")
|
||||||
|
|
||||||
|
# Batch preview dimension reads with asyncio.gather. The previous
|
||||||
|
# loop awaited asyncio.to_thread once per item, so a page_size=100
|
||||||
|
# request submitted 100 sequential thread calls (50-300ms cold-page
|
||||||
|
# latency). gather runs them concurrently while the semaphore caps
|
||||||
|
# disk opens; dimensions stay omitted (not null) when a preview has
|
||||||
|
# no readable size (video, missing file).
|
||||||
|
to_read = [
|
||||||
|
(i, item.get("file_path"))
|
||||||
|
for i, item in enumerate(items)
|
||||||
|
if item.get("file_path")
|
||||||
|
]
|
||||||
|
if to_read:
|
||||||
|
dims_list = await asyncio.gather(
|
||||||
|
*(_read_preview_dims(path) for _, path in to_read)
|
||||||
|
)
|
||||||
|
for (idx, _), dims in zip(to_read, dims_list):
|
||||||
|
if dims:
|
||||||
|
item = items[idx]
|
||||||
|
item["width"], item["height"] = dims
|
||||||
|
|
||||||
return web.json_response(result)
|
return web.json_response(result)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self._logger.error("Error retrieving recipes: %s", exc, exc_info=True)
|
self._logger.error("Error retrieving recipes: %s", exc, exc_info=True)
|
||||||
@@ -538,7 +582,12 @@ class RecipeQueryHandler:
|
|||||||
if recipe_scanner is None:
|
if recipe_scanner is None:
|
||||||
raise RuntimeError("Recipe scanner unavailable")
|
raise RuntimeError("Recipe scanner unavailable")
|
||||||
|
|
||||||
fingerprint_groups = await recipe_scanner.find_all_duplicate_recipes()
|
include_prompt = (
|
||||||
|
request.query.get("include_prompt", "false").lower() in ("1", "true")
|
||||||
|
)
|
||||||
|
fingerprint_groups = await recipe_scanner.find_all_duplicate_recipes(
|
||||||
|
include_prompt=include_prompt
|
||||||
|
)
|
||||||
url_groups = await recipe_scanner.find_duplicate_recipes_by_source()
|
url_groups = await recipe_scanner.find_duplicate_recipes_by_source()
|
||||||
response_data = []
|
response_data = []
|
||||||
|
|
||||||
@@ -571,6 +620,7 @@ class RecipeQueryHandler:
|
|||||||
response_data.append(
|
response_data.append(
|
||||||
{
|
{
|
||||||
"type": "fingerprint",
|
"type": "fingerprint",
|
||||||
|
"key": f"g-{len(response_data) + 1}",
|
||||||
"fingerprint": fingerprint,
|
"fingerprint": fingerprint,
|
||||||
"count": len(recipes),
|
"count": len(recipes),
|
||||||
"recipes": recipes,
|
"recipes": recipes,
|
||||||
@@ -606,6 +656,7 @@ class RecipeQueryHandler:
|
|||||||
response_data.append(
|
response_data.append(
|
||||||
{
|
{
|
||||||
"type": "source_path",
|
"type": "source_path",
|
||||||
|
"key": f"g-{len(response_data) + 1}",
|
||||||
"fingerprint": url,
|
"fingerprint": url,
|
||||||
"count": len(recipes),
|
"count": len(recipes),
|
||||||
"recipes": recipes,
|
"recipes": recipes,
|
||||||
@@ -850,6 +901,159 @@ class RecipeManagementHandler:
|
|||||||
self._logger.error("Error repairing single recipe: %s", exc, exc_info=True)
|
self._logger.error("Error repairing single recipe: %s", exc, exc_info=True)
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def rematch_recipes(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
await self._ensure_dependencies_ready()
|
||||||
|
recipe_scanner = self._recipe_scanner_getter()
|
||||||
|
if recipe_scanner is None:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Recipe scanner unavailable"},
|
||||||
|
status=503,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Mutual exclusion: a global rematch cannot start while a rematch
|
||||||
|
# OR a repair is already running — both mutate recipes under the
|
||||||
|
# same mutation lock.
|
||||||
|
if (
|
||||||
|
self._ws_manager.is_recipe_rematch_running()
|
||||||
|
or self._ws_manager.is_recipe_repair_running()
|
||||||
|
):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Recipe rematch already in progress"},
|
||||||
|
status=409,
|
||||||
|
)
|
||||||
|
|
||||||
|
recipe_scanner.reset_cancellation()
|
||||||
|
|
||||||
|
async def progress_callback(data):
|
||||||
|
await self._ws_manager.broadcast_recipe_rematch_progress(data)
|
||||||
|
|
||||||
|
# Run in background to avoid timeout
|
||||||
|
async def run_rematch():
|
||||||
|
try:
|
||||||
|
await recipe_scanner.rematch_all_recipes(
|
||||||
|
progress_callback=progress_callback
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
self._logger.error(
|
||||||
|
f"Error in recipe rematch task: {e}", exc_info=True
|
||||||
|
)
|
||||||
|
await self._ws_manager.broadcast_recipe_rematch_progress(
|
||||||
|
{"status": "error", "error": str(e)}
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
# Keep the final status for a while so the UI can see it
|
||||||
|
await asyncio.sleep(5)
|
||||||
|
self._ws_manager.cleanup_recipe_rematch_progress()
|
||||||
|
|
||||||
|
asyncio.create_task(run_rematch())
|
||||||
|
|
||||||
|
return web.json_response(
|
||||||
|
{"success": True, "message": "Recipe rematch started"}
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error starting recipe rematch: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def cancel_rematch(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
await self._ensure_dependencies_ready()
|
||||||
|
recipe_scanner = self._recipe_scanner_getter()
|
||||||
|
if recipe_scanner is None:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Recipe scanner unavailable"},
|
||||||
|
status=503,
|
||||||
|
)
|
||||||
|
|
||||||
|
recipe_scanner.cancel_task()
|
||||||
|
return web.json_response(
|
||||||
|
{"success": True, "message": "Cancellation requested"}
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error cancelling recipe rematch: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def rematch_recipes_bulk(self, request: web.Request) -> web.Response:
|
||||||
|
"""Rematch deleted resources for multiple recipes by their IDs.
|
||||||
|
|
||||||
|
Accepts a JSON body with a "recipe_ids" array. The per-recipe loop is
|
||||||
|
delegated to the scanner's rematch_recipes_bulk; this handler only
|
||||||
|
parses the request and returns the scanner's summary.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
await self._ensure_dependencies_ready()
|
||||||
|
recipe_scanner = self._recipe_scanner_getter()
|
||||||
|
if recipe_scanner is None:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Recipe scanner unavailable"},
|
||||||
|
status=503,
|
||||||
|
)
|
||||||
|
|
||||||
|
# A bulk rematch must not queue behind a running global rematch's
|
||||||
|
# mutation lock.
|
||||||
|
if self._ws_manager.is_recipe_rematch_running():
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Recipe rematch already in progress"},
|
||||||
|
status=409,
|
||||||
|
)
|
||||||
|
|
||||||
|
data = await request.json()
|
||||||
|
recipe_ids = data.get("recipe_ids", [])
|
||||||
|
if not recipe_ids:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "recipe_ids are required"},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await recipe_scanner.rematch_recipes_bulk(recipe_ids)
|
||||||
|
return web.json_response(result)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error(
|
||||||
|
"Error performing bulk rematch: %s", exc, exc_info=True
|
||||||
|
)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": str(exc)}, status=500
|
||||||
|
)
|
||||||
|
|
||||||
|
async def rematch_recipe(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
await self._ensure_dependencies_ready()
|
||||||
|
recipe_scanner = self._recipe_scanner_getter()
|
||||||
|
if recipe_scanner is None:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Recipe scanner unavailable"},
|
||||||
|
status=503,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Reject per-recipe rematches while a global run is in progress so
|
||||||
|
# they do not queue behind the mutation lock.
|
||||||
|
if self._ws_manager.is_recipe_rematch_running():
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Recipe rematch already in progress"},
|
||||||
|
status=409,
|
||||||
|
)
|
||||||
|
|
||||||
|
recipe_id = request.match_info["recipe_id"]
|
||||||
|
result = await recipe_scanner.rematch_recipe_by_id(recipe_id)
|
||||||
|
return web.json_response(result)
|
||||||
|
except RecipeNotFoundError as exc:
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=404)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error rematching single recipe: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def get_rematch_progress(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
progress = self._ws_manager.get_recipe_rematch_progress()
|
||||||
|
if progress:
|
||||||
|
return web.json_response({"success": True, "progress": progress})
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "message": "No rematch in progress"}, status=404
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error getting rematch progress: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
async def reimport_recipe(self, request: web.Request) -> web.Response:
|
async def reimport_recipe(self, request: web.Request) -> web.Response:
|
||||||
"""Delete a recipe and re-import it from its source URL.
|
"""Delete a recipe and re-import it from its source URL.
|
||||||
|
|
||||||
@@ -1045,10 +1249,10 @@ class RecipeManagementHandler:
|
|||||||
*,
|
*,
|
||||||
image_url: str,
|
image_url: str,
|
||||||
name: str,
|
name: str,
|
||||||
lora_entries: list,
|
lora_entries: list[Any],
|
||||||
checkpoint_entry: dict,
|
checkpoint_entry: Dict[str, Any] | None,
|
||||||
gen_params_request: dict,
|
gen_params_request: Dict[str, Any] | None,
|
||||||
tags: list,
|
tags: list[Any],
|
||||||
base_model: str,
|
base_model: str,
|
||||||
source_path: str,
|
source_path: str,
|
||||||
) -> web.Response:
|
) -> web.Response:
|
||||||
@@ -1081,6 +1285,12 @@ class RecipeManagementHandler:
|
|||||||
_original_image_url,
|
_original_image_url,
|
||||||
) = await self._download_remote_media(image_url)
|
) = await self._download_remote_media(image_url)
|
||||||
|
|
||||||
|
# Build a version-cached map of local model hashes to cache items so
|
||||||
|
# CivitaiApiMetadataParser can skip CivitAI API calls for models that
|
||||||
|
# exist on disk. Built once and shared by every parse pass below.
|
||||||
|
local_cache = await recipe_scanner.build_local_hash_cache()
|
||||||
|
from ...recipes.parsers.civitai_image import CivitaiApiMetadataParser
|
||||||
|
|
||||||
# Extract embedded EXIF metadata (offloaded to thread pool in this call)
|
# Extract embedded EXIF metadata (offloaded to thread pool in this call)
|
||||||
embedded_gen_params = {}
|
embedded_gen_params = {}
|
||||||
parsed_embedded = None
|
parsed_embedded = None
|
||||||
@@ -1102,9 +1312,16 @@ class RecipeManagementHandler:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if parser:
|
if parser:
|
||||||
parsed_embedded = await parser.parse_metadata(
|
if isinstance(parser, CivitaiApiMetadataParser):
|
||||||
raw_embedded, recipe_scanner=recipe_scanner
|
parsed_embedded = await parser.parse_metadata(
|
||||||
)
|
raw_embedded,
|
||||||
|
recipe_scanner=recipe_scanner,
|
||||||
|
local_cache=local_cache,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
parsed_embedded = await parser.parse_metadata(
|
||||||
|
raw_embedded, recipe_scanner=recipe_scanner
|
||||||
|
)
|
||||||
if parsed_embedded and "gen_params" in parsed_embedded:
|
if parsed_embedded and "gen_params" in parsed_embedded:
|
||||||
embedded_gen_params = parsed_embedded["gen_params"]
|
embedded_gen_params = parsed_embedded["gen_params"]
|
||||||
else:
|
else:
|
||||||
@@ -1135,9 +1352,16 @@ class RecipeManagementHandler:
|
|||||||
civitai_inner_meta
|
civitai_inner_meta
|
||||||
)
|
)
|
||||||
if parser:
|
if parser:
|
||||||
civitai_parsed = await parser.parse_metadata(
|
if isinstance(parser, CivitaiApiMetadataParser):
|
||||||
civitai_inner_meta, recipe_scanner=recipe_scanner
|
civitai_parsed = await parser.parse_metadata(
|
||||||
)
|
civitai_inner_meta,
|
||||||
|
recipe_scanner=recipe_scanner,
|
||||||
|
local_cache=local_cache,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
civitai_parsed = await parser.parse_metadata(
|
||||||
|
civitai_inner_meta, recipe_scanner=recipe_scanner
|
||||||
|
)
|
||||||
if civitai_parsed and "gen_params" in civitai_parsed:
|
if civitai_parsed and "gen_params" in civitai_parsed:
|
||||||
# Merge: API gen_params override EXIF at field level,
|
# Merge: API gen_params override EXIF at field level,
|
||||||
# EXIF fills in fields the API doesn't have.
|
# EXIF fills in fields the API doesn't have.
|
||||||
@@ -1236,6 +1460,33 @@ class RecipeManagementHandler:
|
|||||||
self._logger.error("Error updating recipe: %s", exc, exc_info=True)
|
self._logger.error("Error updating recipe: %s", exc, exc_info=True)
|
||||||
return web.json_response({"error": str(exc)}, status=500)
|
return web.json_response({"error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def record_recipe_open(self, request: web.Request) -> web.Response:
|
||||||
|
"""Record that a recipe's detail modal was opened.
|
||||||
|
|
||||||
|
Lightweight fire-and-forget endpoint backing the "Recently Opened"
|
||||||
|
sort. It only writes the timestamp into the separate open-stats file
|
||||||
|
— recipe JSON and EXIF are never touched.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
await self._ensure_dependencies_ready()
|
||||||
|
recipe_scanner = self._recipe_scanner_getter()
|
||||||
|
if recipe_scanner is None:
|
||||||
|
raise RuntimeError("Recipe scanner unavailable")
|
||||||
|
|
||||||
|
recipe_id = request.match_info["recipe_id"]
|
||||||
|
# Skip recording opens for recipes the scanner no longer knows.
|
||||||
|
recipe_json_path = await recipe_scanner.get_recipe_json_path(recipe_id)
|
||||||
|
if not recipe_json_path:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Recipe not found"}, status=404
|
||||||
|
)
|
||||||
|
|
||||||
|
RecipeOpenStats().record_open(recipe_id)
|
||||||
|
return web.json_response({"success": True})
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error recording recipe open: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
async def move_recipe(self, request: web.Request) -> web.Response:
|
async def move_recipe(self, request: web.Request) -> web.Response:
|
||||||
try:
|
try:
|
||||||
await self._ensure_dependencies_ready()
|
await self._ensure_dependencies_ready()
|
||||||
@@ -1641,7 +1892,7 @@ class RecipeManagementHandler:
|
|||||||
if not provider:
|
if not provider:
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
version_info = await provider.get_model_version_info(version_id)
|
version_info = await provider.get_model_version_info(str(version_id))
|
||||||
if isinstance(version_info, tuple):
|
if isinstance(version_info, tuple):
|
||||||
version_info = version_info[0]
|
version_info = version_info[0]
|
||||||
|
|
||||||
@@ -1761,6 +2012,12 @@ class RecipeManagementHandler:
|
|||||||
await self._download_remote_media(image_url)
|
await self._download_remote_media(image_url)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Build a version-cached map of local model hashes to cache items so
|
||||||
|
# CivitaiApiMetadataParser can skip CivitAI API calls for models that
|
||||||
|
# exist on disk. Built once and shared by every parse pass below.
|
||||||
|
local_cache = await recipe_scanner.build_local_hash_cache()
|
||||||
|
from ...recipes.parsers.civitai_image import CivitaiApiMetadataParser
|
||||||
|
|
||||||
# Extract embedded EXIF metadata
|
# Extract embedded EXIF metadata
|
||||||
embedded_gen_params = {}
|
embedded_gen_params = {}
|
||||||
parsed_embedded = None
|
parsed_embedded = None
|
||||||
@@ -1782,9 +2039,16 @@ class RecipeManagementHandler:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if parser:
|
if parser:
|
||||||
parsed_embedded = await parser.parse_metadata(
|
if isinstance(parser, CivitaiApiMetadataParser):
|
||||||
raw_embedded, recipe_scanner=recipe_scanner
|
parsed_embedded = await parser.parse_metadata(
|
||||||
)
|
raw_embedded,
|
||||||
|
recipe_scanner=recipe_scanner,
|
||||||
|
local_cache=local_cache,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
parsed_embedded = await parser.parse_metadata(
|
||||||
|
raw_embedded, recipe_scanner=recipe_scanner
|
||||||
|
)
|
||||||
if parsed_embedded and "gen_params" in parsed_embedded:
|
if parsed_embedded and "gen_params" in parsed_embedded:
|
||||||
embedded_gen_params = parsed_embedded["gen_params"]
|
embedded_gen_params = parsed_embedded["gen_params"]
|
||||||
finally:
|
finally:
|
||||||
@@ -1822,9 +2086,16 @@ class RecipeManagementHandler:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
if parser:
|
if parser:
|
||||||
parsed_embedded = await parser.parse_metadata(
|
if isinstance(parser, CivitaiApiMetadataParser):
|
||||||
raw_orig, recipe_scanner=recipe_scanner
|
parsed_embedded = await parser.parse_metadata(
|
||||||
)
|
raw_orig,
|
||||||
|
recipe_scanner=recipe_scanner,
|
||||||
|
local_cache=local_cache,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
parsed_embedded = await parser.parse_metadata(
|
||||||
|
raw_orig, recipe_scanner=recipe_scanner
|
||||||
|
)
|
||||||
if (
|
if (
|
||||||
parsed_embedded
|
parsed_embedded
|
||||||
and "gen_params" in parsed_embedded
|
and "gen_params" in parsed_embedded
|
||||||
@@ -1858,9 +2129,16 @@ class RecipeManagementHandler:
|
|||||||
civitai_inner_meta
|
civitai_inner_meta
|
||||||
)
|
)
|
||||||
if parser:
|
if parser:
|
||||||
civitai_parsed = await parser.parse_metadata(
|
if isinstance(parser, CivitaiApiMetadataParser):
|
||||||
civitai_inner_meta, recipe_scanner=recipe_scanner
|
civitai_parsed = await parser.parse_metadata(
|
||||||
)
|
civitai_inner_meta,
|
||||||
|
recipe_scanner=recipe_scanner,
|
||||||
|
local_cache=local_cache,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
civitai_parsed = await parser.parse_metadata(
|
||||||
|
civitai_inner_meta, recipe_scanner=recipe_scanner
|
||||||
|
)
|
||||||
if civitai_parsed and "gen_params" in civitai_parsed:
|
if civitai_parsed and "gen_params" in civitai_parsed:
|
||||||
# Merge: API gen_params override EXIF at field level,
|
# Merge: API gen_params override EXIF at field level,
|
||||||
# EXIF fills in fields the API doesn't have.
|
# EXIF fills in fields the API doesn't have.
|
||||||
@@ -2072,33 +2350,44 @@ class RecipeManagementHandler:
|
|||||||
parsed_input = {**image_data, **inner_meta}
|
parsed_input = {**image_data, **inner_meta}
|
||||||
parsed_input.pop("meta", None)
|
parsed_input.pop("meta", None)
|
||||||
|
|
||||||
# Build a local cache of {hash → cache_item} so the parser can
|
# Build the shared local hash cache so the parser can skip CivitAI
|
||||||
# skip CivitAI API calls for models that exist on disk.
|
# API calls for models that exist on disk.
|
||||||
local_cache: Dict[str, Dict[str, Any]] = {}
|
local_cache: Dict[str, Dict[str, Any]] = (
|
||||||
lora_scanner = getattr(recipe_scanner, "_lora_scanner", None)
|
await recipe_scanner.build_local_hash_cache()
|
||||||
if lora_scanner and model_hash:
|
)
|
||||||
try:
|
|
||||||
parent_cache_data = await lora_scanner.get_cached_data()
|
# Bounded supplement for un-backfilled parents. The shared builder
|
||||||
for item in getattr(parent_cache_data, "raw_data", []):
|
# never computes autov3; when the parent model exists on disk but
|
||||||
if item.get("sha256", "").lower() == model_hash.lower():
|
# its cached entry has no stored AutoV3, compute it for that single
|
||||||
local_cache[model_hash.lower()] = item
|
# file and register the AutoV3 key so the parser can also match on
|
||||||
# Compute AutoV3 so the parser can also match on
|
# that hash type (CivitAI metadata resources use AutoV3). This runs
|
||||||
# that hash type (CivitAI metadata resources use
|
# whenever the parent is found with an empty autov3, independent of
|
||||||
# AutoV3).
|
# whether the sha256 key is already present in the shared cache.
|
||||||
file_path = item.get("file_path")
|
if model_hash:
|
||||||
if file_path and os.path.exists(file_path):
|
lora_scanner = getattr(recipe_scanner, "_lora_scanner", None)
|
||||||
try:
|
if lora_scanner:
|
||||||
from ...utils.file_utils import (
|
try:
|
||||||
calculate_autov3,
|
parent_cache_data = await lora_scanner.get_cached_data()
|
||||||
)
|
for item in getattr(parent_cache_data, "raw_data", []):
|
||||||
autov3 = calculate_autov3(file_path)
|
if item.get("sha256", "").lower() == model_hash.lower():
|
||||||
if autov3:
|
autov3 = (item.get("autov3") or "").lower()
|
||||||
local_cache[autov3.lower()] = item
|
if not autov3:
|
||||||
except Exception:
|
file_path = item.get("file_path")
|
||||||
pass
|
if file_path and os.path.exists(file_path):
|
||||||
break
|
try:
|
||||||
except Exception:
|
from ...utils.file_utils import (
|
||||||
pass
|
calculate_autov3,
|
||||||
|
)
|
||||||
|
autov3 = (
|
||||||
|
calculate_autov3(file_path) or ""
|
||||||
|
).lower()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
if autov3:
|
||||||
|
local_cache[autov3] = item
|
||||||
|
break
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
parser = self._analysis_service._recipe_parser_factory.create_parser(
|
parser = self._analysis_service._recipe_parser_factory.create_parser(
|
||||||
parsed_input
|
parsed_input
|
||||||
@@ -2130,10 +2419,10 @@ class RecipeManagementHandler:
|
|||||||
parent_model_id: int | None = None
|
parent_model_id: int | None = None
|
||||||
parent_version_name: str | None = None
|
parent_version_name: str | None = None
|
||||||
parent_model_name: str | None = None
|
parent_model_name: str | None = None
|
||||||
# Prefer sha256 key; fall back to any cached entry.
|
# Resolve the parent strictly by its sha256 key. There is no
|
||||||
|
# arbitrary fallback: with a full-library cache, picking any entry
|
||||||
|
# would corrupt the isDeleted reconciliation below.
|
||||||
parent_item = local_cache.get(model_hash.lower()) if model_hash else None
|
parent_item = local_cache.get(model_hash.lower()) if model_hash else None
|
||||||
if parent_item is None and local_cache:
|
|
||||||
parent_item = next(iter(local_cache.values()))
|
|
||||||
if parent_item:
|
if parent_item:
|
||||||
civ = parent_item.get("civitai") or {}
|
civ = parent_item.get("civitai") or {}
|
||||||
if isinstance(civ, dict):
|
if isinstance(civ, dict):
|
||||||
@@ -2349,7 +2638,7 @@ class RecipeAnalysisHandler:
|
|||||||
content_type = request.headers.get("Content-Type", "")
|
content_type = request.headers.get("Content-Type", "")
|
||||||
if "multipart/form-data" in content_type:
|
if "multipart/form-data" in content_type:
|
||||||
reader = await request.multipart()
|
reader = await request.multipart()
|
||||||
field = await reader.next()
|
field: Any = await reader.next()
|
||||||
if field is None or field.name != "image":
|
if field is None or field.name != "image":
|
||||||
raise RecipeValidationError("No image field found")
|
raise RecipeValidationError("No image field found")
|
||||||
image_chunks = bytearray()
|
image_chunks = bytearray()
|
||||||
|
|||||||
@@ -1,8 +1,8 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
from typing import Dict
|
from typing import Any, Dict
|
||||||
from server import PromptServer # type: ignore
|
from server import PromptServer # pyright: ignore[reportMissingImports]
|
||||||
|
|
||||||
from .base_model_routes import BaseModelRoutes
|
from .base_model_routes import BaseModelRoutes
|
||||||
from .model_route_registrar import ModelRouteRegistrar
|
from .model_route_registrar import ModelRouteRegistrar
|
||||||
@@ -31,13 +31,13 @@ class LoraRoutes(BaseModelRoutes):
|
|||||||
# Attach service dependencies
|
# Attach service dependencies
|
||||||
self.attach_service(self.service)
|
self.attach_service(self.service)
|
||||||
|
|
||||||
def setup_routes(self, app: web.Application):
|
def setup_routes(self, app: web.Application, prefix: str = "loras"):
|
||||||
"""Setup LoRA routes"""
|
"""Setup LoRA routes"""
|
||||||
# Schedule service initialization on app startup
|
# Schedule service initialization on app startup
|
||||||
app.on_startup.append(lambda _: self.initialize_services())
|
app.on_startup.append(lambda _: self.initialize_services())
|
||||||
|
|
||||||
# Setup common routes with 'loras' prefix (includes page route)
|
# Setup common routes with 'loras' prefix (includes page route)
|
||||||
super().setup_routes(app, "loras")
|
super().setup_routes(app, prefix)
|
||||||
|
|
||||||
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
||||||
"""Setup LoRA-specific routes"""
|
"""Setup LoRA-specific routes"""
|
||||||
@@ -73,7 +73,7 @@ class LoraRoutes(BaseModelRoutes):
|
|||||||
"POST", "/api/lm/{prefix}/get_trigger_words", prefix, self.get_trigger_words
|
"POST", "/api/lm/{prefix}/get_trigger_words", prefix, self.get_trigger_words
|
||||||
)
|
)
|
||||||
|
|
||||||
def _parse_specific_params(self, request: web.Request) -> Dict:
|
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
|
||||||
"""Parse LoRA-specific parameters"""
|
"""Parse LoRA-specific parameters"""
|
||||||
params = {}
|
params = {}
|
||||||
|
|
||||||
@@ -119,25 +119,6 @@ class LoraRoutes(BaseModelRoutes):
|
|||||||
logger.error(f"Error getting letter counts: {e}")
|
logger.error(f"Error getting letter counts: {e}")
|
||||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||||
|
|
||||||
async def get_lora_notes(self, request: web.Request) -> web.Response:
|
|
||||||
"""Get notes for a specific LoRA file"""
|
|
||||||
try:
|
|
||||||
lora_name = request.query.get("name")
|
|
||||||
if not lora_name:
|
|
||||||
return web.Response(text="Lora file name is required", status=400)
|
|
||||||
|
|
||||||
notes = await self.service.get_lora_notes(lora_name)
|
|
||||||
if notes is not None:
|
|
||||||
return web.json_response({"success": True, "notes": notes})
|
|
||||||
else:
|
|
||||||
return web.json_response(
|
|
||||||
{"success": False, "error": "LoRA not found in cache"}, status=404
|
|
||||||
)
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error getting lora notes: {e}", exc_info=True)
|
|
||||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
|
||||||
|
|
||||||
async def get_lora_trigger_words(self, request: web.Request) -> web.Response:
|
async def get_lora_trigger_words(self, request: web.Request) -> web.Response:
|
||||||
"""Get trigger words for a specific LoRA file"""
|
"""Get trigger words for a specific LoRA file"""
|
||||||
try:
|
try:
|
||||||
@@ -168,52 +149,6 @@ class LoraRoutes(BaseModelRoutes):
|
|||||||
logger.error(f"Error getting lora usage tips by path: {e}", exc_info=True)
|
logger.error(f"Error getting lora usage tips by path: {e}", exc_info=True)
|
||||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
return web.json_response({"success": False, "error": str(e)}, status=500)
|
||||||
|
|
||||||
async def get_lora_preview_url(self, request: web.Request) -> web.Response:
|
|
||||||
"""Get the static preview URL for a LoRA file"""
|
|
||||||
try:
|
|
||||||
lora_name = request.query.get("name")
|
|
||||||
if not lora_name:
|
|
||||||
return web.Response(text="Lora file name is required", status=400)
|
|
||||||
|
|
||||||
preview_url = await self.service.get_lora_preview_url(lora_name)
|
|
||||||
if preview_url:
|
|
||||||
return web.json_response({"success": True, "preview_url": preview_url})
|
|
||||||
else:
|
|
||||||
return web.json_response(
|
|
||||||
{
|
|
||||||
"success": False,
|
|
||||||
"error": "No preview URL found for the specified lora",
|
|
||||||
},
|
|
||||||
status=404,
|
|
||||||
)
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error getting lora preview URL: {e}", exc_info=True)
|
|
||||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
|
||||||
|
|
||||||
async def get_lora_civitai_url(self, request: web.Request) -> web.Response:
|
|
||||||
"""Get the Civitai URL for a LoRA file"""
|
|
||||||
try:
|
|
||||||
lora_name = request.query.get("name")
|
|
||||||
if not lora_name:
|
|
||||||
return web.Response(text="Lora file name is required", status=400)
|
|
||||||
|
|
||||||
result = await self.service.get_lora_civitai_url(lora_name)
|
|
||||||
if result["civitai_url"]:
|
|
||||||
return web.json_response({"success": True, **result})
|
|
||||||
else:
|
|
||||||
return web.json_response(
|
|
||||||
{
|
|
||||||
"success": False,
|
|
||||||
"error": "No Civitai data found for the specified lora",
|
|
||||||
},
|
|
||||||
status=404,
|
|
||||||
)
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Error getting lora Civitai URL: {e}", exc_info=True)
|
|
||||||
return web.json_response({"success": False, "error": str(e)}, status=500)
|
|
||||||
|
|
||||||
async def get_random_loras(self, request: web.Request) -> web.Response:
|
async def get_random_loras(self, request: web.Request) -> web.Response:
|
||||||
"""Get random LoRAs based on filters and strength ranges"""
|
"""Get random LoRAs based on filters and strength ranges"""
|
||||||
try:
|
try:
|
||||||
@@ -337,7 +272,7 @@ class LoraRoutes(BaseModelRoutes):
|
|||||||
graph_identifier = entry.get("graph_id")
|
graph_identifier = entry.get("graph_id")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
parsed_node_id = int(node_identifier)
|
parsed_node_id = int(node_identifier) # pyright: ignore[reportArgumentType]
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
parsed_node_id = node_identifier
|
parsed_node_id = node_identifier
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ miscellaneous endpoints share a consistent registration flow.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Callable, Iterable, Mapping
|
from typing import Any, Callable, Iterable, Mapping
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
@@ -147,7 +147,7 @@ class MiscRouteRegistrar:
|
|||||||
handler_lookup[definition.handler_name],
|
handler_lookup[definition.handler_name],
|
||||||
)
|
)
|
||||||
|
|
||||||
def _bind(self, method: str, path: str, handler: Callable) -> None:
|
def _bind(self, method: str, path: str, handler: Callable[..., Any]) -> None:
|
||||||
add_method_name = self._METHOD_MAP[method.upper()]
|
add_method_name = self._METHOD_MAP[method.upper()]
|
||||||
add_method = getattr(self._app.router, add_method_name)
|
add_method = getattr(self._app.router, add_method_name)
|
||||||
add_method(path, handler)
|
add_method(path, handler)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import os
|
|||||||
from typing import Awaitable, Callable, Mapping
|
from typing import Awaitable, Callable, Mapping
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
from server import PromptServer # type: ignore
|
from server import PromptServer # pyright: ignore[reportMissingImports]
|
||||||
|
|
||||||
from ..services.metadata_service import (
|
from ..services.metadata_service import (
|
||||||
get_metadata_archive_manager,
|
get_metadata_archive_manager,
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Callable, Iterable, Mapping
|
from typing import Any, Callable, Iterable, Mapping
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
@@ -174,15 +174,15 @@ class ModelRouteRegistrar:
|
|||||||
handler_lookup[definition.handler_name],
|
handler_lookup[definition.handler_name],
|
||||||
)
|
)
|
||||||
|
|
||||||
def add_route(self, method: str, path: str, handler: Callable) -> None:
|
def add_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
|
||||||
self._bind_route(method, path, handler)
|
self._bind_route(method, path, handler)
|
||||||
|
|
||||||
def add_prefixed_route(
|
def add_prefixed_route(
|
||||||
self, method: str, path_template: str, prefix: str, handler: Callable
|
self, method: str, path_template: str, prefix: str, handler: Callable[..., Any]
|
||||||
) -> None:
|
) -> None:
|
||||||
self._bind_route(method, path_template.replace("{prefix}", prefix), handler)
|
self._bind_route(method, path_template.replace("{prefix}", prefix), handler)
|
||||||
|
|
||||||
def _bind_route(self, method: str, path: str, handler: Callable) -> None:
|
def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
|
||||||
add_method_name = self._METHOD_MAP[method.upper()]
|
add_method_name = self._METHOD_MAP[method.upper()]
|
||||||
add_method = getattr(self._app.router, add_method_name)
|
add_method = getattr(self._app.router, add_method_name)
|
||||||
add_method(path, handler)
|
add_method(path, handler)
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
"""Route controller for the pending-delete undo endpoint."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
from .handlers.pending_delete_handler import PendingDeleteHandler
|
||||||
|
|
||||||
|
|
||||||
|
class PendingDeleteRoutes:
|
||||||
|
"""Shared route controller mirroring MiscRoutes/UpdateRoutes.
|
||||||
|
|
||||||
|
Registered ONCE per mode (py/lora_manager.py, standalone.py); NEVER through
|
||||||
|
the per-model-type ModelRouteRegistrar, which is instantiated per model
|
||||||
|
type and would register this non-prefixed route three times.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def setup_routes(app: web.Application) -> None:
|
||||||
|
"""Register the shared undo-delete endpoint."""
|
||||||
|
handler = PendingDeleteHandler()
|
||||||
|
_ = app.router.add_post("/api/lm/undo-delete", handler.undo_delete)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["PendingDeleteRoutes"]
|
||||||
@@ -3,7 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Callable, Mapping
|
from typing import Any, Callable, Mapping
|
||||||
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
|
|
||||||
@@ -43,6 +43,9 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
),
|
),
|
||||||
RouteDefinition("GET", "/api/lm/recipe/{recipe_id}/syntax", "get_recipe_syntax"),
|
RouteDefinition("GET", "/api/lm/recipe/{recipe_id}/syntax", "get_recipe_syntax"),
|
||||||
RouteDefinition("PUT", "/api/lm/recipe/{recipe_id}/update", "update_recipe"),
|
RouteDefinition("PUT", "/api/lm/recipe/{recipe_id}/update", "update_recipe"),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/recipe/{recipe_id}/opened", "record_recipe_open"
|
||||||
|
),
|
||||||
RouteDefinition("POST", "/api/lm/recipe/move", "move_recipe"),
|
RouteDefinition("POST", "/api/lm/recipe/move", "move_recipe"),
|
||||||
RouteDefinition("POST", "/api/lm/recipes/move-bulk", "move_recipes_bulk"),
|
RouteDefinition("POST", "/api/lm/recipes/move-bulk", "move_recipes_bulk"),
|
||||||
RouteDefinition("POST", "/api/lm/recipe/lora/reconnect", "reconnect_lora"),
|
RouteDefinition("POST", "/api/lm/recipe/lora/reconnect", "reconnect_lora"),
|
||||||
@@ -61,6 +64,11 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
RouteDefinition("POST", "/api/lm/recipe/{recipe_id}/repair", "repair_recipe"),
|
RouteDefinition("POST", "/api/lm/recipe/{recipe_id}/repair", "repair_recipe"),
|
||||||
RouteDefinition("POST", "/api/lm/recipes/repair-bulk", "repair_recipes_bulk"),
|
RouteDefinition("POST", "/api/lm/recipes/repair-bulk", "repair_recipes_bulk"),
|
||||||
RouteDefinition("GET", "/api/lm/recipes/repair-progress", "get_repair_progress"),
|
RouteDefinition("GET", "/api/lm/recipes/repair-progress", "get_repair_progress"),
|
||||||
|
RouteDefinition("POST", "/api/lm/recipes/rematch", "rematch_recipes"),
|
||||||
|
RouteDefinition("POST", "/api/lm/recipes/rematch-bulk", "rematch_recipes_bulk"),
|
||||||
|
RouteDefinition("POST", "/api/lm/recipe/{recipe_id}/rematch", "rematch_recipe"),
|
||||||
|
RouteDefinition("POST", "/api/lm/recipes/cancel-rematch", "cancel_rematch"),
|
||||||
|
RouteDefinition("GET", "/api/lm/recipes/rematch-progress", "get_rematch_progress"),
|
||||||
RouteDefinition("POST", "/api/lm/recipes/batch-import/start", "start_batch_import"),
|
RouteDefinition("POST", "/api/lm/recipes/batch-import/start", "start_batch_import"),
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"GET", "/api/lm/recipes/batch-import/progress", "get_batch_import_progress"
|
"GET", "/api/lm/recipes/batch-import/progress", "get_batch_import_progress"
|
||||||
@@ -105,7 +113,7 @@ class RecipeRouteRegistrar:
|
|||||||
handler = handler_lookup[definition.handler_name]
|
handler = handler_lookup[definition.handler_name]
|
||||||
self._bind_route(definition.method, definition.path, handler)
|
self._bind_route(definition.method, definition.path, handler)
|
||||||
|
|
||||||
def _bind_route(self, method: str, path: str, handler: Callable) -> None:
|
def _bind_route(self, method: str, path: str, handler: Callable[..., Any]) -> None:
|
||||||
add_method_name = self._METHOD_MAP[method.upper()]
|
add_method_name = self._METHOD_MAP[method.upper()]
|
||||||
add_method = getattr(self._app.router, add_method_name)
|
add_method = getattr(self._app.router, add_method_name)
|
||||||
add_method(path, handler)
|
add_method(path, handler)
|
||||||
|
|||||||
+11
-10
@@ -40,10 +40,11 @@ class StatsRoutes:
|
|||||||
"""Route handlers for Statistics page and API endpoints"""
|
"""Route handlers for Statistics page and API endpoints"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.lora_scanner = None
|
self.lora_scanner: Any = None
|
||||||
self.checkpoint_scanner = None
|
self.checkpoint_scanner: Any = None
|
||||||
self.embedding_scanner = None
|
self.embedding_scanner: Any = None
|
||||||
self.usage_stats = None
|
self.usage_stats: Any = None
|
||||||
|
self._i18n_filter_added = False
|
||||||
self.template_env = jinja2.Environment(
|
self.template_env = jinja2.Environment(
|
||||||
loader=jinja2.FileSystemLoader(config.templates_path),
|
loader=jinja2.FileSystemLoader(config.templates_path),
|
||||||
autoescape=True
|
autoescape=True
|
||||||
@@ -95,9 +96,9 @@ class StatsRoutes:
|
|||||||
server_i18n.set_locale(user_language)
|
server_i18n.set_locale(user_language)
|
||||||
|
|
||||||
# 为模板环境添加i18n过滤器
|
# 为模板环境添加i18n过滤器
|
||||||
if not hasattr(self.template_env, '_i18n_filter_added'):
|
if not self._i18n_filter_added:
|
||||||
self.template_env.filters['t'] = server_i18n.create_template_filter()
|
self.template_env.filters['t'] = server_i18n.create_template_filter()
|
||||||
self.template_env._i18n_filter_added = True
|
self._i18n_filter_added = True
|
||||||
|
|
||||||
template = self.template_env.get_template('statistics.html')
|
template = self.template_env.get_template('statistics.html')
|
||||||
rendered = template.render(
|
rendered = template.render(
|
||||||
@@ -549,7 +550,7 @@ class StatsRoutes:
|
|||||||
'error': str(e)
|
'error': str(e)
|
||||||
}, status=500)
|
}, status=500)
|
||||||
|
|
||||||
def _count_unused_models(self, models: List[Dict], usage_data: Dict) -> int:
|
def _count_unused_models(self, models: List[Dict[str, Any]], usage_data: Dict[str, Any]) -> int:
|
||||||
"""Count models that have never been used"""
|
"""Count models that have never been used"""
|
||||||
used_hashes = set(usage_data.keys())
|
used_hashes = set(usage_data.keys())
|
||||||
unused_count = 0
|
unused_count = 0
|
||||||
@@ -560,7 +561,7 @@ class StatsRoutes:
|
|||||||
|
|
||||||
return unused_count
|
return unused_count
|
||||||
|
|
||||||
def _get_top_used_models(self, usage_data: Dict, model_map: Dict, limit: int) -> List[Dict]:
|
def _get_top_used_models(self, usage_data: Dict[str, Any], model_map: Dict[str, Any], limit: int) -> List[Dict[str, Any]]:
|
||||||
"""Get top used models with their metadata"""
|
"""Get top used models with their metadata"""
|
||||||
sorted_usage = sorted(usage_data.items(), key=lambda x: x[1].get('total', 0), reverse=True)
|
sorted_usage = sorted(usage_data.items(), key=lambda x: x[1].get('total', 0), reverse=True)
|
||||||
|
|
||||||
@@ -578,7 +579,7 @@ class StatsRoutes:
|
|||||||
|
|
||||||
return top_models
|
return top_models
|
||||||
|
|
||||||
def _get_usage_timeline(self, usage_data: Dict, days: int) -> List[Dict]:
|
def _get_usage_timeline(self, usage_data: Dict[str, Any], days: int) -> List[Dict[str, Any]]:
|
||||||
"""Get usage timeline for the past N days"""
|
"""Get usage timeline for the past N days"""
|
||||||
timeline = []
|
timeline = []
|
||||||
today = datetime.now()
|
today = datetime.now()
|
||||||
@@ -614,7 +615,7 @@ class StatsRoutes:
|
|||||||
|
|
||||||
return list(reversed(timeline)) # Oldest to newest
|
return list(reversed(timeline)) # Oldest to newest
|
||||||
|
|
||||||
def _format_size(self, size_bytes: int) -> str:
|
def _format_size(self, size_bytes: float) -> str:
|
||||||
"""Format file size in human readable format"""
|
"""Format file size in human readable format"""
|
||||||
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
|
for unit in ['B', 'KB', 'MB', 'GB', 'TB']:
|
||||||
if size_bytes < 1024.0:
|
if size_bytes < 1024.0:
|
||||||
|
|||||||
+324
-53
@@ -6,7 +6,7 @@ import shutil
|
|||||||
import tempfile
|
import tempfile
|
||||||
import asyncio
|
import asyncio
|
||||||
from aiohttp import web, ClientError
|
from aiohttp import web, ClientError
|
||||||
from typing import Dict, List
|
from typing import Any, Dict, List, cast
|
||||||
|
|
||||||
from ..utils.settings_paths import ensure_settings_file
|
from ..utils.settings_paths import ensure_settings_file
|
||||||
from ..services.downloader import get_downloader
|
from ..services.downloader import get_downloader
|
||||||
@@ -38,6 +38,84 @@ def _clean_excludes() -> List[str]:
|
|||||||
return excludes
|
return excludes
|
||||||
|
|
||||||
|
|
||||||
|
def _stage_preserved_items(plugin_root: str) -> tuple[str, list[str]]:
|
||||||
|
"""Move preserved user-data items to a temp directory outside *plugin_root*.
|
||||||
|
|
||||||
|
This ensures that ``git reset --hard``, ``git clean -fd``, and ZIP-based
|
||||||
|
replacement cannot touch these files even when ``-e`` exclusion patterns
|
||||||
|
are mishandled (e.g. on Windows where forward-slash patterns may not
|
||||||
|
match backslash-prefixed paths in some Git builds, or where file locks
|
||||||
|
prevent deletion/recreation).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``(backup_root, staged_names)``: the temp directory path and the
|
||||||
|
list of item names that were successfully moved.
|
||||||
|
"""
|
||||||
|
backup_root = tempfile.mkdtemp(prefix='lora_manager_update_')
|
||||||
|
staged: list[str] = []
|
||||||
|
for name in _PRESERVE_DIRS:
|
||||||
|
src = os.path.join(plugin_root, name)
|
||||||
|
if not os.path.lexists(src):
|
||||||
|
continue
|
||||||
|
dst = os.path.join(backup_root, name)
|
||||||
|
try:
|
||||||
|
shutil.move(src, dst)
|
||||||
|
staged.append(name)
|
||||||
|
logger.debug("Staged '%s' for update safety", name)
|
||||||
|
except OSError:
|
||||||
|
# ``shutil.move`` may fail on Windows if a file handle inside
|
||||||
|
# the directory is still open (e.g. a SQLite WAL file). Fall
|
||||||
|
# back to copy-then-remove.
|
||||||
|
logger.debug("Move failed for '%s', falling back to copy", name)
|
||||||
|
try:
|
||||||
|
if os.path.isdir(src) and not os.path.islink(src):
|
||||||
|
shutil.copytree(src, dst, symlinks=True)
|
||||||
|
shutil.rmtree(src, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
shutil.copy2(src, dst)
|
||||||
|
os.remove(src)
|
||||||
|
staged.append(name)
|
||||||
|
logger.info("Copied (then removed) '%s' for update safety", name)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Could not stage '%s': %s (will rely on git -e / skip lists)", name, exc
|
||||||
|
)
|
||||||
|
return backup_root, staged
|
||||||
|
|
||||||
|
|
||||||
|
def _restore_preserved_items(plugin_root: str, backup_root: str, staged: list[str]) -> None:
|
||||||
|
"""Move staged items back from *backup_root* into *plugin_root*.
|
||||||
|
|
||||||
|
Any leftover placeholder at the destination (created by git checkout or
|
||||||
|
ZIP extraction) is removed before the move.
|
||||||
|
"""
|
||||||
|
for name in staged:
|
||||||
|
src = os.path.join(backup_root, name)
|
||||||
|
dst = os.path.join(plugin_root, name)
|
||||||
|
try:
|
||||||
|
if os.path.lexists(dst):
|
||||||
|
if os.path.isdir(dst) and not os.path.islink(dst):
|
||||||
|
shutil.rmtree(dst, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
os.remove(dst)
|
||||||
|
shutil.move(src, dst)
|
||||||
|
logger.debug("Restored '%s' after update", name)
|
||||||
|
except OSError:
|
||||||
|
logger.debug("Move failed restoring '%s', falling back to copy", name)
|
||||||
|
try:
|
||||||
|
if os.path.isdir(src) and not os.path.islink(src):
|
||||||
|
shutil.copytree(src, dst, symlinks=True, dirs_exist_ok=True)
|
||||||
|
shutil.rmtree(src, ignore_errors=True)
|
||||||
|
else:
|
||||||
|
shutil.copy2(src, dst)
|
||||||
|
os.remove(src)
|
||||||
|
logger.info("Copied '%s' back after update", name)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to restore '%s': %s", name, exc)
|
||||||
|
shutil.rmtree(backup_root, ignore_errors=True)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class UpdateRoutes:
|
class UpdateRoutes:
|
||||||
"""Routes for handling plugin update checks"""
|
"""Routes for handling plugin update checks"""
|
||||||
|
|
||||||
@@ -47,6 +125,7 @@ class UpdateRoutes:
|
|||||||
app.router.add_get('/api/lm/check-updates', UpdateRoutes.check_updates)
|
app.router.add_get('/api/lm/check-updates', UpdateRoutes.check_updates)
|
||||||
app.router.add_get('/api/lm/version-info', UpdateRoutes.get_version_info)
|
app.router.add_get('/api/lm/version-info', UpdateRoutes.get_version_info)
|
||||||
app.router.add_post('/api/lm/perform-update', UpdateRoutes.perform_update)
|
app.router.add_post('/api/lm/perform-update', UpdateRoutes.perform_update)
|
||||||
|
app.router.add_post('/api/lm/switch-channel', UpdateRoutes.switch_channel)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def check_updates(request):
|
async def check_updates(request):
|
||||||
@@ -65,10 +144,17 @@ class UpdateRoutes:
|
|||||||
|
|
||||||
# Fetch remote version from GitHub
|
# Fetch remote version from GitHub
|
||||||
if nightly:
|
if nightly:
|
||||||
remote_version, changelog = await UpdateRoutes._get_nightly_version()
|
local_hash = git_info.get('short_hash', '')
|
||||||
releases = None
|
nightly_version, releases_result = await asyncio.gather(
|
||||||
|
UpdateRoutes._get_nightly_version(local_hash),
|
||||||
|
UpdateRoutes._get_remote_version()
|
||||||
|
)
|
||||||
|
remote_version, _, behind_by, commit_date = nightly_version
|
||||||
|
_, changelog, releases = releases_result
|
||||||
else:
|
else:
|
||||||
remote_version, changelog, releases = await UpdateRoutes._get_remote_version()
|
remote_version, changelog, releases = await UpdateRoutes._get_remote_version()
|
||||||
|
behind_by = 0
|
||||||
|
commit_date = ''
|
||||||
|
|
||||||
# Compare versions
|
# Compare versions
|
||||||
if nightly:
|
if nightly:
|
||||||
@@ -81,6 +167,10 @@ class UpdateRoutes:
|
|||||||
remote_version.replace('v', '')
|
remote_version.replace('v', '')
|
||||||
)
|
)
|
||||||
|
|
||||||
|
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||||
|
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
|
||||||
|
|
||||||
response_data = {
|
response_data = {
|
||||||
'success': True,
|
'success': True,
|
||||||
'current_version': local_version,
|
'current_version': local_version,
|
||||||
@@ -88,13 +178,13 @@ class UpdateRoutes:
|
|||||||
'update_available': update_available,
|
'update_available': update_available,
|
||||||
'changelog': changelog,
|
'changelog': changelog,
|
||||||
'git_info': git_info,
|
'git_info': git_info,
|
||||||
'nightly': nightly
|
'nightly': nightly,
|
||||||
|
'has_git': has_git,
|
||||||
|
'releases': releases,
|
||||||
|
'behind_by': behind_by,
|
||||||
|
'commit_date': commit_date
|
||||||
}
|
}
|
||||||
|
|
||||||
# Include releases list for stable mode
|
|
||||||
if releases is not None:
|
|
||||||
response_data['releases'] = releases
|
|
||||||
|
|
||||||
return web.json_response(response_data)
|
return web.json_response(response_data)
|
||||||
|
|
||||||
except NETWORK_EXCEPTIONS as e:
|
except NETWORK_EXCEPTIONS as e:
|
||||||
@@ -126,9 +216,14 @@ class UpdateRoutes:
|
|||||||
# Format: version-short_hash
|
# Format: version-short_hash
|
||||||
version_string = f"{local_version}-{short_hash}"
|
version_string = f"{local_version}-{short_hash}"
|
||||||
|
|
||||||
|
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||||
|
has_git = os.path.exists(os.path.join(plugin_root, '.git'))
|
||||||
|
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
'success': True,
|
'success': True,
|
||||||
'version': version_string
|
'version': version_string,
|
||||||
|
'has_git': has_git
|
||||||
})
|
})
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -156,20 +251,22 @@ class UpdateRoutes:
|
|||||||
if os.path.exists(settings_path):
|
if os.path.exists(settings_path):
|
||||||
with open(settings_path, 'r', encoding='utf-8') as f:
|
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||||
settings_backup = f.read()
|
settings_backup = f.read()
|
||||||
logger.info("Backed up settings.json")
|
logger.debug("Backed up settings.json (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
git_folder = os.path.join(plugin_root, '.git')
|
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||||
if os.path.exists(git_folder):
|
try:
|
||||||
# Git update
|
git_folder = os.path.join(plugin_root, '.git')
|
||||||
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
if os.path.exists(git_folder):
|
||||||
else:
|
success, new_version = await UpdateRoutes._perform_git_update(plugin_root, nightly)
|
||||||
# Fallback: Download ZIP and replace files
|
else:
|
||||||
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||||
|
finally:
|
||||||
|
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||||
|
|
||||||
if settings_backup and success:
|
if settings_backup and success:
|
||||||
with open(settings_path, 'w', encoding='utf-8') as f:
|
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||||
f.write(settings_backup)
|
f.write(settings_backup)
|
||||||
logger.info("Restored settings.json")
|
logger.debug("Restored settings.json content (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
if success:
|
if success:
|
||||||
return web.json_response({
|
return web.json_response({
|
||||||
@@ -190,6 +287,164 @@ class UpdateRoutes:
|
|||||||
'error': str(e)
|
'error': str(e)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
async def switch_channel(request):
|
||||||
|
"""
|
||||||
|
Switch between release and nightly update channels.
|
||||||
|
|
||||||
|
ZIP/CNR install → Nightly: git init + checkout main (one-way upgrade)
|
||||||
|
Git install → Release: git checkout latest tag (.git preserved)
|
||||||
|
ZIP/CNR install → Release: ZIP download (no .git, stays in ZIP mode)
|
||||||
|
Git install → Nightly: git checkout main + pull
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
body = await request.json() if request.has_body else {}
|
||||||
|
channel = body.get('channel', '')
|
||||||
|
|
||||||
|
if channel not in ('release', 'nightly'):
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'error': f'Invalid channel: {channel}. Must be "release" or "nightly".'
|
||||||
|
})
|
||||||
|
|
||||||
|
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
plugin_root = os.path.dirname(os.path.dirname(current_dir))
|
||||||
|
|
||||||
|
settings_path = ensure_settings_file(logger)
|
||||||
|
settings_backup = None
|
||||||
|
if os.path.exists(settings_path):
|
||||||
|
with open(settings_path, 'r', encoding='utf-8') as f:
|
||||||
|
settings_backup = f.read()
|
||||||
|
logger.debug("Backed up settings.json before channel switch (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
|
staged_backup_dir, staged_items = _stage_preserved_items(plugin_root)
|
||||||
|
try:
|
||||||
|
git_folder = os.path.join(plugin_root, '.git')
|
||||||
|
|
||||||
|
if channel == 'nightly':
|
||||||
|
git_backup = None
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
git_backup = UpdateRoutes._backup_git(git_folder, 'nightly')
|
||||||
|
|
||||||
|
success = False
|
||||||
|
new_version = ''
|
||||||
|
try:
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
success, new_version = await UpdateRoutes._perform_git_update(
|
||||||
|
plugin_root, nightly=True
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
success, new_version = UpdateRoutes._init_git_repo(plugin_root)
|
||||||
|
finally:
|
||||||
|
UpdateRoutes._restore_git(git_backup, git_folder, success, 'nightly')
|
||||||
|
else:
|
||||||
|
success = False
|
||||||
|
new_version = ''
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
success, new_version = await UpdateRoutes._perform_git_update(
|
||||||
|
plugin_root, nightly=False
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||||
|
if os.path.exists(tracking_file):
|
||||||
|
os.remove(tracking_file)
|
||||||
|
success, new_version = await UpdateRoutes._download_and_replace_zip(plugin_root)
|
||||||
|
finally:
|
||||||
|
_restore_preserved_items(plugin_root, staged_backup_dir, staged_items)
|
||||||
|
|
||||||
|
if settings_backup and success:
|
||||||
|
with open(settings_path, 'w', encoding='utf-8') as f:
|
||||||
|
f.write(settings_backup)
|
||||||
|
logger.debug("Restored settings.json content after channel switch (%d bytes)", len(settings_backup))
|
||||||
|
|
||||||
|
if success:
|
||||||
|
return web.json_response({
|
||||||
|
'success': True,
|
||||||
|
'channel': channel,
|
||||||
|
'new_version': new_version,
|
||||||
|
'message': f'Switched to {channel} channel'
|
||||||
|
})
|
||||||
|
else:
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'error': f'Failed to switch to {channel} channel'
|
||||||
|
})
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to switch channel: %s", e, exc_info=True)
|
||||||
|
return web.json_response({
|
||||||
|
'success': False,
|
||||||
|
'error': str(e)
|
||||||
|
})
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _init_git_repo(plugin_root: str) -> tuple[bool, str]:
|
||||||
|
"""
|
||||||
|
Initialize a Git repository in a ZIP-installed plugin folder.
|
||||||
|
Clones the remote history and checks out main branch.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
import git
|
||||||
|
except ImportError:
|
||||||
|
logger.error(
|
||||||
|
"GitPython is not available: cannot initialize git repo. "
|
||||||
|
"Install git or set $GIT_PYTHON_GIT_EXECUTABLE to the git binary path."
|
||||||
|
)
|
||||||
|
return False, ""
|
||||||
|
|
||||||
|
clean_excludes = _clean_excludes()
|
||||||
|
|
||||||
|
try:
|
||||||
|
repo = git.Repo.init(plugin_root)
|
||||||
|
origin = repo.create_remote(
|
||||||
|
'origin',
|
||||||
|
'https://github.com/willmiao/ComfyUI-Lora-Manager.git'
|
||||||
|
)
|
||||||
|
origin.fetch()
|
||||||
|
|
||||||
|
repo.create_head('main', origin.refs.main)
|
||||||
|
repo.git.checkout('main', '--force')
|
||||||
|
repo.git.reset('--hard')
|
||||||
|
repo.git.clean('-fd', *clean_excludes)
|
||||||
|
|
||||||
|
tracking_file = os.path.join(plugin_root, '.tracking')
|
||||||
|
if os.path.exists(tracking_file):
|
||||||
|
os.remove(tracking_file)
|
||||||
|
logger.info("Removed .tracking file (now in git mode)")
|
||||||
|
|
||||||
|
new_version = f"main-{repo.head.commit.hexsha[:7]}"
|
||||||
|
logger.info("Initialized git repo on main branch: %s", new_version)
|
||||||
|
return True, new_version
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to initialize git repo: %s", e, exc_info=True)
|
||||||
|
return False, ""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _backup_git(git_folder, label):
|
||||||
|
try:
|
||||||
|
backup_dir = tempfile.mkdtemp()
|
||||||
|
backup = os.path.join(backup_dir, '.git')
|
||||||
|
shutil.copytree(git_folder, backup)
|
||||||
|
logger.info("Backed up .git before switching to %s", label)
|
||||||
|
return backup
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to backup .git before %s switch: %s", label, e)
|
||||||
|
return None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _restore_git(git_backup, git_folder, success, label):
|
||||||
|
if git_backup and not success:
|
||||||
|
try:
|
||||||
|
if os.path.exists(git_folder):
|
||||||
|
shutil.rmtree(git_folder)
|
||||||
|
shutil.copytree(git_backup, git_folder)
|
||||||
|
logger.info("Restored .git after failed %s switch", label)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Failed to restore .git after %s switch: %s", label, e)
|
||||||
|
if git_backup:
|
||||||
|
shutil.rmtree(os.path.dirname(git_backup), ignore_errors=True)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]:
|
async def _download_and_replace_zip(plugin_root: str) -> tuple[bool, str]:
|
||||||
"""
|
"""
|
||||||
@@ -212,9 +467,10 @@ class UpdateRoutes:
|
|||||||
if not success:
|
if not success:
|
||||||
logger.error(f"Failed to fetch release info: {data}")
|
logger.error(f"Failed to fetch release info: {data}")
|
||||||
return False, ""
|
return False, ""
|
||||||
|
|
||||||
zip_url = data.get("zipball_url")
|
release_payload = cast(dict[str, Any], data)
|
||||||
version = data.get("tag_name", "unknown")
|
zip_url = release_payload.get("zipball_url", "")
|
||||||
|
version = release_payload.get("tag_name", "unknown")
|
||||||
|
|
||||||
# Download ZIP to temporary file
|
# Download ZIP to temporary file
|
||||||
with tempfile.NamedTemporaryFile(delete=False, suffix=".zip") as tmp_zip:
|
with tempfile.NamedTemporaryFile(delete=False, suffix=".zip") as tmp_zip:
|
||||||
@@ -244,8 +500,7 @@ class UpdateRoutes:
|
|||||||
except Exception:
|
except Exception:
|
||||||
logger.debug("Could not close downloaded-version history database", exc_info=True)
|
logger.debug("Could not close downloaded-version history database", exc_info=True)
|
||||||
|
|
||||||
# Skip settings.json, civitai, model cache and runtime cache folders
|
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=list(_PRESERVE_DIRS))
|
||||||
UpdateRoutes._clean_plugin_folder(plugin_root, skip_files=['settings.json', 'civitai', 'model_cache', 'cache', 'wildcards', 'backups', 'stats'])
|
|
||||||
|
|
||||||
# Extract ZIP to temp dir
|
# Extract ZIP to temp dir
|
||||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
with tempfile.TemporaryDirectory() as tmp_dir:
|
||||||
@@ -255,7 +510,7 @@ class UpdateRoutes:
|
|||||||
extracted_root = next(os.scandir(tmp_dir)).path
|
extracted_root = next(os.scandir(tmp_dir)).path
|
||||||
|
|
||||||
# Copy files, skipping user data that should be preserved
|
# Copy files, skipping user data that should be preserved
|
||||||
skip_items = {'settings.json', 'civitai', 'wildcards', 'backups', 'stats'}
|
skip_items = set(_PRESERVE_DIRS)
|
||||||
for item in os.listdir(extracted_root):
|
for item in os.listdir(extracted_root):
|
||||||
if item in skip_items:
|
if item in skip_items:
|
||||||
continue
|
continue
|
||||||
@@ -272,7 +527,7 @@ class UpdateRoutes:
|
|||||||
# for ComfyUI Manager to work properly
|
# for ComfyUI Manager to work properly
|
||||||
tracking_info_file = os.path.join(plugin_root, '.tracking')
|
tracking_info_file = os.path.join(plugin_root, '.tracking')
|
||||||
tracking_files = []
|
tracking_files = []
|
||||||
skip_tracked = {'civitai', 'wildcards', 'backups', 'stats'}
|
skip_tracked = set(_PRESERVE_DIRS) - {'settings.json'}
|
||||||
for root, dirs, files in os.walk(extracted_root):
|
for root, dirs, files in os.walk(extracted_root):
|
||||||
# Skip user data directories and their contents
|
# Skip user data directories and their contents
|
||||||
rel_root = os.path.relpath(root, extracted_root)
|
rel_root = os.path.relpath(root, extracted_root)
|
||||||
@@ -295,7 +550,8 @@ class UpdateRoutes:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"ZIP update failed: {e}", exc_info=True)
|
logger.error(f"ZIP update failed: {e}", exc_info=True)
|
||||||
return False, ""
|
return False, ""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
def _clean_plugin_folder(plugin_root, skip_files=None):
|
def _clean_plugin_folder(plugin_root, skip_files=None):
|
||||||
skip_files = skip_files or []
|
skip_files = skip_files or []
|
||||||
for item in os.listdir(plugin_root):
|
for item in os.listdir(plugin_root):
|
||||||
@@ -308,41 +564,56 @@ class UpdateRoutes:
|
|||||||
os.remove(path)
|
os.remove(path)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _get_nightly_version() -> tuple[str, List[str]]:
|
async def _get_nightly_version(local_hash: str = "") -> tuple[str, List[str], int, str]:
|
||||||
"""
|
|
||||||
Fetch latest commit from main branch
|
|
||||||
"""
|
|
||||||
repo_owner = "willmiao"
|
repo_owner = "willmiao"
|
||||||
repo_name = "ComfyUI-Lora-Manager"
|
repo_name = "ComfyUI-Lora-Manager"
|
||||||
|
|
||||||
# Use GitHub API to fetch the latest commit from main branch
|
|
||||||
github_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/commits/main"
|
github_url = f"https://api.github.com/repos/{repo_owner}/{repo_name}/commits/main"
|
||||||
|
|
||||||
try:
|
try:
|
||||||
downloader = await get_downloader()
|
downloader = await get_downloader()
|
||||||
success, data = await downloader.make_request('GET', github_url, custom_headers={'Accept': 'application/vnd.github+json'})
|
success, data = await downloader.make_request(
|
||||||
|
'GET', github_url,
|
||||||
|
custom_headers={'Accept': 'application/vnd.github+json'}
|
||||||
|
)
|
||||||
|
|
||||||
if not success:
|
if not success:
|
||||||
logger.warning(f"Failed to fetch GitHub commit: {data}")
|
logger.warning("Failed to fetch GitHub commit: %s", data)
|
||||||
return "main", []
|
return "main", [], 0, ""
|
||||||
|
|
||||||
commit_sha = data.get('sha', '')[:7] # Short hash
|
commit_payload = cast(dict[str, Any], data)
|
||||||
commit_message = data.get('commit', {}).get('message', '')
|
commit_sha = commit_payload.get('sha', '')[:7]
|
||||||
|
commit_message = commit_payload.get('commit', {}).get('message', '')
|
||||||
# Format as "main-{short_hash}"
|
commit_date = commit_payload.get('commit', {}).get('committer', {}).get('date', '')[:10]
|
||||||
|
|
||||||
version = f"main-{commit_sha}"
|
version = f"main-{commit_sha}"
|
||||||
|
|
||||||
# Use commit message as changelog
|
|
||||||
changelog = [commit_message] if commit_message else []
|
changelog = [commit_message] if commit_message else []
|
||||||
|
|
||||||
return version, changelog
|
behind_by = 0
|
||||||
|
if local_hash and local_hash not in ('unknown', 'stable'):
|
||||||
|
compare_url = (
|
||||||
|
f"https://api.github.com/repos/{repo_owner}/{repo_name}"
|
||||||
|
f"/compare/{local_hash}...main"
|
||||||
|
)
|
||||||
|
c_ok, c_data = await downloader.make_request(
|
||||||
|
'GET', compare_url,
|
||||||
|
custom_headers={'Accept': 'application/vnd.github+json'}
|
||||||
|
)
|
||||||
|
if c_ok:
|
||||||
|
compare_payload = cast(dict[str, Any], c_data)
|
||||||
|
if compare_payload.get('status') in ('ahead', 'diverged'):
|
||||||
|
behind_by = compare_payload.get('ahead_by', 0)
|
||||||
|
else:
|
||||||
|
behind_by = compare_payload.get('behind_by', 0)
|
||||||
|
|
||||||
|
return version, changelog, behind_by, commit_date
|
||||||
|
|
||||||
except NETWORK_EXCEPTIONS as e:
|
except NETWORK_EXCEPTIONS as e:
|
||||||
logger.warning("Unable to reach GitHub for nightly version: %s", e)
|
logger.warning("Unable to reach GitHub for nightly version: %s", e)
|
||||||
return "main", []
|
return "main", [], 0, ""
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error fetching nightly version: {e}", exc_info=True)
|
logger.error("Error fetching nightly version: %s", e, exc_info=True)
|
||||||
return "main", []
|
return "main", [], 0, ""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _compare_nightly_versions(local_git_info: Dict[str, str], remote_version: str) -> bool:
|
def _compare_nightly_versions(local_git_info: Dict[str, str], remote_version: str) -> bool:
|
||||||
@@ -438,7 +709,7 @@ class UpdateRoutes:
|
|||||||
logger.info(f"Successfully updated to {new_version}")
|
logger.info(f"Successfully updated to {new_version}")
|
||||||
return True, new_version
|
return True, new_version
|
||||||
|
|
||||||
except git.exc.GitError as e:
|
except git.exc.GitError as e: # pyright: ignore[reportAttributeAccessIssue]
|
||||||
logger.error(f"Git error during update: {e}")
|
logger.error(f"Git error during update: {e}")
|
||||||
return False, ""
|
return False, ""
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -499,7 +770,7 @@ class UpdateRoutes:
|
|||||||
return git_info
|
return git_info
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _get_remote_version() -> tuple[str, List[str], List[Dict]]:
|
async def _get_remote_version() -> tuple[str, List[str], List[Dict[str, Any]]]:
|
||||||
"""
|
"""
|
||||||
Fetch remote version from GitHub
|
Fetch remote version from GitHub
|
||||||
Returns:
|
Returns:
|
||||||
@@ -521,7 +792,7 @@ class UpdateRoutes:
|
|||||||
|
|
||||||
# Parse releases
|
# Parse releases
|
||||||
releases = []
|
releases = []
|
||||||
for i, release in enumerate(data):
|
for i, release in enumerate(cast(list[dict[str, Any]], data)):
|
||||||
version = release.get('tag_name', '')
|
version = release.get('tag_name', '')
|
||||||
if not version.startswith('v'):
|
if not version.startswith('v'):
|
||||||
version = f"v{version}"
|
version = f"v{version}"
|
||||||
|
|||||||
@@ -117,7 +117,7 @@ def _render_prompt(template: str, variables: Dict[str, Any]) -> str:
|
|||||||
Uses simple regex substitution — no Jinja2 dependency needed.
|
Uses simple regex substitution — no Jinja2 dependency needed.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def replace(match: re.Match) -> str:
|
def replace(match: re.Match[str]) -> str:
|
||||||
key = match.group(1).strip()
|
key = match.group(1).strip()
|
||||||
value = variables.get(key, "")
|
value = variables.get(key, "")
|
||||||
if isinstance(value, (dict, list)):
|
if isinstance(value, (dict, list)):
|
||||||
|
|||||||
@@ -295,7 +295,7 @@ class PostProcessor:
|
|||||||
normalises every tag to lowercase for case-insensitive dedup.
|
normalises every tag to lowercase for case-insensitive dedup.
|
||||||
"""
|
"""
|
||||||
merged: List[str] = []
|
merged: List[str] = []
|
||||||
seen: set = set()
|
seen: set[str] = set()
|
||||||
for tag in list(existing) + list(new):
|
for tag in list(existing) + list(new):
|
||||||
t = tag.strip().lower()
|
t = tag.strip().lower()
|
||||||
if t and t not in seen:
|
if t and t not in seen:
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ _FRONTMATTER_RE = re.compile(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _parse_skill_file(path: Path) -> tuple[dict, str]:
|
def _parse_skill_file(path: Path) -> tuple[dict[str, Any], str]:
|
||||||
"""Read a prompt definition file (``prompt.md`` or legacy ``SKILL.md``) and
|
"""Read a prompt definition file (``prompt.md`` or legacy ``SKILL.md``) and
|
||||||
return (frontmatter_dict, body_text).
|
return (frontmatter_dict, body_text).
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import html as html_module
|
import html as html_module
|
||||||
import re
|
import re
|
||||||
from typing import List, Tuple
|
from typing import Any, List, Tuple
|
||||||
|
|
||||||
|
|
||||||
_REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
|
_REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
|
||||||
@@ -18,10 +18,10 @@ _REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
|
|||||||
def extract_simple_markdown_images(
|
def extract_simple_markdown_images(
|
||||||
markdown_text: str,
|
markdown_text: str,
|
||||||
repo: str,
|
repo: str,
|
||||||
existing_urls: set | None = None,
|
existing_urls: set[str] | None = None,
|
||||||
default_width: int = 512,
|
default_width: int = 512,
|
||||||
default_height: int = 512,
|
default_height: int = 512,
|
||||||
) -> list[dict]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Extract standalone markdown images from the README body.
|
"""Extract standalone markdown images from the README body.
|
||||||
|
|
||||||
Matches ```` on lines that are NOT part of a markdown table
|
Matches ```` on lines that are NOT part of a markdown table
|
||||||
@@ -36,8 +36,8 @@ def extract_simple_markdown_images(
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
||||||
images: list[dict] = []
|
images: list[dict[str, Any]] = []
|
||||||
seen_urls: set = set(existing_urls) if existing_urls else set()
|
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
||||||
|
|
||||||
# Collect lines that are NOT inside fenced code blocks
|
# Collect lines that are NOT inside fenced code blocks
|
||||||
lines = markdown_text.split("\n")
|
lines = markdown_text.split("\n")
|
||||||
@@ -86,10 +86,10 @@ def extract_simple_markdown_images(
|
|||||||
def extract_html_img_tags(
|
def extract_html_img_tags(
|
||||||
markdown_text: str,
|
markdown_text: str,
|
||||||
repo: str,
|
repo: str,
|
||||||
existing_urls: set | None = None,
|
existing_urls: set[str] | None = None,
|
||||||
default_width: int = 512,
|
default_width: int = 512,
|
||||||
default_height: int = 512,
|
default_height: int = 512,
|
||||||
) -> list[dict]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Extract image URLs from HTML ``<img src=\"...\">`` tags in the README.
|
"""Extract image URLs from HTML ``<img src=\"...\">`` tags in the README.
|
||||||
|
|
||||||
Many HF collection repos (e.g. ``deadman44/Z-Image_LoRA``) use raw HTML
|
Many HF collection repos (e.g. ``deadman44/Z-Image_LoRA``) use raw HTML
|
||||||
@@ -103,8 +103,8 @@ def extract_html_img_tags(
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
||||||
images: list[dict] = []
|
images: list[dict[str, Any]] = []
|
||||||
seen_urls: set = set(existing_urls) if existing_urls else set()
|
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
||||||
|
|
||||||
for m in re.finditer(
|
for m in re.finditer(
|
||||||
r'<img\s[^>]*src=\"([^\"]+)\"',
|
r'<img\s[^>]*src=\"([^\"]+)\"',
|
||||||
@@ -175,7 +175,7 @@ def extract_gallery_images(
|
|||||||
repo: str,
|
repo: str,
|
||||||
default_width: int = 512,
|
default_width: int = 512,
|
||||||
default_height: int = 512,
|
default_height: int = 512,
|
||||||
) -> List[dict]:
|
) -> List[dict[str, Any]]:
|
||||||
"""Extract widget/gallery images from the YAML frontmatter of a HF README.
|
"""Extract widget/gallery images from the YAML frontmatter of a HF README.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -196,7 +196,7 @@ def extract_gallery_images(
|
|||||||
if not frontmatter:
|
if not frontmatter:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
images: List[dict] = []
|
images: List[dict[str, Any]] = []
|
||||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
||||||
w = default_width or 512
|
w = default_width or 512
|
||||||
h = default_height or 512
|
h = default_height or 512
|
||||||
@@ -258,7 +258,7 @@ def extract_gallery_images(
|
|||||||
text = raw_text
|
text = raw_text
|
||||||
|
|
||||||
if url:
|
if url:
|
||||||
image: dict = {
|
image: dict[str, Any] = {
|
||||||
"url": url,
|
"url": url,
|
||||||
"type": "image",
|
"type": "image",
|
||||||
"nsfwLevel": 0,
|
"nsfwLevel": 0,
|
||||||
@@ -276,10 +276,10 @@ def extract_gallery_images(
|
|||||||
def extract_gallery_table_images(
|
def extract_gallery_table_images(
|
||||||
markdown_text: str,
|
markdown_text: str,
|
||||||
repo: str,
|
repo: str,
|
||||||
existing_urls: set | None = None,
|
existing_urls: set[str] | None = None,
|
||||||
default_width: int = 512,
|
default_width: int = 512,
|
||||||
default_height: int = 512,
|
default_height: int = 512,
|
||||||
) -> list[dict]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Extract images from ``| Preview | Prompt |`` markdown gallery tables.
|
"""Extract images from ``| Preview | Prompt |`` markdown gallery tables.
|
||||||
|
|
||||||
Many HF READMEs include a sample-gallery table in the body (outside
|
Many HF READMEs include a sample-gallery table in the body (outside
|
||||||
@@ -295,8 +295,8 @@ def extract_gallery_table_images(
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
||||||
images: list[dict] = []
|
images: list[dict[str, Any]] = []
|
||||||
seen_urls: set = set(existing_urls) if existing_urls else set()
|
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
||||||
lines = markdown_text.split("\n")
|
lines = markdown_text.split("\n")
|
||||||
n = len(lines)
|
n = len(lines)
|
||||||
i = 0
|
i = 0
|
||||||
@@ -514,7 +514,7 @@ def _strip_standalone_images(text: str) -> str:
|
|||||||
URL was stripped entirely, making it impossible for the LLM to return
|
URL was stripped entirely, making it impossible for the LLM to return
|
||||||
a ``preview_url`` for repos that use HTML ``<img>`` tags exclusively.
|
a ``preview_url`` for repos that use HTML ``<img>`` tags exclusively.
|
||||||
"""
|
"""
|
||||||
def _img_to_md(match: re.Match) -> str:
|
def _img_to_md(match: re.Match[str]) -> str:
|
||||||
"""Convert an ``<img>`` tag to markdown image syntax ````."""
|
"""Convert an ``<img>`` tag to markdown image syntax ````."""
|
||||||
tag = match.group(0)
|
tag = match.group(0)
|
||||||
src_m = re.search(r'src="([^"]+)"', tag) or re.search(r"src='([^']+)'", tag)
|
src_m = re.search(r'src="([^"]+)"', tag) or re.search(r"src='([^']+)'", tag)
|
||||||
@@ -942,7 +942,7 @@ def _strip_badge_images(text: str) -> str:
|
|||||||
"twitter", "colab", "gradio", "space",
|
"twitter", "colab", "gradio", "space",
|
||||||
)
|
)
|
||||||
|
|
||||||
def _should_remove(m: re.Match) -> str:
|
def _should_remove(m: re.Match[str]) -> str:
|
||||||
alt = (m.group(1) or "").lower()
|
alt = (m.group(1) or "").lower()
|
||||||
for kw in badge_keywords:
|
for kw in badge_keywords:
|
||||||
if kw in alt:
|
if kw in alt:
|
||||||
|
|||||||
+146
-18
@@ -1,3 +1,7 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
@@ -7,6 +11,7 @@ import os
|
|||||||
import secrets
|
import secrets
|
||||||
import shutil
|
import shutil
|
||||||
import socket
|
import socket
|
||||||
|
import time
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -20,10 +25,43 @@ from .settings_manager import get_settings_manager
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Maximum times the download poll loop will re-schedule a transfer after it
|
||||||
|
# is lost (daemon restart / RPC outage) before failing the download.
|
||||||
|
MAX_TRANSFER_RECOVERY_ATTEMPTS = 2
|
||||||
|
|
||||||
|
# stderr lines matching these markers indicate a disk write failure inside
|
||||||
|
# aria2 (piece cache flush or raw file write). They are promoted to INFO so
|
||||||
|
# the root cause (disk full, permission denied, file locked by another
|
||||||
|
# process, ...) is visible in the default logs; all other stderr output stays
|
||||||
|
# at DEBUG to avoid noise.
|
||||||
|
_DISK_WRITE_ERROR_MARKERS = (
|
||||||
|
# aria2 wrapper messages (write disk cache flush path)
|
||||||
|
"write disk cache flush failure",
|
||||||
|
"error when trying to flush write cache",
|
||||||
|
"failed to write into the file",
|
||||||
|
"failed to open the file",
|
||||||
|
"failed to seek the file",
|
||||||
|
# underlying root-cause phrases reported via "cause: ..." (POSIX + Windows)
|
||||||
|
"no space left on device",
|
||||||
|
"not enough space on the disk",
|
||||||
|
"input/output error",
|
||||||
|
"permission denied",
|
||||||
|
"access is denied",
|
||||||
|
"disk quota exceeded",
|
||||||
|
"used by another process",
|
||||||
|
"sharing violation",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Minimum interval between INFO-level reports of the same stderr line so a
|
||||||
|
# repeated failure (e.g. aria2 retrying against a full disk) does not spam
|
||||||
|
# the log.
|
||||||
|
STDERR_ERROR_REPORT_INTERVAL = 60.0
|
||||||
|
|
||||||
|
|
||||||
def _try_certifi_ca_path() -> str | None:
|
def _try_certifi_ca_path() -> str | None:
|
||||||
"""Return the certifi CA bundle path if available, else None."""
|
"""Return the certifi CA bundle path if available, else None."""
|
||||||
try:
|
try:
|
||||||
import certifi # type: ignore[import-untyped]
|
import certifi # pyright: ignore[reportMissingTypeStubs]
|
||||||
|
|
||||||
path = certifi.where()
|
path = certifi.where()
|
||||||
if os.path.isfile(path):
|
if os.path.isfile(path):
|
||||||
@@ -81,10 +119,12 @@ class Aria2Downloader:
|
|||||||
self._rpc_session: Optional[aiohttp.ClientSession] = None
|
self._rpc_session: Optional[aiohttp.ClientSession] = None
|
||||||
self._rpc_session_lock = asyncio.Lock()
|
self._rpc_session_lock = asyncio.Lock()
|
||||||
self._process_lock = asyncio.Lock()
|
self._process_lock = asyncio.Lock()
|
||||||
|
self._register_lock = asyncio.Lock()
|
||||||
self._transfers: Dict[str, Aria2Transfer] = {}
|
self._transfers: Dict[str, Aria2Transfer] = {}
|
||||||
self._poll_interval = 0.5
|
self._poll_interval = 0.5
|
||||||
self._state_store = Aria2TransferStateStore()
|
self._state_store = Aria2TransferStateStore()
|
||||||
self._stderr_reader_task: Optional[asyncio.Task] = None
|
self._stderr_reader_task: Optional[asyncio.Task[Any]] = None
|
||||||
|
self._stderr_error_report: Dict[str, float] = {}
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_running(self) -> bool:
|
def is_running(self) -> bool:
|
||||||
@@ -99,26 +139,58 @@ class Aria2Downloader:
|
|||||||
progress_callback=None,
|
progress_callback=None,
|
||||||
headers: Optional[Dict[str, str]] = None,
|
headers: Optional[Dict[str, str]] = None,
|
||||||
) -> Tuple[bool, str]:
|
) -> Tuple[bool, str]:
|
||||||
"""Download a file using aria2 RPC and wait for completion."""
|
"""Download a file using aria2 RPC and wait for completion.
|
||||||
|
|
||||||
|
The poll loop is self-healing: when the in-memory transfer entry
|
||||||
|
disappears (e.g. another download restarted the daemon and
|
||||||
|
``close()`` cleared ``_transfers``) or the RPC becomes unreachable,
|
||||||
|
the transfer is re-scheduled with ``continue=true`` so the download
|
||||||
|
resumes from the on-disk ``.aria2`` control file. Recovery is bounded
|
||||||
|
by ``MAX_TRANSFER_RECOVERY_ATTEMPTS``.
|
||||||
|
"""
|
||||||
|
|
||||||
await self._ensure_process()
|
await self._ensure_process()
|
||||||
save_path = os.path.abspath(save_path)
|
save_path = os.path.abspath(save_path)
|
||||||
transfer = self._transfers.get(download_id)
|
|
||||||
if transfer is None or os.path.abspath(transfer.save_path) != save_path:
|
|
||||||
gid = await self._schedule_download(
|
|
||||||
url,
|
|
||||||
save_path,
|
|
||||||
download_id=download_id,
|
|
||||||
headers=headers,
|
|
||||||
)
|
|
||||||
transfer = Aria2Transfer(gid=gid, save_path=save_path)
|
|
||||||
self._transfers[download_id] = transfer
|
|
||||||
|
|
||||||
|
async with self._register_lock:
|
||||||
|
transfer = self._transfers.get(download_id)
|
||||||
|
if transfer is None or os.path.abspath(transfer.save_path) != save_path:
|
||||||
|
transfer = await self._register_transfer(
|
||||||
|
url,
|
||||||
|
save_path,
|
||||||
|
download_id=download_id,
|
||||||
|
headers=headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
recovery_attempts = 0
|
||||||
try:
|
try:
|
||||||
while True:
|
while True:
|
||||||
status = await self._get_status_with_retry(download_id)
|
try:
|
||||||
|
status = await self._get_status_with_retry(download_id)
|
||||||
|
except Aria2Error:
|
||||||
|
status = None
|
||||||
|
|
||||||
if status is None:
|
if status is None:
|
||||||
return False, "aria2 download not found"
|
if recovery_attempts >= MAX_TRANSFER_RECOVERY_ATTEMPTS:
|
||||||
|
return False, "aria2 download not found"
|
||||||
|
recovery_attempts += 1
|
||||||
|
logger.warning(
|
||||||
|
"aria2 transfer %s lost; re-scheduling with resume "
|
||||||
|
"(attempt %d/%d)",
|
||||||
|
download_id,
|
||||||
|
recovery_attempts,
|
||||||
|
MAX_TRANSFER_RECOVERY_ATTEMPTS,
|
||||||
|
)
|
||||||
|
await asyncio.sleep(1.0)
|
||||||
|
await self._ensure_process()
|
||||||
|
async with self._register_lock:
|
||||||
|
transfer = await self._register_transfer(
|
||||||
|
url,
|
||||||
|
save_path,
|
||||||
|
download_id=download_id,
|
||||||
|
headers=headers,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
snapshot = self._build_progress_snapshot(status)
|
snapshot = self._build_progress_snapshot(status)
|
||||||
if progress_callback is not None:
|
if progress_callback is not None:
|
||||||
@@ -135,7 +207,9 @@ class Aria2Downloader:
|
|||||||
|
|
||||||
await asyncio.sleep(self._poll_interval)
|
await asyncio.sleep(self._poll_interval)
|
||||||
finally:
|
finally:
|
||||||
self._transfers.pop(download_id, None)
|
current = self._transfers.get(download_id)
|
||||||
|
if current is not None and current.gid == transfer.gid:
|
||||||
|
self._transfers.pop(download_id, None)
|
||||||
|
|
||||||
async def _get_status_with_retry(
|
async def _get_status_with_retry(
|
||||||
self, download_id: str, *, max_retries: int = 4, retry_delay: float = 3.0
|
self, download_id: str, *, max_retries: int = 4, retry_delay: float = 3.0
|
||||||
@@ -190,7 +264,7 @@ class Aria2Downloader:
|
|||||||
download_id,
|
download_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
options: Dict[str, str] = {
|
options: Dict[str, Any] = {
|
||||||
"dir": save_dir,
|
"dir": save_dir,
|
||||||
"out": out_name,
|
"out": out_name,
|
||||||
"continue": "true",
|
"continue": "true",
|
||||||
@@ -238,6 +312,25 @@ class Aria2Downloader:
|
|||||||
)
|
)
|
||||||
return gid
|
return gid
|
||||||
|
|
||||||
|
async def _register_transfer(
|
||||||
|
self,
|
||||||
|
url: str,
|
||||||
|
save_path: str,
|
||||||
|
*,
|
||||||
|
download_id: str,
|
||||||
|
headers: Optional[Dict[str, str]] = None,
|
||||||
|
) -> Aria2Transfer:
|
||||||
|
"""Schedule a download and track it in the in-memory transfer registry."""
|
||||||
|
gid = await self._schedule_download(
|
||||||
|
url,
|
||||||
|
save_path,
|
||||||
|
download_id=download_id,
|
||||||
|
headers=headers,
|
||||||
|
)
|
||||||
|
transfer = Aria2Transfer(gid=gid, save_path=os.path.abspath(save_path))
|
||||||
|
self._transfers[download_id] = transfer
|
||||||
|
return transfer
|
||||||
|
|
||||||
async def get_status(self, download_id: str) -> Optional[Dict[str, Any]]:
|
async def get_status(self, download_id: str) -> Optional[Dict[str, Any]]:
|
||||||
"""Return the raw aria2 status payload for a known download."""
|
"""Return the raw aria2 status payload for a known download."""
|
||||||
|
|
||||||
@@ -385,16 +478,51 @@ class Aria2Downloader:
|
|||||||
blocks, which freezes the entire ``aria2c`` process — including its
|
blocks, which freezes the entire ``aria2c`` process — including its
|
||||||
RPC handler. This background task reads lines from stderr as they
|
RPC handler. This background task reads lines from stderr as they
|
||||||
arrive and forwards them to Python's logger.
|
arrive and forwards them to Python's logger.
|
||||||
|
|
||||||
|
Lines that indicate a disk write failure (e.g. the "cause: No space
|
||||||
|
left on device" line that follows "Write disk cache flush failure")
|
||||||
|
are promoted to INFO so the root cause is visible without enabling
|
||||||
|
debug logging; every other line stays at DEBUG to avoid noise.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
assert self._process is not None and self._process.stderr is not None
|
assert self._process is not None and self._process.stderr is not None
|
||||||
async for line in self._process.stderr:
|
async for line in self._process.stderr:
|
||||||
text = line.decode("utf-8", errors="replace").rstrip()
|
text = line.decode("utf-8", errors="replace").rstrip()
|
||||||
if text:
|
if text:
|
||||||
logger.debug("aria2 stderr: %s", text)
|
if self._is_disk_write_error(text):
|
||||||
|
self._report_stderr_error(text)
|
||||||
|
else:
|
||||||
|
logger.debug("aria2 stderr: %s", text)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_disk_write_error(text: str) -> bool:
|
||||||
|
lowered = text.lower()
|
||||||
|
return any(marker in lowered for marker in _DISK_WRITE_ERROR_MARKERS)
|
||||||
|
|
||||||
|
def _report_stderr_error(self, text: str) -> None:
|
||||||
|
"""INFO-log a disk write failure line, rate-limited per line text.
|
||||||
|
|
||||||
|
aria2 re-emits the same error chain on every poll/retry while the
|
||||||
|
underlying condition persists; only the first occurrence within
|
||||||
|
``STDERR_ERROR_REPORT_INTERVAL`` seconds is promoted to INFO.
|
||||||
|
"""
|
||||||
|
now = time.monotonic()
|
||||||
|
last = self._stderr_error_report.get(text)
|
||||||
|
if last is not None and now - last < STDERR_ERROR_REPORT_INTERVAL:
|
||||||
|
logger.debug("aria2 stderr (repeated disk write error): %s", text)
|
||||||
|
return
|
||||||
|
# Drop entries older than the window so the map stays bounded even
|
||||||
|
# during a long disk-full episode (piece indexes change per line).
|
||||||
|
self._stderr_error_report = {
|
||||||
|
line: timestamp
|
||||||
|
for line, timestamp in self._stderr_error_report.items()
|
||||||
|
if now - timestamp < STDERR_ERROR_REPORT_INTERVAL
|
||||||
|
}
|
||||||
|
self._stderr_error_report[text] = now
|
||||||
|
logger.info("aria2 disk write failure: %s", text)
|
||||||
|
|
||||||
async def _dispatch_progress(self, callback, snapshot: DownloadProgress) -> None:
|
async def _dispatch_progress(self, callback, snapshot: DownloadProgress) -> None:
|
||||||
try:
|
try:
|
||||||
result = callback(snapshot, snapshot)
|
result = callback(snapshot, snapshot)
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from filename, base_model, and CivitAI version name — no manual tagging requir
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import re
|
import re
|
||||||
from typing import Dict, List, Set
|
from typing import Any, Dict, List, Set
|
||||||
|
|
||||||
# ── Tag category definitions ──────────────────────────────────────────
|
# ── Tag category definitions ──────────────────────────────────────────
|
||||||
# Each category maps a display label to a regex pattern.
|
# Each category maps a display label to a regex pattern.
|
||||||
@@ -52,7 +52,7 @@ AUTO_TAG_GROUPS = {
|
|||||||
DEFAULT_ENABLED_GROUPS = {"mode", "video"}
|
DEFAULT_ENABLED_GROUPS = {"mode", "video"}
|
||||||
|
|
||||||
|
|
||||||
def _collect_sources(model_data: Dict) -> List[str]:
|
def _collect_sources(model_data: Dict[str, Any]) -> List[str]:
|
||||||
"""Collect all text sources from model data for tag matching."""
|
"""Collect all text sources from model data for tag matching."""
|
||||||
sources: List[str] = []
|
sources: List[str] = []
|
||||||
|
|
||||||
@@ -73,7 +73,7 @@ def _collect_sources(model_data: Dict) -> List[str]:
|
|||||||
return sources
|
return sources
|
||||||
|
|
||||||
|
|
||||||
def extract_auto_tags(model_data: Dict) -> List[str]:
|
def extract_auto_tags(model_data: Dict[str, Any]) -> List[str]:
|
||||||
"""Extract auto-detected tags from model metadata.
|
"""Extract auto-detected tags from model metadata.
|
||||||
|
|
||||||
Uses a two-layer approach:
|
Uses a two-layer approach:
|
||||||
|
|||||||
@@ -0,0 +1,144 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
|
"""Backfill the AutoV3 checked state for models loaded from a persisted snapshot.
|
||||||
|
|
||||||
|
The SQLite persistent cache predates the AutoV3 feature, so entries hydrated
|
||||||
|
from it have a NULL ``autov3`` column (the "not checked yet" state). This
|
||||||
|
service computes the embedded AutoV3 hash for each such model — once per
|
||||||
|
process — and persists it through the scanner's single write path
|
||||||
|
(:meth:`ModelScanner.update_autov3_for_model`), marking every visited row so a
|
||||||
|
subsequent run finds nothing left to do.
|
||||||
|
|
||||||
|
Three-state contract honored here:
|
||||||
|
|
||||||
|
- ``NULL`` (sqlite) / absent (dict) = not checked yet → backfill computes it
|
||||||
|
- ``''`` (sqlite/dict) / JSON null = checked, no value available → never recompute
|
||||||
|
- 12-char lowercase hex = value → never recompute
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import threading
|
||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
if TYPE_CHECKING: # pragma: no cover - type-check only; runtime imports are local
|
||||||
|
from .model_scanner import ModelScanner
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_autov3(file_path: str) -> str:
|
||||||
|
"""Resolve the AutoV3 hash for a model file.
|
||||||
|
|
||||||
|
Prefers the Civitai AutoV3 reported for the file whose SHA256 matches
|
||||||
|
(the authoritative value for recipe matching); falls back to the embedded
|
||||||
|
safetensors header hash. Returns ``''`` when neither is available.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
metadata_path = f"{os.path.splitext(file_path)[0]}.metadata.json"
|
||||||
|
if os.path.exists(metadata_path):
|
||||||
|
with open(metadata_path, "r", encoding="utf-8") as handle:
|
||||||
|
payload = json.load(handle)
|
||||||
|
if isinstance(payload, dict):
|
||||||
|
from ..utils.models import autov3_from_civitai_files # local import avoids cycles
|
||||||
|
|
||||||
|
sha256 = (payload.get("sha256") or "").lower()
|
||||||
|
civitai_autov3 = autov3_from_civitai_files(payload.get("civitai"), sha256)
|
||||||
|
if civitai_autov3:
|
||||||
|
return civitai_autov3
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
from ..utils.file_utils import calculate_autov3 # local import avoids cycles
|
||||||
|
|
||||||
|
return calculate_autov3(file_path) or ""
|
||||||
|
|
||||||
|
|
||||||
|
class Autov3BackfillService:
|
||||||
|
"""Compute and persist AutoV3 hashes for models missing a checked state."""
|
||||||
|
|
||||||
|
_instance: Optional["Autov3BackfillService"] = None
|
||||||
|
_instance_lock = threading.Lock()
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
# Re-entrancy guard per model type: scanners for different model types
|
||||||
|
# initialize concurrently (lora_manager.py), so a global guard would
|
||||||
|
# silently skip every type but the first to start. Each model type
|
||||||
|
# runs its own backfill; a duplicate trigger for the same type no-ops.
|
||||||
|
self._running_types: set[str] = set()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_instance(cls) -> "Autov3BackfillService":
|
||||||
|
"""Return the process-wide singleton instance."""
|
||||||
|
if cls._instance is None:
|
||||||
|
with cls._instance_lock:
|
||||||
|
if cls._instance is None:
|
||||||
|
cls._instance = cls()
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
async def backfill(self, scanner: "ModelScanner") -> int:
|
||||||
|
"""Compute AutoV3 for every un-checked model of ``scanner.model_type``.
|
||||||
|
|
||||||
|
Each candidate file is read once via :func:`~py.utils.file_utils.calculate_autov3`
|
||||||
|
(cheap: safetensors header only) and the result is persisted through
|
||||||
|
``scanner.update_autov3_for_model``. Files that no longer exist on
|
||||||
|
disk are skipped — they are intentionally NOT marked, because scanner
|
||||||
|
cleanup removes the stale row later.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The number of models successfully updated. Never raises; on any
|
||||||
|
failure a warning is logged and ``0`` is returned. A duplicate
|
||||||
|
trigger for a model type that is already being backfilled returns
|
||||||
|
``0`` immediately; different model types run concurrently.
|
||||||
|
"""
|
||||||
|
model_type = scanner.model_type
|
||||||
|
if model_type in self._running_types:
|
||||||
|
return 0
|
||||||
|
self._running_types.add(model_type)
|
||||||
|
try:
|
||||||
|
# Local imports avoid import cycles at module load time.
|
||||||
|
from .persistent_model_cache import get_persistent_cache
|
||||||
|
from ..utils.file_utils import calculate_autov3
|
||||||
|
|
||||||
|
persistent = getattr(scanner, "_persistent_cache", None) or get_persistent_cache()
|
||||||
|
paths = persistent.get_models_missing_autov3(model_type)
|
||||||
|
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
count = 0
|
||||||
|
for path in paths:
|
||||||
|
# A file that no longer exists must not be marked; scanner
|
||||||
|
# cleanup removes the stale row later. The existence check and
|
||||||
|
# hash resolution run in the executor so the loop stays
|
||||||
|
# responsive to API requests while the backfill iterates a
|
||||||
|
# large library.
|
||||||
|
if not await loop.run_in_executor(None, os.path.exists, path):
|
||||||
|
continue
|
||||||
|
autov3 = await loop.run_in_executor(None, _resolve_autov3, path)
|
||||||
|
if await scanner.update_autov3_for_model(model_type, path, autov3):
|
||||||
|
count += 1
|
||||||
|
|
||||||
|
if paths:
|
||||||
|
logger.info(
|
||||||
|
"AutoV3 backfill: updated %d/%d models for %s",
|
||||||
|
count,
|
||||||
|
len(paths),
|
||||||
|
model_type,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Steady state after the first run: nothing left to backfill.
|
||||||
|
logger.debug("AutoV3 backfill: nothing to process for %s", model_type)
|
||||||
|
return count
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"AutoV3 backfill failed for %s: %s",
|
||||||
|
getattr(scanner, "model_type", "?"),
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return 0
|
||||||
|
finally:
|
||||||
|
self._running_types.discard(model_type)
|
||||||
@@ -1,3 +1,7 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
import asyncio
|
import asyncio
|
||||||
import re
|
import re
|
||||||
from typing import Any, Dict, List, Optional, Type, TYPE_CHECKING
|
import random
|
||||||
|
from typing import Any, Awaitable, Dict, List, Optional, Type, Union, TYPE_CHECKING, cast
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
@@ -69,24 +70,24 @@ class BaseModelService(ABC):
|
|||||||
page: int,
|
page: int,
|
||||||
page_size: int,
|
page_size: int,
|
||||||
sort_by: str = "name",
|
sort_by: str = "name",
|
||||||
folder: str = None,
|
folder: str | None = None,
|
||||||
folder_include: list = None,
|
folder_include: list[str] | None = None,
|
||||||
folder_exclude: list = None,
|
folder_exclude: list[str] | None = None,
|
||||||
search: str = None,
|
search: str | None = None,
|
||||||
fuzzy_search: bool = False,
|
fuzzy_search: bool = False,
|
||||||
base_models: list = None,
|
base_models: list[str] | None = None,
|
||||||
model_types: list = None,
|
model_types: list[str] | None = None,
|
||||||
tags: Optional[Dict[str, str]] = None,
|
tags: Optional[Dict[str, str]] = None,
|
||||||
auto_tags: Optional[Dict[str, str]] = None,
|
auto_tags: Optional[Dict[str, str]] = None,
|
||||||
search_options: dict = None,
|
search_options: dict[str, Any] | None = None,
|
||||||
hash_filters: dict = None,
|
hash_filters: dict[str, Any] | None = None,
|
||||||
favorites_only: bool = False,
|
favorites_only: bool = False,
|
||||||
update_available_only: bool = False,
|
update_available_only: bool = False,
|
||||||
credit_required: Optional[bool] = None,
|
credit_required: Optional[bool] = None,
|
||||||
allow_selling_generated_content: Optional[bool] = None,
|
allow_selling_generated_content: Optional[bool] = None,
|
||||||
tag_logic: str = "any",
|
tag_logic: str = "any",
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Dict:
|
) -> Dict[str, Any]:
|
||||||
"""Get paginated and filtered model data"""
|
"""Get paginated and filtered model data"""
|
||||||
overall_start = time.perf_counter()
|
overall_start = time.perf_counter()
|
||||||
|
|
||||||
@@ -109,12 +110,15 @@ class BaseModelService(ABC):
|
|||||||
if civitai_model_id is not None:
|
if civitai_model_id is not None:
|
||||||
sorted_data = [
|
sorted_data = [
|
||||||
item for item in sorted_data
|
item for item in sorted_data
|
||||||
if self._extract_model_id(item) == civitai_model_id
|
if self._extract_group_key(item) == civitai_model_id
|
||||||
]
|
]
|
||||||
# VLM mode: always sort by version ID descending (newest version first),
|
# VLM mode: always sort by version ID descending (newest version first),
|
||||||
# regardless of the current sort_by preference.
|
# regardless of the current sort_by preference.
|
||||||
|
# Fall back to modified timestamp for non-CivitAI sources.
|
||||||
sorted_data.sort(
|
sorted_data.sort(
|
||||||
key=lambda x: self._extract_version_id(x) or 0,
|
key=lambda x: self._extract_version_id(x)
|
||||||
|
or x.get("modified", 0)
|
||||||
|
or 0,
|
||||||
reverse=True,
|
reverse=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -129,18 +133,21 @@ class BaseModelService(ABC):
|
|||||||
ufs = self.settings.get("version_grouping", "same_base")
|
ufs = self.settings.get("version_grouping", "same_base")
|
||||||
group_by_base = ufs == "same_base"
|
group_by_base = ufs == "same_base"
|
||||||
|
|
||||||
dedup_map = {} # (modelId [,base_model]) -> (item, version_id)
|
dedup_map = {} # (modelId [,base_model]) -> (item, version_or_modified)
|
||||||
version_counter = {} # same-key -> count
|
version_counter = {} # same-key -> count
|
||||||
standalone = []
|
standalone = []
|
||||||
for item in sorted_data:
|
for item in sorted_data:
|
||||||
mid = self._extract_model_id(item)
|
mid = self._extract_group_key(item)
|
||||||
if mid is None:
|
if mid is None:
|
||||||
standalone.append(item)
|
standalone.append(item)
|
||||||
continue
|
continue
|
||||||
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
||||||
# Count all versions per key
|
# Count all versions per key
|
||||||
version_counter[key] = version_counter.get(key, 0) + 1
|
version_counter[key] = version_counter.get(key, 0) + 1
|
||||||
vid = self._extract_version_id(item) or 0
|
# Prefer CivitAI version_id; fall back to modified timestamp
|
||||||
|
vid = self._extract_version_id(item)
|
||||||
|
if vid is None:
|
||||||
|
vid = item.get("modified", 0) or 0
|
||||||
if key not in dedup_map or vid > dedup_map[key][1]:
|
if key not in dedup_map or vid > dedup_map[key][1]:
|
||||||
dedup_map[key] = (item, vid)
|
dedup_map[key] = (item, vid)
|
||||||
# Attach version_count to each surviving grouped item (shallow copy
|
# Attach version_count to each surviving grouped item (shallow copy
|
||||||
@@ -171,19 +178,22 @@ class BaseModelService(ABC):
|
|||||||
ufs = self.settings.get("version_grouping", "same_base")
|
ufs = self.settings.get("version_grouping", "same_base")
|
||||||
group_by_base = ufs == "same_base"
|
group_by_base = ufs == "same_base"
|
||||||
|
|
||||||
model_groups: Dict[Any, List[Dict]] = {}
|
model_groups: Dict[Any, List[Dict[str, Any]]] = {}
|
||||||
ungrouped_standalone: List[Dict] = []
|
ungrouped_standalone: List[Dict[str, Any]] = []
|
||||||
for item in sorted_data:
|
for item in sorted_data:
|
||||||
mid = self._extract_model_id(item)
|
mid = self._extract_group_key(item)
|
||||||
if mid is None:
|
if mid is None:
|
||||||
ungrouped_standalone.append(item)
|
ungrouped_standalone.append(item)
|
||||||
continue
|
continue
|
||||||
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
key = (mid, item.get("base_model") or "") if group_by_base else mid
|
||||||
model_groups.setdefault(key, []).append(item)
|
model_groups.setdefault(key, []).append(item)
|
||||||
# Sort versions within each group by version id descending
|
# Sort versions within each group by version id (descending);
|
||||||
|
# fall back to modified timestamp for non-CivitAI sources.
|
||||||
for items in model_groups.values():
|
for items in model_groups.values():
|
||||||
items.sort(
|
items.sort(
|
||||||
key=lambda x: self._extract_version_id(x) or 0,
|
key=lambda x: self._extract_version_id(x)
|
||||||
|
or x.get("modified", 0)
|
||||||
|
or 0,
|
||||||
reverse=True,
|
reverse=True,
|
||||||
)
|
)
|
||||||
# Sort groups by version count
|
# Sort groups by version count
|
||||||
@@ -239,7 +249,7 @@ class BaseModelService(ABC):
|
|||||||
filter_duration = time.perf_counter() - t1
|
filter_duration = time.perf_counter() - t1
|
||||||
post_filter_count = len(filtered_data)
|
post_filter_count = len(filtered_data)
|
||||||
|
|
||||||
annotated_for_filter: Optional[List[Dict]] = None
|
annotated_for_filter: Optional[List[Dict[str, Any]]] = None
|
||||||
t2 = time.perf_counter()
|
t2 = time.perf_counter()
|
||||||
if update_available_only:
|
if update_available_only:
|
||||||
annotated_for_filter = await self._annotate_update_flags(filtered_data)
|
annotated_for_filter = await self._annotate_update_flags(filtered_data)
|
||||||
@@ -286,11 +296,11 @@ class BaseModelService(ABC):
|
|||||||
page: int,
|
page: int,
|
||||||
page_size: int,
|
page_size: int,
|
||||||
sort_by: str = "name",
|
sort_by: str = "name",
|
||||||
search: str = None,
|
search: str | None = None,
|
||||||
fuzzy_search: bool = False,
|
fuzzy_search: bool = False,
|
||||||
search_options: dict = None,
|
search_options: dict[str, Any] | None = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Dict:
|
) -> Dict[str, Any]:
|
||||||
"""Get paginated excluded model data."""
|
"""Get paginated excluded model data."""
|
||||||
excluded_paths = list(self.scanner.get_excluded_models())
|
excluded_paths = list(self.scanner.get_excluded_models())
|
||||||
excluded_entries: List[Dict[str, Any]] = []
|
excluded_entries: List[Dict[str, Any]] = []
|
||||||
@@ -316,7 +326,7 @@ class BaseModelService(ABC):
|
|||||||
]
|
]
|
||||||
persist_current_cache = getattr(self.scanner, "_persist_current_cache", None)
|
persist_current_cache = getattr(self.scanner, "_persist_current_cache", None)
|
||||||
if callable(persist_current_cache):
|
if callable(persist_current_cache):
|
||||||
await persist_current_cache()
|
await cast(Awaitable[Any], persist_current_cache())
|
||||||
|
|
||||||
excluded_entries = self._sort_entries(excluded_entries, sort_by)
|
excluded_entries = self._sort_entries(excluded_entries, sort_by)
|
||||||
|
|
||||||
@@ -381,6 +391,12 @@ class BaseModelService(ABC):
|
|||||||
(item.get("model_name") or item.get("file_name") or "").lower(),
|
(item.get("model_name") or item.get("file_name") or "").lower(),
|
||||||
item.get("file_path", "").lower(),
|
item.get("file_path", "").lower(),
|
||||||
)
|
)
|
||||||
|
elif key_name == "random":
|
||||||
|
# Seeded random shuffle: same seed -> same order (stable pagination)
|
||||||
|
rng = random.Random(sort_params.seed or "random")
|
||||||
|
result = list(data)
|
||||||
|
rng.shuffle(result)
|
||||||
|
return result
|
||||||
elif key_name == "size":
|
elif key_name == "size":
|
||||||
key_fn = lambda item: (
|
key_fn = lambda item: (
|
||||||
int(item.get("size", 0) or 0),
|
int(item.get("size", 0) or 0),
|
||||||
@@ -428,39 +444,50 @@ class BaseModelService(ABC):
|
|||||||
return entry
|
return entry
|
||||||
|
|
||||||
async def _apply_hash_filters(
|
async def _apply_hash_filters(
|
||||||
self, data: List[Dict], hash_filters: Dict
|
self, data: List[Dict[str, Any]], hash_filters: Dict[str, Any]
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""Apply hash-based filtering"""
|
"""Apply hash-based filtering (SHA256 and AutoV3)."""
|
||||||
|
|
||||||
|
def matches_hash_set(item: Dict[str, Any], hash_set: set[str]) -> bool:
|
||||||
|
"""Check whether an item matches any hash in the set.
|
||||||
|
|
||||||
|
Compares the item's ``sha256`` field and its non-empty ``autov3``
|
||||||
|
field, both case-insensitively.
|
||||||
|
"""
|
||||||
|
if item.get("sha256", "").lower() in hash_set:
|
||||||
|
return True
|
||||||
|
autov3 = item.get("autov3", "")
|
||||||
|
return bool(autov3) and autov3.lower() in hash_set
|
||||||
|
|
||||||
single_hash = hash_filters.get("single_hash")
|
single_hash = hash_filters.get("single_hash")
|
||||||
multiple_hashes = hash_filters.get("multiple_hashes")
|
multiple_hashes = hash_filters.get("multiple_hashes")
|
||||||
|
|
||||||
if single_hash:
|
if single_hash:
|
||||||
# Filter by single hash
|
# Filter by single hash (SHA256 or AutoV3)
|
||||||
single_hash = single_hash.lower()
|
|
||||||
return [
|
return [
|
||||||
item for item in data if item.get("sha256", "").lower() == single_hash
|
item for item in data if matches_hash_set(item, {single_hash.lower()})
|
||||||
]
|
]
|
||||||
elif multiple_hashes:
|
elif multiple_hashes:
|
||||||
# Filter by multiple hashes
|
# Filter by multiple hashes (SHA256 or AutoV3)
|
||||||
hash_set = set(hash.lower() for hash in multiple_hashes)
|
hash_set = {hash.lower() for hash in multiple_hashes}
|
||||||
return [item for item in data if item.get("sha256", "").lower() in hash_set]
|
return [item for item in data if matches_hash_set(item, hash_set)]
|
||||||
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
async def _apply_common_filters(
|
async def _apply_common_filters(
|
||||||
self,
|
self,
|
||||||
data: List[Dict],
|
data: List[Dict[str, Any]],
|
||||||
folder: str = None,
|
folder: str | None = None,
|
||||||
folder_include: list = None,
|
folder_include: list[str] | None = None,
|
||||||
folder_exclude: list = None,
|
folder_exclude: list[str] | None = None,
|
||||||
base_models: list = None,
|
base_models: list[str] | None = None,
|
||||||
model_types: list = None,
|
model_types: list[str] | None = None,
|
||||||
tags: Optional[Dict[str, str]] = None,
|
tags: Optional[Dict[str, str]] = None,
|
||||||
auto_tags: Optional[Dict[str, str]] = None,
|
auto_tags: Optional[Dict[str, str]] = None,
|
||||||
favorites_only: bool = False,
|
favorites_only: bool = False,
|
||||||
search_options: dict = None,
|
search_options: dict[str, Any] | None = None,
|
||||||
tag_logic: str = "any",
|
tag_logic: str = "any",
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""Apply common filters that work across all model types"""
|
"""Apply common filters that work across all model types"""
|
||||||
normalized_options = self.search_strategy.normalize_options(search_options)
|
normalized_options = self.search_strategy.normalize_options(search_options)
|
||||||
criteria = FilterCriteria(
|
criteria = FilterCriteria(
|
||||||
@@ -479,24 +506,24 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
async def _apply_search_filters(
|
async def _apply_search_filters(
|
||||||
self,
|
self,
|
||||||
data: List[Dict],
|
data: List[Dict[str, Any]],
|
||||||
search: str,
|
search: str,
|
||||||
fuzzy_search: bool,
|
fuzzy_search: bool,
|
||||||
search_options: dict,
|
search_options: dict[str, Any] | None,
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""Apply search filtering"""
|
"""Apply search filtering"""
|
||||||
normalized_options = self.search_strategy.normalize_options(search_options)
|
normalized_options = self.search_strategy.normalize_options(search_options)
|
||||||
return self.search_strategy.apply(
|
return self.search_strategy.apply(
|
||||||
data, search, normalized_options, fuzzy_search
|
data, search, normalized_options, fuzzy_search
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _apply_specific_filters(self, data: List[Dict], **kwargs) -> List[Dict]:
|
async def _apply_specific_filters(self, data: List[Dict[str, Any]], **kwargs) -> List[Dict[str, Any]]:
|
||||||
"""Apply model-specific filters - to be overridden by subclasses if needed"""
|
"""Apply model-specific filters - to be overridden by subclasses if needed"""
|
||||||
return data
|
return data
|
||||||
|
|
||||||
async def _apply_credit_required_filter(
|
async def _apply_credit_required_filter(
|
||||||
self, data: List[Dict], credit_required: bool
|
self, data: List[Dict[str, Any]], credit_required: bool
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""Apply credit required filtering based on license_flags.
|
"""Apply credit required filtering based on license_flags.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -526,8 +553,8 @@ class BaseModelService(ABC):
|
|||||||
return filtered_data
|
return filtered_data
|
||||||
|
|
||||||
async def _apply_allow_selling_filter(
|
async def _apply_allow_selling_filter(
|
||||||
self, data: List[Dict], allow_selling: bool
|
self, data: List[Dict[str, Any]], allow_selling: bool
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""Apply allow selling generated content filtering based on license_flags.
|
"""Apply allow selling generated content filtering based on license_flags.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -559,8 +586,8 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
async def _annotate_update_flags(
|
async def _annotate_update_flags(
|
||||||
self,
|
self,
|
||||||
items: List[Dict],
|
items: List[Dict[str, Any]],
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""Attach an update_available flag to each response item.
|
"""Attach an update_available flag to each response item.
|
||||||
|
|
||||||
Items without a civitai model id default to False.
|
Items without a civitai model id default to False.
|
||||||
@@ -575,7 +602,7 @@ class BaseModelService(ABC):
|
|||||||
item["update_available"] = False
|
item["update_available"] = False
|
||||||
return annotated
|
return annotated
|
||||||
|
|
||||||
id_to_items: Dict[int, List[Dict]] = {}
|
id_to_items: Dict[int, List[Dict[str, Any]]] = {}
|
||||||
ordered_ids: List[int] = []
|
ordered_ids: List[int] = []
|
||||||
for item in annotated:
|
for item in annotated:
|
||||||
model_id = self._extract_model_id(item)
|
model_id = self._extract_model_id(item)
|
||||||
@@ -606,15 +633,25 @@ class BaseModelService(ABC):
|
|||||||
except Exception:
|
except Exception:
|
||||||
hide_early_access = False
|
hide_early_access = False
|
||||||
|
|
||||||
|
# Check user setting for hiding permanent paid updates
|
||||||
|
hide_paid = False
|
||||||
|
try:
|
||||||
|
hide_paid = bool(self.settings.get("hide_paid_updates", False))
|
||||||
|
except Exception:
|
||||||
|
hide_paid = False
|
||||||
|
|
||||||
records = None
|
records = None
|
||||||
resolved: Optional[Dict[int, bool]] = None
|
resolved: Optional[Dict[int, bool]] = None
|
||||||
if same_base_mode:
|
if same_base_mode:
|
||||||
record_method = getattr(self.update_service, "get_records_bulk", None)
|
record_method = getattr(self.update_service, "get_records_bulk", None)
|
||||||
if callable(record_method):
|
if callable(record_method):
|
||||||
try:
|
try:
|
||||||
records = await record_method(self.model_type, ordered_ids)
|
records = await cast(Awaitable[Any], record_method(self.model_type, ordered_ids))
|
||||||
resolved = {
|
resolved = {
|
||||||
model_id: record.has_update(hide_early_access=hide_early_access)
|
model_id: record.has_update(
|
||||||
|
hide_early_access=hide_early_access,
|
||||||
|
hide_paid=hide_paid,
|
||||||
|
)
|
||||||
for model_id, record in records.items()
|
for model_id, record in records.items()
|
||||||
}
|
}
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -632,11 +669,12 @@ class BaseModelService(ABC):
|
|||||||
bulk_method = getattr(self.update_service, "has_updates_bulk", None)
|
bulk_method = getattr(self.update_service, "has_updates_bulk", None)
|
||||||
if callable(bulk_method):
|
if callable(bulk_method):
|
||||||
try:
|
try:
|
||||||
resolved = await bulk_method(
|
resolved = await cast(Awaitable[Any], bulk_method(
|
||||||
self.model_type,
|
self.model_type,
|
||||||
ordered_ids,
|
ordered_ids,
|
||||||
hide_early_access=hide_early_access,
|
hide_early_access=hide_early_access,
|
||||||
)
|
hide_paid=hide_paid,
|
||||||
|
))
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error(
|
logger.error(
|
||||||
"Failed to resolve update status in bulk for %s models (%s): %s",
|
"Failed to resolve update status in bulk for %s models (%s): %s",
|
||||||
@@ -650,7 +688,10 @@ class BaseModelService(ABC):
|
|||||||
if resolved is None:
|
if resolved is None:
|
||||||
tasks = [
|
tasks = [
|
||||||
self.update_service.has_update(
|
self.update_service.has_update(
|
||||||
self.model_type, model_id, hide_early_access=hide_early_access
|
self.model_type,
|
||||||
|
model_id,
|
||||||
|
hide_early_access=hide_early_access,
|
||||||
|
hide_paid=hide_paid,
|
||||||
)
|
)
|
||||||
for model_id in ordered_ids
|
for model_id in ordered_ids
|
||||||
]
|
]
|
||||||
@@ -690,6 +731,7 @@ class BaseModelService(ABC):
|
|||||||
threshold_version,
|
threshold_version,
|
||||||
base_model,
|
base_model,
|
||||||
hide_early_access=hide_early_access,
|
hide_early_access=hide_early_access,
|
||||||
|
hide_paid=hide_paid,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
flag = default_flag
|
flag = default_flag
|
||||||
@@ -698,7 +740,34 @@ class BaseModelService(ABC):
|
|||||||
return annotated
|
return annotated
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _extract_model_id(item: Dict) -> Optional[int]:
|
def _extract_hf_group_key(item: Dict[str, Any]) -> Optional[str]:
|
||||||
|
"""Extract `hf:{owner}/{repo}` from item's ``hf_url``, or None."""
|
||||||
|
hf_url = item.get("hf_url") if isinstance(item, dict) else None
|
||||||
|
if not hf_url or not isinstance(hf_url, str):
|
||||||
|
return None
|
||||||
|
m = re.match(
|
||||||
|
r"https?://huggingface\.co/([^/]+/[^/]+)", hf_url.strip()
|
||||||
|
)
|
||||||
|
if not m:
|
||||||
|
return None
|
||||||
|
return f"hf:{m.group(1)}"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_group_key(item: Dict[str, Any]) -> Union[int, str, None]:
|
||||||
|
"""Return the group identity key: CivitAI modelId (int) or HF repo (str).
|
||||||
|
|
||||||
|
Preference order:
|
||||||
|
1. CivitAI ``modelId`` (int)
|
||||||
|
2. HF repo identity ``hf:{owner}/{repo}`` (str)
|
||||||
|
3. ``None`` (no known grouping source)
|
||||||
|
"""
|
||||||
|
mid = BaseModelService._extract_model_id(item)
|
||||||
|
if mid is not None:
|
||||||
|
return mid
|
||||||
|
return BaseModelService._extract_hf_group_key(item)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _extract_model_id(item: Dict[str, Any]) -> Optional[int]:
|
||||||
civitai = item.get("civitai") if isinstance(item, dict) else None
|
civitai = item.get("civitai") if isinstance(item, dict) else None
|
||||||
if not isinstance(civitai, dict):
|
if not isinstance(civitai, dict):
|
||||||
return None
|
return None
|
||||||
@@ -711,7 +780,7 @@ class BaseModelService(ABC):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _extract_version_id(item: Dict) -> Optional[int]:
|
def _extract_version_id(item: Dict[str, Any]) -> Optional[int]:
|
||||||
civitai = item.get("civitai") if isinstance(item, dict) else None
|
civitai = item.get("civitai") if isinstance(item, dict) else None
|
||||||
if not isinstance(civitai, dict):
|
if not isinstance(civitai, dict):
|
||||||
return None
|
return None
|
||||||
@@ -724,7 +793,7 @@ class BaseModelService(ABC):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _extract_base_model(item: Dict) -> Optional[str]:
|
def _extract_base_model(item: Dict[str, Any]) -> Optional[str]:
|
||||||
value = item.get("base_model")
|
value = item.get("base_model")
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
@@ -776,7 +845,7 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
return highest_by_base
|
return highest_by_base
|
||||||
|
|
||||||
def _paginate(self, data: List[Dict], page: int, page_size: int) -> Dict:
|
def _paginate(self, data: List[Dict[str, Any]], page: int, page_size: int) -> Dict[str, Any]:
|
||||||
"""Apply pagination to filtered data"""
|
"""Apply pagination to filtered data"""
|
||||||
total_items = len(data)
|
total_items = len(data)
|
||||||
start_idx = (page - 1) * page_size
|
start_idx = (page - 1) * page_size
|
||||||
@@ -791,7 +860,7 @@ class BaseModelService(ABC):
|
|||||||
}
|
}
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def format_response(self, model_data: Dict) -> Optional[Dict]:
|
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||||
"""Format model data for API response - must be implemented by subclasses.
|
"""Format model data for API response - must be implemented by subclasses.
|
||||||
|
|
||||||
Subclasses should return None for corrupted entries so the handler
|
Subclasses should return None for corrupted entries so the handler
|
||||||
@@ -800,17 +869,17 @@ class BaseModelService(ABC):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
# Common service methods that delegate to scanner
|
# Common service methods that delegate to scanner
|
||||||
async def get_top_tags(self, limit: int = 20) -> List[Dict]:
|
async def get_top_tags(self, limit: int = 20) -> List[Dict[str, Any]]:
|
||||||
"""Get top tags sorted by frequency"""
|
"""Get top tags sorted by frequency"""
|
||||||
return await self.scanner.get_top_tags(limit)
|
return await self.scanner.get_top_tags(limit)
|
||||||
|
|
||||||
async def search_tags(
|
async def search_tags(
|
||||||
self, query: str, limit: int = 50
|
self, query: str, limit: int = 50
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""Search tags by substring, sorted by frequency"""
|
"""Search tags by substring, sorted by frequency"""
|
||||||
return await self.scanner.search_tags(query, limit)
|
return await self.scanner.search_tags(query, limit)
|
||||||
|
|
||||||
async def get_base_models(self, limit: int = 20) -> List[Dict]:
|
async def get_base_models(self, limit: int = 20) -> List[Dict[str, Any]]:
|
||||||
"""Get base models sorted by frequency"""
|
"""Get base models sorted by frequency"""
|
||||||
return await self.scanner.get_base_models(limit)
|
return await self.scanner.get_base_models(limit)
|
||||||
|
|
||||||
@@ -877,7 +946,7 @@ class BaseModelService(ABC):
|
|||||||
"""Get model root directories"""
|
"""Get model root directories"""
|
||||||
return self.scanner.get_model_roots()
|
return self.scanner.get_model_roots()
|
||||||
|
|
||||||
def filter_civitai_data(self, data: Dict, minimal: bool = False) -> Dict:
|
def filter_civitai_data(self, data: Dict[str, Any], minimal: bool = False) -> Dict[str, Any]:
|
||||||
"""Filter relevant fields from CivitAI data"""
|
"""Filter relevant fields from CivitAI data"""
|
||||||
if not data:
|
if not data:
|
||||||
return {}
|
return {}
|
||||||
@@ -903,7 +972,7 @@ class BaseModelService(ABC):
|
|||||||
)
|
)
|
||||||
return {k: data[k] for k in fields if k in data}
|
return {k: data[k] for k in fields if k in data}
|
||||||
|
|
||||||
async def get_folder_tree(self, model_root: str) -> Dict:
|
async def get_folder_tree(self, model_root: str) -> Dict[str, Any]:
|
||||||
"""Get hierarchical folder tree for a specific model root"""
|
"""Get hierarchical folder tree for a specific model root"""
|
||||||
cache = await self.scanner.get_cached_data()
|
cache = await self.scanner.get_cached_data()
|
||||||
|
|
||||||
@@ -932,7 +1001,7 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
return tree
|
return tree
|
||||||
|
|
||||||
async def get_unified_folder_tree(self) -> Dict:
|
async def get_unified_folder_tree(self) -> Dict[str, Any]:
|
||||||
"""Get unified folder tree across all model roots"""
|
"""Get unified folder tree across all model roots"""
|
||||||
cache = await self.scanner.get_cached_data()
|
cache = await self.scanner.get_cached_data()
|
||||||
|
|
||||||
@@ -961,7 +1030,7 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
return unified_tree
|
return unified_tree
|
||||||
|
|
||||||
async def get_model_notes(self, model_name: str) -> Optional[dict]:
|
async def get_model_notes(self, model_name: str) -> Optional[dict[str, Any]]:
|
||||||
"""Get notes and file_path for a specific model file.
|
"""Get notes and file_path for a specific model file.
|
||||||
|
|
||||||
Supports both simple names (``OWSMianne_ANIMA_V1``) and full-path
|
Supports both simple names (``OWSMianne_ANIMA_V1``) and full-path
|
||||||
@@ -1093,7 +1162,7 @@ class BaseModelService(ABC):
|
|||||||
|
|
||||||
return {"civitai_url": None, "model_id": None, "version_id": None}
|
return {"civitai_url": None, "model_id": None, "version_id": None}
|
||||||
|
|
||||||
async def get_model_metadata(self, file_path: str) -> Optional[Dict]:
|
async def get_model_metadata(self, file_path: str) -> Optional[Dict[str, Any]]:
|
||||||
"""Load full metadata for a single model.
|
"""Load full metadata for a single model.
|
||||||
|
|
||||||
Listing/search endpoints return lightweight cache entries; this method performs
|
Listing/search endpoints return lightweight cache entries; this method performs
|
||||||
@@ -1189,7 +1258,7 @@ class BaseModelService(ABC):
|
|||||||
return True
|
return True
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _relative_path_sort_key(relative_path: str, include_terms: List[str]) -> tuple:
|
def _relative_path_sort_key(relative_path: str, include_terms: List[str]) -> tuple[int, int, int, str]:
|
||||||
"""Sort paths by how well they satisfy the include tokens.
|
"""Sort paths by how well they satisfy the include tokens.
|
||||||
|
|
||||||
Sorts based on path without extension for consistent ordering.
|
Sorts based on path without extension for consistent ordering.
|
||||||
@@ -1216,19 +1285,87 @@ class BaseModelService(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def search_relative_paths(
|
async def search_relative_paths(
|
||||||
self, search_term: str, limit: int = 15, offset: int = 0
|
self,
|
||||||
|
search_term: str,
|
||||||
|
limit: int = 15,
|
||||||
|
offset: int = 0,
|
||||||
|
*,
|
||||||
|
folder: Optional[str] = None,
|
||||||
|
folder_include: Optional[list[str]] = None,
|
||||||
|
folder_exclude: Optional[list[str]] = None,
|
||||||
|
base_models: Optional[list[str]] = None,
|
||||||
|
model_types: Optional[list[str]] = None,
|
||||||
|
tags: Optional[dict[str, str]] = None,
|
||||||
|
auto_tags: Optional[dict[str, str]] = None,
|
||||||
|
tag_logic: str = "any",
|
||||||
|
credit_required: Optional[bool] = None,
|
||||||
|
allow_selling_generated_content: Optional[bool] = None,
|
||||||
|
recursive: bool = True,
|
||||||
|
apply_filters: bool = False,
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
"""Search model relative file paths for autocomplete functionality"""
|
"""Search model relative file paths for autocomplete functionality.
|
||||||
|
|
||||||
|
Optional filter kwargs mirror the filters used by the list endpoint
|
||||||
|
(/api/lm/{prefix}/list). When no filter kwargs are provided the
|
||||||
|
behavior is identical to plain token-based path matching.
|
||||||
|
"""
|
||||||
cache = await self.scanner.get_cached_data()
|
cache = await self.scanner.get_cached_data()
|
||||||
include_terms, exclude_terms = self._parse_search_tokens(search_term)
|
include_terms, exclude_terms = self._parse_search_tokens(search_term)
|
||||||
|
|
||||||
|
data = cache.raw_data
|
||||||
|
has_filters = any(
|
||||||
|
[
|
||||||
|
apply_filters,
|
||||||
|
folder is not None,
|
||||||
|
folder_include,
|
||||||
|
folder_exclude,
|
||||||
|
base_models,
|
||||||
|
model_types,
|
||||||
|
tags,
|
||||||
|
auto_tags,
|
||||||
|
credit_required is not None,
|
||||||
|
allow_selling_generated_content is not None,
|
||||||
|
]
|
||||||
|
)
|
||||||
|
if has_filters:
|
||||||
|
# Auto-tags are not stored in the scanner cache — they are computed
|
||||||
|
# on the fly. Pre-compute them only when an auto-tag filter is
|
||||||
|
# active to avoid mutating cache entries unnecessarily.
|
||||||
|
if auto_tags:
|
||||||
|
from .auto_tag_service import extract_auto_tags
|
||||||
|
|
||||||
|
for item in data:
|
||||||
|
if not item.get("auto_tags"):
|
||||||
|
item["auto_tags"] = extract_auto_tags(item)
|
||||||
|
|
||||||
|
criteria = FilterCriteria(
|
||||||
|
folder=folder,
|
||||||
|
folder_include=folder_include,
|
||||||
|
folder_exclude=folder_exclude,
|
||||||
|
base_models=base_models,
|
||||||
|
model_types=model_types,
|
||||||
|
tags=tags,
|
||||||
|
auto_tags=auto_tags,
|
||||||
|
search_options={"recursive": recursive},
|
||||||
|
tag_logic=tag_logic,
|
||||||
|
)
|
||||||
|
data = self.filter_set.apply(data, criteria)
|
||||||
|
if credit_required is not None:
|
||||||
|
data = await self._apply_credit_required_filter(
|
||||||
|
data, credit_required
|
||||||
|
)
|
||||||
|
if allow_selling_generated_content is not None:
|
||||||
|
data = await self._apply_allow_selling_filter(
|
||||||
|
data, allow_selling_generated_content
|
||||||
|
)
|
||||||
|
|
||||||
matching_paths = []
|
matching_paths = []
|
||||||
|
|
||||||
# Get model roots for path calculation
|
# Get model roots for path calculation
|
||||||
model_roots = self.scanner.get_model_roots()
|
model_roots = self.scanner.get_model_roots()
|
||||||
|
|
||||||
# Collect all matching paths first (needed for proper sorting and offset)
|
# Collect all matching paths first (needed for proper sorting and offset)
|
||||||
for model in cache.raw_data:
|
for model in data:
|
||||||
file_path = model.get("file_path", "")
|
file_path = model.get("file_path", "")
|
||||||
if not file_path:
|
if not file_path:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -59,6 +59,7 @@ class CacheEntryValidator:
|
|||||||
'notes': ('', False),
|
'notes': ('', False),
|
||||||
'usage_tips': ('', False),
|
'usage_tips': ('', False),
|
||||||
'hash_status': ('completed', False),
|
'hash_status': ('completed', False),
|
||||||
|
'autov3': (None, False),
|
||||||
}
|
}
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -119,8 +120,13 @@ class CacheEntryValidator:
|
|||||||
if is_required:
|
if is_required:
|
||||||
errors.append(f"Required field '{field_name}' is missing or None")
|
errors.append(f"Required field '{field_name}' is missing or None")
|
||||||
if auto_repair:
|
if auto_repair:
|
||||||
working_entry[field_name] = cls._get_default_copy(default_value)
|
# A missing optional field whose default is None is already
|
||||||
repaired = True
|
# semantically equal to its default (e.g. autov3: absent
|
||||||
|
# means "not checked") — writing None back is a no-op, not
|
||||||
|
# a repair.
|
||||||
|
if default_value is not None:
|
||||||
|
working_entry[field_name] = cls._get_default_copy(default_value)
|
||||||
|
repaired = True
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Validate field type and value
|
# Validate field type and value
|
||||||
@@ -175,6 +181,15 @@ class CacheEntryValidator:
|
|||||||
# that invalidates the entry, but we also don't mark it repaired.
|
# that invalidates the entry, but we also don't mark it repaired.
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
# Normalize autov3 to lowercase if needed (optional field, never stripped).
|
||||||
|
autov3 = working_entry.get('autov3')
|
||||||
|
if isinstance(autov3, str) and autov3:
|
||||||
|
normalized_autov3 = autov3.lower()
|
||||||
|
if normalized_autov3 != autov3:
|
||||||
|
if auto_repair:
|
||||||
|
working_entry['autov3'] = normalized_autov3
|
||||||
|
repaired = True
|
||||||
|
|
||||||
# Determine if entry is valid
|
# Determine if entry is valid
|
||||||
# Entry is valid if no critical required field errors remain after repair
|
# Entry is valid if no critical required field errors remain after repair
|
||||||
# Critical fields are file_path and sha256
|
# Critical fields are file_path and sha256
|
||||||
@@ -242,6 +257,19 @@ class CacheEntryValidator:
|
|||||||
"""
|
"""
|
||||||
expected_type = type(default_value)
|
expected_type = type(default_value)
|
||||||
|
|
||||||
|
# Special case: autov3 is optional with a three-state contract.
|
||||||
|
# None = not checked, "" = checked but unavailable, otherwise a
|
||||||
|
# 12-character hex string (case-insensitive here; normalized to
|
||||||
|
# lowercase separately).
|
||||||
|
if field_name == 'autov3':
|
||||||
|
if value is None or value == "":
|
||||||
|
return None
|
||||||
|
if not isinstance(value, str):
|
||||||
|
return f"Field 'autov3' should be string or None, got {type(value).__name__}"
|
||||||
|
if len(value) != 12 or any(c not in '0123456789abcdefABCDEF' for c in value):
|
||||||
|
return "Field 'autov3' should be a 12-character hex string"
|
||||||
|
return None
|
||||||
|
|
||||||
# Special handling for numeric types
|
# Special handling for numeric types
|
||||||
if expected_type == int:
|
if expected_type == int:
|
||||||
if not isinstance(value, (int, float)):
|
if not isinstance(value, (int, float)):
|
||||||
|
|||||||
@@ -1,3 +1,7 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
@@ -6,10 +10,10 @@ from datetime import datetime
|
|||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
from ..utils.models import CheckpointMetadata
|
from ..utils.models import CheckpointMetadata
|
||||||
from ..utils.file_utils import find_preview_file, normalize_path
|
from ..utils.file_utils import find_preview_file, normalize_path, calculate_autov3
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from ..config import config
|
from ..config import config
|
||||||
from .model_scanner import ModelScanner
|
from .model_scanner import ModelScanner, _is_excluded_dir
|
||||||
from .model_hash_index import ModelHashIndex
|
from .model_hash_index import ModelHashIndex
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -62,6 +66,11 @@ class CheckpointScanner(ModelScanner):
|
|||||||
# Find preview image
|
# Find preview image
|
||||||
preview_url = find_preview_file(base_name, dir_path)
|
preview_url = find_preview_file(base_name, dir_path)
|
||||||
|
|
||||||
|
# AutoV3 reads only the safetensors header, so it is cheap even for
|
||||||
|
# large checkpoints; record the checked state at creation time ("" =
|
||||||
|
# checked but unavailable).
|
||||||
|
autov3 = calculate_autov3(real_path)
|
||||||
|
|
||||||
# Create metadata WITHOUT calculating hash
|
# Create metadata WITHOUT calculating hash
|
||||||
metadata = CheckpointMetadata(
|
metadata = CheckpointMetadata(
|
||||||
file_name=base_name,
|
file_name=base_name,
|
||||||
@@ -77,6 +86,7 @@ class CheckpointScanner(ModelScanner):
|
|||||||
sub_type="checkpoint",
|
sub_type="checkpoint",
|
||||||
from_civitai=False, # Mark as local model since no hash yet
|
from_civitai=False, # Mark as local model since no hash yet
|
||||||
hash_status="pending", # Mark hash as pending
|
hash_status="pending", # Mark hash as pending
|
||||||
|
autov3=autov3 or "",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Save the created metadata
|
# Save the created metadata
|
||||||
@@ -120,7 +130,11 @@ class CheckpointScanner(ModelScanner):
|
|||||||
# that queries get_hash_by_filename first) will miss on every
|
# that queries get_hash_by_filename first) will miss on every
|
||||||
# lookup and keep calling back into this method, creating a
|
# lookup and keep calling back into this method, creating a
|
||||||
# tight loop that never populates the index.
|
# tight loop that never populates the index.
|
||||||
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
|
self._hash_index.add_entry(
|
||||||
|
metadata.sha256.lower(),
|
||||||
|
file_path,
|
||||||
|
getattr(metadata, "autov3", None) or None,
|
||||||
|
)
|
||||||
return metadata.sha256
|
return metadata.sha256
|
||||||
|
|
||||||
async with self._hash_calculation_lock:
|
async with self._hash_calculation_lock:
|
||||||
@@ -132,7 +146,11 @@ class CheckpointScanner(ModelScanner):
|
|||||||
and metadata.hash_status == "completed"
|
and metadata.hash_status == "completed"
|
||||||
and metadata.sha256
|
and metadata.sha256
|
||||||
):
|
):
|
||||||
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
|
self._hash_index.add_entry(
|
||||||
|
metadata.sha256.lower(),
|
||||||
|
file_path,
|
||||||
|
getattr(metadata, "autov3", None) or None,
|
||||||
|
)
|
||||||
return metadata.sha256
|
return metadata.sha256
|
||||||
|
|
||||||
task = self._hash_calculation_tasks.get(real_path)
|
task = self._hash_calculation_tasks.get(real_path)
|
||||||
@@ -185,7 +203,11 @@ class CheckpointScanner(ModelScanner):
|
|||||||
if metadata.hash_status == "completed" and metadata.sha256:
|
if metadata.hash_status == "completed" and metadata.sha256:
|
||||||
# Populate the in-memory hash index even for pre-computed
|
# Populate the in-memory hash index even for pre-computed
|
||||||
# hashes, mirroring the fix in calculate_hash_for_model.
|
# hashes, mirroring the fix in calculate_hash_for_model.
|
||||||
self._hash_index.add_entry(metadata.sha256.lower(), file_path)
|
self._hash_index.add_entry(
|
||||||
|
metadata.sha256.lower(),
|
||||||
|
file_path,
|
||||||
|
getattr(metadata, "autov3", None) or None,
|
||||||
|
)
|
||||||
return metadata.sha256
|
return metadata.sha256
|
||||||
|
|
||||||
# Update status to calculating
|
# Update status to calculating
|
||||||
@@ -202,7 +224,11 @@ class CheckpointScanner(ModelScanner):
|
|||||||
await MetadataManager.save_metadata(file_path, metadata)
|
await MetadataManager.save_metadata(file_path, metadata)
|
||||||
|
|
||||||
# Update hash index
|
# Update hash index
|
||||||
self._hash_index.add_entry(sha256.lower(), file_path)
|
self._hash_index.add_entry(
|
||||||
|
sha256.lower(),
|
||||||
|
file_path,
|
||||||
|
getattr(metadata, "autov3", None) or None,
|
||||||
|
)
|
||||||
|
|
||||||
# Update the in-memory cache entry so that subsequent
|
# Update the in-memory cache entry so that subsequent
|
||||||
# _persist_current_cache / _save_persistent_cache calls
|
# _persist_current_cache / _save_persistent_cache calls
|
||||||
@@ -216,6 +242,7 @@ class CheckpointScanner(ModelScanner):
|
|||||||
if entry.get("file_path") == file_path:
|
if entry.get("file_path") == file_path:
|
||||||
entry["sha256"] = sha256.lower()
|
entry["sha256"] = sha256.lower()
|
||||||
entry["hash_status"] = "completed"
|
entry["hash_status"] = "completed"
|
||||||
|
self.bump_cache_version()
|
||||||
break
|
break
|
||||||
|
|
||||||
logger.info(f"Hash calculated for checkpoint: {file_path}")
|
logger.info(f"Hash calculated for checkpoint: {file_path}")
|
||||||
@@ -301,7 +328,8 @@ class CheckpointScanner(ModelScanner):
|
|||||||
if not os.path.exists(root_path):
|
if not os.path.exists(root_path):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
for dirpath, _dirnames, filenames in os.walk(root_path):
|
for dirpath, dirnames, filenames in os.walk(root_path):
|
||||||
|
dirnames[:] = [d for d in dirnames if not _is_excluded_dir(d)]
|
||||||
for filename in filenames:
|
for filename in filenames:
|
||||||
if not filename.endswith(".metadata.json"):
|
if not filename.endswith(".metadata.json"):
|
||||||
continue
|
continue
|
||||||
@@ -405,7 +433,7 @@ class CheckpointScanner(ModelScanner):
|
|||||||
roots.extend(config.extra_checkpoints_roots or [])
|
roots.extend(config.extra_checkpoints_roots or [])
|
||||||
roots.extend(config.extra_unet_roots or [])
|
roots.extend(config.extra_unet_roots or [])
|
||||||
# Remove duplicates while preserving order
|
# Remove duplicates while preserving order
|
||||||
seen: set = set()
|
seen: set[str] = set()
|
||||||
unique_roots: List[str] = []
|
unique_roots: List[str] = []
|
||||||
for root in roots:
|
for root in roots:
|
||||||
if root not in seen:
|
if root not in seen:
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
import os
|
import os
|
||||||
import logging
|
import logging
|
||||||
from typing import Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
from .base_model_service import BaseModelService
|
from .base_model_service import BaseModelService
|
||||||
from .auto_tag_service import extract_auto_tags
|
from .auto_tag_service import extract_auto_tags
|
||||||
@@ -21,58 +21,58 @@ class CheckpointService(BaseModelService):
|
|||||||
"""
|
"""
|
||||||
super().__init__("checkpoint", scanner, CheckpointMetadata, update_service=update_service)
|
super().__init__("checkpoint", scanner, CheckpointMetadata, update_service=update_service)
|
||||||
|
|
||||||
async def format_response(self, checkpoint_data: Dict) -> Optional[Dict]:
|
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||||
"""Format Checkpoint data for API response.
|
"""Format Checkpoint data for API response.
|
||||||
|
|
||||||
Returns None when the entry is missing critical fields (corrupted cache
|
Returns None when the entry is missing critical fields (corrupted cache
|
||||||
row), so the handler layer can filter it out. See issue #730.
|
row), so the handler layer can filter it out. See issue #730.
|
||||||
"""
|
"""
|
||||||
# Guard against corrupted cache entries missing critical fields
|
# Guard against corrupted cache entries missing critical fields
|
||||||
file_path = checkpoint_data.get("file_path")
|
file_path = model_data.get("file_path")
|
||||||
if not file_path or not isinstance(file_path, str):
|
if not file_path or not isinstance(file_path, str):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Skipping corrupted checkpoint entry (missing file_path): %s",
|
"Skipping corrupted checkpoint entry (missing file_path): %s",
|
||||||
checkpoint_data.get("file_name", "<unknown>"),
|
model_data.get("file_name", "<unknown>"),
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Get sub_type from cache entry (new canonical field)
|
# Get sub_type from cache entry (new canonical field)
|
||||||
sub_type = checkpoint_data.get("sub_type", "checkpoint")
|
sub_type = model_data.get("sub_type", "checkpoint")
|
||||||
|
|
||||||
file_name = checkpoint_data.get("file_name") or ""
|
file_name = model_data.get("file_name") or ""
|
||||||
model_name = checkpoint_data.get("model_name") or file_name
|
model_name = model_data.get("model_name") or file_name
|
||||||
folder = checkpoint_data.get("folder") or ""
|
folder = model_data.get("folder") or ""
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"model_name": model_name,
|
"model_name": model_name,
|
||||||
"file_name": file_name,
|
"file_name": file_name,
|
||||||
"preview_url": config.get_preview_static_url(checkpoint_data.get("preview_url", "")),
|
"preview_url": config.get_preview_static_url(model_data.get("preview_url", "")),
|
||||||
"preview_nsfw_level": checkpoint_data.get("preview_nsfw_level", 0),
|
"preview_nsfw_level": model_data.get("preview_nsfw_level", 0),
|
||||||
"base_model": checkpoint_data.get("base_model", ""),
|
"base_model": model_data.get("base_model", ""),
|
||||||
"folder": folder,
|
"folder": folder,
|
||||||
"sha256": checkpoint_data.get("sha256", ""),
|
"sha256": model_data.get("sha256", ""),
|
||||||
"file_path": file_path.replace(os.sep, "/"),
|
"file_path": file_path.replace(os.sep, "/"),
|
||||||
"file_size": checkpoint_data.get("size", 0),
|
"file_size": model_data.get("size", 0),
|
||||||
"modified": checkpoint_data.get("modified", ""),
|
"modified": model_data.get("modified", ""),
|
||||||
"tags": checkpoint_data.get("tags", []),
|
"tags": model_data.get("tags", []),
|
||||||
"from_civitai": checkpoint_data.get("from_civitai", True),
|
"from_civitai": model_data.get("from_civitai", True),
|
||||||
"usage_count": checkpoint_data.get("usage_count", 0),
|
"usage_count": model_data.get("usage_count", 0),
|
||||||
"notes": checkpoint_data.get("notes", ""),
|
"notes": model_data.get("notes", ""),
|
||||||
"sub_type": sub_type,
|
"sub_type": sub_type,
|
||||||
"favorite": checkpoint_data.get("favorite", False),
|
"favorite": model_data.get("favorite", False),
|
||||||
"exclude": bool(checkpoint_data.get("exclude", False)),
|
"exclude": bool(model_data.get("exclude", False)),
|
||||||
"update_available": bool(checkpoint_data.get("update_available", False)),
|
"update_available": bool(model_data.get("update_available", False)),
|
||||||
"skip_metadata_refresh": bool(checkpoint_data.get("skip_metadata_refresh", False)),
|
"skip_metadata_refresh": bool(model_data.get("skip_metadata_refresh", False)),
|
||||||
"civitai": self.filter_civitai_data(checkpoint_data.get("civitai", {}), minimal=True),
|
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
|
||||||
"auto_tags": checkpoint_data.get("auto_tags") or extract_auto_tags(checkpoint_data),
|
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
||||||
"version_count": checkpoint_data.get("version_count"),
|
"version_count": model_data.get("version_count"),
|
||||||
"hf_url": checkpoint_data.get("hf_url", ""),
|
"hf_url": model_data.get("hf_url", ""),
|
||||||
}
|
}
|
||||||
|
|
||||||
def find_duplicate_hashes(self) -> Dict:
|
def find_duplicate_hashes(self) -> Dict[str, Any]:
|
||||||
"""Find Checkpoints with duplicate SHA256 hashes"""
|
"""Find Checkpoints with duplicate SHA256 hashes"""
|
||||||
return self.scanner._hash_index.get_duplicate_hashes()
|
return self.scanner._hash_index.get_duplicate_hashes()
|
||||||
|
|
||||||
def find_duplicate_filenames(self) -> Dict:
|
def find_duplicate_filenames(self) -> Dict[str, Any]:
|
||||||
"""Find Checkpoints with conflicting filenames"""
|
"""Find Checkpoints with conflicting filenames"""
|
||||||
return self.scanner._hash_index.get_duplicate_filenames()
|
return self.scanner._hash_index.get_duplicate_filenames()
|
||||||
|
|||||||
@@ -1,8 +1,12 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import asyncio
|
import asyncio
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from typing import Optional, Dict, Tuple, List
|
from typing import Any, Optional, Dict, Tuple, List, cast
|
||||||
from .model_metadata_provider import CivArchiveModelMetadataProvider, ModelMetadataProviderManager
|
from .model_metadata_provider import CivArchiveModelMetadataProvider, ModelMetadataProviderManager
|
||||||
from .downloader import get_downloader
|
from .downloader import get_downloader
|
||||||
from .errors import RateLimitError
|
from .errors import RateLimitError
|
||||||
@@ -37,8 +41,8 @@ class CivArchiveClient:
|
|||||||
async def _request_json(
|
async def _request_json(
|
||||||
self,
|
self,
|
||||||
path: str,
|
path: str,
|
||||||
params: Optional[Dict[str, str]] = None
|
params: Optional[Dict[str, Any]] = None
|
||||||
) -> Tuple[Optional[Dict], Optional[str]]:
|
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||||
"""Call CivArchive API and return JSON payload"""
|
"""Call CivArchive API and return JSON payload"""
|
||||||
success, payload = await self._make_request(path, params=params)
|
success, payload = await self._make_request(path, params=params)
|
||||||
if not success:
|
if not success:
|
||||||
@@ -52,12 +56,12 @@ class CivArchiveClient:
|
|||||||
self,
|
self,
|
||||||
path: str,
|
path: str,
|
||||||
*,
|
*,
|
||||||
params: Optional[Dict[str, str]] = None,
|
params: Optional[Dict[str, Any]] = None,
|
||||||
) -> Tuple[bool, Dict | str]:
|
) -> Tuple[bool, Dict[str, Any] | str]:
|
||||||
"""Wrapper around downloader.make_request that surfaces rate limits."""
|
"""Wrapper around downloader.make_request that surfaces rate limits."""
|
||||||
|
|
||||||
downloader = await get_downloader()
|
downloader = await get_downloader()
|
||||||
kwargs: Dict[str, Dict[str, str]] = {}
|
kwargs: Dict[str, Dict[str, Any]] = {}
|
||||||
if params:
|
if params:
|
||||||
safe_params = {str(key): str(value) for key, value in params.items() if value is not None}
|
safe_params = {str(key): str(value) for key, value in params.items() if value is not None}
|
||||||
if safe_params:
|
if safe_params:
|
||||||
@@ -73,10 +77,11 @@ class CivArchiveClient:
|
|||||||
if payload.provider is None:
|
if payload.provider is None:
|
||||||
payload.provider = "civarchive_api"
|
payload.provider = "civarchive_api"
|
||||||
raise payload
|
raise payload
|
||||||
return success, payload
|
# RateLimitError is always raised above, so the returned payload is a dict or str.
|
||||||
|
return success, cast(Dict[str, Any] | str, payload)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _normalize_payload(payload: Dict) -> Dict:
|
def _normalize_payload(payload: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
"""Unwrap CivArchive responses that wrap content under a data key"""
|
"""Unwrap CivArchive responses that wrap content under a data key"""
|
||||||
if not isinstance(payload, dict):
|
if not isinstance(payload, dict):
|
||||||
return {}
|
return {}
|
||||||
@@ -86,12 +91,12 @@ class CivArchiveClient:
|
|||||||
return payload
|
return payload
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _split_context(payload: Dict) -> Tuple[Dict, Dict, List[Dict]]:
|
def _split_context(payload: Dict[str, Any]) -> Tuple[Dict[str, Any], Dict[str, Any], List[Dict[str, Any]]]:
|
||||||
"""Separate version payload from surrounding model context"""
|
"""Separate version payload from surrounding model context"""
|
||||||
data = CivArchiveClient._normalize_payload(payload)
|
data = CivArchiveClient._normalize_payload(payload)
|
||||||
context: Dict = {}
|
context: Dict[str, Any] = {}
|
||||||
fallback_files: List[Dict] = []
|
fallback_files: List[Dict[str, Any]] = []
|
||||||
version: Dict = {}
|
version: Dict[str, Any] = {}
|
||||||
|
|
||||||
for key, value in data.items():
|
for key, value in data.items():
|
||||||
if key in {"version", "model"}:
|
if key in {"version", "model"}:
|
||||||
@@ -115,7 +120,7 @@ class CivArchiveClient:
|
|||||||
return context, version, fallback_files
|
return context, version, fallback_files
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _ensure_list(value) -> List:
|
def _ensure_list(value: Any) -> List[Any]:
|
||||||
if isinstance(value, list):
|
if isinstance(value, list):
|
||||||
return value
|
return value
|
||||||
if value is None:
|
if value is None:
|
||||||
@@ -123,7 +128,7 @@ class CivArchiveClient:
|
|||||||
return [value]
|
return [value]
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _build_model_info(context: Dict) -> Dict:
|
def _build_model_info(context: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
tags = context.get("tags")
|
tags = context.get("tags")
|
||||||
if not isinstance(tags, list):
|
if not isinstance(tags, list):
|
||||||
tags = list(tags) if isinstance(tags, (set, tuple)) else ([] if tags is None else [tags])
|
tags = list(tags) if isinstance(tags, (set, tuple)) else ([] if tags is None else [tags])
|
||||||
@@ -136,7 +141,7 @@ class CivArchiveClient:
|
|||||||
}
|
}
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _build_creator_info(context: Dict) -> Dict:
|
def _build_creator_info(context: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
username = context.get("creator_username") or context.get("username") or ""
|
username = context.get("creator_username") or context.get("username") or ""
|
||||||
image = context.get("creator_image") or context.get("creator_avatar") or ""
|
image = context.get("creator_image") or context.get("creator_avatar") or ""
|
||||||
creator: Dict[str, Optional[str]] = {
|
creator: Dict[str, Optional[str]] = {
|
||||||
@@ -150,7 +155,7 @@ class CivArchiveClient:
|
|||||||
return creator
|
return creator
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _transform_file_entry(file_data: Dict) -> Dict:
|
def _transform_file_entry(file_data: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
mirrors = file_data.get("mirrors") or []
|
mirrors = file_data.get("mirrors") or []
|
||||||
if not isinstance(mirrors, list):
|
if not isinstance(mirrors, list):
|
||||||
mirrors = [mirrors]
|
mirrors = [mirrors]
|
||||||
@@ -165,7 +170,7 @@ class CivArchiveClient:
|
|||||||
if not name and available_mirror:
|
if not name and available_mirror:
|
||||||
name = available_mirror.get("filename")
|
name = available_mirror.get("filename")
|
||||||
|
|
||||||
transformed: Dict = {
|
transformed: Dict[str, Any] = {
|
||||||
"id": file_data.get("id"),
|
"id": file_data.get("id"),
|
||||||
"sizeKB": file_data.get("sizeKB"),
|
"sizeKB": file_data.get("sizeKB"),
|
||||||
"name": name,
|
"name": name,
|
||||||
@@ -216,23 +221,23 @@ class CivArchiveClient:
|
|||||||
|
|
||||||
def _transform_files(
|
def _transform_files(
|
||||||
self,
|
self,
|
||||||
files: Optional[List[Dict]],
|
files: Optional[List[Dict[str, Any]]],
|
||||||
fallback_files: Optional[List[Dict]] = None
|
fallback_files: Optional[List[Dict[str, Any]]] = None
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
candidates: List[Dict] = []
|
candidates: List[Dict[str, Any]] = []
|
||||||
if isinstance(files, list) and files:
|
if isinstance(files, list) and files:
|
||||||
candidates = files
|
candidates = files
|
||||||
elif isinstance(fallback_files, list):
|
elif isinstance(fallback_files, list):
|
||||||
candidates = fallback_files
|
candidates = fallback_files
|
||||||
|
|
||||||
transformed_files: List[Dict] = []
|
transformed_files: List[Dict[str, Any]] = []
|
||||||
for file_data in candidates:
|
for file_data in candidates:
|
||||||
if isinstance(file_data, dict):
|
if isinstance(file_data, dict):
|
||||||
transformed_files.append(self._transform_file_entry(file_data))
|
transformed_files.append(self._transform_file_entry(file_data))
|
||||||
|
|
||||||
# Sort: .safetensors first, .ckpt second, others last
|
# Sort: .safetensors first, .ckpt second, others last
|
||||||
# so the backend fallback (no file_params) prefers safetensors
|
# so the backend fallback (no file_params) prefers safetensors
|
||||||
def _sort_key(f: Dict) -> int:
|
def _sort_key(f: Dict[str, Any]) -> int:
|
||||||
fname = f.get("name") or ""
|
fname = f.get("name") or ""
|
||||||
if isinstance(fname, str):
|
if isinstance(fname, str):
|
||||||
lower = fname.lower()
|
lower = fname.lower()
|
||||||
@@ -247,10 +252,10 @@ class CivArchiveClient:
|
|||||||
|
|
||||||
def _transform_version(
|
def _transform_version(
|
||||||
self,
|
self,
|
||||||
context: Dict,
|
context: Dict[str, Any],
|
||||||
version: Dict,
|
version: Dict[str, Any],
|
||||||
fallback_files: Optional[List[Dict]] = None
|
fallback_files: Optional[List[Dict[str, Any]]] = None
|
||||||
) -> Optional[Dict]:
|
) -> Optional[Dict[str, Any]]:
|
||||||
if not version:
|
if not version:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -291,7 +296,7 @@ class CivArchiveClient:
|
|||||||
|
|
||||||
return version_copy
|
return version_copy
|
||||||
|
|
||||||
async def _resolve_version_from_files(self, payload: Dict) -> Optional[Dict]:
|
async def _resolve_version_from_files(self, payload: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||||
"""Fallback to fetch version data when only file metadata is available"""
|
"""Fallback to fetch version data when only file metadata is available"""
|
||||||
data = self._normalize_payload(payload)
|
data = self._normalize_payload(payload)
|
||||||
files = data.get("files") or payload.get("files") or []
|
files = data.get("files") or payload.get("files") or []
|
||||||
@@ -323,7 +328,7 @@ class CivArchiveClient:
|
|||||||
return resolved
|
return resolved
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict], Optional[str]]:
|
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||||
"""Find model by SHA256 hash value using CivArchive API"""
|
"""Find model by SHA256 hash value using CivArchive API"""
|
||||||
try:
|
try:
|
||||||
payload, error = await self._request_json(f"/sha256/{model_hash.lower()}")
|
payload, error = await self._request_json(f"/sha256/{model_hash.lower()}")
|
||||||
@@ -332,12 +337,12 @@ class CivArchiveClient:
|
|||||||
return None, "Model not found"
|
return None, "Model not found"
|
||||||
return None, error
|
return None, error
|
||||||
|
|
||||||
context, version_data, fallback_files = self._split_context(payload)
|
context, version_data, fallback_files = self._split_context(cast(Dict[str, Any], payload))
|
||||||
transformed = self._transform_version(context, version_data, fallback_files)
|
transformed = self._transform_version(context, version_data, fallback_files)
|
||||||
if transformed:
|
if transformed:
|
||||||
return transformed, None
|
return transformed, None
|
||||||
|
|
||||||
resolved = await self._resolve_version_from_files(payload)
|
resolved = await self._resolve_version_from_files(cast(Dict[str, Any], payload))
|
||||||
if resolved:
|
if resolved:
|
||||||
return resolved, None
|
return resolved, None
|
||||||
|
|
||||||
@@ -350,7 +355,7 @@ class CivArchiveClient:
|
|||||||
logger.error(f"Error fetching CivArchive model by hash {model_hash[:10]}: {e}")
|
logger.error(f"Error fetching CivArchive model by hash {model_hash[:10]}: {e}")
|
||||||
return None, str(e)
|
return None, str(e)
|
||||||
|
|
||||||
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
|
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
|
||||||
"""Get all versions of a model using CivArchive API"""
|
"""Get all versions of a model using CivArchive API"""
|
||||||
try:
|
try:
|
||||||
payload, error = await self._request_json(f"/models/{model_id}")
|
payload, error = await self._request_json(f"/models/{model_id}")
|
||||||
@@ -364,7 +369,7 @@ class CivArchiveClient:
|
|||||||
context, version_data, fallback_files = self._split_context(payload)
|
context, version_data, fallback_files = self._split_context(payload)
|
||||||
|
|
||||||
versions_meta = data.get("versions") or []
|
versions_meta = data.get("versions") or []
|
||||||
transformed_versions: List[Dict] = []
|
transformed_versions: List[Dict[str, Any]] = []
|
||||||
for meta in versions_meta:
|
for meta in versions_meta:
|
||||||
if not isinstance(meta, dict):
|
if not isinstance(meta, dict):
|
||||||
continue
|
continue
|
||||||
@@ -381,7 +386,7 @@ class CivArchiveClient:
|
|||||||
if primary_version:
|
if primary_version:
|
||||||
transformed_versions.insert(0, primary_version)
|
transformed_versions.insert(0, primary_version)
|
||||||
|
|
||||||
ordered_versions: List[Dict] = []
|
ordered_versions: List[Dict[str, Any]] = []
|
||||||
seen_ids = set()
|
seen_ids = set()
|
||||||
for version in transformed_versions:
|
for version in transformed_versions:
|
||||||
version_id = version.get("id")
|
version_id = version.get("id")
|
||||||
@@ -402,7 +407,7 @@ class CivArchiveClient:
|
|||||||
logger.error(f"Error fetching CivArchive model versions for {model_id}: {e}")
|
logger.error(f"Error fetching CivArchive model versions for {model_id}: {e}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_model_version(self, model_id: int = None, version_id: int = None) -> Optional[Dict]:
|
async def get_model_version(self, model_id: int | str | None = None, version_id: int | str | None = None) -> Optional[Dict[str, Any]]:
|
||||||
"""Get specific model version using CivArchive API
|
"""Get specific model version using CivArchive API
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -459,7 +464,7 @@ class CivArchiveClient:
|
|||||||
logger.error(f"Error fetching CivArchive model version via API {model_id}/{version_id}: {e}")
|
logger.error(f"Error fetching CivArchive model version via API {model_id}/{version_id}: {e}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict], Optional[str]]:
|
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||||
""" Fetch model version metadata using a known bogus model lookup
|
""" Fetch model version metadata using a known bogus model lookup
|
||||||
CivArchive lacks a direct version lookup API, this uses a workaround (which we handle in the main model request now)
|
CivArchive lacks a direct version lookup API, this uses a workaround (which we handle in the main model request now)
|
||||||
|
|
||||||
|
|||||||
@@ -283,7 +283,7 @@ class CivitaiBaseModelService:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
if isinstance(result, str):
|
if isinstance(result, str):
|
||||||
data = json.loads(result)
|
data: Any = json.loads(result)
|
||||||
else:
|
else:
|
||||||
data = result
|
data = result
|
||||||
|
|
||||||
|
|||||||
+137
-42
@@ -1,9 +1,14 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
import asyncio
|
import asyncio
|
||||||
import copy
|
import copy
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import time
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import Any, Optional, Dict, Tuple, List, Sequence
|
from typing import Any, Optional, Dict, Tuple, List, Sequence, cast
|
||||||
from .connectivity_guard import (
|
from .connectivity_guard import (
|
||||||
OFFLINE_FRIENDLY_MESSAGE,
|
OFFLINE_FRIENDLY_MESSAGE,
|
||||||
is_expected_offline_error,
|
is_expected_offline_error,
|
||||||
@@ -16,9 +21,16 @@ from .model_metadata_provider import (
|
|||||||
from .downloader import get_downloader
|
from .downloader import get_downloader
|
||||||
from .errors import RateLimitError, ResourceNotFoundError
|
from .errors import RateLimitError, ResourceNotFoundError
|
||||||
from ..utils.civitai_utils import resolve_license_payload
|
from ..utils.civitai_utils import resolve_license_payload
|
||||||
|
from ..utils.constants import MODEL_WEIGHT_FILE_TYPES
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Best-effort cache for creator model counts, keyed by lowercase username.
|
||||||
|
# Values are (monotonic timestamp, count or None); None results are cached
|
||||||
|
# too so repeated failures don't hammer the API.
|
||||||
|
_CREATOR_COUNT_CACHE_TTL_SECONDS = 600
|
||||||
|
_creator_model_count_cache: Dict[str, Tuple[float, Optional[int]]] = {}
|
||||||
|
|
||||||
|
|
||||||
class CivitaiClient:
|
class CivitaiClient:
|
||||||
_instance = None
|
_instance = None
|
||||||
@@ -51,7 +63,7 @@ class CivitaiClient:
|
|||||||
# Uses OrderedDict with LRU eviction at MAX_CACHE_ENTRIES to prevent
|
# Uses OrderedDict with LRU eviction at MAX_CACHE_ENTRIES to prevent
|
||||||
# unbounded growth in long-running server processes.
|
# unbounded growth in long-running server processes.
|
||||||
self._version_info_cache: OrderedDict[
|
self._version_info_cache: OrderedDict[
|
||||||
str, Tuple[Optional[Dict], Optional[str]]
|
str, Tuple[Optional[Dict[str, Any]], Optional[str]]
|
||||||
] = OrderedDict()
|
] = OrderedDict()
|
||||||
self._MAX_CACHE_ENTRIES = 500
|
self._MAX_CACHE_ENTRIES = 500
|
||||||
|
|
||||||
@@ -65,7 +77,7 @@ class CivitaiClient:
|
|||||||
*,
|
*,
|
||||||
use_auth: bool = False,
|
use_auth: bool = False,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Tuple[bool, Dict | str]:
|
) -> Tuple[bool, Dict[str, Any] | str]:
|
||||||
"""Wrapper around downloader.make_request that surfaces rate limits,
|
"""Wrapper around downloader.make_request that surfaces rate limits,
|
||||||
with retry for transient server errors (5xx, Cloudflare 524, network flakiness)."""
|
with retry for transient server errors (5xx, Cloudflare 524, network flakiness)."""
|
||||||
|
|
||||||
@@ -79,7 +91,8 @@ class CivitaiClient:
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
if success:
|
if success:
|
||||||
return True, result
|
# RateLimitError is raised below; a successful result is dict or str.
|
||||||
|
return True, cast(Dict[str, Any] | str, result)
|
||||||
|
|
||||||
if isinstance(result, RateLimitError):
|
if isinstance(result, RateLimitError):
|
||||||
if result.provider is None:
|
if result.provider is None:
|
||||||
@@ -119,7 +132,7 @@ class CivitaiClient:
|
|||||||
return False, "Unexpected error in _make_request"
|
return False, "Unexpected error in _make_request"
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _remove_comfy_metadata(model_version: Optional[Dict]) -> None:
|
def _remove_comfy_metadata(model_version: Optional[Dict[str, Any]]) -> None:
|
||||||
"""Remove Comfy-specific metadata from model version images."""
|
"""Remove Comfy-specific metadata from model version images."""
|
||||||
if not isinstance(model_version, dict):
|
if not isinstance(model_version, dict):
|
||||||
return
|
return
|
||||||
@@ -166,7 +179,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
async def get_model_by_hash(
|
async def get_model_by_hash(
|
||||||
self, model_hash: str
|
self, model_hash: str
|
||||||
) -> Tuple[Optional[Dict], Optional[str]]:
|
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||||
try:
|
try:
|
||||||
success, version = await self._make_request(
|
success, version = await self._make_request(
|
||||||
"GET",
|
"GET",
|
||||||
@@ -213,7 +226,7 @@ class CivitaiClient:
|
|||||||
# Ensure directory exists
|
# Ensure directory exists
|
||||||
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
||||||
with open(save_path, "wb") as f:
|
with open(save_path, "wb") as f:
|
||||||
f.write(content)
|
f.write(content if isinstance(content, bytes) else content.encode("utf-8"))
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -268,7 +281,7 @@ class CivitaiClient:
|
|||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
async def get_model_versions(self, model_id: str) -> Optional[Dict]:
|
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
|
||||||
"""Get all versions of a model with local availability info"""
|
"""Get all versions of a model with local availability info"""
|
||||||
try:
|
try:
|
||||||
success, result = await self._make_request(
|
success, result = await self._make_request(
|
||||||
@@ -276,7 +289,7 @@ class CivitaiClient:
|
|||||||
f"{self.base_url}/models/{model_id}",
|
f"{self.base_url}/models/{model_id}",
|
||||||
use_auth=True,
|
use_auth=True,
|
||||||
)
|
)
|
||||||
if success:
|
if success and isinstance(result, dict):
|
||||||
# Also return model type along with versions
|
# Also return model type along with versions
|
||||||
return {
|
return {
|
||||||
"modelVersions": result.get("modelVersions", []),
|
"modelVersions": result.get("modelVersions", []),
|
||||||
@@ -310,7 +323,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
async def get_model_versions_bulk(
|
async def get_model_versions_bulk(
|
||||||
self, model_ids: Sequence[int]
|
self, model_ids: Sequence[int]
|
||||||
) -> Optional[Dict[int, Dict]]:
|
) -> Optional[Dict[int, Dict[str, Any]]]:
|
||||||
"""Fetch model metadata for multiple ids using the batch API."""
|
"""Fetch model metadata for multiple ids using the batch API."""
|
||||||
|
|
||||||
deduped: Dict[int, None] = {}
|
deduped: Dict[int, None] = {}
|
||||||
@@ -340,13 +353,13 @@ class CivitaiClient:
|
|||||||
if not isinstance(items, list):
|
if not isinstance(items, list):
|
||||||
return {}
|
return {}
|
||||||
|
|
||||||
payload: Dict[int, Dict] = {}
|
payload: Dict[int, Dict[str, Any]] = {}
|
||||||
for item in items:
|
for item in items:
|
||||||
if not isinstance(item, dict):
|
if not isinstance(item, dict):
|
||||||
continue
|
continue
|
||||||
model_id = item.get("id")
|
model_id = item.get("id")
|
||||||
try:
|
try:
|
||||||
normalized_id = int(model_id)
|
normalized_id = int(cast(Any, model_id))
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
continue
|
continue
|
||||||
payload[normalized_id] = {
|
payload[normalized_id] = {
|
||||||
@@ -366,8 +379,8 @@ class CivitaiClient:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_model_version(
|
async def get_model_version(
|
||||||
self, model_id: int = None, version_id: int = None
|
self, model_id: int | None = None, version_id: int | None = None
|
||||||
) -> Optional[Dict]:
|
) -> Optional[Dict[str, Any]]:
|
||||||
"""Get specific model version with additional metadata."""
|
"""Get specific model version with additional metadata."""
|
||||||
try:
|
try:
|
||||||
if model_id is None and version_id is not None:
|
if model_id is None and version_id is not None:
|
||||||
@@ -385,7 +398,7 @@ class CivitaiClient:
|
|||||||
logger.error(f"Error fetching model version: {e}")
|
logger.error(f"Error fetching model version: {e}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _get_version_by_id_only(self, version_id: int) -> Optional[Dict]:
|
async def _get_version_by_id_only(self, version_id: int) -> Optional[Dict[str, Any]]:
|
||||||
version = await self._fetch_version_by_id(version_id)
|
version = await self._fetch_version_by_id(version_id)
|
||||||
if version is None:
|
if version is None:
|
||||||
return None
|
return None
|
||||||
@@ -404,7 +417,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
async def _get_version_with_model_id(
|
async def _get_version_with_model_id(
|
||||||
self, model_id: int, version_id: Optional[int]
|
self, model_id: int, version_id: Optional[int]
|
||||||
) -> Optional[Dict]:
|
) -> Optional[Dict[str, Any]]:
|
||||||
model_data = await self._fetch_model_data(model_id)
|
model_data = await self._fetch_model_data(model_id)
|
||||||
if not model_data:
|
if not model_data:
|
||||||
return None
|
return None
|
||||||
@@ -457,20 +470,20 @@ class CivitaiClient:
|
|||||||
self._remove_comfy_metadata(version)
|
self._remove_comfy_metadata(version)
|
||||||
return version
|
return version
|
||||||
|
|
||||||
async def _fetch_model_data(self, model_id: int) -> Optional[Dict]:
|
async def _fetch_model_data(self, model_id: int) -> Optional[Dict[str, Any]]:
|
||||||
success, data = await self._make_request(
|
success, data = await self._make_request(
|
||||||
"GET",
|
"GET",
|
||||||
f"{self.base_url}/models/{model_id}",
|
f"{self.base_url}/models/{model_id}",
|
||||||
use_auth=True,
|
use_auth=True,
|
||||||
)
|
)
|
||||||
if success:
|
if success and isinstance(data, dict):
|
||||||
return data
|
return data
|
||||||
if is_expected_offline_error(data):
|
if is_expected_offline_error(data):
|
||||||
return None
|
return None
|
||||||
logger.warning(f"Failed to fetch model data for model {model_id}")
|
logger.warning(f"Failed to fetch model data for model {model_id}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _fetch_version_by_id(self, version_id: Optional[int]) -> Optional[Dict]:
|
async def _fetch_version_by_id(self, version_id: Optional[int]) -> Optional[Dict[str, Any]]:
|
||||||
if version_id is None:
|
if version_id is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -479,7 +492,7 @@ class CivitaiClient:
|
|||||||
f"{self.base_url}/model-versions/{version_id}",
|
f"{self.base_url}/model-versions/{version_id}",
|
||||||
use_auth=True,
|
use_auth=True,
|
||||||
)
|
)
|
||||||
if success:
|
if success and isinstance(version, dict):
|
||||||
return version
|
return version
|
||||||
if is_expected_offline_error(version):
|
if is_expected_offline_error(version):
|
||||||
return None
|
return None
|
||||||
@@ -487,7 +500,7 @@ class CivitaiClient:
|
|||||||
logger.warning(f"Failed to fetch version by id {version_id}")
|
logger.warning(f"Failed to fetch version by id {version_id}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict]:
|
async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict[str, Any]]:
|
||||||
if not model_hash:
|
if not model_hash:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -496,7 +509,7 @@ class CivitaiClient:
|
|||||||
f"{self.base_url}/model-versions/by-hash/{model_hash}",
|
f"{self.base_url}/model-versions/by-hash/{model_hash}",
|
||||||
use_auth=True,
|
use_auth=True,
|
||||||
)
|
)
|
||||||
if success:
|
if success and isinstance(version, dict):
|
||||||
return version
|
return version
|
||||||
if is_expected_offline_error(version):
|
if is_expected_offline_error(version):
|
||||||
return None
|
return None
|
||||||
@@ -505,8 +518,8 @@ class CivitaiClient:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
def _select_target_version(
|
def _select_target_version(
|
||||||
self, model_data: Dict, model_id: int, version_id: Optional[int]
|
self, model_data: Dict[str, Any], model_id: int, version_id: Optional[int]
|
||||||
) -> Optional[Dict]:
|
) -> Optional[Dict[str, Any]]:
|
||||||
model_versions = model_data.get("modelVersions", [])
|
model_versions = model_data.get("modelVersions", [])
|
||||||
if not model_versions:
|
if not model_versions:
|
||||||
logger.warning(f"No model versions found for model {model_id}")
|
logger.warning(f"No model versions found for model {model_id}")
|
||||||
@@ -525,18 +538,24 @@ class CivitaiClient:
|
|||||||
|
|
||||||
return model_versions[0]
|
return model_versions[0]
|
||||||
|
|
||||||
def _extract_primary_model_hash(self, version_entry: Dict) -> Optional[str]:
|
def _extract_primary_model_hash(self, version_entry: Dict[str, Any]) -> Optional[str]:
|
||||||
|
# Prefer the generic "Model" file (most reliable version identity);
|
||||||
|
# fall back to any other weights-type primary.
|
||||||
for file_info in version_entry.get("files", []):
|
for file_info in version_entry.get("files", []):
|
||||||
if file_info.get("type") == "Model" and file_info.get("primary"):
|
if file_info.get("type") == "Model" and file_info.get("primary"):
|
||||||
hashes = file_info.get("hashes", {})
|
model_hash = (file_info.get("hashes", {}) or {}).get("SHA256")
|
||||||
model_hash = hashes.get("SHA256")
|
if model_hash:
|
||||||
|
return model_hash
|
||||||
|
for file_info in version_entry.get("files", []):
|
||||||
|
if file_info.get("type") in MODEL_WEIGHT_FILE_TYPES and file_info.get("primary"):
|
||||||
|
model_hash = (file_info.get("hashes", {}) or {}).get("SHA256")
|
||||||
if model_hash:
|
if model_hash:
|
||||||
return model_hash
|
return model_hash
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _build_version_from_model_data(
|
def _build_version_from_model_data(
|
||||||
self, version_entry: Dict, model_id: int, model_data: Dict
|
self, version_entry: Dict[str, Any], model_id: int, model_data: Dict[str, Any]
|
||||||
) -> Dict:
|
) -> Dict[str, Any]:
|
||||||
version = copy.deepcopy(version_entry)
|
version = copy.deepcopy(version_entry)
|
||||||
version.pop("index", None)
|
version.pop("index", None)
|
||||||
version["modelId"] = model_id
|
version["modelId"] = model_id
|
||||||
@@ -548,7 +567,7 @@ class CivitaiClient:
|
|||||||
}
|
}
|
||||||
return version
|
return version
|
||||||
|
|
||||||
def _enrich_version_with_model_data(self, version: Dict, model_data: Dict) -> None:
|
def _enrich_version_with_model_data(self, version: Dict[str, Any], model_data: Dict[str, Any]) -> None:
|
||||||
model_info = version.get("model")
|
model_info = version.get("model")
|
||||||
if not isinstance(model_info, dict):
|
if not isinstance(model_info, dict):
|
||||||
model_info = {}
|
model_info = {}
|
||||||
@@ -564,7 +583,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
async def get_model_version_info(
|
async def get_model_version_info(
|
||||||
self, version_id: str
|
self, version_id: str
|
||||||
) -> Tuple[Optional[Dict], Optional[str]]:
|
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||||
"""Fetch model version metadata from Civitai
|
"""Fetch model version metadata from Civitai
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -589,7 +608,7 @@ class CivitaiClient:
|
|||||||
logger.debug("Resolving Civitai model version info: %s", url)
|
logger.debug("Resolving Civitai model version info: %s", url)
|
||||||
success, result = await self._make_request("GET", url, use_auth=True)
|
success, result = await self._make_request("GET", url, use_auth=True)
|
||||||
|
|
||||||
if success:
|
if success and isinstance(result, dict):
|
||||||
logger.debug("Successfully fetched model version info for: %s", version_id)
|
logger.debug("Successfully fetched model version info for: %s", version_id)
|
||||||
self._remove_comfy_metadata(result)
|
self._remove_comfy_metadata(result)
|
||||||
self._version_info_cache[version_id] = (result, None)
|
self._version_info_cache[version_id] = (result, None)
|
||||||
@@ -619,7 +638,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
async def get_image_info(
|
async def get_image_info(
|
||||||
self, image_id: str, source_url: str | None = None
|
self, image_id: str, source_url: str | None = None
|
||||||
) -> Optional[Dict]:
|
) -> Optional[Dict[str, Any]]:
|
||||||
"""Fetch image information from Civitai API
|
"""Fetch image information from Civitai API
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -652,7 +671,7 @@ class CivitaiClient:
|
|||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
if result and "items" in result and isinstance(result["items"], list):
|
if isinstance(result, dict) and "items" in result and isinstance(result["items"], list):
|
||||||
items = result["items"]
|
items = result["items"]
|
||||||
|
|
||||||
for item in items:
|
for item in items:
|
||||||
@@ -692,7 +711,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
async def get_model_versions_by_hashes(
|
async def get_model_versions_by_hashes(
|
||||||
self, hashes: List[str]
|
self, hashes: List[str]
|
||||||
) -> Optional[List[Dict]]:
|
) -> Optional[List[Dict[str, Any]]]:
|
||||||
"""Fetch full version details for up to 100 SHA256 hashes via the batch endpoint.
|
"""Fetch full version details for up to 100 SHA256 hashes via the batch endpoint.
|
||||||
|
|
||||||
Uses POST /api/v1/model-versions/by-hash which returns full version
|
Uses POST /api/v1/model-versions/by-hash which returns full version
|
||||||
@@ -709,7 +728,7 @@ class CivitaiClient:
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
BATCH_SIZE = 100
|
BATCH_SIZE = 100
|
||||||
all_versions: List[Dict] = []
|
all_versions: List[Dict[str, Any]] = []
|
||||||
|
|
||||||
for start in range(0, len(hashes), BATCH_SIZE):
|
for start in range(0, len(hashes), BATCH_SIZE):
|
||||||
batch = hashes[start : start + BATCH_SIZE]
|
batch = hashes[start : start + BATCH_SIZE]
|
||||||
@@ -729,7 +748,7 @@ class CivitaiClient:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
if isinstance(result, list):
|
if isinstance(result, list):
|
||||||
all_versions.extend(result)
|
all_versions.extend(cast(Any, result))
|
||||||
else:
|
else:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Unexpected by-hash response type: %s", type(result)
|
"Unexpected by-hash response type: %s", type(result)
|
||||||
@@ -743,17 +762,34 @@ class CivitaiClient:
|
|||||||
|
|
||||||
return all_versions if all_versions else None
|
return all_versions if all_versions else None
|
||||||
|
|
||||||
async def get_user_models(self, username: str) -> Optional[List[Dict]]:
|
async def get_user_models(
|
||||||
"""Fetch all models for a specific Civitai user."""
|
self, username: str, cursor: Optional[str] = None
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Fetch one page (up to 100 models) for a specific Civitai user.
|
||||||
|
|
||||||
|
Returns ``{"items": [...], "nextCursor": <str|None>}`` on success,
|
||||||
|
or None on failure. Pass ``cursor`` (from a previous response's
|
||||||
|
``nextCursor``) to fetch subsequent pages.
|
||||||
|
"""
|
||||||
if not username:
|
if not username:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
params: Dict[str, Any] = {
|
||||||
|
"username": username,
|
||||||
|
"nsfw": "true",
|
||||||
|
"limit": 100,
|
||||||
|
"sort": "Newest",
|
||||||
|
"period": "AllTime",
|
||||||
|
}
|
||||||
|
if cursor:
|
||||||
|
params["cursor"] = cursor
|
||||||
|
|
||||||
try:
|
try:
|
||||||
success, result = await self._make_request(
|
success, result = await self._make_request(
|
||||||
"GET",
|
"GET",
|
||||||
f"{self.base_url}/models",
|
f"{self.base_url}/models",
|
||||||
use_auth=True,
|
use_auth=True,
|
||||||
params={"username": username, "nsfw": "true"},
|
params=params,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not success:
|
if not success:
|
||||||
@@ -765,7 +801,7 @@ class CivitaiClient:
|
|||||||
|
|
||||||
items = result.get("items") if isinstance(result, dict) else None
|
items = result.get("items") if isinstance(result, dict) else None
|
||||||
if not isinstance(items, list):
|
if not isinstance(items, list):
|
||||||
return []
|
items = []
|
||||||
|
|
||||||
for model in items:
|
for model in items:
|
||||||
versions = model.get("modelVersions")
|
versions = model.get("modelVersions")
|
||||||
@@ -774,9 +810,68 @@ class CivitaiClient:
|
|||||||
for version in versions:
|
for version in versions:
|
||||||
self._remove_comfy_metadata(version)
|
self._remove_comfy_metadata(version)
|
||||||
|
|
||||||
return items
|
next_cursor: Optional[str] = None
|
||||||
|
metadata = result.get("metadata") if isinstance(result, dict) else None
|
||||||
|
if isinstance(metadata, dict):
|
||||||
|
raw_cursor = metadata.get("nextCursor")
|
||||||
|
if raw_cursor is not None:
|
||||||
|
next_cursor = str(raw_cursor)
|
||||||
|
|
||||||
|
return {"items": items, "nextCursor": next_cursor}
|
||||||
except RateLimitError:
|
except RateLimitError:
|
||||||
raise
|
raise
|
||||||
except Exception as exc: # pragma: no cover - defensive logging
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
logger.error("Error fetching models for %s: %s", username, exc)
|
logger.error("Error fetching models for %s: %s", username, exc)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
|
"""Best-effort lookup of a creator's published model count.
|
||||||
|
|
||||||
|
Uses the ``/creators`` endpoint (a contains-match query), picking the
|
||||||
|
entry whose username matches exactly (case-insensitive). Returns None
|
||||||
|
on any failure; never raises. Results (including None) are cached
|
||||||
|
for ``_CREATOR_COUNT_CACHE_TTL_SECONDS``.
|
||||||
|
"""
|
||||||
|
if not username:
|
||||||
|
return None
|
||||||
|
|
||||||
|
cache_key = username.lower()
|
||||||
|
cached = _creator_model_count_cache.get(cache_key)
|
||||||
|
if cached is not None:
|
||||||
|
cached_at, cached_count = cached
|
||||||
|
if time.monotonic() - cached_at < _CREATOR_COUNT_CACHE_TTL_SECONDS:
|
||||||
|
return cached_count
|
||||||
|
|
||||||
|
count: Optional[int] = None
|
||||||
|
try:
|
||||||
|
success, result = await self._make_request(
|
||||||
|
"GET",
|
||||||
|
f"{self.base_url}/creators",
|
||||||
|
use_auth=True,
|
||||||
|
params={"query": username, "limit": 10},
|
||||||
|
)
|
||||||
|
|
||||||
|
if success and isinstance(result, dict):
|
||||||
|
creators = result.get("items")
|
||||||
|
if isinstance(creators, list):
|
||||||
|
for creator in creators:
|
||||||
|
if not isinstance(creator, dict):
|
||||||
|
continue
|
||||||
|
creator_name = creator.get("username")
|
||||||
|
if not isinstance(creator_name, str):
|
||||||
|
continue
|
||||||
|
if creator_name.lower() != cache_key:
|
||||||
|
continue
|
||||||
|
model_count = creator.get("modelCount")
|
||||||
|
if isinstance(model_count, (int, float)) and not isinstance(
|
||||||
|
model_count, bool
|
||||||
|
):
|
||||||
|
count = int(model_count)
|
||||||
|
break
|
||||||
|
except Exception as exc: # best-effort only, never propagate
|
||||||
|
logger.debug(
|
||||||
|
"Failed to fetch creator model count for %s: %s", username, exc
|
||||||
|
)
|
||||||
|
|
||||||
|
_creator_model_count_cache[cache_key] = (time.monotonic(), count)
|
||||||
|
return count
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ class DownloadCoordinator:
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
ws_manager,
|
ws_manager,
|
||||||
download_manager_factory: Callable[[], Awaitable],
|
download_manager_factory: Callable[[], Awaitable[Any]],
|
||||||
) -> None:
|
) -> None:
|
||||||
self._ws_manager = ws_manager
|
self._ws_manager = ws_manager
|
||||||
self._download_manager_factory = download_manager_factory
|
self._download_manager_factory = download_manager_factory
|
||||||
@@ -83,6 +83,7 @@ class DownloadCoordinator:
|
|||||||
save_dir=payload.get("model_root"),
|
save_dir=payload.get("model_root"),
|
||||||
relative_path=payload.get("relative_path", ""),
|
relative_path=payload.get("relative_path", ""),
|
||||||
use_default_paths=payload.get("use_default_paths", False),
|
use_default_paths=payload.get("use_default_paths", False),
|
||||||
|
use_save_dir_as_root=payload.get("use_save_dir_as_root", False),
|
||||||
progress_callback=progress_callback,
|
progress_callback=progress_callback,
|
||||||
download_id=download_id,
|
download_id=download_id,
|
||||||
source=payload.get("source"),
|
source=payload.get("source"),
|
||||||
|
|||||||
+240
-126
@@ -1,4 +1,9 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
import copy
|
import copy
|
||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import asyncio
|
import asyncio
|
||||||
@@ -8,17 +13,18 @@ import zipfile
|
|||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Dict, List, Optional, Set, Tuple
|
from typing import Any, Dict, List, Optional, Set, Tuple, cast
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
from ..utils.models import LoraMetadata, CheckpointMetadata, EmbeddingMetadata
|
from ..utils.models import LoraMetadata, CheckpointMetadata, EmbeddingMetadata
|
||||||
from ..utils.constants import (
|
from ..utils.constants import (
|
||||||
CARD_PREVIEW_WIDTH,
|
CARD_PREVIEW_WIDTH,
|
||||||
DIFFUSION_MODEL_BASE_MODELS,
|
DIFFUSION_MODEL_BASE_MODELS,
|
||||||
|
MODEL_WEIGHT_FILE_TYPES,
|
||||||
SUPPORTED_DOWNLOAD_SKIP_BASE_MODELS,
|
SUPPORTED_DOWNLOAD_SKIP_BASE_MODELS,
|
||||||
VALID_LORA_TYPES,
|
VALID_LORA_TYPES,
|
||||||
)
|
)
|
||||||
from ..utils.civitai_utils import normalize_civitai_download_url, rewrite_preview_url
|
from ..utils.civitai_utils import normalize_civitai_download_url, rewrite_preview_url
|
||||||
from ..utils.file_utils import calculate_sha256
|
from ..utils.file_utils import calculate_sha256, calculate_autov3
|
||||||
from ..utils.preview_selection import resolve_mature_threshold, select_preview_media
|
from ..utils.preview_selection import resolve_mature_threshold, select_preview_media
|
||||||
from ..utils.utils import sanitize_folder_name
|
from ..utils.utils import sanitize_folder_name
|
||||||
from ..utils.exif_utils import ExifUtils
|
from ..utils.exif_utils import ExifUtils
|
||||||
@@ -42,6 +48,11 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# File types that are never the intended download target even when CivitAI
|
||||||
|
# marks them primary — configs/archives/workflows are auxiliary artifacts.
|
||||||
|
NON_DOWNLOADABLE_PRIMARY_TYPES = ("Config", "Archive", "Workflow", "Training Data")
|
||||||
|
|
||||||
|
|
||||||
class DownloadManager:
|
class DownloadManager:
|
||||||
_instance = None
|
_instance = None
|
||||||
_lock = asyncio.Lock()
|
_lock = asyncio.Lock()
|
||||||
@@ -121,7 +132,7 @@ class DownloadManager:
|
|||||||
"delay": 0,
|
"delay": 0,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
except DownloadInProgressError:
|
except DownloadInProgressError: # pyright: ignore[reportPossiblyUnboundVariable]
|
||||||
logger.info(
|
logger.info(
|
||||||
"Skipping automatic example images download for %s; another example images download is already running",
|
"Skipping automatic example images download for %s; another example images download is already running",
|
||||||
model_hash,
|
model_hash,
|
||||||
@@ -170,7 +181,7 @@ class DownloadManager:
|
|||||||
logger.error("aria2 download failed for %s: %s", download_url, exc)
|
logger.error("aria2 download failed for %s: %s", download_url, exc)
|
||||||
return False, str(exc)
|
return False, str(exc)
|
||||||
|
|
||||||
download_kwargs = {
|
download_kwargs: Dict[str, Any] = {
|
||||||
"progress_callback": progress_callback,
|
"progress_callback": progress_callback,
|
||||||
"use_auth": use_auth,
|
"use_auth": use_auth,
|
||||||
}
|
}
|
||||||
@@ -204,16 +215,17 @@ class DownloadManager:
|
|||||||
|
|
||||||
async def download_from_civitai(
|
async def download_from_civitai(
|
||||||
self,
|
self,
|
||||||
model_id: int = None,
|
model_id: int | None = None,
|
||||||
model_version_id: int = None,
|
model_version_id: int | None = None,
|
||||||
save_dir: str = None,
|
save_dir: str | None = None,
|
||||||
relative_path: str = "",
|
relative_path: str = "",
|
||||||
progress_callback=None,
|
progress_callback=None,
|
||||||
use_default_paths: bool = False,
|
use_default_paths: bool = False,
|
||||||
download_id: str = None,
|
download_id: str | None = None,
|
||||||
source: str = None,
|
source: str | None = None,
|
||||||
file_params: Dict = None,
|
file_params: Dict[str, Any] | None = None,
|
||||||
) -> Dict:
|
use_save_dir_as_root: bool = False,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
"""Download model from Civitai with task tracking and concurrency control
|
"""Download model from Civitai with task tracking and concurrency control
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -253,6 +265,7 @@ class DownloadManager:
|
|||||||
"save_dir": save_dir,
|
"save_dir": save_dir,
|
||||||
"relative_path": relative_path,
|
"relative_path": relative_path,
|
||||||
"use_default_paths": bool(use_default_paths),
|
"use_default_paths": bool(use_default_paths),
|
||||||
|
"use_save_dir_as_root": bool(use_save_dir_as_root),
|
||||||
"source": source,
|
"source": source,
|
||||||
"file_params": copy.deepcopy(file_params) if file_params is not None else None,
|
"file_params": copy.deepcopy(file_params) if file_params is not None else None,
|
||||||
"progress": 0,
|
"progress": 0,
|
||||||
@@ -283,6 +296,7 @@ class DownloadManager:
|
|||||||
use_default_paths,
|
use_default_paths,
|
||||||
source,
|
source,
|
||||||
file_params,
|
file_params,
|
||||||
|
use_save_dir_as_root,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -309,14 +323,15 @@ class DownloadManager:
|
|||||||
async def _download_with_semaphore(
|
async def _download_with_semaphore(
|
||||||
self,
|
self,
|
||||||
task_id: str,
|
task_id: str,
|
||||||
model_id: int,
|
model_id: int | None,
|
||||||
model_version_id: int,
|
model_version_id: int | None,
|
||||||
save_dir: str,
|
save_dir: str | None,
|
||||||
relative_path: str,
|
relative_path: str,
|
||||||
progress_callback=None,
|
progress_callback=None,
|
||||||
use_default_paths: bool = False,
|
use_default_paths: bool = False,
|
||||||
source: str = None,
|
source: str | None = None,
|
||||||
file_params: Dict = None,
|
file_params: Dict[str, Any] | None = None,
|
||||||
|
use_save_dir_as_root: bool = False,
|
||||||
):
|
):
|
||||||
"""Execute download with semaphore to limit concurrency"""
|
"""Execute download with semaphore to limit concurrency"""
|
||||||
# Update status to waiting
|
# Update status to waiting
|
||||||
@@ -380,7 +395,8 @@ class DownloadManager:
|
|||||||
# Use original download implementation
|
# Use original download implementation
|
||||||
try:
|
try:
|
||||||
# Check for cancellation before starting
|
# Check for cancellation before starting
|
||||||
if asyncio.current_task().cancelled():
|
current_task = asyncio.current_task()
|
||||||
|
if current_task is not None and current_task.cancelled():
|
||||||
raise asyncio.CancelledError()
|
raise asyncio.CancelledError()
|
||||||
|
|
||||||
result = await self._execute_original_download(
|
result = await self._execute_original_download(
|
||||||
@@ -396,6 +412,7 @@ class DownloadManager:
|
|||||||
),
|
),
|
||||||
source,
|
source,
|
||||||
file_params,
|
file_params,
|
||||||
|
use_save_dir_as_root=use_save_dir_as_root,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Update status based on result
|
# Update status based on result
|
||||||
@@ -484,11 +501,11 @@ class DownloadManager:
|
|||||||
# Schedule cleanup of download record after delay
|
# Schedule cleanup of download record after delay
|
||||||
asyncio.create_task(self._cleanup_download_record(task_id))
|
asyncio.create_task(self._cleanup_download_record(task_id))
|
||||||
|
|
||||||
def _start_background_download_task(self, download_id: str, coroutine) -> asyncio.Task:
|
def _start_background_download_task(self, download_id: str, coroutine) -> asyncio.Task[Any]:
|
||||||
task = asyncio.create_task(coroutine)
|
task = asyncio.create_task(coroutine)
|
||||||
self._download_tasks[download_id] = task
|
self._download_tasks[download_id] = task
|
||||||
|
|
||||||
def _cleanup_done_task(done_task: asyncio.Task) -> None:
|
def _cleanup_done_task(done_task: asyncio.Task[Any]) -> None:
|
||||||
current_task = self._download_tasks.get(download_id)
|
current_task = self._download_tasks.get(download_id)
|
||||||
if current_task is done_task:
|
if current_task is done_task:
|
||||||
self._download_tasks.pop(download_id, None)
|
self._download_tasks.pop(download_id, None)
|
||||||
@@ -530,7 +547,7 @@ class DownloadManager:
|
|||||||
async def _cleanup_cancelled_download_files(
|
async def _cleanup_cancelled_download_files(
|
||||||
self,
|
self,
|
||||||
download_id: str,
|
download_id: str,
|
||||||
download_info: Optional[Dict],
|
download_info: Optional[Dict[str, Any]],
|
||||||
) -> None:
|
) -> None:
|
||||||
target_files = set()
|
target_files = set()
|
||||||
persisted = await self._aria2_state_store.get(download_id)
|
persisted = await self._aria2_state_store.get(download_id)
|
||||||
@@ -603,19 +620,20 @@ class DownloadManager:
|
|||||||
self,
|
self,
|
||||||
download_id: str,
|
download_id: str,
|
||||||
*,
|
*,
|
||||||
extra: Optional[Dict] = None,
|
extra: Optional[Dict[str, Any]] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
info = self._active_downloads.get(download_id)
|
info = self._active_downloads.get(download_id)
|
||||||
if not info:
|
if not info:
|
||||||
return
|
return
|
||||||
|
|
||||||
payload = {
|
payload: Dict[str, Any] = {
|
||||||
"download_id": download_id,
|
"download_id": download_id,
|
||||||
"model_id": info.get("model_id"),
|
"model_id": info.get("model_id"),
|
||||||
"model_version_id": info.get("model_version_id"),
|
"model_version_id": info.get("model_version_id"),
|
||||||
"save_dir": info.get("save_dir"),
|
"save_dir": info.get("save_dir"),
|
||||||
"relative_path": info.get("relative_path", ""),
|
"relative_path": info.get("relative_path", ""),
|
||||||
"use_default_paths": bool(info.get("use_default_paths", False)),
|
"use_default_paths": bool(info.get("use_default_paths", False)),
|
||||||
|
"use_save_dir_as_root": bool(info.get("use_save_dir_as_root", False)),
|
||||||
"source": info.get("source"),
|
"source": info.get("source"),
|
||||||
"file_params": copy.deepcopy(info.get("file_params")),
|
"file_params": copy.deepcopy(info.get("file_params")),
|
||||||
"transfer_backend": info.get("transfer_backend", "aria2"),
|
"transfer_backend": info.get("transfer_backend", "aria2"),
|
||||||
@@ -631,13 +649,14 @@ class DownloadManager:
|
|||||||
|
|
||||||
await self._aria2_state_store.upsert(download_id, payload)
|
await self._aria2_state_store.upsert(download_id, payload)
|
||||||
|
|
||||||
def _build_restored_download_info(self, record: Dict, save_path: str) -> Dict:
|
def _build_restored_download_info(self, record: Dict[str, Any], save_path: str) -> Dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
"model_id": record.get("model_id"),
|
"model_id": record.get("model_id"),
|
||||||
"model_version_id": record.get("model_version_id"),
|
"model_version_id": record.get("model_version_id"),
|
||||||
"save_dir": record.get("save_dir"),
|
"save_dir": record.get("save_dir"),
|
||||||
"relative_path": record.get("relative_path", ""),
|
"relative_path": record.get("relative_path", ""),
|
||||||
"use_default_paths": bool(record.get("use_default_paths", False)),
|
"use_default_paths": bool(record.get("use_default_paths", False)),
|
||||||
|
"use_save_dir_as_root": bool(record.get("use_save_dir_as_root", False)),
|
||||||
"source": record.get("source"),
|
"source": record.get("source"),
|
||||||
"file_params": copy.deepcopy(record.get("file_params")),
|
"file_params": copy.deepcopy(record.get("file_params")),
|
||||||
"progress": record.get("progress", 0),
|
"progress": record.get("progress", 0),
|
||||||
@@ -653,8 +672,8 @@ class DownloadManager:
|
|||||||
|
|
||||||
def _is_same_aria2_download_request(
|
def _is_same_aria2_download_request(
|
||||||
self,
|
self,
|
||||||
current_info: Optional[Dict],
|
current_info: Optional[Dict[str, Any]],
|
||||||
persisted_record: Dict,
|
persisted_record: Dict[str, Any],
|
||||||
) -> bool:
|
) -> bool:
|
||||||
if not isinstance(current_info, dict):
|
if not isinstance(current_info, dict):
|
||||||
return False
|
return False
|
||||||
@@ -666,13 +685,15 @@ class DownloadManager:
|
|||||||
|
|
||||||
return current_version_id == persisted_version_id
|
return current_version_id == persisted_version_id
|
||||||
|
|
||||||
def _build_download_urls_from_file_info(self, file_info: Dict, source: str = None) -> List[str]:
|
def _build_download_urls_from_file_info(self, file_info: Dict[str, Any], source: str | None = None) -> List[str]:
|
||||||
mirrors = file_info.get("mirrors") or []
|
mirrors = file_info.get("mirrors") or []
|
||||||
download_urls: List[str] = []
|
download_urls: List[str] = []
|
||||||
if mirrors:
|
if mirrors:
|
||||||
for mirror in mirrors:
|
for mirror in mirrors:
|
||||||
if mirror.get("deletedAt") is None and mirror.get("url"):
|
if mirror.get("deletedAt") is None and mirror.get("url"):
|
||||||
download_urls.append(normalize_civitai_download_url(mirror["url"]))
|
normalized_url = normalize_civitai_download_url(mirror["url"])
|
||||||
|
if normalized_url:
|
||||||
|
download_urls.append(normalized_url)
|
||||||
|
|
||||||
if source == "civarchive" and len(download_urls) > 1:
|
if source == "civarchive" and len(download_urls) > 1:
|
||||||
civitai_urls = [
|
civitai_urls = [
|
||||||
@@ -688,7 +709,9 @@ class DownloadManager:
|
|||||||
if not download_urls:
|
if not download_urls:
|
||||||
download_url = file_info.get("downloadUrl")
|
download_url = file_info.get("downloadUrl")
|
||||||
if download_url:
|
if download_url:
|
||||||
download_urls.append(normalize_civitai_download_url(download_url))
|
normalized_url = normalize_civitai_download_url(download_url)
|
||||||
|
if normalized_url:
|
||||||
|
download_urls.append(normalized_url)
|
||||||
|
|
||||||
return download_urls
|
return download_urls
|
||||||
|
|
||||||
@@ -696,8 +719,8 @@ class DownloadManager:
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
model_type: str,
|
model_type: str,
|
||||||
version_info: Dict,
|
version_info: Dict[str, Any],
|
||||||
file_info: Dict,
|
file_info: Dict[str, Any],
|
||||||
save_path: str,
|
save_path: str,
|
||||||
):
|
):
|
||||||
if model_type == "checkpoint":
|
if model_type == "checkpoint":
|
||||||
@@ -706,7 +729,7 @@ class DownloadManager:
|
|||||||
return EmbeddingMetadata.from_civitai_info(version_info, file_info, save_path)
|
return EmbeddingMetadata.from_civitai_info(version_info, file_info, save_path)
|
||||||
return LoraMetadata.from_civitai_info(version_info, file_info, save_path)
|
return LoraMetadata.from_civitai_info(version_info, file_info, save_path)
|
||||||
|
|
||||||
def _resolve_save_path_from_persisted_record(self, record: Dict) -> Optional[str]:
|
def _resolve_save_path_from_persisted_record(self, record: Dict[str, Any]) -> Optional[str]:
|
||||||
save_path = record.get("save_path") or record.get("file_path")
|
save_path = record.get("save_path") or record.get("file_path")
|
||||||
if isinstance(save_path, str) and save_path:
|
if isinstance(save_path, str) and save_path:
|
||||||
return os.path.abspath(save_path)
|
return os.path.abspath(save_path)
|
||||||
@@ -728,7 +751,7 @@ class DownloadManager:
|
|||||||
|
|
||||||
return os.path.abspath(os.path.join(save_dir, file_name))
|
return os.path.abspath(os.path.join(save_dir, file_name))
|
||||||
|
|
||||||
async def _resume_restored_aria2_download(self, download_id: str, record: Dict) -> Dict:
|
async def _resume_restored_aria2_download(self, download_id: str, record: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
try:
|
try:
|
||||||
if download_id in self._active_downloads:
|
if download_id in self._active_downloads:
|
||||||
self._active_downloads[download_id]["status"] = "downloading"
|
self._active_downloads[download_id]["status"] = "downloading"
|
||||||
@@ -842,7 +865,7 @@ class DownloadManager:
|
|||||||
self,
|
self,
|
||||||
previous_download_id: str,
|
previous_download_id: str,
|
||||||
new_download_id: str,
|
new_download_id: str,
|
||||||
persisted_record: Dict,
|
persisted_record: Dict[str, Any],
|
||||||
save_path: str,
|
save_path: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
aria2_downloader = await get_aria2_downloader()
|
aria2_downloader = await get_aria2_downloader()
|
||||||
@@ -938,7 +961,7 @@ class DownloadManager:
|
|||||||
except Exception:
|
except Exception:
|
||||||
status_payload = None
|
status_payload = None
|
||||||
|
|
||||||
if status_payload is not None:
|
if status_payload is not None and isinstance(gid, str):
|
||||||
remote_status = status_payload.get("status", "")
|
remote_status = status_payload.get("status", "")
|
||||||
if remote_status in {"active", "waiting", "paused"}:
|
if remote_status in {"active", "waiting", "paused"}:
|
||||||
await aria2_downloader.restore_transfer(download_id, gid, save_path)
|
await aria2_downloader.restore_transfer(download_id, gid, save_path)
|
||||||
@@ -992,6 +1015,7 @@ class DownloadManager:
|
|||||||
bool(restored.get("use_default_paths", False)),
|
bool(restored.get("use_default_paths", False)),
|
||||||
restored.get("source"),
|
restored.get("source"),
|
||||||
restored.get("file_params"),
|
restored.get("file_params"),
|
||||||
|
bool(restored.get("use_save_dir_as_root", False)),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
@@ -1115,17 +1139,18 @@ class DownloadManager:
|
|||||||
|
|
||||||
async def _execute_original_download(
|
async def _execute_original_download(
|
||||||
self,
|
self,
|
||||||
model_id,
|
model_id: int | None,
|
||||||
model_version_id,
|
model_version_id: int | None,
|
||||||
save_dir,
|
save_dir: str | None,
|
||||||
relative_path,
|
relative_path: str,
|
||||||
progress_callback,
|
progress_callback,
|
||||||
use_default_paths,
|
use_default_paths: bool,
|
||||||
download_id=None,
|
download_id: str | None = None,
|
||||||
transfer_backend="python",
|
transfer_backend: str = "python",
|
||||||
source=None,
|
source: str | None = None,
|
||||||
file_params=None,
|
file_params: Dict[str, Any] | None = None,
|
||||||
):
|
use_save_dir_as_root: bool = False,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
"""Wrapper for original download_from_civitai implementation"""
|
"""Wrapper for original download_from_civitai implementation"""
|
||||||
try:
|
try:
|
||||||
# Check if model version already exists in library
|
# Check if model version already exists in library
|
||||||
@@ -1172,7 +1197,7 @@ class DownloadManager:
|
|||||||
|
|
||||||
# Get version info based on the provided identifier
|
# Get version info based on the provided identifier
|
||||||
version_info = await metadata_provider.get_model_version(
|
version_info = await metadata_provider.get_model_version(
|
||||||
model_id, model_version_id
|
cast(int, model_id), cast(int, model_version_id)
|
||||||
)
|
)
|
||||||
|
|
||||||
if not version_info:
|
if not version_info:
|
||||||
@@ -1183,7 +1208,7 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
metadata_provider = await get_default_metadata_provider()
|
metadata_provider = await get_default_metadata_provider()
|
||||||
version_info = await metadata_provider.get_model_version(
|
version_info = await metadata_provider.get_model_version(
|
||||||
model_id, model_version_id
|
cast(int, model_id), cast(int, model_version_id)
|
||||||
)
|
)
|
||||||
|
|
||||||
if not version_info:
|
if not version_info:
|
||||||
@@ -1353,47 +1378,54 @@ class DownloadManager:
|
|||||||
# Handle use_default_paths
|
# Handle use_default_paths
|
||||||
if use_default_paths:
|
if use_default_paths:
|
||||||
settings_manager = get_settings_manager()
|
settings_manager = get_settings_manager()
|
||||||
# Set save_dir based on model type
|
# With use_save_dir_as_root, an explicitly provided save_dir is kept
|
||||||
if model_type == "checkpoint":
|
# as the base root and the path template is resolved underneath it.
|
||||||
if is_diffusion_model:
|
# Otherwise fall back to the configured default root, which keeps the
|
||||||
default_path = settings_manager.get("default_unet_root")
|
# classic "download to default root" behavior for regular downloads.
|
||||||
error_msg = "Default unet root path not set in settings"
|
if not save_dir or not use_save_dir_as_root:
|
||||||
else:
|
# Set save_dir based on model type
|
||||||
default_path = settings_manager.get("default_checkpoint_root")
|
if model_type == "checkpoint":
|
||||||
error_msg = "Default checkpoint root path not set in settings"
|
if is_diffusion_model:
|
||||||
if not default_path:
|
default_path = settings_manager.get("default_unet_root")
|
||||||
return {
|
error_msg = "Default unet root path not set in settings"
|
||||||
"success": False,
|
else:
|
||||||
"error": error_msg,
|
default_path = settings_manager.get("default_checkpoint_root")
|
||||||
}
|
error_msg = "Default checkpoint root path not set in settings"
|
||||||
save_dir = default_path
|
if not default_path:
|
||||||
elif model_type == "lora":
|
return {
|
||||||
default_path = settings_manager.get("default_lora_root")
|
"success": False,
|
||||||
if not default_path:
|
"error": error_msg,
|
||||||
return {
|
}
|
||||||
"success": False,
|
save_dir = default_path
|
||||||
"error": "Default lora root path not set in settings",
|
elif model_type == "lora":
|
||||||
}
|
default_path = settings_manager.get("default_lora_root")
|
||||||
save_dir = default_path
|
if not default_path:
|
||||||
elif model_type == "embedding":
|
return {
|
||||||
default_path = settings_manager.get("default_embedding_root")
|
"success": False,
|
||||||
if not default_path:
|
"error": "Default lora root path not set in settings",
|
||||||
return {
|
}
|
||||||
"success": False,
|
save_dir = default_path
|
||||||
"error": "Default embedding root path not set in settings",
|
elif model_type == "embedding":
|
||||||
}
|
default_path = settings_manager.get("default_embedding_root")
|
||||||
save_dir = default_path
|
if not default_path:
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": "Default embedding root path not set in settings",
|
||||||
|
}
|
||||||
|
save_dir = default_path
|
||||||
|
|
||||||
# Calculate relative path using template
|
# Calculate relative path using template
|
||||||
relative_path = self._calculate_relative_path(version_info, model_type)
|
relative_path = self._calculate_relative_path(version_info, model_type)
|
||||||
|
|
||||||
# Update save directory with relative path if provided
|
# Update save directory with relative path if provided
|
||||||
|
if not save_dir:
|
||||||
|
return {"success": False, "error": "No save directory specified"}
|
||||||
if relative_path:
|
if relative_path:
|
||||||
base_save_dir = save_dir
|
base_save_dir = save_dir
|
||||||
save_dir = os.path.join(save_dir, relative_path)
|
save_dir = os.path.join(save_dir, relative_path)
|
||||||
# Security: validate path containment after joining
|
# Security: validate path containment after joining
|
||||||
resolved_dir = os.path.realpath(os.path.normpath(save_dir))
|
resolved_dir = os.path.abspath(os.path.normpath(save_dir))
|
||||||
base_dir = os.path.realpath(os.path.normpath(base_save_dir))
|
base_dir = os.path.abspath(os.path.normpath(base_save_dir))
|
||||||
if not resolved_dir.startswith(base_dir + os.sep) and resolved_dir != base_dir:
|
if not resolved_dir.startswith(base_dir + os.sep) and resolved_dir != base_dir:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Path traversal detected: %s escapes %s",
|
"Path traversal detected: %s escapes %s",
|
||||||
@@ -1403,24 +1435,48 @@ class DownloadManager:
|
|||||||
# Create directory if it doesn't exist
|
# Create directory if it doesn't exist
|
||||||
os.makedirs(save_dir, exist_ok=True)
|
os.makedirs(save_dir, exist_ok=True)
|
||||||
|
|
||||||
# Check if this is an early access model
|
# Check if this is a paid or early access model
|
||||||
if version_info.get("earlyAccessEndsAt"):
|
paid_access = version_info.get("paidAccess")
|
||||||
early_access_date = version_info.get("earlyAccessEndsAt", "")
|
if isinstance(paid_access, str):
|
||||||
# Convert to a readable date if possible
|
# Some providers (e.g. CivArchive fallback) carry the DTO as JSON text
|
||||||
try:
|
try:
|
||||||
from datetime import datetime
|
parsed = json.loads(paid_access)
|
||||||
|
paid_access = parsed if isinstance(parsed, dict) else None
|
||||||
date_obj = datetime.fromisoformat(
|
except (TypeError, ValueError):
|
||||||
early_access_date.replace("Z", "+00:00")
|
paid_access = None
|
||||||
)
|
if not isinstance(paid_access, dict):
|
||||||
formatted_date = date_obj.strftime("%Y-%m-%d")
|
paid_access = None
|
||||||
|
# An empty DTO ({"permanent": false, "endsAt": null}) is not a gate
|
||||||
|
if paid_access and not paid_access.get("permanent") and not paid_access.get("endsAt"):
|
||||||
|
paid_access = None
|
||||||
|
if version_info.get("earlyAccessEndsAt") or paid_access:
|
||||||
|
permanent_paid = bool(paid_access.get("permanent")) if paid_access else False
|
||||||
|
if permanent_paid:
|
||||||
early_access_msg = (
|
early_access_msg = (
|
||||||
f"This model requires payment (until {formatted_date}). "
|
"This model requires payment. Please ensure you have "
|
||||||
|
"purchased access and are logged in to Civitai."
|
||||||
)
|
)
|
||||||
except:
|
else:
|
||||||
early_access_msg = "This model requires payment. "
|
early_access_date = version_info.get("earlyAccessEndsAt")
|
||||||
|
if not early_access_date and paid_access:
|
||||||
|
early_access_date = paid_access.get("endsAt")
|
||||||
|
if not early_access_date:
|
||||||
|
early_access_date = ""
|
||||||
|
# Convert to a readable date if possible
|
||||||
|
try:
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
early_access_msg += "Please ensure you have purchased early access and are logged in to Civitai."
|
date_obj = datetime.fromisoformat(
|
||||||
|
early_access_date.replace("Z", "+00:00")
|
||||||
|
)
|
||||||
|
formatted_date = date_obj.strftime("%Y-%m-%d")
|
||||||
|
early_access_msg = (
|
||||||
|
f"This model requires payment (until {formatted_date}). "
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
early_access_msg = "This model requires payment. "
|
||||||
|
|
||||||
|
early_access_msg += "Please ensure you have purchased early access and are logged in to Civitai."
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Early access model detected: {version_info.get('name', 'Unknown')}"
|
f"Early access model detected: {version_info.get('name', 'Unknown')}"
|
||||||
)
|
)
|
||||||
@@ -1475,7 +1531,7 @@ class DownloadManager:
|
|||||||
f
|
f
|
||||||
for f in files
|
for f in files
|
||||||
if f.get("primary")
|
if f.get("primary")
|
||||||
and f.get("type") in ("Model", "Negative", "Diffusion Model", "UNet")
|
and f.get("type") in MODEL_WEIGHT_FILE_TYPES
|
||||||
),
|
),
|
||||||
None,
|
None,
|
||||||
)
|
)
|
||||||
@@ -1515,21 +1571,52 @@ class DownloadManager:
|
|||||||
# Fallback to primary file if no match found
|
# Fallback to primary file if no match found
|
||||||
if not file_info:
|
if not file_info:
|
||||||
logger.debug("[download] Looking for primary file as fallback")
|
logger.debug("[download] Looking for primary file as fallback")
|
||||||
|
# Prefer a weights-type file CivitAI marked primary; then any
|
||||||
|
# weights-type file (providers without primary flags, e.g.
|
||||||
|
# civarchive); then trust CivitAI's primary flag regardless of
|
||||||
|
# type — newer types like 'Enhancement LoRA' are valid primary
|
||||||
|
# files. Weights files are preferred over non-weights primary
|
||||||
|
# files so a Config/Archive primary never replaces a Model.
|
||||||
file_info = next(
|
file_info = next(
|
||||||
(
|
(
|
||||||
f
|
f
|
||||||
for f in files
|
for f in files
|
||||||
if f.get("primary") and f.get("type") in ("Model", "Negative", "Diffusion Model", "UNet")
|
if f.get("primary") and f.get("type") in MODEL_WEIGHT_FILE_TYPES
|
||||||
),
|
),
|
||||||
None,
|
None,
|
||||||
)
|
)
|
||||||
if file_info:
|
if file_info:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"[download] Fallback primary file selected: id=%s, name=%s",
|
"[download] Fallback primary file selected (primary + weights): id=%s, name=%s",
|
||||||
file_info.get("id"), file_info.get("name"),
|
file_info.get("id"), file_info.get("name"),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.debug("[download] No primary file found in fallback lookup")
|
file_info = next(
|
||||||
|
(f for f in files if f.get("type") in MODEL_WEIGHT_FILE_TYPES),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if file_info:
|
||||||
|
logger.debug(
|
||||||
|
"[download] Fallback primary file selected (weights type, no primary flag): id=%s, name=%s",
|
||||||
|
file_info.get("id"), file_info.get("name"),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
file_info = next(
|
||||||
|
(
|
||||||
|
f
|
||||||
|
for f in files
|
||||||
|
if f.get("primary")
|
||||||
|
and f.get("type") not in NON_DOWNLOADABLE_PRIMARY_TYPES
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if file_info:
|
||||||
|
logger.debug(
|
||||||
|
"[download] Fallback primary file selected (trusting CivitAI primary flag): id=%s, name=%s, type=%s",
|
||||||
|
file_info.get("id"), file_info.get("name"), file_info.get("type"),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.debug("[download] No primary file found in fallback lookup")
|
||||||
|
|
||||||
if not file_info:
|
if not file_info:
|
||||||
return {"success": False, "error": "No suitable file found in metadata"}
|
return {"success": False, "error": "No suitable file found in metadata"}
|
||||||
@@ -1561,6 +1648,11 @@ class DownloadManager:
|
|||||||
version_info, file_info, save_path
|
version_info, file_info, save_path
|
||||||
)
|
)
|
||||||
logger.info(f"Creating EmbeddingMetadata for {file_name}")
|
logger.info(f"Creating EmbeddingMetadata for {file_name}")
|
||||||
|
else:
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": f'Unsupported model type "{model_type}"',
|
||||||
|
}
|
||||||
|
|
||||||
# 6. Start download process
|
# 6. Start download process
|
||||||
if transfer_backend == "aria2" and download_id:
|
if transfer_backend == "aria2" and download_id:
|
||||||
@@ -1580,7 +1672,7 @@ class DownloadManager:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
execute_kwargs = {
|
execute_kwargs: Dict[str, Any] = {
|
||||||
"download_urls": download_urls,
|
"download_urls": download_urls,
|
||||||
"save_dir": save_dir,
|
"save_dir": save_dir,
|
||||||
"metadata": metadata,
|
"metadata": metadata,
|
||||||
@@ -1627,7 +1719,8 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# If early_access_msg exists and download failed, replace error message
|
# If early_access_msg exists and download failed, replace error message
|
||||||
if "early_access_msg" in locals() and not result.get("success", False):
|
early_access_msg = locals().get("early_access_msg")
|
||||||
|
if early_access_msg and not result.get("success", False):
|
||||||
result["error"] = early_access_msg
|
result["error"] = early_access_msg
|
||||||
|
|
||||||
return result
|
return result
|
||||||
@@ -1652,7 +1745,7 @@ class DownloadManager:
|
|||||||
self,
|
self,
|
||||||
model_type: str,
|
model_type: str,
|
||||||
model_id_value,
|
model_id_value,
|
||||||
version_info: Dict,
|
version_info: Dict[str, Any],
|
||||||
fallback_version_id=None,
|
fallback_version_id=None,
|
||||||
file_path: str | None = None,
|
file_path: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -1683,8 +1776,8 @@ class DownloadManager:
|
|||||||
try:
|
try:
|
||||||
await history_service.mark_downloaded(
|
await history_service.mark_downloaded(
|
||||||
model_type,
|
model_type,
|
||||||
int(version_id),
|
int(cast(Any, version_id)),
|
||||||
model_id=int(resolved_model_id) if resolved_model_id is not None else None,
|
model_id=int(cast(Any, resolved_model_id)) if resolved_model_id is not None else None,
|
||||||
source="download",
|
source="download",
|
||||||
file_path=file_path,
|
file_path=file_path,
|
||||||
)
|
)
|
||||||
@@ -1701,7 +1794,7 @@ class DownloadManager:
|
|||||||
self,
|
self,
|
||||||
model_type: str,
|
model_type: str,
|
||||||
model_id_value,
|
model_id_value,
|
||||||
version_info: Dict,
|
version_info: Dict[str, Any],
|
||||||
fallback_version_id=None,
|
fallback_version_id=None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Ensure update tracking reflects a newly downloaded version."""
|
"""Ensure update tracking reflects a newly downloaded version."""
|
||||||
@@ -1725,7 +1818,7 @@ class DownloadManager:
|
|||||||
if isinstance(model_info, dict):
|
if isinstance(model_info, dict):
|
||||||
resolved_model_id = model_info.get("id")
|
resolved_model_id = model_info.get("id")
|
||||||
try:
|
try:
|
||||||
resolved_model_id = int(resolved_model_id)
|
resolved_model_id = int(cast(Any, resolved_model_id))
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Skipping update sync; invalid model id: %s", resolved_model_id
|
"Skipping update sync; invalid model id: %s", resolved_model_id
|
||||||
@@ -1736,7 +1829,7 @@ class DownloadManager:
|
|||||||
if version_id is None:
|
if version_id is None:
|
||||||
version_id = fallback_version_id
|
version_id = fallback_version_id
|
||||||
try:
|
try:
|
||||||
version_id = int(version_id)
|
version_id = int(cast(Any, version_id))
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Skipping update sync; invalid version id for model %s: %s",
|
"Skipping update sync; invalid version id for model %s: %s",
|
||||||
@@ -1773,7 +1866,7 @@ class DownloadManager:
|
|||||||
for entry in local_versions or []:
|
for entry in local_versions or []:
|
||||||
vid = entry.get("versionId")
|
vid = entry.get("versionId")
|
||||||
try:
|
try:
|
||||||
version_ids.add(int(vid))
|
version_ids.add(int(cast(Any, vid)))
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
@@ -1795,7 +1888,7 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def _calculate_relative_path(
|
def _calculate_relative_path(
|
||||||
self, version_info: Dict, model_type: str = "lora"
|
self, version_info: Dict[str, Any], model_type: str = "lora"
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Calculate relative path using template from settings
|
"""Calculate relative path using template from settings
|
||||||
|
|
||||||
@@ -1871,21 +1964,22 @@ class DownloadManager:
|
|||||||
download_urls: List[str],
|
download_urls: List[str],
|
||||||
save_dir: str,
|
save_dir: str,
|
||||||
metadata,
|
metadata,
|
||||||
version_info: Dict,
|
version_info: Dict[str, Any],
|
||||||
relative_path: str,
|
relative_path: str,
|
||||||
progress_callback=None,
|
progress_callback=None,
|
||||||
model_type: str = "lora",
|
model_type: str = "lora",
|
||||||
download_id: str = None,
|
download_id: str | None = None,
|
||||||
transfer_backend: Optional[str] = None,
|
transfer_backend: Optional[str] = None,
|
||||||
) -> Dict:
|
) -> Dict[str, Any]:
|
||||||
"""Execute the actual download process including preview images and model files"""
|
"""Execute the actual download process including preview images and model files"""
|
||||||
metadata_entries: List = []
|
metadata_entries: List[Any] = []
|
||||||
metadata_files_for_cleanup: List[str] = []
|
metadata_files_for_cleanup: List[str] = []
|
||||||
extracted_paths: List[str] = []
|
extracted_paths: List[str] = []
|
||||||
metadata_path = ""
|
metadata_path = ""
|
||||||
preview_targets: List[str] = []
|
preview_targets: List[str] = []
|
||||||
preview_path: str | None = None
|
preview_path: str | None = None
|
||||||
preview_nsfw_level = 0
|
preview_nsfw_level = 0
|
||||||
|
save_path: str | None = None
|
||||||
transfer_backend = (transfer_backend or self._get_model_download_backend()).lower()
|
transfer_backend = (transfer_backend or self._get_model_download_backend()).lower()
|
||||||
try:
|
try:
|
||||||
resolved, save_path = await self._resolve_download_target_path(
|
resolved, save_path = await self._resolve_download_target_path(
|
||||||
@@ -1933,9 +2027,9 @@ class DownloadManager:
|
|||||||
mature_threshold=mature_threshold,
|
mature_threshold=mature_threshold,
|
||||||
)
|
)
|
||||||
|
|
||||||
preview_url = selected_image.get("url") if selected_image else None
|
preview_url = cast(Optional[str], selected_image.get("url")) if selected_image else None
|
||||||
media_type = (
|
media_type = (
|
||||||
(selected_image.get("type") or "").lower() if selected_image else ""
|
cast(str, selected_image.get("type") or "").lower() if selected_image else ""
|
||||||
)
|
)
|
||||||
|
|
||||||
def _extension_from_url(url: str, fallback: str) -> str:
|
def _extension_from_url(url: str, fallback: str) -> str:
|
||||||
@@ -1959,9 +2053,10 @@ class DownloadManager:
|
|||||||
preview_url, media_type="video"
|
preview_url, media_type="video"
|
||||||
)
|
)
|
||||||
attempt_urls: List[str] = []
|
attempt_urls: List[str] = []
|
||||||
if rewritten:
|
if rewritten and rewritten_url:
|
||||||
attempt_urls.append(rewritten_url)
|
attempt_urls.append(rewritten_url)
|
||||||
attempt_urls.append(preview_url)
|
if preview_url:
|
||||||
|
attempt_urls.append(preview_url)
|
||||||
|
|
||||||
seen_attempts = set()
|
seen_attempts = set()
|
||||||
for attempt in attempt_urls:
|
for attempt in attempt_urls:
|
||||||
@@ -1978,7 +2073,7 @@ class DownloadManager:
|
|||||||
rewritten_url, rewritten = rewrite_preview_url(
|
rewritten_url, rewritten = rewrite_preview_url(
|
||||||
preview_url, media_type="image"
|
preview_url, media_type="image"
|
||||||
)
|
)
|
||||||
if rewritten:
|
if rewritten and rewritten_url:
|
||||||
preview_ext = _extension_from_url(preview_url, ".png")
|
preview_ext = _extension_from_url(preview_url, ".png")
|
||||||
preview_path = os.path.splitext(save_path)[0] + preview_ext
|
preview_path = os.path.splitext(save_path)[0] + preview_ext
|
||||||
success, _ = await downloader.download_file(
|
success, _ = await downloader.download_file(
|
||||||
@@ -2004,7 +2099,9 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
if success:
|
if success:
|
||||||
with open(temp_path, "wb") as temp_file_handle:
|
with open(temp_path, "wb") as temp_file_handle:
|
||||||
temp_file_handle.write(content)
|
temp_file_handle.write(
|
||||||
|
content if isinstance(content, bytes) else content.encode("utf-8")
|
||||||
|
)
|
||||||
preview_path = (
|
preview_path = (
|
||||||
os.path.splitext(save_path)[0] + ".webp"
|
os.path.splitext(save_path)[0] + ".webp"
|
||||||
)
|
)
|
||||||
@@ -2056,6 +2153,8 @@ class DownloadManager:
|
|||||||
last_error = None
|
last_error = None
|
||||||
for download_url in download_urls:
|
for download_url in download_urls:
|
||||||
download_url = normalize_civitai_download_url(download_url)
|
download_url = normalize_civitai_download_url(download_url)
|
||||||
|
if download_url is None:
|
||||||
|
continue
|
||||||
use_auth = download_url.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
|
use_auth = download_url.startswith(CIVITAI_DOWNLOAD_URL_PREFIXES)
|
||||||
if transfer_backend == "aria2" and download_id:
|
if transfer_backend == "aria2" and download_id:
|
||||||
await self._persist_aria2_state(
|
await self._persist_aria2_state(
|
||||||
@@ -2160,6 +2259,10 @@ class DownloadManager:
|
|||||||
"error": f"Zip archive does not contain any supported model files ({supported_text})",
|
"error": f"Zip archive does not contain any supported model files ({supported_text})",
|
||||||
}
|
}
|
||||||
actual_file_paths = extracted_paths
|
actual_file_paths = extracted_paths
|
||||||
|
# The archive entry's AutoV3 (if any) describes the zip itself,
|
||||||
|
# not the extracted models; clear it so per-file header
|
||||||
|
# resolution applies to every extracted model.
|
||||||
|
metadata.autov3 = None
|
||||||
try:
|
try:
|
||||||
os.remove(save_path)
|
os.remove(save_path)
|
||||||
except OSError as exc:
|
except OSError as exc:
|
||||||
@@ -2235,7 +2338,7 @@ class DownloadManager:
|
|||||||
entry, normalized_file_path, adjust_root
|
entry, normalized_file_path, adjust_root
|
||||||
)
|
)
|
||||||
if adjusted_entry is not None:
|
if adjusted_entry is not None:
|
||||||
entry = adjusted_entry
|
entry = cast(Any, adjusted_entry)
|
||||||
metadata_entries[index] = entry
|
metadata_entries[index] = entry
|
||||||
|
|
||||||
metadata_file_path = (
|
metadata_file_path = (
|
||||||
@@ -2355,11 +2458,11 @@ class DownloadManager:
|
|||||||
|
|
||||||
async def _build_metadata_entries(
|
async def _build_metadata_entries(
|
||||||
self, base_metadata, file_paths: List[str]
|
self, base_metadata, file_paths: List[str]
|
||||||
) -> List:
|
) -> List[Any]:
|
||||||
if not file_paths:
|
if not file_paths:
|
||||||
return []
|
return []
|
||||||
|
|
||||||
entries: List = []
|
entries: List[Any] = []
|
||||||
for index, file_path in enumerate(file_paths):
|
for index, file_path in enumerate(file_paths):
|
||||||
entry = base_metadata if index == 0 else copy.deepcopy(base_metadata)
|
entry = base_metadata if index == 0 else copy.deepcopy(base_metadata)
|
||||||
# Update file paths without modifying size and modified timestamps
|
# Update file paths without modifying size and modified timestamps
|
||||||
@@ -2374,6 +2477,16 @@ class DownloadManager:
|
|||||||
sha256 = await calculate_sha256(file_path)
|
sha256 = await calculate_sha256(file_path)
|
||||||
if sha256:
|
if sha256:
|
||||||
entry.sha256 = sha256.lower()
|
entry.sha256 = sha256.lower()
|
||||||
|
# AutoV3: the Civitai-reported value for the downloaded file (set
|
||||||
|
# by from_civitai_info) takes precedence. Only the un-checked
|
||||||
|
# state (None) triggers a header read; '' (checked-unavailable)
|
||||||
|
# is never re-read, honoring the three-state contract so rows
|
||||||
|
# marked at download time stay untouched by later passes.
|
||||||
|
if entry.autov3 is None:
|
||||||
|
autov3 = await asyncio.get_running_loop().run_in_executor(
|
||||||
|
None, calculate_autov3, file_path
|
||||||
|
)
|
||||||
|
entry.autov3 = (autov3 or "").lower()
|
||||||
entries.append(entry)
|
entries.append(entry)
|
||||||
|
|
||||||
return entries
|
return entries
|
||||||
@@ -2392,7 +2505,7 @@ class DownloadManager:
|
|||||||
return destination
|
return destination
|
||||||
|
|
||||||
def _distribute_preview_to_entries(
|
def _distribute_preview_to_entries(
|
||||||
self, preview_path: str, entries: List
|
self, preview_path: str, entries: List[Any]
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
if not preview_path or not entries:
|
if not preview_path or not entries:
|
||||||
return []
|
return []
|
||||||
@@ -2451,7 +2564,7 @@ class DownloadManager:
|
|||||||
progress_callback, normalized_snapshot, rounded_progress
|
progress_callback, normalized_snapshot, rounded_progress
|
||||||
)
|
)
|
||||||
|
|
||||||
async def cancel_download(self, download_id: str) -> Dict:
|
async def cancel_download(self, download_id: str) -> Dict[str, Any]:
|
||||||
"""Cancel an active download by download_id
|
"""Cancel an active download by download_id
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -2533,7 +2646,7 @@ class DownloadManager:
|
|||||||
self._download_tasks.pop(download_id, None)
|
self._download_tasks.pop(download_id, None)
|
||||||
await self._aria2_state_store.remove(download_id)
|
await self._aria2_state_store.remove(download_id)
|
||||||
|
|
||||||
async def skip_download(self, download_id: str) -> Dict:
|
async def skip_download(self, download_id: str) -> Dict[str, Any]:
|
||||||
"""Skip a download while preserving all partial files on disk.
|
"""Skip a download while preserving all partial files on disk.
|
||||||
|
|
||||||
Removes all in-memory tracking (asyncio task, semaphore, active/pause
|
Removes all in-memory tracking (asyncio task, semaphore, active/pause
|
||||||
@@ -2616,7 +2729,7 @@ class DownloadManager:
|
|||||||
# Preserve aria2 state store entry so the partial download
|
# Preserve aria2 state store entry so the partial download
|
||||||
# info survives restarts and can be resumed later
|
# info survives restarts and can be resumed later
|
||||||
|
|
||||||
async def pause_download(self, download_id: str) -> Dict:
|
async def pause_download(self, download_id: str) -> Dict[str, Any]:
|
||||||
"""Pause an active download without losing progress."""
|
"""Pause an active download without losing progress."""
|
||||||
|
|
||||||
await self._restore_persisted_downloads()
|
await self._restore_persisted_downloads()
|
||||||
@@ -2663,7 +2776,7 @@ class DownloadManager:
|
|||||||
|
|
||||||
return {"success": True, "message": "Download paused successfully"}
|
return {"success": True, "message": "Download paused successfully"}
|
||||||
|
|
||||||
async def resume_download(self, download_id: str) -> Dict:
|
async def resume_download(self, download_id: str) -> Dict[str, Any]:
|
||||||
"""Resume a previously paused download."""
|
"""Resume a previously paused download."""
|
||||||
|
|
||||||
await self._restore_persisted_downloads()
|
await self._restore_persisted_downloads()
|
||||||
@@ -2680,7 +2793,7 @@ class DownloadManager:
|
|||||||
self._pause_events[download_id] = pause_control
|
self._pause_events[download_id] = pause_control
|
||||||
self._active_downloads[download_id] = self._build_restored_download_info(
|
self._active_downloads[download_id] = self._build_restored_download_info(
|
||||||
persisted,
|
persisted,
|
||||||
os.path.abspath(save_path),
|
os.path.abspath(cast(str, save_path)),
|
||||||
)
|
)
|
||||||
|
|
||||||
if pause_control.is_set():
|
if pause_control.is_set():
|
||||||
@@ -2724,6 +2837,7 @@ class DownloadManager:
|
|||||||
bool(persisted.get("use_default_paths", False)),
|
bool(persisted.get("use_default_paths", False)),
|
||||||
persisted.get("source"),
|
persisted.get("source"),
|
||||||
persisted.get("file_params"),
|
persisted.get("file_params"),
|
||||||
|
bool(persisted.get("use_save_dir_as_root", False)),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -2807,7 +2921,7 @@ class DownloadManager:
|
|||||||
elif asyncio.iscoroutine(result):
|
elif asyncio.iscoroutine(result):
|
||||||
await result
|
await result
|
||||||
|
|
||||||
async def get_active_downloads(self) -> Dict:
|
async def get_active_downloads(self) -> Dict[str, Any]:
|
||||||
"""Get information about all active downloads
|
"""Get information about all active downloads
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ class DownloadQueueService:
|
|||||||
_instance: Optional[DownloadQueueService] = None
|
_instance: Optional[DownloadQueueService] = None
|
||||||
_class_lock: asyncio.Lock = asyncio.Lock()
|
_class_lock: asyncio.Lock = asyncio.Lock()
|
||||||
|
|
||||||
_SCHEMA = """
|
_SCHEMA_TABLES = """
|
||||||
CREATE TABLE IF NOT EXISTS download_queue (
|
CREATE TABLE IF NOT EXISTS download_queue (
|
||||||
download_id TEXT PRIMARY KEY,
|
download_id TEXT PRIMARY KEY,
|
||||||
model_id INTEGER,
|
model_id INTEGER,
|
||||||
@@ -74,6 +74,9 @@ class DownloadQueueService:
|
|||||||
);
|
);
|
||||||
CREATE INDEX IF NOT EXISTS idx_dh_completed ON download_history(completed_at DESC);
|
CREATE INDEX IF NOT EXISTS idx_dh_completed ON download_history(completed_at DESC);
|
||||||
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status);
|
CREATE INDEX IF NOT EXISTS idx_dh_status ON download_history(status);
|
||||||
|
"""
|
||||||
|
|
||||||
|
_CREATE_UNIQUE_INDEX = """
|
||||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_dh_download_id
|
CREATE UNIQUE INDEX IF NOT EXISTS idx_dh_download_id
|
||||||
ON download_history(download_id) WHERE download_id IS NOT NULL;
|
ON download_history(download_id) WHERE download_id IS NOT NULL;
|
||||||
"""
|
"""
|
||||||
@@ -115,10 +118,39 @@ class DownloadQueueService:
|
|||||||
if self._schema_initialized:
|
if self._schema_initialized:
|
||||||
return
|
return
|
||||||
with self._connect() as conn:
|
with self._connect() as conn:
|
||||||
conn.executescript(self._SCHEMA)
|
conn.executescript(self._SCHEMA_TABLES)
|
||||||
|
|
||||||
|
# Creating the unique index on download_history.download_id can
|
||||||
|
# fail if pre-existing rows have duplicate values (e.g. from a
|
||||||
|
# previous version that lacked the index). Deduplicate first so
|
||||||
|
# that the migration does not crash on startup.
|
||||||
|
if not self._index_exists(conn, "idx_dh_download_id"):
|
||||||
|
self._remove_duplicate_download_ids(conn)
|
||||||
|
conn.executescript(self._CREATE_UNIQUE_INDEX)
|
||||||
|
|
||||||
conn.commit()
|
conn.commit()
|
||||||
self._schema_initialized = True
|
self._schema_initialized = True
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _index_exists(conn: sqlite3.Connection, name: str) -> bool:
|
||||||
|
return conn.execute(
|
||||||
|
"SELECT 1 FROM sqlite_master WHERE type='index' AND name=?",
|
||||||
|
(name,),
|
||||||
|
).fetchone() is not None
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _remove_duplicate_download_ids(conn: sqlite3.Connection) -> None:
|
||||||
|
conn.execute("""
|
||||||
|
DELETE FROM download_history
|
||||||
|
WHERE id NOT IN (
|
||||||
|
SELECT MIN(id)
|
||||||
|
FROM download_history
|
||||||
|
WHERE download_id IS NOT NULL
|
||||||
|
GROUP BY download_id
|
||||||
|
)
|
||||||
|
AND download_id IS NOT NULL
|
||||||
|
""")
|
||||||
|
|
||||||
def get_database_path(self) -> str:
|
def get_database_path(self) -> str:
|
||||||
"""Return the resolved database file path."""
|
"""Return the resolved database file path."""
|
||||||
return self._db_path
|
return self._db_path
|
||||||
|
|||||||
@@ -1,3 +1,7 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
|||||||
@@ -1,3 +1,7 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
"""
|
"""
|
||||||
Unified download manager for all HTTP/HTTPS downloads in the application.
|
Unified download manager for all HTTP/HTTPS downloads in the application.
|
||||||
|
|
||||||
@@ -20,7 +24,7 @@ from dataclasses import dataclass
|
|||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
from email.utils import parsedate_to_datetime
|
from email.utils import parsedate_to_datetime
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
from typing import Optional, Dict, Tuple, Callable, Union, Awaitable
|
from typing import Optional, Dict, Tuple, Callable, Union, Awaitable, Any, cast
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
from .connectivity_guard import (
|
from .connectivity_guard import (
|
||||||
OFFLINE_COOLDOWN_ERROR,
|
OFFLINE_COOLDOWN_ERROR,
|
||||||
@@ -204,6 +208,7 @@ class Downloader:
|
|||||||
# Double check after acquiring lock
|
# Double check after acquiring lock
|
||||||
if self._session is None or self._should_refresh_session():
|
if self._session is None or self._should_refresh_session():
|
||||||
await self._create_session()
|
await self._create_session()
|
||||||
|
assert self._session is not None
|
||||||
return self._session
|
return self._session
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -231,7 +236,7 @@ class Downloader:
|
|||||||
)
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
timeout_value = float(raw_value)
|
timeout_value = float(cast(Any, raw_value))
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
timeout_value = default_timeout
|
timeout_value = default_timeout
|
||||||
|
|
||||||
@@ -243,7 +248,7 @@ class Downloader:
|
|||||||
raw_value = os.environ.get("COMFYUI_DOWNLOAD_MAX_RETRIES")
|
raw_value = os.environ.get("COMFYUI_DOWNLOAD_MAX_RETRIES")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
retries = int(raw_value)
|
retries = int(cast(Any, raw_value))
|
||||||
except (TypeError, ValueError):
|
except (TypeError, ValueError):
|
||||||
retries = default_retries
|
retries = default_retries
|
||||||
|
|
||||||
@@ -320,7 +325,7 @@ class Downloader:
|
|||||||
# CA coverage across different Python environments (especially
|
# CA coverage across different Python environments (especially
|
||||||
# embedded/compatibility Python builds).
|
# embedded/compatibility Python builds).
|
||||||
try:
|
try:
|
||||||
import certifi # type: ignore[import-untyped]
|
import certifi # pyright: ignore[reportMissingTypeStubs]
|
||||||
|
|
||||||
ca_path = certifi.where()
|
ca_path = certifi.where()
|
||||||
ssl_context = ssl.create_default_context(cafile=ca_path)
|
ssl_context = ssl.create_default_context(cafile=ca_path)
|
||||||
@@ -330,7 +335,7 @@ class Downloader:
|
|||||||
logger.debug("SSL: certifi unavailable; using system default CA bundle")
|
logger.debug("SSL: certifi unavailable; using system default CA bundle")
|
||||||
|
|
||||||
# Optimize TCP connection parameters
|
# Optimize TCP connection parameters
|
||||||
connector_kwargs = dict(
|
connector_kwargs: Dict[str, Any] = dict(
|
||||||
ssl=ssl_context,
|
ssl=ssl_context,
|
||||||
limit=8, # Concurrent connections
|
limit=8, # Concurrent connections
|
||||||
ttl_dns_cache=300, # DNS cache timeout
|
ttl_dns_cache=300, # DNS cache timeout
|
||||||
@@ -890,7 +895,7 @@ class Downloader:
|
|||||||
use_auth: bool = False,
|
use_auth: bool = False,
|
||||||
custom_headers: Optional[Dict[str, str]] = None,
|
custom_headers: Optional[Dict[str, str]] = None,
|
||||||
return_headers: bool = False,
|
return_headers: bool = False,
|
||||||
) -> Tuple[bool, Union[bytes, str], Optional[Dict]]:
|
) -> Tuple[bool, Union[bytes, str], Optional[Dict[str, Any]]]:
|
||||||
"""
|
"""
|
||||||
Download a file to memory (for small files like preview images)
|
Download a file to memory (for small files like preview images)
|
||||||
|
|
||||||
@@ -976,7 +981,7 @@ class Downloader:
|
|||||||
url: str,
|
url: str,
|
||||||
use_auth: bool = False,
|
use_auth: bool = False,
|
||||||
custom_headers: Optional[Dict[str, str]] = None,
|
custom_headers: Optional[Dict[str, str]] = None,
|
||||||
) -> Tuple[bool, Union[Dict, str]]:
|
) -> Tuple[bool, Union[Dict[str, Any], str]]:
|
||||||
"""
|
"""
|
||||||
Get response headers without downloading the full content
|
Get response headers without downloading the full content
|
||||||
|
|
||||||
@@ -1036,7 +1041,7 @@ class Downloader:
|
|||||||
use_auth: bool = False,
|
use_auth: bool = False,
|
||||||
custom_headers: Optional[Dict[str, str]] = None,
|
custom_headers: Optional[Dict[str, str]] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Tuple[bool, Union[Dict, str]]:
|
) -> Tuple[bool, Union[Dict[str, Any], str, RateLimitError]]:
|
||||||
"""
|
"""
|
||||||
Make a generic HTTP request and return JSON response
|
Make a generic HTTP request and return JSON response
|
||||||
|
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ class EmbeddingScanner(ModelScanner):
|
|||||||
roots.extend(config.embeddings_roots or [])
|
roots.extend(config.embeddings_roots or [])
|
||||||
roots.extend(config.extra_embeddings_roots or [])
|
roots.extend(config.extra_embeddings_roots or [])
|
||||||
# Remove duplicates while preserving order
|
# Remove duplicates while preserving order
|
||||||
seen: set = set()
|
seen: set[str] = set()
|
||||||
unique_roots: List[str] = []
|
unique_roots: List[str] = []
|
||||||
for root in roots:
|
for root in roots:
|
||||||
if root and root not in seen:
|
if root and root not in seen:
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
import os
|
import os
|
||||||
import logging
|
import logging
|
||||||
from typing import Dict, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
from .base_model_service import BaseModelService
|
from .base_model_service import BaseModelService
|
||||||
from .auto_tag_service import extract_auto_tags
|
from .auto_tag_service import extract_auto_tags
|
||||||
@@ -21,58 +21,58 @@ class EmbeddingService(BaseModelService):
|
|||||||
"""
|
"""
|
||||||
super().__init__("embedding", scanner, EmbeddingMetadata, update_service=update_service)
|
super().__init__("embedding", scanner, EmbeddingMetadata, update_service=update_service)
|
||||||
|
|
||||||
async def format_response(self, embedding_data: Dict) -> Optional[Dict]:
|
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||||
"""Format Embedding data for API response.
|
"""Format Embedding data for API response.
|
||||||
|
|
||||||
Returns None when the entry is missing critical fields (corrupted cache
|
Returns None when the entry is missing critical fields (corrupted cache
|
||||||
row), so the handler layer can filter it out. See issue #730.
|
row), so the handler layer can filter it out. See issue #730.
|
||||||
"""
|
"""
|
||||||
# Guard against corrupted cache entries missing critical fields
|
# Guard against corrupted cache entries missing critical fields
|
||||||
file_path = embedding_data.get("file_path")
|
file_path = model_data.get("file_path")
|
||||||
if not file_path or not isinstance(file_path, str):
|
if not file_path or not isinstance(file_path, str):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Skipping corrupted embedding entry (missing file_path): %s",
|
"Skipping corrupted embedding entry (missing file_path): %s",
|
||||||
embedding_data.get("file_name", "<unknown>"),
|
model_data.get("file_name", "<unknown>"),
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Get sub_type from cache entry (new canonical field)
|
# Get sub_type from cache entry (new canonical field)
|
||||||
sub_type = embedding_data.get("sub_type", "embedding")
|
sub_type = model_data.get("sub_type", "embedding")
|
||||||
|
|
||||||
file_name = embedding_data.get("file_name") or ""
|
file_name = model_data.get("file_name") or ""
|
||||||
model_name = embedding_data.get("model_name") or file_name
|
model_name = model_data.get("model_name") or file_name
|
||||||
folder = embedding_data.get("folder") or ""
|
folder = model_data.get("folder") or ""
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"model_name": model_name,
|
"model_name": model_name,
|
||||||
"file_name": file_name,
|
"file_name": file_name,
|
||||||
"preview_url": config.get_preview_static_url(embedding_data.get("preview_url", "")),
|
"preview_url": config.get_preview_static_url(model_data.get("preview_url", "")),
|
||||||
"preview_nsfw_level": embedding_data.get("preview_nsfw_level", 0),
|
"preview_nsfw_level": model_data.get("preview_nsfw_level", 0),
|
||||||
"base_model": embedding_data.get("base_model", ""),
|
"base_model": model_data.get("base_model", ""),
|
||||||
"folder": folder,
|
"folder": folder,
|
||||||
"sha256": embedding_data.get("sha256", ""),
|
"sha256": model_data.get("sha256", ""),
|
||||||
"file_path": file_path.replace(os.sep, "/"),
|
"file_path": file_path.replace(os.sep, "/"),
|
||||||
"file_size": embedding_data.get("size", 0),
|
"file_size": model_data.get("size", 0),
|
||||||
"modified": embedding_data.get("modified", ""),
|
"modified": model_data.get("modified", ""),
|
||||||
"tags": embedding_data.get("tags", []),
|
"tags": model_data.get("tags", []),
|
||||||
"from_civitai": embedding_data.get("from_civitai", True),
|
"from_civitai": model_data.get("from_civitai", True),
|
||||||
# "usage_count": embedding_data.get("usage_count", 0), # TODO: Enable when embedding usage tracking is implemented
|
# "usage_count": model_data.get("usage_count", 0), # TODO: Enable when embedding usage tracking is implemented
|
||||||
"notes": embedding_data.get("notes", ""),
|
"notes": model_data.get("notes", ""),
|
||||||
"sub_type": sub_type,
|
"sub_type": sub_type,
|
||||||
"favorite": embedding_data.get("favorite", False),
|
"favorite": model_data.get("favorite", False),
|
||||||
"exclude": bool(embedding_data.get("exclude", False)),
|
"exclude": bool(model_data.get("exclude", False)),
|
||||||
"update_available": bool(embedding_data.get("update_available", False)),
|
"update_available": bool(model_data.get("update_available", False)),
|
||||||
"skip_metadata_refresh": bool(embedding_data.get("skip_metadata_refresh", False)),
|
"skip_metadata_refresh": bool(model_data.get("skip_metadata_refresh", False)),
|
||||||
"civitai": self.filter_civitai_data(embedding_data.get("civitai", {}), minimal=True),
|
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
|
||||||
"auto_tags": embedding_data.get("auto_tags") or extract_auto_tags(embedding_data),
|
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
||||||
"version_count": embedding_data.get("version_count"),
|
"version_count": model_data.get("version_count"),
|
||||||
"hf_url": embedding_data.get("hf_url", ""),
|
"hf_url": model_data.get("hf_url", ""),
|
||||||
}
|
}
|
||||||
|
|
||||||
def find_duplicate_hashes(self) -> Dict:
|
def find_duplicate_hashes(self) -> Dict[str, Any]:
|
||||||
"""Find Embeddings with duplicate SHA256 hashes"""
|
"""Find Embeddings with duplicate SHA256 hashes"""
|
||||||
return self.scanner._hash_index.get_duplicate_hashes()
|
return self.scanner._hash_index.get_duplicate_hashes()
|
||||||
|
|
||||||
def find_duplicate_filenames(self) -> Dict:
|
def find_duplicate_filenames(self) -> Dict[str, Any]:
|
||||||
"""Find Embeddings with conflicting filenames"""
|
"""Find Embeddings with conflicting filenames"""
|
||||||
return self.scanner._hash_index.get_duplicate_filenames()
|
return self.scanner._hash_index.get_duplicate_filenames()
|
||||||
|
|||||||
@@ -35,7 +35,7 @@ class CleanupResult:
|
|||||||
def to_dict(self) -> Dict[str, object]:
|
def to_dict(self) -> Dict[str, object]:
|
||||||
"""Convert the dataclass to a serialisable dictionary."""
|
"""Convert the dataclass to a serialisable dictionary."""
|
||||||
|
|
||||||
data = {
|
data: Dict[str, object] = {
|
||||||
"success": self.success,
|
"success": self.success,
|
||||||
"checked_folders": self.checked_folders,
|
"checked_folders": self.checked_folders,
|
||||||
"moved_empty_folders": self.moved_empty_folders,
|
"moved_empty_folders": self.moved_empty_folders,
|
||||||
|
|||||||
@@ -201,6 +201,11 @@ PROVIDER_PRESETS: Dict[str, Dict[str, Any]] = {
|
|||||||
"api_base": "https://openrouter.ai/api/v1",
|
"api_base": "https://openrouter.ai/api/v1",
|
||||||
"requires_key": True,
|
"requires_key": True,
|
||||||
},
|
},
|
||||||
|
"google": {
|
||||||
|
"name": "Gemini",
|
||||||
|
"api_base": "https://generativelanguage.googleapis.com/v1beta/openai",
|
||||||
|
"requires_key": True,
|
||||||
|
},
|
||||||
"opencode-go": {
|
"opencode-go": {
|
||||||
"name": "OpenCode Go",
|
"name": "OpenCode Go",
|
||||||
"api_base": "https://opencode.ai/zen/go/v1",
|
"api_base": "https://opencode.ai/zen/go/v1",
|
||||||
@@ -566,18 +571,52 @@ class LLMService:
|
|||||||
if effective_max is None:
|
if effective_max is None:
|
||||||
effective_max = 4096
|
effective_max = 4096
|
||||||
|
|
||||||
result = await self.chat_completion(
|
# Use json_schema (not json_object) for broader provider compatibility:
|
||||||
messages=messages,
|
# LM Studio and some other OpenAI-compatible servers reject
|
||||||
model=model,
|
# json_object but accept json_schema. {"type": "object"} is
|
||||||
temperature=temperature,
|
# functionally equivalent — it accepts any JSON object without
|
||||||
response_format={"type": "json_object"},
|
# constraining specific fields.
|
||||||
max_tokens=effective_max,
|
response_format = {
|
||||||
)
|
"type": "json_schema",
|
||||||
|
"json_schema": {
|
||||||
|
"name": "metadata",
|
||||||
|
"schema": {"type": "object"},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = await self.chat_completion(
|
||||||
|
messages=messages,
|
||||||
|
model=model,
|
||||||
|
temperature=temperature,
|
||||||
|
response_format=response_format,
|
||||||
|
max_tokens=effective_max,
|
||||||
|
)
|
||||||
|
except LLMResponseError as e:
|
||||||
|
# Only fall back when the provider rejects the response_format
|
||||||
|
# type value (e.g. "'response_format.type' must be..."). Avoid
|
||||||
|
# catching unrelated 400 errors whose body happens to mention
|
||||||
|
# "response_format" (e.g. "model does not support
|
||||||
|
# response_format restrictions on this endpoint").
|
||||||
|
if "'response_format.type'" not in str(e).lower():
|
||||||
|
raise
|
||||||
|
logger.info(
|
||||||
|
"Provider rejected response_format, retrying without it. "
|
||||||
|
"Falling back to prompt-only JSON mode. Error: %s",
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
result = await self.chat_completion(
|
||||||
|
messages=messages,
|
||||||
|
model=model,
|
||||||
|
temperature=temperature,
|
||||||
|
response_format=None,
|
||||||
|
max_tokens=effective_max,
|
||||||
|
)
|
||||||
|
|
||||||
content = result.get("content", "") or ""
|
content = result.get("content", "") or ""
|
||||||
if not content:
|
if not content:
|
||||||
raise LLMResponseError(
|
raise LLMResponseError(
|
||||||
"LLM returned empty content in json_object mode. "
|
"LLM returned empty content. "
|
||||||
f"Raw response: {json.dumps(result)[:500]}"
|
f"Raw response: {json.dumps(result)[:500]}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +1,12 @@
|
|||||||
|
# pyright: reportImportCycles=false
|
||||||
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
import logging
|
import logging
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
from ..utils.models import LoraMetadata
|
from ..utils.models import LoraMetadata
|
||||||
from ..config import config
|
|
||||||
from .model_scanner import ModelScanner
|
from .model_scanner import ModelScanner
|
||||||
from .model_hash_index import ModelHashIndex # Changed from LoraHashIndex to ModelHashIndex
|
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -15,8 +17,10 @@ class LoraScanner(ModelScanner):
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
# Define supported file extensions
|
# Define supported file extensions
|
||||||
file_extensions = {'.safetensors'}
|
file_extensions = {'.safetensors'}
|
||||||
|
|
||||||
# Initialize parent class with ModelHashIndex
|
# Initialize parent class with ModelHashIndex
|
||||||
|
from .model_hash_index import ModelHashIndex
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
model_type="lora",
|
model_type="lora",
|
||||||
model_class=LoraMetadata,
|
model_class=LoraMetadata,
|
||||||
@@ -26,11 +30,13 @@ class LoraScanner(ModelScanner):
|
|||||||
|
|
||||||
def get_model_roots(self) -> List[str]:
|
def get_model_roots(self) -> List[str]:
|
||||||
"""Get lora root directories (including extra paths)"""
|
"""Get lora root directories (including extra paths)"""
|
||||||
|
from ..config import config
|
||||||
|
|
||||||
roots: List[str] = []
|
roots: List[str] = []
|
||||||
roots.extend(config.loras_roots or [])
|
roots.extend(config.loras_roots or [])
|
||||||
roots.extend(config.extra_loras_roots or [])
|
roots.extend(config.extra_loras_roots or [])
|
||||||
# Remove duplicates while preserving order
|
# Remove duplicates while preserving order
|
||||||
seen: set = set()
|
seen: set[str] = set()
|
||||||
unique_roots: List[str] = []
|
unique_roots: List[str] = []
|
||||||
for root in roots:
|
for root in roots:
|
||||||
if root and root not in seen:
|
if root and root not in seen:
|
||||||
@@ -68,8 +74,12 @@ class LoraScanner(ModelScanner):
|
|||||||
test_hash = next(iter(self._hash_index._hash_to_path.keys()))
|
test_hash = next(iter(self._hash_index._hash_to_path.keys()))
|
||||||
test_path = self._hash_index.get_path(test_hash)
|
test_path = self._hash_index.get_path(test_hash)
|
||||||
logger.debug(f"\nTest lookup by hash: {test_hash[:8]}... -> {test_path}")
|
logger.debug(f"\nTest lookup by hash: {test_hash[:8]}... -> {test_path}")
|
||||||
|
if test_path is None:
|
||||||
|
return
|
||||||
|
|
||||||
# Also test reverse lookup
|
# Also test reverse lookup
|
||||||
test_hash_result = self._hash_index.get_hash(test_path)
|
test_hash_result = self._hash_index.get_hash(test_path)
|
||||||
|
if test_hash_result is None:
|
||||||
|
return
|
||||||
logger.debug(f"Test reverse lookup: {test_path} -> {test_hash_result[:8]}...\n\n")
|
logger.debug(f"Test reverse lookup: {test_path} -> {test_hash_result[:8]}...\n\n")
|
||||||
|
|
||||||
|
|||||||
+41
-41
@@ -1,7 +1,7 @@
|
|||||||
import logging
|
import logging
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
from typing import Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
from .base_model_service import BaseModelService
|
from .base_model_service import BaseModelService
|
||||||
from .model_query import resolve_sub_type
|
from .model_query import resolve_sub_type
|
||||||
@@ -24,7 +24,7 @@ class LoraService(BaseModelService):
|
|||||||
"""
|
"""
|
||||||
super().__init__("lora", scanner, LoraMetadata, update_service=update_service)
|
super().__init__("lora", scanner, LoraMetadata, update_service=update_service)
|
||||||
|
|
||||||
async def format_response(self, lora_data: Dict) -> Optional[Dict]:
|
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||||
"""Format LoRA data for API response.
|
"""Format LoRA data for API response.
|
||||||
|
|
||||||
Returns None when the entry is missing critical fields (corrupted cache
|
Returns None when the entry is missing critical fields (corrupted cache
|
||||||
@@ -32,56 +32,56 @@ class LoraService(BaseModelService):
|
|||||||
whole listing request. See issue #730.
|
whole listing request. See issue #730.
|
||||||
"""
|
"""
|
||||||
# Guard against corrupted cache entries missing critical fields
|
# Guard against corrupted cache entries missing critical fields
|
||||||
file_path = lora_data.get("file_path")
|
file_path = model_data.get("file_path")
|
||||||
if not file_path or not isinstance(file_path, str):
|
if not file_path or not isinstance(file_path, str):
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Skipping corrupted LoRA entry (missing file_path): %s",
|
"Skipping corrupted LoRA entry (missing file_path): %s",
|
||||||
lora_data.get("file_name", "<unknown>"),
|
model_data.get("file_name", "<unknown>"),
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Resolve sub_type using priority: sub_type > model_type > civitai.model.type > default
|
# Resolve sub_type using priority: sub_type > model_type > civitai.model.type > default
|
||||||
# Normalize to lowercase for consistent API responses
|
# Normalize to lowercase for consistent API responses
|
||||||
sub_type = resolve_sub_type(lora_data).lower()
|
sub_type = resolve_sub_type(model_data).lower()
|
||||||
|
|
||||||
file_name = lora_data.get("file_name") or ""
|
file_name = model_data.get("file_name") or ""
|
||||||
model_name = lora_data.get("model_name") or file_name
|
model_name = model_data.get("model_name") or file_name
|
||||||
folder = lora_data.get("folder") or ""
|
folder = model_data.get("folder") or ""
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"model_name": model_name,
|
"model_name": model_name,
|
||||||
"file_name": file_name,
|
"file_name": file_name,
|
||||||
"preview_url": config.get_preview_static_url(
|
"preview_url": config.get_preview_static_url(
|
||||||
lora_data.get("preview_url", "")
|
model_data.get("preview_url", "")
|
||||||
),
|
),
|
||||||
"preview_nsfw_level": lora_data.get("preview_nsfw_level", 0),
|
"preview_nsfw_level": model_data.get("preview_nsfw_level", 0),
|
||||||
"base_model": lora_data.get("base_model", ""),
|
"base_model": model_data.get("base_model", ""),
|
||||||
"folder": folder,
|
"folder": folder,
|
||||||
"sha256": lora_data.get("sha256", ""),
|
"sha256": model_data.get("sha256", ""),
|
||||||
"file_path": file_path.replace(os.sep, "/"),
|
"file_path": file_path.replace(os.sep, "/"),
|
||||||
"file_size": lora_data.get("size", 0),
|
"file_size": model_data.get("size", 0),
|
||||||
"modified": lora_data.get("modified", ""),
|
"modified": model_data.get("modified", ""),
|
||||||
"tags": lora_data.get("tags", []),
|
"tags": model_data.get("tags", []),
|
||||||
"from_civitai": lora_data.get("from_civitai", True),
|
"from_civitai": model_data.get("from_civitai", True),
|
||||||
"usage_count": lora_data.get("usage_count", 0),
|
"usage_count": model_data.get("usage_count", 0),
|
||||||
"usage_tips": lora_data.get("usage_tips", ""),
|
"usage_tips": model_data.get("usage_tips", ""),
|
||||||
"notes": lora_data.get("notes", ""),
|
"notes": model_data.get("notes", ""),
|
||||||
"favorite": lora_data.get("favorite", False),
|
"favorite": model_data.get("favorite", False),
|
||||||
"exclude": bool(lora_data.get("exclude", False)),
|
"exclude": bool(model_data.get("exclude", False)),
|
||||||
"update_available": bool(lora_data.get("update_available", False)),
|
"update_available": bool(model_data.get("update_available", False)),
|
||||||
"skip_metadata_refresh": bool(
|
"skip_metadata_refresh": bool(
|
||||||
lora_data.get("skip_metadata_refresh", False)
|
model_data.get("skip_metadata_refresh", False)
|
||||||
),
|
),
|
||||||
"sub_type": sub_type,
|
"sub_type": sub_type,
|
||||||
"civitai": self.filter_civitai_data(
|
"civitai": self.filter_civitai_data(
|
||||||
lora_data.get("civitai", {}), minimal=True
|
model_data.get("civitai", {}), minimal=True
|
||||||
),
|
),
|
||||||
"auto_tags": lora_data.get("auto_tags") or extract_auto_tags(lora_data),
|
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
||||||
"version_count": lora_data.get("version_count"),
|
"version_count": model_data.get("version_count"),
|
||||||
"hf_url": lora_data.get("hf_url", ""),
|
"hf_url": model_data.get("hf_url", ""),
|
||||||
}
|
}
|
||||||
|
|
||||||
async def _apply_specific_filters(self, data: List[Dict], **kwargs) -> List[Dict]:
|
async def _apply_specific_filters(self, data: List[Dict[str, Any]], **kwargs) -> List[Dict[str, Any]]:
|
||||||
"""Apply LoRA-specific filters"""
|
"""Apply LoRA-specific filters"""
|
||||||
# Handle first_letter filter for LoRAs
|
# Handle first_letter filter for LoRAs
|
||||||
first_letter = kwargs.get("first_letter")
|
first_letter = kwargs.get("first_letter")
|
||||||
@@ -152,7 +152,7 @@ class LoraService(BaseModelService):
|
|||||||
|
|
||||||
return data
|
return data
|
||||||
|
|
||||||
def _filter_by_first_letter(self, data: List[Dict], letter: str) -> List[Dict]:
|
def _filter_by_first_letter(self, data: List[Dict[str, Any]], letter: str) -> List[Dict[str, Any]]:
|
||||||
"""Filter data by first letter of model name
|
"""Filter data by first letter of model name
|
||||||
|
|
||||||
Special handling:
|
Special handling:
|
||||||
@@ -307,7 +307,7 @@ class LoraService(BaseModelService):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_recommended_strength_from_lora_data(lora_data: Dict) -> Optional[float]:
|
def get_recommended_strength_from_lora_data(lora_data: Dict[str, Any]) -> Optional[float]:
|
||||||
"""Parse usage_tips JSON and extract recommended model strength."""
|
"""Parse usage_tips JSON and extract recommended model strength."""
|
||||||
try:
|
try:
|
||||||
usage_tips = lora_data.get("usage_tips", "")
|
usage_tips = lora_data.get("usage_tips", "")
|
||||||
@@ -320,7 +320,7 @@ class LoraService(BaseModelService):
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_recommended_clip_strength_from_lora_data(
|
def get_recommended_clip_strength_from_lora_data(
|
||||||
lora_data: Dict,
|
lora_data: Dict[str, Any],
|
||||||
) -> Optional[float]:
|
) -> Optional[float]:
|
||||||
"""Parse usage_tips JSON and extract recommended clip strength."""
|
"""Parse usage_tips JSON and extract recommended clip strength."""
|
||||||
try:
|
try:
|
||||||
@@ -332,7 +332,7 @@ class LoraService(BaseModelService):
|
|||||||
except (json.JSONDecodeError, TypeError, AttributeError):
|
except (json.JSONDecodeError, TypeError, AttributeError):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_lora_metadata_by_filename(self, filename: str) -> Optional[Dict]:
|
async def get_lora_metadata_by_filename(self, filename: str) -> Optional[Dict[str, Any]]:
|
||||||
"""Return cached raw metadata for a LoRA matching the given filename."""
|
"""Return cached raw metadata for a LoRA matching the given filename."""
|
||||||
cache = await self.scanner.get_cached_data(force_refresh=False)
|
cache = await self.scanner.get_cached_data(force_refresh=False)
|
||||||
|
|
||||||
@@ -357,11 +357,11 @@ class LoraService(BaseModelService):
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def find_duplicate_hashes(self) -> Dict:
|
def find_duplicate_hashes(self) -> Dict[str, Any]:
|
||||||
"""Find LoRAs with duplicate SHA256 hashes"""
|
"""Find LoRAs with duplicate SHA256 hashes"""
|
||||||
return self.scanner._hash_index.get_duplicate_hashes()
|
return self.scanner._hash_index.get_duplicate_hashes()
|
||||||
|
|
||||||
def find_duplicate_filenames(self) -> Dict:
|
def find_duplicate_filenames(self) -> Dict[str, Any]:
|
||||||
"""Find LoRAs with conflicting filenames"""
|
"""Find LoRAs with conflicting filenames"""
|
||||||
return self.scanner._hash_index.get_duplicate_filenames()
|
return self.scanner._hash_index.get_duplicate_filenames()
|
||||||
|
|
||||||
@@ -373,8 +373,8 @@ class LoraService(BaseModelService):
|
|||||||
use_same_clip_strength: bool = True,
|
use_same_clip_strength: bool = True,
|
||||||
clip_strength_min: float = 0.0,
|
clip_strength_min: float = 0.0,
|
||||||
clip_strength_max: float = 1.0,
|
clip_strength_max: float = 1.0,
|
||||||
locked_loras: Optional[List[Dict]] = None,
|
locked_loras: Optional[List[Dict[str, Any]]] = None,
|
||||||
pool_config: Optional[Dict] = None,
|
pool_config: Optional[Dict[str, Any]] = None,
|
||||||
count_mode: str = "fixed",
|
count_mode: str = "fixed",
|
||||||
count_min: int = 3,
|
count_min: int = 3,
|
||||||
count_max: int = 7,
|
count_max: int = 7,
|
||||||
@@ -382,7 +382,7 @@ class LoraService(BaseModelService):
|
|||||||
recommended_strength_scale_min: float = 0.5,
|
recommended_strength_scale_min: float = 0.5,
|
||||||
recommended_strength_scale_max: float = 1.0,
|
recommended_strength_scale_max: float = 1.0,
|
||||||
seed: Optional[int] = None,
|
seed: Optional[int] = None,
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
Get random LoRAs with specified strength ranges.
|
Get random LoRAs with specified strength ranges.
|
||||||
|
|
||||||
@@ -513,8 +513,8 @@ class LoraService(BaseModelService):
|
|||||||
return result_loras
|
return result_loras
|
||||||
|
|
||||||
async def _apply_pool_filters(
|
async def _apply_pool_filters(
|
||||||
self, available_loras: List[Dict], pool_config: Dict
|
self, available_loras: List[Dict[str, Any]], pool_config: Dict[str, Any]
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
Apply pool_config filters to available LoRAs.
|
Apply pool_config filters to available LoRAs.
|
||||||
|
|
||||||
@@ -671,8 +671,8 @@ class LoraService(BaseModelService):
|
|||||||
return available_loras
|
return available_loras
|
||||||
|
|
||||||
async def get_cycler_list(
|
async def get_cycler_list(
|
||||||
self, pool_config: Optional[Dict] = None, sort_by: str = "filename"
|
self, pool_config: Optional[Dict[str, Any]] = None, sort_by: str = "filename"
|
||||||
) -> List[Dict]:
|
) -> List[Dict[str, Any]]:
|
||||||
"""
|
"""
|
||||||
Get filtered and sorted LoRA list for cycling.
|
Get filtered and sorted LoRA list for cycling.
|
||||||
|
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user