mirror of
https://github.com/willmiao/ComfyUI-Lora-Manager.git
synced 2026-09-21 11:11:26 -03:00
Compare commits
190 Commits
da071e8452
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
| 521531111a | |||
| 474da1b264 | |||
| 78d38b449e | |||
| 2bc9860b24 | |||
| 327da0465b | |||
| c8c84bfc54 | |||
| 3b9e8efb3d | |||
| d45a523fb5 | |||
| 6dc9f34f7d | |||
| 5adfa3be36 | |||
| d4b82d98b2 | |||
| 8c1c1691e3 | |||
| e14a084f0d | |||
| c55c6f0a41 | |||
| 7d963b27b5 | |||
| eba03800b9 | |||
| bf497d5144 | |||
| 369613f811 | |||
| 9eeebac40b | |||
| b9a516c9f8 | |||
| ef7fa7d3dd | |||
| 9c67dbbf15 | |||
| 16b0bdf70a | |||
| e09fe5888b | |||
| f67689b0f9 | |||
| 1d6da1787a | |||
| d572292142 | |||
| 1b1a8d63db | |||
| a0a5b13ab0 | |||
| 779bd18e75 | |||
| 01137eed88 | |||
| c6c44b741a | |||
| 5095b23eb2 | |||
| cc25bb3dc2 | |||
| 3b54a13cae | |||
| 9bbe57ee85 | |||
| 4938faa049 | |||
| cc8eedcff7 | |||
| 9734df15b4 | |||
| 2ceb1e2850 | |||
| 942717f0b6 | |||
| 0f160e157f | |||
| e9e9ee20c6 | |||
| f0ee30fc68 | |||
| 51de85a6ca | |||
| 4064ea7d3a | |||
| 35b291ab19 | |||
| e711e643f1 | |||
| db38ad80e6 | |||
| 326df32933 | |||
| 31ef9ffa06 | |||
| 38d4c59b4c | |||
| b9bf006998 | |||
| 5ab0e88abc | |||
| 84146b62fd | |||
| adeb40bfff | |||
| 8a21837ca2 | |||
| 3302147a43 | |||
| 4d87ae7637 | |||
| b1a653f18f | |||
| 6fe0543d2e | |||
| 931dfbe1d3 | |||
| b5c1331911 | |||
| 37f2cba72d | |||
| 0e789cb38c | |||
| f3b3393a16 | |||
| 480a3f4ea5 | |||
| 69a62d739c | |||
| 28fbb86dce | |||
| f88fe2665c | |||
| 3592eab48c | |||
| 1dbdf5b00c | |||
| fc3b2d7c13 | |||
| f2a7297cb9 | |||
| 57729375b6 | |||
| fa7ce725c1 | |||
| 27da7b3ca3 | |||
| 3070838a42 | |||
| fe160134d0 | |||
| 6d3f82976f | |||
| 91b2735dad | |||
| 3112869a21 | |||
| aa630bf85b | |||
| e0052cd237 | |||
| 3cdc5ba7a2 | |||
| 04485e384f | |||
| a03dc4002f | |||
| cc9d3bff42 | |||
| 2672b3331b | |||
| 4963bf2b2e | |||
| 51cad6f852 | |||
| 1b5cbbbaa0 | |||
| e747946f7a | |||
| 53fa22f39c | |||
| 82b34097fb | |||
| a7995db009 | |||
| 5ae4aef30e | |||
| 08023f0cd9 | |||
| 6e2185c182 | |||
| 41302e75ba | |||
| a17399d667 | |||
| e2d85a0a21 | |||
| 303833bbae | |||
| f86b7b55d6 | |||
| 782bb53784 | |||
| 139231e225 | |||
| 121d8d5cea | |||
| ec147bd677 | |||
| 93fc28b499 | |||
| 7afed1a14b | |||
| e6f5142e48 | |||
| 87f05fb66c | |||
| cf64e5baa8 | |||
| 634ea7f299 | |||
| 6ba64ebb3c | |||
| 03569c62df | |||
| a61840b366 | |||
| 726fc178f1 | |||
| 8260bd022d | |||
| b309becdf9 | |||
| 1e375bb8d9 | |||
| 14da8a6f17 | |||
| da71985c3e | |||
| 7c4c8b8f30 | |||
| 77109b3cf8 | |||
| 00095a5398 | |||
| 6b41c3bbb4 | |||
| b37238d790 | |||
| bc33e32c6f | |||
| 49704d801c | |||
| 34ca14d7fc | |||
| f7b247f9e8 | |||
| 3005d2877e | |||
| ed2a17970f | |||
| 9584fa85c9 | |||
| 1fd7cc0123 | |||
| 39e7c1376c | |||
| 2a3c632dc5 | |||
| 8d46d26abe | |||
| d761ac77f7 | |||
| c8b9db5bf4 | |||
| bce7d1d30c | |||
| bccd494a56 | |||
| 3fd29f6943 | |||
| 838a374a56 | |||
| 6e31da7a70 | |||
| fc9088bfd6 | |||
| 675421ea84 | |||
| 2ff98ae089 | |||
| c972c755fc | |||
| ebe3df7d22 | |||
| be44a75b74 | |||
| fd1227d3b8 | |||
| 3a9e02137d | |||
| d8a2be8edc | |||
| 1c46b2e8c3 | |||
| 3c3ac49f2f | |||
| 1a1be95a64 | |||
| 7a36659a20 | |||
| cb18281b14 | |||
| 856c9a87ac | |||
| a7d65fe84a | |||
| 15bf079af2 | |||
| 65ba750634 | |||
| 17dcbd3d4f | |||
| e914a0e19d | |||
| 2ba04bb1bd | |||
| 1b7314591a | |||
| 2bfb987312 | |||
| df34efafbc | |||
| c2a2048c8b | |||
| 1e1921cabb | |||
| ee233548e5 | |||
| 574dfbbe55 | |||
| 1d3bcdfe47 | |||
| 74369940bf | |||
| d188cec306 | |||
| 641a61f804 | |||
| 3025c64fea | |||
| c52cfc7e7a | |||
| 4ed9f775f6 | |||
| 0b08ad283a | |||
| 08895f77ff | |||
| 74f889f160 | |||
| c51090ab16 | |||
| cdb044cb45 | |||
| c83b26b556 | |||
| a202c666bc | |||
| e05046af10 | |||
| 41ed03e5c6 |
@@ -1,373 +0,0 @@
|
|||||||
---
|
|
||||||
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 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
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
- LoRa Manager project cloned and dependencies installed (`pip install -r requirements.txt`) — run everything from `<repo-root>`
|
|
||||||
- Chrome browser available for debugging
|
|
||||||
- Chrome DevTools MCP connected
|
|
||||||
- `ss` (or `lsof`/`netstat`) available for port checks: `ss -tlnp`
|
|
||||||
|
|
||||||
## Port Selection
|
|
||||||
|
|
||||||
`8188` is only the *default candidate*. Verify it is actually free before every run:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# Is anything listening on 8188?
|
|
||||||
ss -tlnp | grep ':8188' || echo "8188 is free"
|
|
||||||
```
|
|
||||||
|
|
||||||
- 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).
|
|
||||||
|
|
||||||
## Quick Start Workflow (sandboxed)
|
|
||||||
|
|
||||||
### 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:
|
|
||||||
- Take snapshots: `take_snapshot`
|
|
||||||
- Click elements: `click`
|
|
||||||
- Fill forms: `fill` or `fill_form`
|
|
||||||
- Evaluate scripts: `evaluate_script`
|
|
||||||
- Wait for elements: `wait_for`
|
|
||||||
|
|
||||||
## Common E2E Test Patterns
|
|
||||||
|
|
||||||
### Pattern: Full Page Load Verification
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Navigate to LoRA list page
|
|
||||||
navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
|
||||||
|
|
||||||
# Wait for page to load
|
|
||||||
wait_for(text="LoRAs", timeout=10000)
|
|
||||||
|
|
||||||
# Take snapshot to verify UI state
|
|
||||||
snapshot = take_snapshot()
|
|
||||||
```
|
|
||||||
|
|
||||||
### Pattern: Restart Server for Configuration Changes
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Stop current server (if running), start with new configuration.
|
|
||||||
# --restart only kills the E2E server this script started before (via its pidfile);
|
|
||||||
# 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
|
|
||||||
navigate_page(type="reload", ignoreCache=True)
|
|
||||||
wait_for(text="LoRAs", timeout=15000)
|
|
||||||
```
|
|
||||||
|
|
||||||
### Pattern: Verify Backend API via Frontend
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Execute script in browser to call backend API
|
|
||||||
result = evaluate_script(function="""
|
|
||||||
async () => {
|
|
||||||
const response = await fetch('/loras/api/list');
|
|
||||||
const data = await response.json();
|
|
||||||
return { count: data.length, firstItem: data[0]?.name };
|
|
||||||
}
|
|
||||||
""")
|
|
||||||
```
|
|
||||||
|
|
||||||
### Pattern: Form Submission Flow
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Fill a form (e.g., search or filter)
|
|
||||||
fill_form(elements=[
|
|
||||||
{"uid": "search-input", "value": "character"},
|
|
||||||
])
|
|
||||||
|
|
||||||
# Click submit button
|
|
||||||
click(uid="search-button")
|
|
||||||
|
|
||||||
# Wait for results
|
|
||||||
wait_for(text="Results", timeout=5000)
|
|
||||||
|
|
||||||
# Verify results via snapshot
|
|
||||||
snapshot = take_snapshot()
|
|
||||||
```
|
|
||||||
|
|
||||||
### Pattern: Modal Dialog Interaction
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Open modal (e.g., add LoRA)
|
|
||||||
click(uid="add-lora-button")
|
|
||||||
|
|
||||||
# Wait for modal to appear
|
|
||||||
wait_for(text="Add LoRA", timeout=3000)
|
|
||||||
|
|
||||||
# Fill modal form
|
|
||||||
fill_form(elements=[
|
|
||||||
{"uid": "lora-name", "value": "Test LoRA"},
|
|
||||||
{"uid": "lora-path", "value": "/path/to/lora.safetensors"},
|
|
||||||
])
|
|
||||||
|
|
||||||
# Submit
|
|
||||||
click(uid="modal-submit-button")
|
|
||||||
|
|
||||||
# Wait for success message or close
|
|
||||||
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
|
|
||||||
|
|
||||||
### scripts/start_server.py
|
|
||||||
|
|
||||||
Starts or restarts the LoRa Manager standalone server for E2E testing.
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python scripts/start_server.py [--port PORT] [--restart] [--wait] [--timeout SECONDS] [--detach]
|
|
||||||
```
|
|
||||||
|
|
||||||
Options:
|
|
||||||
- `--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 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 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
|
|
||||||
|
|
||||||
Polls the server until ready or timeout.
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python scripts/wait_for_server.py [--port PORT] [--timeout SECONDS]
|
|
||||||
```
|
|
||||||
|
|
||||||
## Test Scenarios Reference
|
|
||||||
|
|
||||||
See [references/test-scenarios.md](references/test-scenarios.md) for detailed test scenarios including:
|
|
||||||
- LoRA list display and filtering
|
|
||||||
- Model metadata editing
|
|
||||||
- Recipe creation and management
|
|
||||||
- Settings configuration
|
|
||||||
- Import/export functionality
|
|
||||||
|
|
||||||
## Network Request Verification
|
|
||||||
|
|
||||||
Use `list_network_requests` and `get_network_request` to verify API calls:
|
|
||||||
|
|
||||||
```python
|
|
||||||
# List recent XHR/fetch requests
|
|
||||||
requests = list_network_requests(resourceTypes=["xhr", "fetch"])
|
|
||||||
|
|
||||||
# Get details of specific request
|
|
||||||
details = get_network_request(reqid=123)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Console Message Monitoring
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Check for errors or warnings
|
|
||||||
messages = list_console_messages(types=["error", "warn"])
|
|
||||||
```
|
|
||||||
|
|
||||||
## Performance Testing
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Start performance trace
|
|
||||||
performance_start_trace(reload=True, autoStop=False)
|
|
||||||
|
|
||||||
# Perform actions...
|
|
||||||
|
|
||||||
# Stop and analyze
|
|
||||||
results = performance_stop_trace()
|
|
||||||
```
|
|
||||||
|
|
||||||
## Cleanup
|
|
||||||
|
|
||||||
Always ensure proper cleanup after tests:
|
|
||||||
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).
|
|
||||||
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.
|
|
||||||
@@ -1,360 +0,0 @@
|
|||||||
# Chrome DevTools MCP Cheatsheet for LoRa Manager
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Navigate to LoRA list page
|
|
||||||
navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
|
||||||
|
|
||||||
# Reload page with cache clear
|
|
||||||
navigate_page(type="reload", ignoreCache=True)
|
|
||||||
|
|
||||||
# Go back/forward
|
|
||||||
navigate_page(type="back")
|
|
||||||
navigate_page(type="forward")
|
|
||||||
```
|
|
||||||
|
|
||||||
## Waiting
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Wait for text to appear
|
|
||||||
wait_for(text="LoRAs", timeout=10000)
|
|
||||||
|
|
||||||
# Wait for specific element (via evaluate_script)
|
|
||||||
evaluate_script(function="""
|
|
||||||
() => {
|
|
||||||
return new Promise((resolve) => {
|
|
||||||
const check = () => {
|
|
||||||
if (document.querySelector('.lora-card')) {
|
|
||||||
resolve(true);
|
|
||||||
} else {
|
|
||||||
setTimeout(check, 100);
|
|
||||||
}
|
|
||||||
};
|
|
||||||
check();
|
|
||||||
});
|
|
||||||
}
|
|
||||||
""")
|
|
||||||
```
|
|
||||||
|
|
||||||
## Taking Snapshots
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Full page snapshot
|
|
||||||
snapshot = take_snapshot()
|
|
||||||
|
|
||||||
# Verbose snapshot (more details)
|
|
||||||
snapshot = take_snapshot(verbose=True)
|
|
||||||
|
|
||||||
# Save to file
|
|
||||||
take_snapshot(filePath="test-snapshots/page-load.json")
|
|
||||||
```
|
|
||||||
|
|
||||||
## Element Interaction
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Click element
|
|
||||||
click(uid="element-uid-from-snapshot")
|
|
||||||
|
|
||||||
# Double click
|
|
||||||
click(uid="element-uid", dblClick=True)
|
|
||||||
|
|
||||||
# Fill input
|
|
||||||
fill(uid="search-input", value="test query")
|
|
||||||
|
|
||||||
# Fill multiple inputs
|
|
||||||
fill_form(elements=[
|
|
||||||
{"uid": "input-1", "value": "value 1"},
|
|
||||||
{"uid": "input-2", "value": "value 2"},
|
|
||||||
])
|
|
||||||
|
|
||||||
# Hover
|
|
||||||
hover(uid="lora-card-1")
|
|
||||||
|
|
||||||
# Upload file
|
|
||||||
upload_file(uid="file-input", filePath="/path/to/file.safetensors")
|
|
||||||
```
|
|
||||||
|
|
||||||
## Keyboard Input
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Press key
|
|
||||||
press_key(key="Enter")
|
|
||||||
press_key(key="Escape")
|
|
||||||
press_key(key="Tab")
|
|
||||||
|
|
||||||
# Keyboard shortcuts
|
|
||||||
press_key(key="Control+A") # Select all
|
|
||||||
press_key(key="Control+F") # Find
|
|
||||||
```
|
|
||||||
|
|
||||||
## JavaScript Evaluation
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Simple evaluation
|
|
||||||
result = evaluate_script(function="() => document.title")
|
|
||||||
|
|
||||||
# Async evaluation
|
|
||||||
result = evaluate_script(function="""
|
|
||||||
async () => {
|
|
||||||
const response = await fetch('/loras/api/list');
|
|
||||||
return await response.json();
|
|
||||||
}
|
|
||||||
""")
|
|
||||||
|
|
||||||
# Check element existence
|
|
||||||
exists = evaluate_script(function="""
|
|
||||||
() => document.querySelector('.lora-card') !== null
|
|
||||||
""")
|
|
||||||
|
|
||||||
# Get element count
|
|
||||||
count = evaluate_script(function="""
|
|
||||||
() => document.querySelectorAll('.lora-card').length
|
|
||||||
""")
|
|
||||||
```
|
|
||||||
|
|
||||||
## Network Monitoring
|
|
||||||
|
|
||||||
```python
|
|
||||||
# List all network requests
|
|
||||||
requests = list_network_requests()
|
|
||||||
|
|
||||||
# Filter by resource type
|
|
||||||
xhr_requests = list_network_requests(resourceTypes=["xhr", "fetch"])
|
|
||||||
|
|
||||||
# Get specific request details
|
|
||||||
details = get_network_request(reqid=123)
|
|
||||||
|
|
||||||
# Include preserved requests from previous navigations
|
|
||||||
all_requests = list_network_requests(includePreservedRequests=True)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Console Monitoring
|
|
||||||
|
|
||||||
```python
|
|
||||||
# List all console messages
|
|
||||||
messages = list_console_messages()
|
|
||||||
|
|
||||||
# Filter by type
|
|
||||||
errors = list_console_messages(types=["error", "warn"])
|
|
||||||
|
|
||||||
# Include preserved messages
|
|
||||||
all_messages = list_console_messages(includePreservedMessages=True)
|
|
||||||
|
|
||||||
# Get specific message
|
|
||||||
details = get_console_message(msgid=1)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Performance Testing
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Start trace with page reload
|
|
||||||
performance_start_trace(reload=True, autoStop=False)
|
|
||||||
|
|
||||||
# Start trace without reload
|
|
||||||
performance_start_trace(reload=False, autoStop=True, filePath="trace.json.gz")
|
|
||||||
|
|
||||||
# Stop trace
|
|
||||||
results = performance_stop_trace()
|
|
||||||
|
|
||||||
# Stop and save
|
|
||||||
performance_stop_trace(filePath="trace-results.json.gz")
|
|
||||||
|
|
||||||
# Analyze specific insight
|
|
||||||
insight = performance_analyze_insight(
|
|
||||||
insightSetId="results.insightSets[0].id",
|
|
||||||
insightName="LCPBreakdown"
|
|
||||||
)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Page Management
|
|
||||||
|
|
||||||
```python
|
|
||||||
# List open pages
|
|
||||||
pages = list_pages()
|
|
||||||
|
|
||||||
# Select a page
|
|
||||||
select_page(pageId=0, bringToFront=True)
|
|
||||||
|
|
||||||
# Create new page
|
|
||||||
new_page(url="http://127.0.0.1:{PORT}/loras")
|
|
||||||
|
|
||||||
# Close page (keep at least one open!)
|
|
||||||
close_page(pageId=1)
|
|
||||||
|
|
||||||
# Resize page
|
|
||||||
resize_page(width=1920, height=1080)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Screenshots
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Full page screenshot
|
|
||||||
take_screenshot(fullPage=True)
|
|
||||||
|
|
||||||
# Viewport screenshot
|
|
||||||
take_screenshot()
|
|
||||||
|
|
||||||
# Element screenshot
|
|
||||||
take_screenshot(uid="lora-card-1")
|
|
||||||
|
|
||||||
# Save to file
|
|
||||||
take_screenshot(filePath="screenshots/page.png", format="png")
|
|
||||||
|
|
||||||
# JPEG with quality
|
|
||||||
take_screenshot(filePath="screenshots/page.jpg", format="jpeg", quality=90)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Dialog Handling
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Accept dialog
|
|
||||||
handle_dialog(action="accept")
|
|
||||||
|
|
||||||
# Accept with text input
|
|
||||||
handle_dialog(action="accept", promptText="user input")
|
|
||||||
|
|
||||||
# Dismiss dialog
|
|
||||||
handle_dialog(action="dismiss")
|
|
||||||
```
|
|
||||||
|
|
||||||
## Device Emulation
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Mobile viewport
|
|
||||||
emulate(viewport={"width": 375, "height": 667, "isMobile": True, "hasTouch": True})
|
|
||||||
|
|
||||||
# Tablet viewport
|
|
||||||
emulate(viewport={"width": 768, "height": 1024, "isMobile": True, "hasTouch": True})
|
|
||||||
|
|
||||||
# Desktop viewport
|
|
||||||
emulate(viewport={"width": 1920, "height": 1080})
|
|
||||||
|
|
||||||
# Network throttling
|
|
||||||
emulate(networkConditions="Slow 3G")
|
|
||||||
emulate(networkConditions="Fast 4G")
|
|
||||||
|
|
||||||
# CPU throttling
|
|
||||||
emulate(cpuThrottlingRate=4) # 4x slowdown
|
|
||||||
|
|
||||||
# Geolocation
|
|
||||||
emulate(geolocation={"latitude": 37.7749, "longitude": -122.4194})
|
|
||||||
|
|
||||||
# User agent
|
|
||||||
emulate(userAgent="Mozilla/5.0 (Custom)")
|
|
||||||
|
|
||||||
# Reset emulation
|
|
||||||
emulate(viewport=None, networkConditions="No emulation", userAgent=None)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Drag and Drop
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Drag element to another
|
|
||||||
drag(from_uid="draggable-item", to_uid="drop-zone")
|
|
||||||
```
|
|
||||||
|
|
||||||
## Common LoRa Manager Test Patterns
|
|
||||||
|
|
||||||
### Verify LoRA Cards Loaded
|
|
||||||
|
|
||||||
```python
|
|
||||||
navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
|
||||||
wait_for(text="LoRAs", timeout=10000)
|
|
||||||
|
|
||||||
# Check if cards loaded
|
|
||||||
result = evaluate_script(function="""
|
|
||||||
() => {
|
|
||||||
const cards = document.querySelectorAll('.lora-card');
|
|
||||||
return {
|
|
||||||
count: cards.length,
|
|
||||||
hasData: cards.length > 0
|
|
||||||
};
|
|
||||||
}
|
|
||||||
""")
|
|
||||||
```
|
|
||||||
|
|
||||||
### Search and Verify Results
|
|
||||||
|
|
||||||
```python
|
|
||||||
fill(uid="search-input", value="character")
|
|
||||||
press_key(key="Enter")
|
|
||||||
wait_for(timeout=2000) # Wait for debounce
|
|
||||||
|
|
||||||
# Check results
|
|
||||||
result = evaluate_script(function="""
|
|
||||||
() => {
|
|
||||||
const cards = document.querySelectorAll('.lora-card');
|
|
||||||
const names = Array.from(cards).map(c => c.dataset.name || c.textContent);
|
|
||||||
return { count: cards.length, names };
|
|
||||||
}
|
|
||||||
""")
|
|
||||||
```
|
|
||||||
|
|
||||||
### Check API Response
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Trigger API call
|
|
||||||
evaluate_script(function="""
|
|
||||||
() => window.loraApiCallPromise = fetch('/loras/api/list').then(r => r.json())
|
|
||||||
""")
|
|
||||||
|
|
||||||
# Wait and get result
|
|
||||||
import time
|
|
||||||
time.sleep(1)
|
|
||||||
|
|
||||||
result = evaluate_script(function="""
|
|
||||||
async () => await window.loraApiCallPromise
|
|
||||||
""")
|
|
||||||
```
|
|
||||||
|
|
||||||
### Monitor Console for Errors
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Before test: clear console (navigate reloads)
|
|
||||||
navigate_page(type="reload")
|
|
||||||
|
|
||||||
# ... perform actions ...
|
|
||||||
|
|
||||||
# Check for errors
|
|
||||||
errors = list_console_messages(types=["error"])
|
|
||||||
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.
|
|
||||||
@@ -1,280 +0,0 @@
|
|||||||
# LoRa Manager E2E Test Scenarios
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
1. [LoRA List Page](#lora-list-page)
|
|
||||||
2. [Model Details](#model-details)
|
|
||||||
3. [Recipes](#recipes)
|
|
||||||
4. [Settings](#settings)
|
|
||||||
5. [Import/Export](#importexport)
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## LoRA List Page
|
|
||||||
|
|
||||||
### Scenario: Page Load and Display
|
|
||||||
|
|
||||||
**Objective**: Verify the LoRA list page loads correctly and displays models.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to `http://127.0.0.1:{PORT}/loras`
|
|
||||||
2. Wait for page title "LoRAs" to appear
|
|
||||||
3. Take snapshot to verify:
|
|
||||||
- Header with "LoRAs" title is visible
|
|
||||||
- Search/filter controls are present
|
|
||||||
- Grid/list view toggle exists
|
|
||||||
- LoRA cards are displayed (if models exist)
|
|
||||||
- Pagination controls (if applicable)
|
|
||||||
|
|
||||||
**Expected Result**: Page loads without errors, UI elements are present.
|
|
||||||
|
|
||||||
### Scenario: Search Functionality
|
|
||||||
|
|
||||||
**Objective**: Verify search filters LoRA models correctly.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Ensure at least one LoRA exists with known name (e.g., "test-character")
|
|
||||||
2. Navigate to LoRA list page
|
|
||||||
3. Enter search term in search box: "test"
|
|
||||||
4. Press Enter or click search button
|
|
||||||
5. Wait for results to update
|
|
||||||
|
|
||||||
**Expected Result**: Only LoRAs matching search term are displayed.
|
|
||||||
|
|
||||||
**Verification Script**:
|
|
||||||
```python
|
|
||||||
# After search, verify filtered results
|
|
||||||
evaluate_script(function="""
|
|
||||||
() => {
|
|
||||||
const cards = document.querySelectorAll('.lora-card');
|
|
||||||
const names = Array.from(cards).map(c => c.dataset.name);
|
|
||||||
return { count: cards.length, names };
|
|
||||||
}
|
|
||||||
""")
|
|
||||||
```
|
|
||||||
|
|
||||||
### Scenario: Filter by Tags
|
|
||||||
|
|
||||||
**Objective**: Verify tag filtering works correctly.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to LoRA list page
|
|
||||||
2. Click on a tag (e.g., "character", "style")
|
|
||||||
3. Wait for filtered results
|
|
||||||
|
|
||||||
**Expected Result**: Only LoRAs with selected tag are displayed.
|
|
||||||
|
|
||||||
### Scenario: View Mode Toggle
|
|
||||||
|
|
||||||
**Objective**: Verify grid/list view toggle works.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to LoRA list page
|
|
||||||
2. Click list view button
|
|
||||||
3. Verify list layout
|
|
||||||
4. Click grid view button
|
|
||||||
5. Verify grid layout
|
|
||||||
|
|
||||||
**Expected Result**: View mode changes correctly, layout updates.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Model Details
|
|
||||||
|
|
||||||
### Scenario: Open Model Details
|
|
||||||
|
|
||||||
**Objective**: Verify clicking a LoRA opens its details.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to LoRA list page
|
|
||||||
2. Click on a LoRA card
|
|
||||||
3. Wait for details panel/modal to open
|
|
||||||
|
|
||||||
**Expected Result**: Details panel shows:
|
|
||||||
- Model name
|
|
||||||
- Preview image
|
|
||||||
- Metadata (trigger words, tags, etc.)
|
|
||||||
- Action buttons (edit, delete, etc.)
|
|
||||||
|
|
||||||
### Scenario: Edit Model Metadata
|
|
||||||
|
|
||||||
**Objective**: Verify metadata editing works end-to-end.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Open a LoRA's details
|
|
||||||
2. Click "Edit" button
|
|
||||||
3. Modify trigger words field
|
|
||||||
4. Add/remove tags
|
|
||||||
5. Save changes
|
|
||||||
6. Refresh page
|
|
||||||
7. Reopen the same LoRA
|
|
||||||
|
|
||||||
**Expected Result**: Changes persist after refresh.
|
|
||||||
|
|
||||||
### Scenario: Delete Model
|
|
||||||
|
|
||||||
**Objective**: Verify model deletion works.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Open a LoRA's details
|
|
||||||
2. Click "Delete" button
|
|
||||||
3. Confirm deletion in dialog
|
|
||||||
4. Wait for removal
|
|
||||||
|
|
||||||
**Expected Result**: Model removed from list, success message shown.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Recipes
|
|
||||||
|
|
||||||
### Scenario: Recipe List Display
|
|
||||||
|
|
||||||
**Objective**: Verify recipes page loads and displays recipes.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to `http://127.0.0.1:{PORT}/recipes`
|
|
||||||
2. Wait for "Recipes" title
|
|
||||||
3. Take snapshot
|
|
||||||
|
|
||||||
**Expected Result**: Recipe list displayed with cards/items.
|
|
||||||
|
|
||||||
### Scenario: Create New Recipe
|
|
||||||
|
|
||||||
**Objective**: Verify recipe creation workflow.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to recipes page
|
|
||||||
2. Click "New Recipe" button
|
|
||||||
3. Fill recipe form:
|
|
||||||
- Name: "Test Recipe"
|
|
||||||
- Description: "E2E test recipe"
|
|
||||||
- Add LoRA models
|
|
||||||
4. Save recipe
|
|
||||||
5. Verify recipe appears in list
|
|
||||||
|
|
||||||
**Expected Result**: New recipe created and displayed.
|
|
||||||
|
|
||||||
### Scenario: Apply Recipe
|
|
||||||
|
|
||||||
**Objective**: Verify applying a recipe to ComfyUI.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Open a recipe
|
|
||||||
2. Click "Apply" or "Load in ComfyUI"
|
|
||||||
3. Verify action completes
|
|
||||||
|
|
||||||
**Expected Result**: Recipe applied successfully.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Settings
|
|
||||||
|
|
||||||
### Scenario: Settings Page Load
|
|
||||||
|
|
||||||
**Objective**: Verify settings page displays correctly.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to `http://127.0.0.1:{PORT}/settings`
|
|
||||||
2. Wait for "Settings" title
|
|
||||||
3. Take snapshot
|
|
||||||
|
|
||||||
**Expected Result**: Settings form with various options displayed.
|
|
||||||
|
|
||||||
### Scenario: Change Setting and Restart
|
|
||||||
|
|
||||||
**Objective**: Verify settings persist after restart.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to settings page
|
|
||||||
2. Change a setting (e.g., default view mode)
|
|
||||||
3. Save settings
|
|
||||||
4. Restart server: `python scripts/start_server.py --port {PORT} --restart --wait --timeout 30 --detach`
|
|
||||||
5. Refresh browser page
|
|
||||||
6. Navigate to settings
|
|
||||||
|
|
||||||
**Expected Result**: Changed setting value persists.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Import/Export
|
|
||||||
|
|
||||||
### Scenario: Export Models List
|
|
||||||
|
|
||||||
**Objective**: Verify export functionality.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Navigate to LoRA list
|
|
||||||
2. Click "Export" button
|
|
||||||
3. Select format (JSON/CSV)
|
|
||||||
4. Download file
|
|
||||||
|
|
||||||
**Expected Result**: File downloaded with correct data.
|
|
||||||
|
|
||||||
### Scenario: Import Models
|
|
||||||
|
|
||||||
**Objective**: Verify import functionality.
|
|
||||||
|
|
||||||
**Steps**:
|
|
||||||
1. Prepare import file
|
|
||||||
2. Navigate to import page
|
|
||||||
3. Upload file
|
|
||||||
4. Verify import results
|
|
||||||
|
|
||||||
**Expected Result**: Models imported successfully, confirmation shown.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## API Integration Tests
|
|
||||||
|
|
||||||
### Scenario: Verify API Endpoints
|
|
||||||
|
|
||||||
**Objective**: Verify backend API responds correctly.
|
|
||||||
|
|
||||||
**Test via browser console**:
|
|
||||||
```javascript
|
|
||||||
// List LoRAs
|
|
||||||
fetch('/loras/api/list').then(r => r.json()).then(console.log)
|
|
||||||
|
|
||||||
// Get LoRA details
|
|
||||||
fetch('/loras/api/detail/<id>').then(r => r.json()).then(console.log)
|
|
||||||
|
|
||||||
// Search LoRAs
|
|
||||||
fetch('/loras/api/search?q=test').then(r => r.json()).then(console.log)
|
|
||||||
```
|
|
||||||
|
|
||||||
**Expected Result**: APIs return valid JSON with expected structure.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Console Error Monitoring
|
|
||||||
|
|
||||||
During all tests, monitor browser console for errors:
|
|
||||||
|
|
||||||
```python
|
|
||||||
# Check for JavaScript errors
|
|
||||||
messages = list_console_messages(types=["error"])
|
|
||||||
assert len(messages) == 0, f"Console errors found: {messages}"
|
|
||||||
```
|
|
||||||
|
|
||||||
## Network Request Verification
|
|
||||||
|
|
||||||
Verify key API calls are made:
|
|
||||||
|
|
||||||
```python
|
|
||||||
# List XHR requests
|
|
||||||
requests = list_network_requests(resourceTypes=["xhr", "fetch"])
|
|
||||||
|
|
||||||
# Look for specific endpoints
|
|
||||||
lora_list_requests = [r for r in requests if "/api/list" in r.get("url", "")]
|
|
||||||
assert len(lora_list_requests) > 0, "LoRA list API not called"
|
|
||||||
```
|
|
||||||
@@ -1,215 +0,0 @@
|
|||||||
#!/usr/bin/env python3
|
|
||||||
"""
|
|
||||||
Example E2E test demonstrating LoRa Manager testing workflow.
|
|
||||||
|
|
||||||
This script shows how to:
|
|
||||||
1. Start the standalone server
|
|
||||||
2. Use Chrome DevTools MCP to interact with the UI
|
|
||||||
3. Verify functionality end-to-end
|
|
||||||
|
|
||||||
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 sys
|
|
||||||
|
|
||||||
# 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():
|
|
||||||
"""Run example E2E test flow."""
|
|
||||||
|
|
||||||
print("=" * 60)
|
|
||||||
print("LoRa Manager E2E Test Example")
|
|
||||||
print("=" * 60)
|
|
||||||
|
|
||||||
# Step 1: Start server (detached so it survives the shell)
|
|
||||||
print("\n[1/5] Starting LoRa Manager standalone server...")
|
|
||||||
result = subprocess.run(
|
|
||||||
[sys.executable, "start_server.py", "--port", PORT, "--wait", "--timeout", "30", "--detach"],
|
|
||||||
capture_output=True,
|
|
||||||
text=True,
|
|
||||||
)
|
|
||||||
if result.returncode != 0:
|
|
||||||
print(f"Failed to start server: {result.stderr}")
|
|
||||||
return 1
|
|
||||||
print("Server ready!")
|
|
||||||
|
|
||||||
# Step 2: Open Chrome (manual step - show command)
|
|
||||||
print("\n[2/5] Open Chrome with debug mode:")
|
|
||||||
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)")
|
|
||||||
|
|
||||||
# Step 3: Navigate and verify page load
|
|
||||||
print("\n[3/5] Page Load Verification:")
|
|
||||||
print(
|
|
||||||
f"""
|
|
||||||
MCP Commands to execute:
|
|
||||||
1. navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
|
||||||
2. wait_for(text="LoRAs", timeout=10000)
|
|
||||||
3. snapshot = take_snapshot()
|
|
||||||
"""
|
|
||||||
)
|
|
||||||
|
|
||||||
# Step 4: Test search functionality
|
|
||||||
print("\n[4/5] Search Functionality Test:")
|
|
||||||
print(
|
|
||||||
"""
|
|
||||||
MCP Commands to execute:
|
|
||||||
1. fill(uid="search-input", value="test")
|
|
||||||
2. press_key(key="Enter")
|
|
||||||
3. wait_for(text="Results", timeout=5000)
|
|
||||||
4. result = evaluate_script(function=`
|
|
||||||
() => {
|
|
||||||
const cards = document.querySelectorAll('.lora-card');
|
|
||||||
return { count: cards.length };
|
|
||||||
}
|
|
||||||
`)
|
|
||||||
"""
|
|
||||||
)
|
|
||||||
|
|
||||||
# Step 5: Verify API
|
|
||||||
print("\n[5/5] API Verification:")
|
|
||||||
print(
|
|
||||||
"""
|
|
||||||
MCP Commands to execute:
|
|
||||||
1. api_result = evaluate_script(function=`
|
|
||||||
async () => {
|
|
||||||
const response = await fetch('/loras/api/list');
|
|
||||||
const data = await response.json();
|
|
||||||
return { count: data.length, status: response.status };
|
|
||||||
}
|
|
||||||
`)
|
|
||||||
2. Verify api_result['status'] == 200
|
|
||||||
"""
|
|
||||||
)
|
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
|
||||||
print("Test flow completed!")
|
|
||||||
print("=" * 60)
|
|
||||||
|
|
||||||
return 0
|
|
||||||
|
|
||||||
|
|
||||||
def example_restart_flow():
|
|
||||||
"""Example: Testing configuration change that requires restart."""
|
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
|
||||||
print("Example: Server Restart Flow")
|
|
||||||
print("=" * 60)
|
|
||||||
|
|
||||||
print(
|
|
||||||
f"""
|
|
||||||
Scenario: Change setting and verify after restart
|
|
||||||
|
|
||||||
Steps:
|
|
||||||
1. Navigate to settings page
|
|
||||||
- navigate_page(type="url", url="http://127.0.0.1:{PORT}/settings")
|
|
||||||
|
|
||||||
2. Change a setting (e.g., theme)
|
|
||||||
- fill(uid="theme-select", value="dark")
|
|
||||||
- click(uid="save-settings-button")
|
|
||||||
|
|
||||||
3. Restart server
|
|
||||||
- subprocess.run([python, "start_server.py", "--port", "{PORT}", "--restart", "--wait", "--detach"])
|
|
||||||
|
|
||||||
4. Refresh browser
|
|
||||||
- navigate_page(type="reload", ignoreCache=True)
|
|
||||||
- wait_for(text="LoRAs", timeout=15000)
|
|
||||||
|
|
||||||
5. Verify setting persisted
|
|
||||||
- navigate_page(type="url", url="http://127.0.0.1:{PORT}/settings")
|
|
||||||
- theme = evaluate_script(function="() => document.querySelector('#theme-select').value")
|
|
||||||
- assert theme == "dark"
|
|
||||||
"""
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def example_modal_interaction():
|
|
||||||
"""Example: Testing modal dialog interaction."""
|
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
|
||||||
print("Example: Modal Dialog Interaction")
|
|
||||||
print("=" * 60)
|
|
||||||
|
|
||||||
print(
|
|
||||||
"""
|
|
||||||
Scenario: Add new LoRA via modal
|
|
||||||
|
|
||||||
Steps:
|
|
||||||
1. Open modal
|
|
||||||
- click(uid="add-lora-button")
|
|
||||||
- wait_for(text="Add LoRA", timeout=3000)
|
|
||||||
|
|
||||||
2. Fill form
|
|
||||||
- fill_form(elements=[
|
|
||||||
{"uid": "lora-name", "value": "Test Character"},
|
|
||||||
{"uid": "lora-path", "value": "/models/test.safetensors"},
|
|
||||||
])
|
|
||||||
|
|
||||||
3. Submit
|
|
||||||
- click(uid="modal-submit-button")
|
|
||||||
|
|
||||||
4. Verify success
|
|
||||||
- wait_for(text="Successfully added", timeout=5000)
|
|
||||||
- snapshot = take_snapshot()
|
|
||||||
"""
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def example_network_monitoring():
|
|
||||||
"""Example: Network request monitoring."""
|
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
|
||||||
print("Example: Network Request Monitoring")
|
|
||||||
print("=" * 60)
|
|
||||||
|
|
||||||
print(
|
|
||||||
f"""
|
|
||||||
Scenario: Verify API calls during user interaction
|
|
||||||
|
|
||||||
Steps:
|
|
||||||
1. Clear network log (implicit on navigation)
|
|
||||||
- navigate_page(type="url", url="http://127.0.0.1:{PORT}/loras")
|
|
||||||
|
|
||||||
2. Perform action that triggers API call
|
|
||||||
- fill(uid="search-input", value="character")
|
|
||||||
- press_key(key="Enter")
|
|
||||||
|
|
||||||
3. List network requests
|
|
||||||
- requests = list_network_requests(resourceTypes=["xhr", "fetch"])
|
|
||||||
|
|
||||||
4. Find search API call
|
|
||||||
- search_requests = [r for r in requests if "/api/search" in r.get("url", "")]
|
|
||||||
- assert len(search_requests) > 0, "Search API was not called"
|
|
||||||
|
|
||||||
5. Get request details
|
|
||||||
- if search_requests:
|
|
||||||
details = get_network_request(reqid=search_requests[0]["reqid"])
|
|
||||||
- Verify request method, response status, etc.
|
|
||||||
"""
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
print("LoRa Manager E2E Test Examples\n")
|
|
||||||
print("This script demonstrates E2E testing patterns.\n")
|
|
||||||
print("Note: Actual execution requires Chrome DevTools MCP connection.\n")
|
|
||||||
|
|
||||||
run_test()
|
|
||||||
example_restart_flow()
|
|
||||||
example_modal_interaction()
|
|
||||||
example_network_monitoring()
|
|
||||||
|
|
||||||
print("\n" + "=" * 60)
|
|
||||||
print("All examples shown!")
|
|
||||||
print("=" * 60)
|
|
||||||
@@ -9,7 +9,10 @@ description: Inspect ComfyUI LoRA Manager runtime configuration and local diagno
|
|||||||
|
|
||||||
- Treat runtime state as local user data. Prefer read-only inspection unless the user explicitly asks for mutation.
|
- Treat runtime state as local user data. Prefer read-only inspection unless the user explicitly asks for mutation.
|
||||||
- Never print secret-like settings values. Redact keys containing `key`, `token`, `secret`, `password`, `auth`, or `credential`, including `civitai_api_key`.
|
- Never print secret-like settings values. Redact keys containing `key`, `token`, `secret`, `password`, `auth`, or `credential`, including `civitai_api_key`.
|
||||||
- Resolve paths from the runtime configuration before guessing. In this environment the settings file is normally `/home/miao/.config/ComfyUI-LoRA-Manager/settings.json`, but portable settings can override this through the repository `settings.json`.
|
- Resolve paths from the runtime configuration before guessing. Settings-directory precedence (highest first):
|
||||||
|
1. **Explicit override** — env `LORA_MANAGER_SETTINGS_DIR` or standalone `--settings-path` (also accepted by the inspect script as `--settings-path DIR`). Pins EVERYTHING (`settings.json`, `cache/`, `wildcards/`, `backups/`, `logs/`, `stats/`) under the given directory; bypasses portable mode and the user config dir. Common when inspecting a sandboxed/E2E instance.
|
||||||
|
2. **Portable** — repository `<repo-root>/settings.json` with `"use_portable_settings": true` (or `LORA_MANAGER_PORTABLE=1`): settings dir = `<repo-root>`.
|
||||||
|
3. **Default** — `~/.config/ComfyUI-LoRA-Manager` on this machine (`platformdirs.user_config_dir("ComfyUI-LoRA-Manager", appauthor=False)`).
|
||||||
- Use the active library when selecting per-library caches and paths. Read `active_library` from settings; fall back to `default` if missing.
|
- Use the active library when selecting per-library caches and paths. Read `active_library` from settings; fall back to `default` if missing.
|
||||||
- Normalize and expand `~` before comparing paths. Symlinks are common in this repo.
|
- Normalize and expand `~` before comparing paths. Symlinks are common in this repo.
|
||||||
|
|
||||||
@@ -32,9 +35,17 @@ python .agents/skills/lora-manager-runtime-context/scripts/inspect_runtime_conte
|
|||||||
python .agents/skills/lora-manager-runtime-context/scripts/inspect_runtime_context.py sqlite --db /path/to/cache.sqlite --limit 3
|
python .agents/skills/lora-manager-runtime-context/scripts/inspect_runtime_context.py sqlite --db /path/to/cache.sqlite --limit 3
|
||||||
```
|
```
|
||||||
|
|
||||||
|
To inspect a sandboxed/E2E instance that pins its settings directory:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# --settings-path DIR (or LORA_MANAGER_SETTINGS_DIR) works with every subcommand:
|
||||||
|
python .agents/skills/lora-manager-runtime-context/scripts/inspect_runtime_context.py \
|
||||||
|
--settings-path /tmp/opencode/<plan>-e2e/settings summary
|
||||||
|
```
|
||||||
|
|
||||||
## Runtime Path Rules
|
## Runtime Path Rules
|
||||||
|
|
||||||
- Settings directory: use `py/utils/settings_paths.py`. Default platform path is `platformdirs.user_config_dir("ComfyUI-LoRA-Manager", appauthor=False)`.
|
- Settings directory: resolve via `py/utils/settings_paths.py` — `get_settings_dir()` honors the `LORA_MANAGER_SETTINGS_DIR` / programmatic override first, then portable mode, then `platformdirs.user_config_dir("ComfyUI-LoRA-Manager", appauthor=False)`. The inspect script mirrors this precedence in `resolve_settings_path()`.
|
||||||
- Settings file: `<settings_dir>/settings.json`.
|
- Settings file: `<settings_dir>/settings.json`.
|
||||||
- Cache root: `<settings_dir>/cache`.
|
- Cache root: `<settings_dir>/cache`.
|
||||||
- Canonical cache files:
|
- Canonical cache files:
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from typing import Any
|
|||||||
|
|
||||||
SECRET_PATTERN = re.compile(r"(key|token|secret|password|auth|credential)", re.IGNORECASE)
|
SECRET_PATTERN = re.compile(r"(key|token|secret|password|auth|credential)", re.IGNORECASE)
|
||||||
APP_NAME = "ComfyUI-LoRA-Manager"
|
APP_NAME = "ComfyUI-LoRA-Manager"
|
||||||
|
SETTINGS_DIR_ENV = "LORA_MANAGER_SETTINGS_DIR"
|
||||||
CACHE_SQLITE = {
|
CACHE_SQLITE = {
|
||||||
"model": ("model", "{library}.sqlite"),
|
"model": ("model", "{library}.sqlite"),
|
||||||
"recipe": ("recipe", "{library}.sqlite"),
|
"recipe": ("recipe", "{library}.sqlite"),
|
||||||
@@ -30,6 +31,15 @@ CACHE_JSON = {
|
|||||||
|
|
||||||
def main() -> int:
|
def main() -> int:
|
||||||
parser = argparse.ArgumentParser(description="Inspect LoRA Manager runtime state read-only.")
|
parser = argparse.ArgumentParser(description="Inspect LoRA Manager runtime state read-only.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--settings-path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
metavar="DIR",
|
||||||
|
help="Explicit settings directory (same as LORA_MANAGER_SETTINGS_DIR / "
|
||||||
|
"standalone --settings-path). Overrides portable mode and the default "
|
||||||
|
"user config dir.",
|
||||||
|
)
|
||||||
subparsers = parser.add_subparsers(dest="command", required=True)
|
subparsers = parser.add_subparsers(dest="command", required=True)
|
||||||
|
|
||||||
subparsers.add_parser("summary", help="Print redacted settings and resolved paths.")
|
subparsers.add_parser("summary", help="Print redacted settings and resolved paths.")
|
||||||
@@ -44,6 +54,8 @@ def main() -> int:
|
|||||||
sqlite_parser.add_argument("--limit", type=int, default=3, help="Rows to sample from each user table.")
|
sqlite_parser.add_argument("--limit", type=int, default=3, help="Rows to sample from each user table.")
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
if args.settings_path:
|
||||||
|
os.environ[SETTINGS_DIR_ENV] = args.settings_path
|
||||||
context = build_context()
|
context = build_context()
|
||||||
|
|
||||||
if args.command == "summary":
|
if args.command == "summary":
|
||||||
@@ -78,6 +90,11 @@ def build_context() -> dict[str, Any]:
|
|||||||
|
|
||||||
|
|
||||||
def resolve_settings_path() -> Path:
|
def resolve_settings_path() -> Path:
|
||||||
|
# Explicit override: LORA_MANAGER_SETTINGS_DIR env or --settings-path.
|
||||||
|
explicit = os.environ.get(SETTINGS_DIR_ENV)
|
||||||
|
if explicit:
|
||||||
|
return Path(explicit).expanduser() / "settings.json"
|
||||||
|
|
||||||
repo_root = find_repo_root()
|
repo_root = find_repo_root()
|
||||||
portable = repo_root / "settings.json"
|
portable = repo_root / "settings.json"
|
||||||
if portable.exists():
|
if portable.exists():
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ node_modules/
|
|||||||
coverage/
|
coverage/
|
||||||
.coverage
|
.coverage
|
||||||
model_cache/
|
model_cache/
|
||||||
|
recipe_cache/
|
||||||
|
|
||||||
# agent / dev tooling
|
# agent / dev tooling
|
||||||
.opencode/
|
.opencode/
|
||||||
|
|||||||
@@ -2,6 +2,10 @@
|
|||||||
|
|
||||||
This file provides guidance for agentic coding assistants working in this repository.
|
This file provides guidance for agentic coding assistants working in this repository.
|
||||||
|
|
||||||
|
## Overview
|
||||||
|
|
||||||
|
ComfyUI LoRA Manager is a comprehensive LoRA management system for ComfyUI that combines a Python backend with browser-based widgets. It provides model organization, downloading from CivitAI/CivArchive, recipe management, and one-click workflow integration.
|
||||||
|
|
||||||
## Development Commands
|
## Development Commands
|
||||||
|
|
||||||
### Backend Development
|
### Backend Development
|
||||||
@@ -28,16 +32,21 @@ COVERAGE_FILE=coverage/backend/.coverage pytest \
|
|||||||
--cov=py --cov=standalone \
|
--cov=py --cov=standalone \
|
||||||
--cov-report=term-missing \
|
--cov-report=term-missing \
|
||||||
--cov-report=html:coverage/backend/html \
|
--cov-report=html:coverage/backend/html \
|
||||||
--cov-report=xml:coverage/backend/coverage.xml
|
--cov-report=xml:coverage/backend/coverage.xml \
|
||||||
|
--cov-report=json:coverage/backend/coverage.json
|
||||||
```
|
```
|
||||||
|
|
||||||
### Frontend Development (LoRA Manager Web UI)
|
### Frontend Development (LoRA Manager Web UI)
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
# Install dependencies (root and Vue widgets)
|
||||||
npm install
|
npm install
|
||||||
|
cd vue-widgets && npm install && cd ..
|
||||||
|
|
||||||
npm test # Run all tests (JS + Vue)
|
npm test # Run all tests (JS + Vue)
|
||||||
npm run test:js # Run JS tests only
|
npm run test:js # Run JS tests only
|
||||||
npm run test:watch # Watch mode
|
npm run test:vue # Run Vue widget tests only
|
||||||
|
npm run test:watch # Watch mode (JS tests only)
|
||||||
npm run test:coverage # Generate coverage report
|
npm run test:coverage # Generate coverage report
|
||||||
```
|
```
|
||||||
|
|
||||||
@@ -54,88 +63,210 @@ npm run test:watch # Watch mode
|
|||||||
npm run test:coverage # Generate coverage report
|
npm run test:coverage # Generate coverage report
|
||||||
```
|
```
|
||||||
|
|
||||||
## Python Code Style
|
### Localization
|
||||||
|
|
||||||
### Imports & Formatting
|
```bash
|
||||||
|
# Sync translation keys after UI string updates
|
||||||
|
python scripts/sync_translation_keys.py
|
||||||
|
```
|
||||||
|
|
||||||
|
Locale files are in `locales/` (en, zh-CN, zh-TW, ja, ko, fr, de, es, ru, he).
|
||||||
|
|
||||||
|
After adding keys to `en.json` and syncing, **stop**: the `[TODO: Translate]` placeholders in
|
||||||
|
the other locales are the expected end state during feature development. Do NOT translate
|
||||||
|
proactively — translate only when the feature owner explicitly asks (see
|
||||||
|
`docs/i18n-translation-guidelines.md` §7).
|
||||||
|
|
||||||
|
**Before translating anything, read `docs/i18n-translation-guidelines.md`** — it defines the
|
||||||
|
term conventions (e.g. "Recipe" stays untranslated in French, 配方 in Chinese; model-type and
|
||||||
|
brand names are never translated), per-locale preferred renderings, placeholder rules, and
|
||||||
|
the known confusion hot-spots.
|
||||||
|
|
||||||
|
## Code Style
|
||||||
|
|
||||||
|
### Python
|
||||||
|
|
||||||
|
#### Imports & Formatting
|
||||||
|
|
||||||
- Use `from __future__ import annotations` for forward references
|
- Use `from __future__ import annotations` for forward references
|
||||||
- Group imports: standard library, third-party, local (blank line separated)
|
- Group imports: standard library, third-party, local (blank line separated)
|
||||||
|
- Use `TYPE_CHECKING` guard for type-checking-only imports
|
||||||
- Absolute imports within `py/`: `from ..services import X`
|
- Absolute imports within `py/`: `from ..services import X`
|
||||||
- PEP 8 with 4-space indentation, type hints required
|
- PEP 8 with 4-space indentation, type hints required
|
||||||
|
|
||||||
### Naming Conventions
|
#### Naming Conventions
|
||||||
|
|
||||||
- Files: `snake_case.py`, Classes: `PascalCase`, Functions/vars: `snake_case`
|
- Files: `snake_case.py`, Classes: `PascalCase`, Functions/vars: `snake_case`
|
||||||
- Constants: `UPPER_SNAKE_CASE`, Private: `_protected`, `__mangled`
|
- Constants: `UPPER_SNAKE_CASE`, Private: `_protected`, `__mangled`
|
||||||
|
|
||||||
### Error Handling & Async
|
#### Error Handling & Async
|
||||||
|
|
||||||
- Use `logging.getLogger(__name__)`, define custom exceptions in `py/services/errors.py`
|
- Use `logging.getLogger(__name__)`, define custom exceptions in `py/services/errors.py`
|
||||||
- `async def` for I/O, `@pytest.mark.asyncio` for async tests
|
- `async def` for I/O, `@pytest.mark.asyncio` for async tests
|
||||||
- Singleton with `asyncio.Lock`: see `ModelScanner.get_instance()`
|
- Singleton with `asyncio.Lock`: see `ModelScanner.get_instance()`
|
||||||
- Return `aiohttp.web.json_response` or `web.Response`
|
- Return `aiohttp.web.json_response` or `web.Response`
|
||||||
|
|
||||||
### Testing
|
### JavaScript/TypeScript
|
||||||
|
|
||||||
- `pytest` with `--import-mode=importlib`
|
#### Imports & Modules
|
||||||
- Fixtures in `tests/conftest.py`, use `tmp_path_factory` for isolation
|
|
||||||
- Mark tests needing real paths: `@pytest.mark.no_settings_dir_isolation`
|
|
||||||
- Mock ComfyUI dependencies via conftest patterns
|
|
||||||
|
|
||||||
## JavaScript/TypeScript Code Style
|
|
||||||
|
|
||||||
### Imports & Modules
|
|
||||||
|
|
||||||
- ES modules: `import { app } from "../../scripts/app.js"` for ComfyUI
|
- ES modules: `import { app } from "../../scripts/app.js"` for ComfyUI
|
||||||
- Vue: `import { ref, computed } from 'vue'`, type imports: `import type { Foo }`
|
- Vue: `import { ref, computed } from 'vue'`, type imports: `import type { Foo }`
|
||||||
- Export named functions: `export function foo() {}`
|
- Export named functions: `export function foo() {}`
|
||||||
|
|
||||||
### Naming & Formatting
|
#### Naming & Formatting
|
||||||
|
|
||||||
- camelCase for functions/vars/props, PascalCase for classes
|
- camelCase for functions/vars/props, PascalCase for classes
|
||||||
- Constants: `UPPER_SNAKE_CASE`, Files: `snake_case.js` or `kebab-case.js`
|
- Constants: `UPPER_SNAKE_CASE`, Files: `snake_case.js` or `kebab-case.js`
|
||||||
- 2-space indentation preferred (follow existing file conventions)
|
- 2-space indentation preferred (follow existing file conventions)
|
||||||
- Vue Single File Components: `<script setup lang="ts">` preferred
|
- Vue Single File Components: `<script setup lang="ts">` preferred
|
||||||
|
|
||||||
### Widget Development
|
#### Widget Development
|
||||||
|
|
||||||
|
- Prefer vanilla JS for `web/comfyui/` widgets; avoid framework dependencies (except the Vue widgets in `vue-widgets/`)
|
||||||
- ComfyUI: `app.registerExtension()`, `node.addDOMWidget(name, type, element, options)`
|
- ComfyUI: `app.registerExtension()`, `node.addDOMWidget(name, type, element, options)`
|
||||||
- Event handlers via `addEventListener` or widget callbacks
|
- Event handlers via `addEventListener` or widget callbacks
|
||||||
- Shared utilities: `web/comfyui/utils.js`
|
- Shared utilities: `web/comfyui/utils.js`
|
||||||
- Dual-mode rendering patterns (canvas vs Vue): see `docs/comfyui-dual-mode-widgets.md`
|
- Dual-mode rendering patterns (canvas vs Vue): see `docs/comfyui-dual-mode-widgets.md`
|
||||||
|
|
||||||
### Vue Composables Pattern
|
#### Vue Composables Pattern
|
||||||
|
|
||||||
- Use composition API: `useXxxState(widget)`, return reactive refs and methods
|
- Use composition API: `useXxxState(widget)`, return reactive refs and methods
|
||||||
- Guard restoration loops with flag: `let isRestoring = false`
|
- Guard restoration loops with flag: `let isRestoring = false`
|
||||||
- Build config from state: `const buildConfig = (): Config => { ... }`
|
- Build config from state: `const buildConfig = (): Config => { ... }`
|
||||||
|
|
||||||
## Architecture Patterns
|
## Architecture
|
||||||
|
|
||||||
|
### Dual Mode Operation
|
||||||
|
|
||||||
|
The system runs in two modes:
|
||||||
|
- **ComfyUI plugin mode**: Integrates with ComfyUI's PromptServer, uses `folder_paths` for model discovery
|
||||||
|
- **Standalone mode**: `standalone.py` mocks ComfyUI dependencies, reads paths from `settings.json`
|
||||||
|
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
||||||
|
|
||||||
|
### Backend Entry Points
|
||||||
|
|
||||||
|
- `__init__.py` — ComfyUI plugin entry: registers nodes via `NODE_CLASS_MAPPINGS`, sets `WEB_DIRECTORY`, calls `LoraManager.add_routes()`
|
||||||
|
- `standalone.py` — Standalone server: mocks `folder_paths` and node modules, starts aiohttp server
|
||||||
|
- `py/lora_manager.py` — Main `LoraManager` class that registers all HTTP routes
|
||||||
|
|
||||||
### Service Layer
|
### Service Layer
|
||||||
|
|
||||||
- `ServiceRegistry` singleton for DI, services use `get_instance()` classmethod
|
- `ServiceRegistry` singleton for DI, services use `get_instance()` classmethod
|
||||||
|
- `BaseModelService` abstract base → `LoraService`, `CheckpointService`, `EmbeddingService`
|
||||||
|
- `ModelScanner` base → `LoraScanner`, `CheckpointScanner`, `EmbeddingScanner` for file discovery with hash-based deduplication
|
||||||
|
- `PersistentModelCache` (SQLite) for metadata persistence
|
||||||
|
- `MetadataSyncService` — background sync from CivitAI/CivArchive APIs
|
||||||
|
- `SettingsManager` — settings with schema migration support
|
||||||
|
- `WebSocketManager` — real-time progress broadcasting
|
||||||
|
- `ModelServiceFactory` — creates the right service for each model type
|
||||||
|
- Use cases in `py/services/use_cases/` orchestrate complex business logic (auto-organize, bulk refresh, downloads)
|
||||||
- Separate scanners (discovery) from services (business logic)
|
- Separate scanners (discovery) from services (business logic)
|
||||||
- Handlers in `py/routes/handlers/` are pure functions with deps as params
|
- Handlers in `py/routes/handlers/` are pure functions with deps as params
|
||||||
|
|
||||||
### Model Types & Routes
|
### Model Types & Routes
|
||||||
|
|
||||||
- `BaseModelService` base for LoRA, Checkpoint, Embedding
|
- API endpoints follow `/loras/*`, `/checkpoints/*`, `/embeddings/*`, `/other/*` patterns
|
||||||
- `ModelScanner` for file discovery, hash deduplication
|
- Route registrars organize endpoints by domain: `ModelRouteRegistrar`, `RecipeRouteRegistrar`, etc.
|
||||||
- `PersistentModelCache` (SQLite) for persistence
|
- Request handlers in `py/routes/handlers/` implement route logic
|
||||||
- Route registrars: `ModelRouteRegistrar`, endpoints: `/loras/*`, `/checkpoints/*`, `/embeddings/*`
|
- All routes use aiohttp, return `web.json_response` or `web.Response`
|
||||||
- WebSocket via `WebSocketManager` for real-time updates
|
- Endpoints consumed by the companion browser extension (lm-civitai-extension)
|
||||||
|
MUST also accept `GET` with query-string params: the extension is GET-only by
|
||||||
|
convention (see its AGENTS.md), even for state-changing operations such as
|
||||||
|
`GET /api/lm/recipe/{recipe_id}/reimport`
|
||||||
|
|
||||||
### Recipe System
|
### Recipe System
|
||||||
|
|
||||||
- Base: `py/recipes/base.py`, Enrichment: `RecipeEnrichmentService`
|
- Base: `py/recipes/base.py`, Enrichment: `RecipeEnrichmentService` in `py/recipes/enrichment.py`
|
||||||
- Parsers: `py/recipes/parsers/`
|
- Parsers: `py/recipes/parsers/` for PNG metadata, JSON, and workflow formats
|
||||||
|
|
||||||
|
### Custom Nodes
|
||||||
|
|
||||||
|
- Location: `py/nodes/`, all nodes registered in `__init__.py`
|
||||||
|
- Each node class has a `NAME` class attribute used as key in `NODE_CLASS_MAPPINGS`
|
||||||
|
- Standard ComfyUI node pattern: `INPUT_TYPES()` classmethod, `RETURN_TYPES`, `FUNCTION`
|
||||||
|
|
||||||
|
### Configuration
|
||||||
|
|
||||||
|
- `py/config.py` manages folder paths for models and handles symlink mappings
|
||||||
|
- Auto-saves paths to `settings.json` in ComfyUI mode
|
||||||
|
- `settings.json.example` is intentionally minimal (see Important Notes); all
|
||||||
|
other defaults live in `DEFAULT_SETTINGS` (`py/services/settings_manager.py`)
|
||||||
|
- **`folder_paths` vs `extra_folder_paths` — different purposes, do not conflate:**
|
||||||
|
- `folder_paths` (primary model roots): in ComfyUI plugin mode these come
|
||||||
|
from the ComfyUI host; in standalone mode they are the ONLY source of
|
||||||
|
model library paths and are currently edited by hand in `settings.json`.
|
||||||
|
- `extra_folder_paths` is a **ComfyUI-plugin-mode feature**: paths visible
|
||||||
|
ONLY to LoRA Manager, not to ComfyUI. Its motivation is that a very large
|
||||||
|
model library slows ComfyUI itself down, while LoRA Manager handles large
|
||||||
|
libraries without performance issues — so users keep ComfyUI's library
|
||||||
|
small and add the bulk via `extra_folder_paths`.
|
||||||
|
|
||||||
|
### Frontend UI Architecture
|
||||||
|
|
||||||
|
#### 1. LoRA Manager Web UI
|
||||||
|
- Location: `./static/` (JS/CSS) and `./templates/` (HTML)
|
||||||
|
- Tech: Vanilla JS + CSS, served by the hosting server (ComfyUI app in plugin mode, `standalone.py` in standalone mode)
|
||||||
|
- Tests: `tests/frontend/**/*.test.js` (vitest + jsdom)
|
||||||
|
|
||||||
|
#### 2. ComfyUI Custom Node Widgets
|
||||||
|
- Location: `./web/comfyui/` (Vanilla JS) + `./vue-widgets/` (Vue)
|
||||||
|
- Primary styles: `./web/comfyui/lm_styles.css` (NOT `./static/css/`)
|
||||||
|
- Vue widgets: Vue 3 + TypeScript + PrimeVue + vue-i18n, e.g. `LoraPoolWidget`, `LoraRandomizerWidget`, `LoraCyclerWidget`, `AutocompleteTextWidget`
|
||||||
|
- Vue builds to `./web/comfyui/vue-widgets/`; auto-built on ComfyUI startup via `py/vue_widget_builder.py`, typecheck via `vue-tsc`
|
||||||
|
- Widget registration: `app.registerExtension()` and `getCustomWidgets` hooks; `node.addDOMWidget(...)` embeds HTML in LiteGraph nodes
|
||||||
|
- See `docs/dom_widget_dev_guide.md` for the DOMWidget development guide
|
||||||
|
|
||||||
|
## Testing
|
||||||
|
|
||||||
|
### Backend (pytest)
|
||||||
|
|
||||||
|
- Config in `pytest.ini`: `--import-mode=importlib`, testpaths=`tests`
|
||||||
|
- Fixtures in `tests/conftest.py` mock ComfyUI dependencies; use `tmp_path_factory` for isolation
|
||||||
|
- Markers: `@pytest.mark.asyncio`, `@pytest.mark.no_settings_dir_isolation` (tests needing real settings paths)
|
||||||
|
|
||||||
|
### Frontend (vitest)
|
||||||
|
|
||||||
|
- Vanilla JS tests: `tests/frontend/**/*.test.js` with jsdom; setup in `tests/frontend/setup.js`
|
||||||
|
- Vue widget tests: `vue-widgets/tests/**/*.test.ts` with jsdom + `@vue/test-utils`
|
||||||
|
|
||||||
|
### UI Verification (manual default)
|
||||||
|
|
||||||
|
UI/layout changes are verified by the user by eye — do NOT spin up a sandbox,
|
||||||
|
standalone server, or browser automation to "prove" a visual fix. Ask the user to
|
||||||
|
look instead. The full browser E2E ceremony (server + Chrome DevTools MCP +
|
||||||
|
screenshots) is slow, token-heavy, and fragile; reserve it for genuine
|
||||||
|
server+browser integration bugs, and only when the user explicitly agrees.
|
||||||
|
|
||||||
|
If a cross-layer issue ever needs a live server, the sandboxed helpers live in
|
||||||
|
`scripts/e2e/` (`start_server.py`, `wait_for_server.py`). Non-negotiable rules:
|
||||||
|
|
||||||
|
- Always launch with `--settings-path <sandbox>/settings` and sandboxed
|
||||||
|
`folder_paths` under `/tmp` — the repo folder is the real plugin folder and a
|
||||||
|
`settings.json` there is read by the live instance. Never touch real config or
|
||||||
|
real model libraries.
|
||||||
|
- Never kill a process you did not start; `start_server.py` tracks its own PIDs
|
||||||
|
via pidfile and refuses to touch unrelated processes on the port.
|
||||||
|
- Abort after ~30 minutes or 3 consecutive tool failures; report `BLOCKED` with
|
||||||
|
observed state instead of retrying blindly. Clean up sandbox and server after.
|
||||||
|
|
||||||
|
## Key Integration Points
|
||||||
|
|
||||||
|
- **Settings:** Stored in the user config directory (via `platformdirs`) or portable mode (`"use_portable_settings": true`)
|
||||||
|
- **CivitAI/CivArchive:** API clients for metadata sync and model downloads; CivitAI API key stored in settings
|
||||||
|
- **Symlinks:** Config scans symlinks to map virtual→physical paths; fingerprinting prevents redundant rescans
|
||||||
|
- **WebSocket:** Broadcasts real-time progress for downloads, scans, and metadata sync
|
||||||
|
- **Model scanning flow:** Walk folders → compute hashes → deduplicate → extract safetensors metadata → cache in SQLite → background CivitAI sync → WebSocket broadcast
|
||||||
|
|
||||||
## Important Notes
|
## Important Notes
|
||||||
|
|
||||||
- ALWAYS use English for comments (per copilot-instructions.md)
|
- ALWAYS use English for comments (per copilot-instructions.md)
|
||||||
- Dual mode: ComfyUI plugin (folder_paths) vs standalone (settings.json)
|
- **`settings.json.example` must stay minimal**: only `use_portable_settings`,
|
||||||
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
`civitai_api_key`, and the four core `folder_paths` keys (`loras`,
|
||||||
|
`checkpoints`, `unet`, `embeddings`). Do NOT add optional/default keys
|
||||||
|
(model-category folders, `default_*_root`, `auto_organize_exclusions`, etc.)
|
||||||
|
to this file unless the user explicitly asks for it. Defaults belong in
|
||||||
|
`DEFAULT_SETTINGS` in `py/services/settings_manager.py`.
|
||||||
- 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
|
**Business paths vs real paths**: All stored paths and operation routing use the
|
||||||
@@ -143,23 +274,4 @@ npm run test:coverage # Generate coverage report
|
|||||||
resolved. `os.path.realpath` is only for scanner dedup and the symlink cache.
|
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
|
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
|
containment check MUST use the business path (i.e. `os.path.abspath`, not
|
||||||
`realpath`).
|
`realpath`).
|
||||||
|
|
||||||
## Git / Commit Messages
|
|
||||||
|
|
||||||
- Follow the style of recent repository commits when writing commit messages
|
|
||||||
- Prefer the repo's existing `feat(...)`, `fix(...)`, `chore:` style where applicable
|
|
||||||
- If the user has provided a GitHub issue link or issue ID for the task, mention that issue in the commit message, for example `(#871)`
|
|
||||||
- When unrelated local changes exist, stage and commit only the files relevant to the requested task
|
|
||||||
|
|
||||||
## Frontend UI Architecture
|
|
||||||
|
|
||||||
### 1. LoRA Manager Web UI
|
|
||||||
- Location: `./static/` and `./templates/`
|
|
||||||
- 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
|
|
||||||
|
|
||||||
### 2. ComfyUI Custom Node Widgets
|
|
||||||
- Location: `./web/comfyui/` (Vanilla JS) + `./vue-widgets/` (Vue)
|
|
||||||
- Primary styles: `./web/comfyui/lm_styles.css` (NOT `./static/css/`)
|
|
||||||
- Vue builds to `./web/comfyui/vue-widgets/`, typecheck via `vue-tsc`
|
|
||||||
@@ -1,189 +0,0 @@
|
|||||||
# CLAUDE.md
|
|
||||||
|
|
||||||
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
|
|
||||||
|
|
||||||
## Overview
|
|
||||||
|
|
||||||
ComfyUI LoRA Manager is a comprehensive LoRA management system for ComfyUI that combines a Python backend with browser-based widgets. It provides model organization, downloading from CivitAI/CivArchive, recipe management, and one-click workflow integration.
|
|
||||||
|
|
||||||
## Development Commands
|
|
||||||
|
|
||||||
### Backend
|
|
||||||
|
|
||||||
```bash
|
|
||||||
pip install -r requirements.txt
|
|
||||||
pip install -r requirements-dev.txt
|
|
||||||
|
|
||||||
# Run standalone server (port 8188 by default)
|
|
||||||
python standalone.py --port 8188
|
|
||||||
|
|
||||||
# Run all backend tests
|
|
||||||
pytest
|
|
||||||
|
|
||||||
# Run specific test file or function
|
|
||||||
pytest tests/test_recipes.py
|
|
||||||
pytest tests/test_recipes.py::test_function_name
|
|
||||||
|
|
||||||
# Run backend tests with coverage
|
|
||||||
COVERAGE_FILE=coverage/backend/.coverage pytest \
|
|
||||||
--cov=py \
|
|
||||||
--cov=standalone \
|
|
||||||
--cov-report=term-missing \
|
|
||||||
--cov-report=html:coverage/backend/html \
|
|
||||||
--cov-report=xml:coverage/backend/coverage.xml \
|
|
||||||
--cov-report=json:coverage/backend/coverage.json
|
|
||||||
```
|
|
||||||
|
|
||||||
### Frontend
|
|
||||||
|
|
||||||
There are three test suites run by `npm test`: vanilla JS tests (vitest at root) and Vue widget tests (`vue-widgets/` vitest).
|
|
||||||
|
|
||||||
```bash
|
|
||||||
npm install
|
|
||||||
cd vue-widgets && npm install && cd ..
|
|
||||||
|
|
||||||
# Run all frontend tests (JS + Vue)
|
|
||||||
npm test
|
|
||||||
|
|
||||||
# Run only vanilla JS tests
|
|
||||||
npm run test:js
|
|
||||||
|
|
||||||
# Run only Vue widget tests
|
|
||||||
npm run test:vue
|
|
||||||
|
|
||||||
# Watch mode (JS tests only)
|
|
||||||
npm run test:watch
|
|
||||||
|
|
||||||
# Frontend coverage
|
|
||||||
npm run test:coverage
|
|
||||||
|
|
||||||
# Build Vue widgets (output to web/comfyui/vue-widgets/)
|
|
||||||
cd vue-widgets && npm run build
|
|
||||||
|
|
||||||
# Vue widget dev mode (watch + rebuild)
|
|
||||||
cd vue-widgets && npm run dev
|
|
||||||
|
|
||||||
# Typecheck Vue widgets
|
|
||||||
cd vue-widgets && npm run typecheck
|
|
||||||
```
|
|
||||||
|
|
||||||
### Localization
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# Sync translation keys after UI string updates
|
|
||||||
python scripts/sync_translation_keys.py
|
|
||||||
```
|
|
||||||
|
|
||||||
Locale files are in `locales/` (en, zh-CN, zh-TW, ja, ko, fr, de, es, ru, he).
|
|
||||||
|
|
||||||
## Architecture
|
|
||||||
|
|
||||||
### Dual Mode Operation
|
|
||||||
|
|
||||||
The system runs in two modes:
|
|
||||||
- **ComfyUI plugin mode**: Integrates with ComfyUI's PromptServer, uses `folder_paths` for model discovery
|
|
||||||
- **Standalone mode**: `standalone.py` mocks ComfyUI dependencies, reads paths from `settings.json`
|
|
||||||
- Detection: `os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"`
|
|
||||||
|
|
||||||
### Backend (Python)
|
|
||||||
|
|
||||||
**Entry points:**
|
|
||||||
- `__init__.py` — ComfyUI plugin entry: registers nodes via `NODE_CLASS_MAPPINGS`, sets `WEB_DIRECTORY`, calls `LoraManager.add_routes()`
|
|
||||||
- `standalone.py` — Standalone server: mocks `folder_paths` and node modules, starts aiohttp server
|
|
||||||
- `py/lora_manager.py` — Main `LoraManager` class that registers all HTTP routes
|
|
||||||
|
|
||||||
**Service layer** (`py/services/`):
|
|
||||||
- `ServiceRegistry` singleton for dependency injection; services follow `get_instance()` singleton pattern
|
|
||||||
- `BaseModelService` abstract base → `LoraService`, `CheckpointService`, `EmbeddingService`
|
|
||||||
- `ModelScanner` base → `LoraScanner`, `CheckpointScanner`, `EmbeddingScanner` for file discovery with hash-based deduplication
|
|
||||||
- `PersistentModelCache` — SQLite-based metadata cache
|
|
||||||
- `MetadataSyncService` — Background sync from CivitAI/CivArchive APIs
|
|
||||||
- `SettingsManager` — Settings with schema migration support
|
|
||||||
- `WebSocketManager` — Real-time progress broadcasting
|
|
||||||
- `ModelServiceFactory` — Creates the right service for each model type
|
|
||||||
- Use cases in `py/services/use_cases/` orchestrate complex business logic (auto-organize, bulk refresh, downloads)
|
|
||||||
|
|
||||||
**Routes** (`py/routes/`):
|
|
||||||
- Route registrars organize endpoints by domain: `ModelRouteRegistrar`, `RecipeRouteRegistrar`, etc.
|
|
||||||
- Request handlers in `py/routes/handlers/` implement route logic
|
|
||||||
- API endpoints follow `/loras/*`, `/checkpoints/*`, `/embeddings/*` patterns
|
|
||||||
- All routes use aiohttp, return `web.json_response` or `web.Response`
|
|
||||||
|
|
||||||
**Recipe system** (`py/recipes/`):
|
|
||||||
- `base.py` — Recipe metadata structure
|
|
||||||
- `enrichment.py` — Enriches recipes with model metadata
|
|
||||||
- `parsers/` — Parsers for PNG metadata, JSON, and workflow formats
|
|
||||||
|
|
||||||
**Custom nodes** (`py/nodes/`):
|
|
||||||
- Each node class has a `NAME` class attribute used as key in `NODE_CLASS_MAPPINGS`
|
|
||||||
- Standard ComfyUI node pattern: `INPUT_TYPES()` classmethod, `RETURN_TYPES`, `FUNCTION`
|
|
||||||
- All nodes registered in `__init__.py`
|
|
||||||
|
|
||||||
**Configuration** (`py/config.py`):
|
|
||||||
- Manages folder paths for models, handles symlink mappings
|
|
||||||
- Auto-saves paths to settings.json in ComfyUI mode
|
|
||||||
|
|
||||||
### Frontend — Two Distinct UI Systems
|
|
||||||
|
|
||||||
#### 1. Standalone Manager Web UI
|
|
||||||
- **Location:** `static/` (JS/CSS) and `templates/` (HTML)
|
|
||||||
- **Tech:** Vanilla JS + CSS, served by standalone server
|
|
||||||
- **Structure:** `static/js/core.js` (shared), `loras.js`, `checkpoints.js`, `embeddings.js`, `recipes.js`, `statistics.js`
|
|
||||||
- **Tests:** `tests/frontend/**/*.test.js` (vitest + jsdom)
|
|
||||||
|
|
||||||
#### 2. ComfyUI Custom Node Widgets
|
|
||||||
- **Vanilla JS widgets:** `web/comfyui/*.js` — ES modules extending ComfyUI's LiteGraph UI
|
|
||||||
- `loras_widget.js` / `loras_widget_events.js` — Main LoRA selection widget
|
|
||||||
- `autocomplete.js` — Trigger word and embedding autocomplete
|
|
||||||
- `preview_tooltip.js` — Model card preview tooltips
|
|
||||||
- `top_menu_extension.js` — "Launch LoRA Manager" menu item
|
|
||||||
- `utils.js` — Shared utilities and API helpers
|
|
||||||
- Widget styling in `web/comfyui/lm_styles.css` (NOT `static/css/`)
|
|
||||||
- **Vue widgets:** `vue-widgets/src/` → built to `web/comfyui/vue-widgets/`
|
|
||||||
- Vue 3 + TypeScript + PrimeVue + vue-i18n
|
|
||||||
- Vite build with CSS-injected-by-JS plugin
|
|
||||||
- Components: `LoraPoolWidget`, `LoraRandomizerWidget`, `LoraCyclerWidget`, `AutocompleteTextWidget`
|
|
||||||
- Auto-built on ComfyUI startup via `py/vue_widget_builder.py`
|
|
||||||
- Tests: `vue-widgets/tests/**/*.test.ts` (vitest)
|
|
||||||
|
|
||||||
**Widget registration pattern:**
|
|
||||||
- Widgets use `app.registerExtension()` and `getCustomWidgets` hooks
|
|
||||||
- `node.addDOMWidget(name, type, element, options)` embeds HTML in LiteGraph nodes
|
|
||||||
- See `docs/dom_widget_dev_guide.md` for DOMWidget development guide
|
|
||||||
|
|
||||||
## Code Style
|
|
||||||
|
|
||||||
**Python:**
|
|
||||||
- PEP 8, 4-space indentation, English comments only
|
|
||||||
- Use `from __future__ import annotations` for forward references
|
|
||||||
- Use `TYPE_CHECKING` guard for type-checking-only imports
|
|
||||||
- Loggers via `logging.getLogger(__name__)`
|
|
||||||
- Custom exceptions in `py/services/errors.py`
|
|
||||||
- Async patterns: `async def` for I/O, `@pytest.mark.asyncio` for async tests
|
|
||||||
- Singleton pattern with class-level `asyncio.Lock` (see `ModelScanner.get_instance()`)
|
|
||||||
|
|
||||||
**JavaScript:**
|
|
||||||
- ES modules, camelCase functions/variables, PascalCase classes
|
|
||||||
- Widget files use `*_widget.js` suffix
|
|
||||||
- Prefer vanilla JS for `web/comfyui/` widgets, avoid framework dependencies (except Vue widgets)
|
|
||||||
|
|
||||||
## Testing
|
|
||||||
|
|
||||||
**Backend (pytest):**
|
|
||||||
- Config in `pytest.ini`: `--import-mode=importlib`, testpaths=`tests`
|
|
||||||
- Fixtures in `tests/conftest.py` handle ComfyUI dependency mocking
|
|
||||||
- Markers: `@pytest.mark.asyncio`, `@pytest.mark.no_settings_dir_isolation`
|
|
||||||
- Uses `tmp_path_factory` for directory isolation
|
|
||||||
|
|
||||||
**Frontend (vitest):**
|
|
||||||
- Vanilla JS tests: `tests/frontend/**/*.test.js` with jsdom
|
|
||||||
- Vue widget tests: `vue-widgets/tests/**/*.test.ts` with jsdom + @vue/test-utils
|
|
||||||
- Setup in `tests/frontend/setup.js`
|
|
||||||
|
|
||||||
## Key Integration Points
|
|
||||||
|
|
||||||
- **Settings:** Stored in user directory (via `platformdirs`) or portable mode (`"use_portable_settings": true`)
|
|
||||||
- **CivitAI/CivArchive:** API clients for metadata sync and model downloads; CivitAI API key in settings
|
|
||||||
- **Symlink handling:** Config scans symlinks to map virtual→physical paths; fingerprinting prevents redundant rescans
|
|
||||||
- **WebSocket:** Broadcasts real-time progress for downloads, scans, and metadata sync
|
|
||||||
- **Model scanning flow:** Walk folders → compute hashes → deduplicate → extract safetensors metadata → cache in SQLite → background CivitAI sync → WebSocket broadcast
|
|
||||||
-10
@@ -3,8 +3,6 @@ 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
|
||||||
@@ -42,12 +40,6 @@ 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
|
||||||
@@ -87,8 +79,6 @@ 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,
|
||||||
|
|||||||
+326
-283
File diff suppressed because it is too large
Load Diff
+124
-8
@@ -62,27 +62,143 @@ Environment variable overrides: `LLM_API_KEY`, `LLM_MODEL`, `LLM_API_BASE`, `LLM
|
|||||||
|
|
||||||
### enrich_hf_metadata
|
### enrich_hf_metadata
|
||||||
|
|
||||||
Enriches HuggingFace-downloaded models with metadata extracted by an LLM from the HF model card.
|
Enriches models linked to an external model site with metadata extracted by an LLM from the site's model card (README).
|
||||||
|
|
||||||
**Entry point**: Right-click context menu → "Enrich Metadata (Agent)"
|
**Entry point**: Right-click context menu → "Enrich Metadata with AI"
|
||||||
|
|
||||||
|
**Supported model sources**:
|
||||||
|
|
||||||
|
| Platform | Link | AI enrichment | Direct download |
|
||||||
|
| --- | --- | --- | --- |
|
||||||
|
| Hugging Face | yes | yes | yes |
|
||||||
|
| ModelScope (`modelscope.cn`) | yes | yes | yes |
|
||||||
|
| ModelScope International (`modelscope.ai`) | yes | yes | yes |
|
||||||
|
| TensorArt | yes | no (see below) | no |
|
||||||
|
|
||||||
|
`modelscope.cn` and `modelscope.ai` are **separate catalogues, not mirrors** — a
|
||||||
|
repository published on one is routinely absent from the other — so each is
|
||||||
|
registered as its own source (`ModelScopeSource` / `ModelScopeIntlSource` in
|
||||||
|
`py/services/model_sources/modelscope.py`). The host therefore decides which
|
||||||
|
API and CDN a model resolves against, and the two deployments get separate
|
||||||
|
version groups (`ms:` / `msai:`) and default download directories. Keep the two
|
||||||
|
tables in `modelSourceHelpers.js` and `registry.py` in step when adding a site.
|
||||||
|
|
||||||
|
TensorArt is link-only: `tensor.art` sits behind a Cloudflare managed challenge and its internal API requires session authorization, so the backend cannot read its model pages. Linking still stores the canonical page URL and the "View on TensorArt" link works.
|
||||||
|
|
||||||
**What it does**:
|
**What it does**:
|
||||||
1. Reads the model's `.metadata.json` to get the `hf_url`
|
1. Reads the model's `.metadata.json` to get the source (`source_platform` + `source_url`, or the legacy `hf_url`)
|
||||||
2. Fetches the README.md from the HuggingFace repository
|
2. Fetches the model card through the provider in `py/services/model_sources/` — the README via `fetch_model_card()`, plus any extras the site keeps outside it via `fetch_model_card_context()`
|
||||||
3. Sends the README + local metadata to the LLM for structured extraction
|
3. Sends the README + site-provided extras + local metadata to the LLM for structured extraction
|
||||||
4. Writes extracted fields to `.metadata.json`:
|
4. Writes extracted fields to `.metadata.json`:
|
||||||
- `base_model` — only if current value is empty
|
- `base_model` — only if current value is empty
|
||||||
- `trainedWords` — trigger words (LoRA only, if none exist)
|
- `trainedWords` — trigger words (LoRA only, if none exist)
|
||||||
- `modelDescription` — concise summary (if none exists)
|
- `modelDescription` — the site's author description (if any) followed by the README rendered as HTML
|
||||||
- `tags` — merged with existing tags, deduplicated
|
- `tags` — merged with existing tags, deduplicated
|
||||||
|
- `civitai.images` — example images
|
||||||
- `metadata_source` — audit trail: `agent:enrich_hf_metadata`
|
- `metadata_source` — audit trail: `agent:enrich_hf_metadata`
|
||||||
- `llm_enriched_at` — ISO timestamp
|
- `llm_enriched_at` — ISO timestamp
|
||||||
5. Downloads and optimizes preview image (if LLM found one in the README)
|
5. Downloads and optimizes a preview image, using the per-file example image the
|
||||||
|
site publishes when the README has none
|
||||||
6. Updates the scanner cache
|
6. Updates the scanner cache
|
||||||
7. Broadcasts WebSocket progress events
|
7. Broadcasts WebSocket progress events
|
||||||
|
|
||||||
|
#### Site-provided card extras (`fetch_model_card_context`)
|
||||||
|
|
||||||
|
A model card is not always just `README.md`. ModelScope keeps the author's
|
||||||
|
summary (`Description`), the site-curated tags (`OfficialTags`), and — per
|
||||||
|
published version — the model filenames together with that file's example
|
||||||
|
images (`MuseInfo.versions[].coverImages`) and trigger words in its
|
||||||
|
model-detail API. AIGC repositories there often ship an auto-generated
|
||||||
|
boilerplate README and put everything useful in `Description`, so reading only
|
||||||
|
the README yields almost nothing.
|
||||||
|
|
||||||
|
Providers opt in by overriding `ModelSource.fetch_model_card_context()`, which
|
||||||
|
returns a `ModelCardContext`. The wanted file is identified by its sha256 when
|
||||||
|
the caller knows it (the scanner already records one) and by **basename**
|
||||||
|
otherwise, so each checkpoint in a collection repo gets its own images — and
|
||||||
|
keeps getting them after the user renames the weights, which is the only
|
||||||
|
identifier a rename cannot invalidate. Sites with no such extras inherit an
|
||||||
|
empty context, and the pipeline behaves exactly as before.
|
||||||
|
|
||||||
|
The README and the repository metadata describe the whole repository, not one
|
||||||
|
file, so `execute_skill()` creates a `ModelSourceCache` for the duration of a
|
||||||
|
run and passes it down. Enriching the eight checkpoints of one ModelScope
|
||||||
|
repository costs two HTTP requests instead of sixteen; only the per-file
|
||||||
|
selection is redone for each file. Nothing is cached across runs, and download
|
||||||
|
URLs never go through it.
|
||||||
|
|
||||||
|
#### Deterministic data is applied whether or not an LLM is configured
|
||||||
|
|
||||||
|
`AgentService._load_source_card()` runs for every source-backed enrichment, and
|
||||||
|
the post-processor applies what it returns before the LLM output is merged. A
|
||||||
|
user with **no** provider configured therefore still gets the author summary,
|
||||||
|
the example images, the preview, the site-curated tags, the trigger words and
|
||||||
|
the README rendered as the model description.
|
||||||
|
|
||||||
|
The LLM is always consulted when one is configured — invoking **Enrich Metadata
|
||||||
|
with AI** must call the provider every time, and the site data is never treated
|
||||||
|
as a reason to skip it. The deterministic values act as fallbacks that fill
|
||||||
|
gaps the LLM leaves behind:
|
||||||
|
|
||||||
|
| Field | Deterministic source | LLM role |
|
||||||
|
| --- | --- | --- |
|
||||||
|
| `model_name` | site display name (`Name`), written only while the value is still the file stem | — |
|
||||||
|
| `modelDescription` | author summary + README as HTML | — |
|
||||||
|
| `civitai.name` | the matched version's label (`modelVersion.showName`) | — |
|
||||||
|
| `civitai.images` | site example images, then README images | — |
|
||||||
|
| `preview_url` | first available example image | may propose one from the README |
|
||||||
|
| `tags` | site-curated tags, always merged in | proposes additional content tags |
|
||||||
|
| `civitai.description` | author summary | richer 1-2 sentence summary wins |
|
||||||
|
| `base_model` | site hints resolved against the canonical vocabulary (`py/services/agent/base_model_resolver.py`) | mapping it is the LLM's job; the resolver only fills in when the LLM returns nothing |
|
||||||
|
| `trainedWords` | per-file site trigger words, then YAML `instance_prompt` | primary extraction |
|
||||||
|
| `usage_tips` | regex over an explicitly stated strength range | primary extraction |
|
||||||
|
| `notes` | — | LLM-only |
|
||||||
|
|
||||||
|
Models with no source, an unknown source, or a source without model-card access (TensorArt) are skipped with an explicit reason and counted in the run summary.
|
||||||
|
|
||||||
**Model types**: LoRA, Checkpoint, Embedding
|
**Model types**: LoRA, Checkpoint, Embedding
|
||||||
|
|
||||||
|
### Download-time hydration
|
||||||
|
|
||||||
|
The same deterministic mapping runs automatically when a model is downloaded
|
||||||
|
from a model source, so a ModelScope or Hugging Face download lands with the
|
||||||
|
populated card a CivitAI download produces instead of a bare filename and
|
||||||
|
hash. Nothing needs to be triggered by hand and no provider is called.
|
||||||
|
|
||||||
|
`py/services/model_sources/hydration.py` owns this path:
|
||||||
|
|
||||||
|
* `_save_source_metadata()` in `py/routes/handlers/model_source_handlers.py`
|
||||||
|
creates the sidecar (hash, source link, scanner-cache entry) and then calls
|
||||||
|
`hydrate_from_source()`. It also runs for a file that was already on disk, so
|
||||||
|
models downloaded before this existed get topped up on the next attempt.
|
||||||
|
* Metadata is created through the **owning scanner**
|
||||||
|
(`scanner._create_default_metadata()`) rather than
|
||||||
|
`MetadataManager.create_default_metadata()`, so the per-type lazy-hash rule
|
||||||
|
applies: `CheckpointScanner` and `OtherScanner` store
|
||||||
|
`hash_status="pending"` with an empty `sha256` for their multi-GB files, and
|
||||||
|
the generic helper would read a 10 GB checkpoint end to end inside the
|
||||||
|
download request. Hydration copes with the empty hash — `_matching_versions()`
|
||||||
|
falls back to the repository basename, which the download just wrote.
|
||||||
|
* Hydration reuses `PostProcessor` with an empty `llm_output`, so the two paths
|
||||||
|
cannot drift apart. It reports `metadata_source = "source:<platform>"` rather
|
||||||
|
than the skill's `agent:enrich_hf_metadata`, and — because no provider ran —
|
||||||
|
it does not stamp `llm_enriched_at`.
|
||||||
|
* `model_name` is only written while it still equals the file stem: once a user
|
||||||
|
renames a model, that choice is kept.
|
||||||
|
* Only a model whose stored `source_platform`/`source_url` match the repository
|
||||||
|
being downloaded is updated; a local file that merely shares a name must not
|
||||||
|
receive another model's card.
|
||||||
|
* The README and repository payload describe the *repository*, so a short-lived
|
||||||
|
process-wide `ModelSourceCache` (`shared_source_cache`, 300 s, 32 entries)
|
||||||
|
keeps a batch over one repository to two HTTP requests.
|
||||||
|
* Every failure — unreachable site, changed payload shape, broken post-processor
|
||||||
|
— is logged and swallowed. Metadata hydration can never fail a download.
|
||||||
|
* Neither stage advances the byte counter, so both are announced to the
|
||||||
|
progress UI (`_report_phase()` → `{"status": "metadata", "stage": ...}`) as
|
||||||
|
they start. Without that the bar sits at 100% reporting `0 B/s` for several
|
||||||
|
seconds and the download looks stuck. `stage` and `platform` are
|
||||||
|
machine-readable; the wording is localised in `LoadingManager`.
|
||||||
|
|
||||||
## Adding a New Skill
|
## Adding a New Skill
|
||||||
|
|
||||||
### 1. Create the skill directory
|
### 1. Create the skill directory
|
||||||
@@ -129,7 +245,7 @@ Use `{{variable}}` placeholders that will be replaced with data from the `prepar
|
|||||||
```markdown
|
```markdown
|
||||||
You are an expert assistant...
|
You are an expert assistant...
|
||||||
|
|
||||||
Model URL: {{hf_url}}
|
Model URL: {{source_url}}
|
||||||
README content:
|
README content:
|
||||||
{{readme_content}}
|
{{readme_content}}
|
||||||
|
|
||||||
|
|||||||
@@ -54,7 +54,7 @@ The dedicated services encapsulate long-running work so handlers stay thin.
|
|||||||
| Use case | Entry point | Dependencies | Guarantees |
|
| Use case | Entry point | Dependencies | Guarantees |
|
||||||
| --- | --- | --- | --- |
|
| --- | --- | --- | --- |
|
||||||
| `RecipeAnalysisService` | `analyze_uploaded_image`, `analyze_remote_image`, `analyze_local_image`, `analyze_widget_metadata` | `ExifUtils`, `RecipeParserFactory`, downloader factory, optional metadata collector/processor | Normalises missing/invalid payloads into `RecipeValidationError`; generates consistent fingerprint data to keep duplicate detection stable; temporary files are cleaned up after every analysis path. |
|
| `RecipeAnalysisService` | `analyze_uploaded_image`, `analyze_remote_image`, `analyze_local_image`, `analyze_widget_metadata` | `ExifUtils`, `RecipeParserFactory`, downloader factory, optional metadata collector/processor | Normalises missing/invalid payloads into `RecipeValidationError`; generates consistent fingerprint data to keep duplicate detection stable; temporary files are cleaned up after every analysis path. |
|
||||||
| `RecipePersistenceService` | `save_recipe`, `delete_recipe`, `update_recipe`, `reconnect_lora`, `bulk_delete`, `save_recipe_from_widget` | `ExifUtils`, recipe scanner, card preview sizing constants | Writes images/JSON metadata atomically; updates scanner caches and hash indices before returning; recalculates fingerprints whenever LoRA assignments change. |
|
| `RecipePersistenceService` | `save_recipe`, `delete_recipe`, `update_recipe`, `reconnect_lora`, `get_reconnect_suggestions`, `bulk_delete`, `save_recipe_from_widget` | `ExifUtils`, recipe scanner, card preview sizing constants | Writes images/JSON metadata atomically; updates scanner caches and hash indices before returning; recalculates fingerprints whenever LoRA assignments change. |
|
||||||
| `RecipeSharingService` | `share_recipe`, `prepare_download` | `tempfile`, recipe scanner | Copies originals to TTL-managed temp files; metadata lookups re-use the scanner; expired shares trigger cleanup and `RecipeNotFoundError`. |
|
| `RecipeSharingService` | `share_recipe`, `prepare_download` | `tempfile`, recipe scanner | Copies originals to TTL-managed temp files; metadata lookups re-use the scanner; expired shares trigger cleanup and `RecipeNotFoundError`. |
|
||||||
|
|
||||||
## Maintaining critical invariants
|
## Maintaining critical invariants
|
||||||
|
|||||||
@@ -0,0 +1,590 @@
|
|||||||
|
# i18n Translation Guidelines
|
||||||
|
|
||||||
|
This document is the canonical set of conventions for translating LoRA Manager UI strings.
|
||||||
|
It applies to **human translators and AI agents** alike. Read it before editing anything in
|
||||||
|
`locales/`.
|
||||||
|
|
||||||
|
Source of truth: `locales/en.json` (10 locales, 2025 leaf keys; all locales share the exact
|
||||||
|
same key structure).
|
||||||
|
|
||||||
|
Locales: `en`, `zh-CN`, `zh-TW`, `ja`, `ko`, `fr`, `de`, `es`, `ru`, `he` (RTL).
|
||||||
|
|
||||||
|
> **Status (2026-08 sweep):** a full audit was executed and the terminology, placeholder,
|
||||||
|
> stale-text, and untranslated-block fixes described in §2–§6 were applied across all locales
|
||||||
|
> (commits `3c3ac49f` … `fd1227d3`). The tables below are now the **normative target state**,
|
||||||
|
> not a to-do list — future edits should preserve these renderings and only add what is new.
|
||||||
|
>
|
||||||
|
> **Status (2026-09, Other Models):** the `other` model type (VAE / Upscaler / Text Encoder /
|
||||||
|
> CLIP Vision / ControlNet) and the Other Models opt-in toggles added 36 new keys; all of them
|
||||||
|
> are now translated in all 9 locales (terminology in §2 "Other Models feature"). There are no
|
||||||
|
> remaining `[TODO: Translate]` placeholders in any locale.
|
||||||
|
>
|
||||||
|
> **Status (2026-09, revision):** `other.disabled.description`, `banners.otherModels.content` and
|
||||||
|
> `settings.folderSettings.enableOtherModelsHelp` were refreshed in `en.json` to name all five
|
||||||
|
> sub_types (they had listed four, which read as "these are what enabling manages") and
|
||||||
|
> re-translated in all 9 locales in the same pass. `clip_vision` and `controlnet` are now both
|
||||||
|
> opt-in, so the first two describe **capability** and the third the **master switch**, not the
|
||||||
|
> default set — keep all three enumerating the full five (`VAE / upscaler / text encoder /
|
||||||
|
> CLIP vision / ControlNet` in `en`; locale slash-list casing follows each file's existing
|
||||||
|
> `VAE / Upscaler / Text Encoder / …` style, de compounds as `CLIP-Vision- und ControlNet-Ordner`).
|
||||||
|
>
|
||||||
|
> **Status (2026-09, "no folders found" state):** the Other Models page gained an *enabled but
|
||||||
|
> nothing to scan* empty state with 6 new keys (`other.noPaths.*`); translated in all 9 locales
|
||||||
|
> in the same pass. The `folder_paths` JSON snippet shown in that state lives in
|
||||||
|
> `templates/other.html`, **not** in the locale files, so it is never translated — only the
|
||||||
|
> surrounding prose is. Terminology added in §2.
|
||||||
|
>
|
||||||
|
> **Status (2026-09, model sources):** models can now be linked to ModelScope and TensorArt
|
||||||
|
> alongside Hugging Face, which added 15 keys (`modelCard.actions.viewOnSource`,
|
||||||
|
> `loras.contextMenu.linkModelSource`, `modals.linkModelSource.*`,
|
||||||
|
> `modals.model.versions.sourceGroupInfo`, `toast.contextMenu.enrichNeedsSource`,
|
||||||
|
> `toast.contextMenu.enrichUnsupportedSource`) and refreshed the two `enrichHfAgent` labels,
|
||||||
|
> which had hardcoded "HF" for a button that now also enriches ModelScope models. The
|
||||||
|
> `modals.linkModelSource.urlPlaceholder` value stays byte-identical to `en.json` (it is a URL,
|
||||||
|
> the §6 exception). Terminology in §2, "Model source feature".
|
||||||
|
>
|
||||||
|
> **Status (2026-09, folder sidebar):** the model-root sidebar gained on-disk folder management
|
||||||
|
> (create / rename / delete folders, show empty folders, tree vs list view) plus its `...`
|
||||||
|
> view-options menu, adding 35 `sidebar.*` keys. Those were the only `[TODO: Translate]`
|
||||||
|
> placeholders left behind by the feature series, and all 35 are now translated in all 9
|
||||||
|
> locales, so the "no remaining placeholders" claim above holds again. Terminology in §2,
|
||||||
|
> "Folder sidebar feature".
|
||||||
|
>
|
||||||
|
> **Status (2026-09, chip reordering):** model tags and trigger words now share one drag/`⠿`
|
||||||
|
> grip reorder affordance, which added the single `common.reorder.dragHandle` key (it lives
|
||||||
|
> under `common` because both editors render it). All 9 locales are translated (renderings in
|
||||||
|
> §2, "Chip reordering"). Reordering is pointer-only by design: an `Alt + Arrow` shortcut was
|
||||||
|
> prototyped and removed because it collided with the browser's Alt + Arrow handling and the
|
||||||
|
> modal's arrow-key navigation.
|
||||||
|
|
||||||
|
> **Status (2026-09, standalone no-paths guidance):** the standalone branch of the
|
||||||
|
> `other.noPaths` empty state now shows the real `settings.json` path plus an
|
||||||
|
> `other.noPaths.openSettingsFolder` button (each locale reuses its
|
||||||
|
> `settings.openSettingsFileLocation.label` rendering), and `descriptionStandalone` was
|
||||||
|
> reworded in `en.json` — from "none of the configured folders exist on disk" to "no
|
||||||
|
> other-model folders were found; add the folder keys you need to the `folder_paths`
|
||||||
|
> section" — and re-translated in all 9 locales. The `on disk` phrase now survives only in
|
||||||
|
> the ComfyUI variant (`descriptionComfyUI`).
|
||||||
|
|
||||||
|
> **Status (2026-09, settings Organization tab):** the settings modal split its overloaded
|
||||||
|
> Library tab, adding the single `settings.nav.organization` key (renderings in §2,
|
||||||
|
> "Settings Organization tab"). All 9 locales are translated, so the "no remaining
|
||||||
|
> placeholders" claim holds again.
|
||||||
|
|
||||||
|
> **Status (2026-09, filename templates):** the Filename Templates feature (per-model-type
|
||||||
|
> download filename templates + bulk "Apply to Library Now" rename, with an empty template
|
||||||
|
> restoring recorded original filenames) added 26 keys across `settings.filenameTemplates.*`,
|
||||||
|
> `loras.bulkOperations.filenameTemplateProgress.*`, `modals.filenameTemplateConfirm.*` and
|
||||||
|
> the `toast.loras.filenameTemplate*` / `toast.settings.filenameTemplates*` toasts. All 9
|
||||||
|
> locales are translated (terminology in §2, "Filename Templates feature").
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. Hard rules (do not violate)
|
||||||
|
|
||||||
|
### R1 — Key structure is sacred
|
||||||
|
- Only `locales/en.json` may add/remove/rename keys. All other locales must keep the exact
|
||||||
|
same nested key set. `tests/i18n/test_i18n.py` enforces this.
|
||||||
|
- When a new UI string is added to `en.json`, run
|
||||||
|
`python scripts/sync_translation_keys.py` (adds the missing keys to all locales with
|
||||||
|
`[TODO: Translate]` placeholder copies) — **then stop**. Do NOT translate proactively:
|
||||||
|
placeholders are the expected end state during feature development, and translations are
|
||||||
|
filled in only when the feature owner explicitly asks (workflow details in §7).
|
||||||
|
- Never reorder, re-indent, or reformat a locale file "for tidiness". The sync script
|
||||||
|
preserves formatting; manual reformatting creates noisy diffs.
|
||||||
|
|
||||||
|
### R2 — Placeholders and HTML must be preserved verbatim
|
||||||
|
- `{name}`-style placeholders must appear in the translation exactly as in `en.json`.
|
||||||
|
Do not invent placeholders the source string does not have — the caller may not pass them
|
||||||
|
(example bug: `zh-CN recipes.controls.import.downloadLocationPreview` added `{path}`; the
|
||||||
|
template renders this key with no parameters, so the literal text `{path}` shows in the UI).
|
||||||
|
- `{{...}}` in a locale value is an escaped literal brace — keep it identical.
|
||||||
|
- Keep embedded HTML tags (e.g. `<strong>...</strong>`, `<code>...</code>`) intact.
|
||||||
|
You may move the tag around the sentence if the target language needs different word order.
|
||||||
|
|
||||||
|
### R3 — Never translate or transliterate these
|
||||||
|
- Model types: **LoRA, Checkpoint, Embedding, Diffusion Model**
|
||||||
|
- Products/brands: **LoRA Manager, ComfyUI, CivitAI, CivArchive, HuggingFace, Ko-fi**
|
||||||
|
- Ecosystem names: **LyCORIS, DoRA**, trigger-adjacent jargon **Prompt, Workflow**
|
||||||
|
(these are used as-is in the target-language SD community; see §2 per-language policy)
|
||||||
|
- Theme names: **Nord, Midnight, Monokai, Dracula, Solarized**
|
||||||
|
|
||||||
|
### R4 — The "Recipe" convention (the most important domain term)
|
||||||
|
Product intent: a *Recipe* records a **LoRA combination + generation parameters**
|
||||||
|
(prompt, seed, sampler, …) that reproduces an image style. The metaphor is a **cooking
|
||||||
|
recipe** — "follow it and you get a similar dish". It is **not** a menu, not a dish list,
|
||||||
|
not a prescription.
|
||||||
|
|
||||||
|
Decision per language — translate only into a word whose everyday primary meaning is a
|
||||||
|
cooking recipe; where that word would mislead users, **keep the English "Recipe(s)"**:
|
||||||
|
|
||||||
|
| Locale | Use | Never use |
|
||||||
|
|---|---|---|
|
||||||
|
| fr | **Recipe / Recipes** (keep English) | recette(s) — cooking reading is secondary and it was explicitly judged misleading |
|
||||||
|
| zh-CN / zh-TW | 配方 | 食谱 (reads as "food cookbook") |
|
||||||
|
| ja | レシピ | — (leftover English "Recipe" in `initialization.recipes.title` / `toast.recipes.recipeSaved` → translate) |
|
||||||
|
| ko | 레시피 | — |
|
||||||
|
| de | Rezept / Rezepte | — (cooking meaning dominant; prescription reading acceptable) |
|
||||||
|
| es | receta / recetas | — (cooking meaning dominant) |
|
||||||
|
| ru | рецепт / рецепты | — (leftover English "Recipe" in `initialization.recipes.title` / `toast.recipes.recipeSaved` → translate) |
|
||||||
|
| he | מתכון / מתכונים | — (cooking meaning dominant) |
|
||||||
|
|
||||||
|
Whatever the choice, **one concept = one noun within a locale**. Currently violated in:
|
||||||
|
- `fr` — "Recipe" (~97 keys, incl. nav) mixed with "recette" (~58 keys)
|
||||||
|
- `zh-CN` / `zh-TW` — 配方 (126/122 keys) mixed with 食谱 / 食譜 (14/17 keys, all in the
|
||||||
|
*rematch* flow: `globalContextMenu.rematchRecipes.*`, `toast.recipes.rematch*`)
|
||||||
|
- `de` — "Rezept" (136 keys) mixed with leftover English "Recipe" (5 keys)
|
||||||
|
- `ja` / `ru` — leftover English "Recipe" in `initialization.recipes.title` ("Recipe Manager
|
||||||
|
zu initialisieren" / «Инициализация Recipe Manager») and `toast.recipes.recipeSaved`
|
||||||
|
|
||||||
|
### R5 — One term, one rendering (within each locale)
|
||||||
|
Same source word must not be translated several ways in one file. Known offender areas
|
||||||
|
(see §5 for the full fix list): recipe, Checkpoint, Embedding, prompt, base model, preset,
|
||||||
|
workflow, hash, metadata, tags, bulk. Every locale currently mixes variants of at least one
|
||||||
|
of these — pick the preferred form in the §2 tables and normalize.
|
||||||
|
|
||||||
|
### R6 — Register consistency
|
||||||
|
- `zh-CN` / `zh-TW`: pick 你 or 您 once. Do not mix (zh-CN has 44×你 + 5×您; zh-TW has
|
||||||
|
27×您 + 18×你).
|
||||||
|
- `de`: pick "du" or "Sie" once (currently 143×Sie + ~7×du).
|
||||||
|
- `es`: pick "tú" or "usted" once.
|
||||||
|
|
||||||
|
### R7 — Punctuation per script
|
||||||
|
- Full-width punctuation `:()` is correct **only in CJK locales** (zh-CN, zh-TW, ja, ko).
|
||||||
|
- Latin/Cyrillic/Hebrew locales must use ASCII `: ()` — full-width colons leaked in there
|
||||||
|
are machine-translation artifacts. Known: `fr toast.recipes.createError/createFailed`,
|
||||||
|
`es toast.recipes.createError/createFailed` (e.g. "…de la receta:" should be "…de la receta:").
|
||||||
|
- `fr` apostrophes must be U+2019 `'` / ASCII `'`, never a straight double quote:
|
||||||
|
`fr header.filter.allowSellingGeneratedContentTooltip` currently reads
|
||||||
|
`vendre d"images` → fix to `d'images`. Do not mix `'` and `'` in one file (fr has 299 vs 15).
|
||||||
|
- Ellipsis: use ASCII `...` (project style). Don't introduce `…`.
|
||||||
|
- Keep the sentence-ending period/omission consistent with the source string where the
|
||||||
|
language allows it.
|
||||||
|
- `he` is RTL: mix of Hebrew and Latin scripts is normal; keep Latin term ordering natural.
|
||||||
|
|
||||||
|
### R8 — No untranslated English leftovers
|
||||||
|
Full sentences left byte-identical to `en.json` are bugs (brand names and URL placeholders
|
||||||
|
are the exception). Every locale has them; see §6 for the per-locale checklist.
|
||||||
|
`[TODO: Translate]` placeholders are the sanctioned intermediate state during feature
|
||||||
|
development (see §7) — do not "fix" them unless the feature owner asked for translations.
|
||||||
|
|
||||||
|
### R9 — Mirror the source even when the source is wrong
|
||||||
|
If `en.json` itself contains an inconsistency (e.g. the `Civitai` vs `CivitAI` casing split,
|
||||||
|
or the `CivitArchive` typo in `modals.relinkCivitai.helpText.format4`), translate/transcribe
|
||||||
|
it as-is in your locale and instead **fix the source** in `en.json` (then propagate by
|
||||||
|
re-syncing and re-translating affected keys). Do not silently diverge in one locale only.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. Per-language term maps
|
||||||
|
|
||||||
|
Preferred rendering per term. "Fix" means the locale currently contains the wrong variant
|
||||||
|
and must be normalized. `en` = keep the English word as-is.
|
||||||
|
|
||||||
|
### fr
|
||||||
|
|
||||||
|
| Term | Use | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| recipe | Recipe(s) | Replace all "recette(s)" (58 keys, e.g. `recipes.actions.deleteRecipeWithShortcut`, `toast.recipes.rematchComplete`) with "Recipe(s)" |
|
||||||
|
| Checkpoint | Checkpoint | `statistics.modelTypes.checkpoint` = "Point de contrôle" → "Checkpoint" |
|
||||||
|
| trigger words | mot(s)-clé(s) | unify: `modals.model.triggerWords.editWord` uses "mot déclencheur" — pick one |
|
||||||
|
| prompt / negative prompt | Prompt / prompt négatif | — |
|
||||||
|
| base model | modèle(s) de base | — |
|
||||||
|
| preset | préréglage | unify: `modals.model.usageTips.addPresetParameter` "prédéfini", `toast.presets.restored` "par défaut" |
|
||||||
|
| hash | hash | `conflictConfirm.message` "hachage" → "hash" |
|
||||||
|
| tags | tags | `settings.sections.priorityTags` "Étiquettes" → "Tags" |
|
||||||
|
| metadata | métadonnées | `loras.controls.refresh.fullTooltip` keeps English "metadata" |
|
||||||
|
| duplicates | doublon(s) | unify with "dupliqué(e)s" |
|
||||||
|
| bulk | groupé(e) | unify with "par lot / mode lot" variants |
|
||||||
|
|
||||||
|
### de
|
||||||
|
|
||||||
|
| Term | Use | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| recipe | Rezept/Rezepte | leftover English "Recipe" keys → Rezept (e.g. `toast.recipes.recipeSaved`) |
|
||||||
|
| base model | pick Basis-Modell or Basismodell | currently 27× hyphenated vs 15× closed |
|
||||||
|
| metadata | Metadaten | 4 keys use "Modelldaten" (`onboarding.steps.fetch.title/content`) → Metadaten |
|
||||||
|
| bulk | pick Massen- or Sammelmodus | `loras.controls.bulk.action` = "Massen" reads as "crowds" — use "Massenbearbeitung"/"Mehrfachauswahl" |
|
||||||
|
| register | Sie (formal) | 7 keys use "du/dein" (`settings.backup.managementHelp`, `modals.checkUpdates.message/tip`, `doctor.footer`, …) |
|
||||||
|
|
||||||
|
### es
|
||||||
|
|
||||||
|
| Term | Use | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| recipe | receta(s) | — |
|
||||||
|
| Checkpoint | Checkpoint | 5 statistics keys "Punto(s) de control" → "Checkpoints" (`statistics.metrics.checkpoints`, `statistics.insights.unusedCheckpoints.*`, `statistics.modelTypes.checkpoint`) |
|
||||||
|
| trigger words | palabra(s) de activación | 2 keys already use it; ~15 keys "palabra(s) clave" (reads as search keyword) → unify |
|
||||||
|
| base model | modelo base | — |
|
||||||
|
| preset | preajuste | 3 keys keep English "preset", 1 "preestablecido" → preajuste |
|
||||||
|
| workflow | pick flujo de trabajo or workflow | currently 21× "flujo de trabajo" vs 10× "workflow" |
|
||||||
|
| bulk | masivo / por lotes | unify; "Batch Import" → traducción |
|
||||||
|
| tags | etiquetas | — |
|
||||||
|
|
||||||
|
### ru
|
||||||
|
|
||||||
|
| Term | Use | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| recipe | рецепт(ы) | English leftovers: `initialization.recipes.title`, `recipes.batchImport.*`, `toast.recipes.recipeSaved` → translate |
|
||||||
|
| Checkpoint | Checkpoint (recommended) | 3 variants today: "Checkpoint" (17 keys), «Чекпойнт», «Контрольная точка» (statistics, 6 keys) — statistics MUST drop «Контрольная точка» |
|
||||||
|
| Embedding | Embedding | «Эмбеддинг» variant exists in `settings.priorityTags.modelTypes.embedding` — unify |
|
||||||
|
| prompt | промпт | 8 keys use «запрос» (reads as "database/HTTP request") → «промпт» |
|
||||||
|
| base model | базовая модель | — |
|
||||||
|
| preset | пресет | `header.theme.presets` "Предустановки" → пресеты |
|
||||||
|
| workflow | Workflow (recommended) | «рабочий процесс» used in 4 keys — unify |
|
||||||
|
| hash | pick хеш or хэш | both spellings co-occur |
|
||||||
|
| tag(s) | тег(и) | — |
|
||||||
|
| typos | — | `settings.misc.loraSyntaxFormatHelp`: «безпотерьного» → «беспотерьного» |
|
||||||
|
|
||||||
|
### he
|
||||||
|
|
||||||
|
| Term | Use | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| recipe | מתכון / מתכונים | — |
|
||||||
|
| Checkpoint | Checkpoint | 5 statistics keys «נקודת/נקודות ביקורת» (road/security checkpoint) → "Checkpoint(s)" (`statistics.metrics.checkpoints`, `statistics.modelTypes.checkpoint`, `statistics.insights.unusedCheckpoints.*`) |
|
||||||
|
| Embedding | Embedding | `statistics` keys use הטמעות → Embedding |
|
||||||
|
| prompt | pick הנחיה or פרומפט | 9 keys הנחיה vs 3 פרומפט — unify (recommend פרומפט, SD-community loanword) |
|
||||||
|
| preset | קביעה מראש | `header.filter.presetOverwriteConfirm` uses פריסט → unify |
|
||||||
|
| hash | pick one of האש / גיבוב / hash | 3 variants co-occur — unify (recommend hash or גיבוב) |
|
||||||
|
| metadata | pick מטא-דאטה or מטא-נתונים | 38 vs 17 keys — unify |
|
||||||
|
| model | מודל | 13 keys use דגם/דגמים — unify |
|
||||||
|
| bulk | pick one of 5 variants | 5 different renderings ("כמות גדולה", "המוני", "קבוצתי", "אצווה", …) — unify; `loras.controls.bulk.action` "כמות גדולה" reads as "large quantity" |
|
||||||
|
|
||||||
|
### ja
|
||||||
|
|
||||||
|
| Term | Use | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| recipe | レシピ | `initialization.recipes.title` keeps English "Recipe Manager" — translate to レシピマネージャー |
|
||||||
|
| Checkpoint | Checkpoint or チェックポイント (pick one) | 3 variants: Checkpoint (~14), checkpoint lowercase (4), チェックポイント (4, e.g. `settings.priorityTags.modelTypes.checkpoint`) |
|
||||||
|
| Embedding | Embedding | 4 keys lowercase "embedding" mid-sentence |
|
||||||
|
| bulk | 一括 | `modals.checkUpdates.tip` "バルクモード" → 一括モード |
|
||||||
|
| recipe counter | 件 or 個 | `globalContextMenu.rematchRecipes.success` uses 件, `.cancelled` uses 個 — unify |
|
||||||
|
|
||||||
|
### ko
|
||||||
|
|
||||||
|
| Term | Use | Fix |
|
||||||
|
|---|---|---|
|
||||||
|
| recipe | 레시피 | — |
|
||||||
|
| Checkpoint | Checkpoint (recommended) | 4 keys transliterate 체크포인트 (`settings.priorityTags.modelTypes.checkpoint`, `toast.recipes.missingCheckpointPath/missingCheckpointInfo/downloadCheckpointFailed`) |
|
||||||
|
| Embedding | Embedding | 3 keys 임베딩 (`settings.priorityTags.modelTypes.embedding`, `uiHelpers.nodeSelector.embedding`) |
|
||||||
|
| base model | 베이스 모델 | 6 keys «기본 모델» read as "default model" → 베이스 모델 (`settings.downloadSkipBaseModels.*`, `toast.loras.downloadSkippedByBaseModel`) |
|
||||||
|
| workflow | pick 워크플로 or 워크플로우 | 26 vs 6 keys — unify |
|
||||||
|
| bulk | 일괄 | `modals.checkUpdates.tip` "벌크 모드" → 일괄 모드 |
|
||||||
|
| tag logic | — | `header.filter.tagLogicAny` = "모든 태그 일치 (OR)" is **inverted** (should be "하나 이상의 태그 일치") and identical to `tagLogicAll` |
|
||||||
|
| particle | — | `modelCard.sendToWorkflow.checkpointNotImplemented`: "Checkpoint을" → "Checkpoint를" |
|
||||||
|
|
||||||
|
### zh-CN / zh-TW
|
||||||
|
|
||||||
|
| Term | zh-CN | zh-TW |
|
||||||
|
|---|---|---|
|
||||||
|
| recipe | 配方 (fix 食谱 → 配方, 14 keys in rematch flow) | 配方 (fix 食譜 → 配方, 17 keys in rematch flow) |
|
||||||
|
| Checkpoint | Checkpoint (fix 检查点 → Checkpoint, 5 keys: `toast.recipes.missingCheckpointPath/missingCheckpointInfo/downloadCheckpointFailed`, `modelCard.actions.checkpointNameCopied`, `modelCard.sendToWorkflow.checkpointNotImplemented`) | Checkpoint (fix 檢查點 → Checkpoint, 4 keys: `modelCard.actions.copyCheckpointName`, `toast.recipes.missing*`×2, `toast.recipes.downloadCheckpointFailed`) |
|
||||||
|
| base model | 基础模型 (fix 基模型 → 基础模型, 3 keys in `modals.model.versions.filters.*`) | 基礎模型 ✓ consistent |
|
||||||
|
| prompt | 提示词 ✓ | 提示詞 ✓ |
|
||||||
|
| preset | 预设 ✓ | 預設 ✓ |
|
||||||
|
| workflow | 工作流 ✓ | 工作流 ✓ |
|
||||||
|
| trigger words | 触发词 ✓ | 觸發詞 ✓ |
|
||||||
|
| hash | 哈希 (哈希值 variant OK) | 雜湊 ✓ |
|
||||||
|
| register | 你 (fix 5×您 → 你) | 您 (fix 18×你 → 您) |
|
||||||
|
|
||||||
|
### Other Models feature (VAE / Upscaler / Text Encoder / CLIP Vision / ControlNet)
|
||||||
|
|
||||||
|
The `other` model type exposes five sub_types. They are **model-type names**, so they follow
|
||||||
|
R3 and stay in Latin in every locale. The `settings.folderSettings.subType*` values are
|
||||||
|
therefore **intentionally byte-identical to `en.json`** (same precedent as
|
||||||
|
`settings.priorityTags.modelTypes` / `checkpoints.modelTypes.checkpoint`) — a §6 sweep must
|
||||||
|
not "fix" them.
|
||||||
|
|
||||||
|
| Term | Rendering | Note |
|
||||||
|
|---|---|---|
|
||||||
|
| VAE | `VAE` everywhere | acronym, always upper-case |
|
||||||
|
| Upscaler | `Upscaler` everywhere | CivitAI `ModelType` name |
|
||||||
|
| Text Encoder | `Text Encoder` everywhere | de compounds as `Text-Encoder-Stammordner` |
|
||||||
|
| CLIP Vision | `CLIP Vision` everywhere | de compounds as `CLIP-Vision-Stammordner` |
|
||||||
|
| ControlNet | `ControlNet` everywhere | brand casing, capital N |
|
||||||
|
|
||||||
|
In prose these names sit next to localized nouns the same way `Diffusion Model` does
|
||||||
|
(zh `VAE 根目录`, ja `VAEルート`, ko `VAE 루트`, ru `Корневая папка VAE`).
|
||||||
|
|
||||||
|
**"Other Models" is the page/feature name, not a model type — translate it:**
|
||||||
|
|
||||||
|
| Locale | `other.title` | `header.navigation.other` |
|
||||||
|
|---|---|---|
|
||||||
|
| fr | Autres modèles | Autres |
|
||||||
|
| zh-CN | 其他模型 | 其他 |
|
||||||
|
| zh-TW | 其他模型 | 其他 |
|
||||||
|
| ja | その他のモデル | その他 |
|
||||||
|
| ko | 기타 모델 | 기타 |
|
||||||
|
| de | Weitere Modelle | Andere |
|
||||||
|
| es | Otros modelos | Otros |
|
||||||
|
| ru | Другие модели | Другое |
|
||||||
|
| he | מודלים אחרים | אחרים |
|
||||||
|
|
||||||
|
`settings.folderSettings.otherSubTypes` ("Managed Types") must name **model** types, matching
|
||||||
|
each locale's `header.filter.modelTypes` rendering (zh `管理的模型类型`, ja `管理するモデルタイプ`,
|
||||||
|
de `Verwaltete Modelltypen`, …).
|
||||||
|
|
||||||
|
The "no folders found" empty state (`other.noPaths.*`) uses two phrases that must stay
|
||||||
|
consistent whenever that copy is edited. `folder key` means the `folder_paths` key name
|
||||||
|
(`vae`, `upscale_models`, … — Latin per the table above); `on disk` means the folder must
|
||||||
|
physically exist:
|
||||||
|
|
||||||
|
| Phrase | Rendering |
|
||||||
|
|---|---|
|
||||||
|
| folder key | zh-CN 文件夹键 · zh-TW 資料夾鍵 · ja フォルダーキー · ko 폴더 키 · fr clé de dossier · de Ordnerschlüssel · es clave de carpeta · ru ключ папки · he מפתח תיקייה |
|
||||||
|
| on disk | zh-CN 在磁盘上 · zh-TW 在磁碟上 · ja ディスク上 · ko 디스크에 · fr sur le disque · de auf dem Datenträger · es en el disco · ru на диске · he בדיסק |
|
||||||
|
|
||||||
|
`settings.json` and `ComfyUI` stay verbatim in every locale; "reload this page" / "restart
|
||||||
|
LoRA Manager" reuse each locale's existing restart wording (`settings.extraFolderPaths.*`).
|
||||||
|
|
||||||
|
### Model source feature (Hugging Face / ModelScope / TensorArt)
|
||||||
|
|
||||||
|
A model file can be linked to the page of an external model site. **Hugging Face**,
|
||||||
|
**ModelScope** and **TensorArt** are brand names and stay Latin in every locale (R3); the
|
||||||
|
generic nouns around them are translated:
|
||||||
|
|
||||||
|
| Term | Rendering |
|
||||||
|
|---|---|
|
||||||
|
| model source | zh-CN 模型来源 · zh-TW 模型來源 · ja モデルソース · ko 모델 소스 · fr source de modèle · de Modellquelle · es fuente de modelo · ru источник модели · he מקור מודל |
|
||||||
|
| model page | zh-CN 模型页面 · zh-TW 模型頁面 · ja モデルページ · ko 모델 페이지 · fr page du modèle · de Modellseite · es página del modelo · ru страница модели · he עמוד המודל |
|
||||||
|
| model card | zh-CN 模型卡 · zh-TW 模型卡 · ja モデルカード · ko 모델 카드 · fr fiche de modèle · de Modellkarte · es ficha de modelo · ru карточка модели · he כרטיס מודל |
|
||||||
|
| AI enrichment (noun) | reuse the existing pair per locale: zh-CN 增强 · zh-TW 增強 · ja 補完 · ko 보강 · fr enrichissement (par IA) · de Anreicherung (KI-) · es enriquecimiento (con IA) · ru обогащение (с помощью ИИ) · he העשרה (AI) |
|
||||||
|
|
||||||
|
`modelCard.actions.viewOnSource` ("View on {source}") follows each locale's existing
|
||||||
|
`viewOnHuggingFace` pattern — de `Auf … ansehen`, ru `Открыть …`, he `צפייה ב-…`,
|
||||||
|
ja `… で見る`, ko `…에서 보기`, zh `在 … 查看`, fr `Voir sur …`, es `Ver en …`. `{source}` is
|
||||||
|
replaced at runtime with the untranslated platform name, so the brand never appears inside the
|
||||||
|
translated text.
|
||||||
|
|
||||||
|
`modals.linkModelSource.enrichNote` states the rule that only sites exposing a readable model
|
||||||
|
card can be enriched and names TensorArt as the current exception. Keep the parenthetical
|
||||||
|
exception in sync if another link-only source is ever added — the sentence is deliberately
|
||||||
|
phrased as a rule, not as an apology for one site.
|
||||||
|
|
||||||
|
The context-menu and bulk-operation enrichment entry points read **"Enrich Metadata with AI"**
|
||||||
|
in `en`, not "Enrich HF Metadata": they cover ModelScope as well, so no locale may reintroduce
|
||||||
|
an `HF` qualifier in `loras.contextMenu.enrichHfAgent` / `loras.bulkOperations.enrichHfAgent`
|
||||||
|
(the key names keep the historical `Hf`; only the values changed).
|
||||||
|
|
||||||
|
### Folder sidebar feature (create / rename / delete folders, empty folders, view options)
|
||||||
|
|
||||||
|
The model-root sidebar manages on-disk folders. "Folder" reuses the noun already fixed in §2
|
||||||
|
(the `folder key` row); the rest is new surface:
|
||||||
|
|
||||||
|
| Term | Rendering |
|
||||||
|
|---|---|
|
||||||
|
| folder | zh-CN 文件夹 · zh-TW 資料夾 · ja フォルダ · ko 폴더 · fr dossier · de Ordner · es carpeta · ru папка · he תיקייה |
|
||||||
|
| model root (as in "no model root is configured") | zh-CN 模型根目录 · zh-TW 模型根目錄 · ja モデルルート · ko 모델 루트 · fr racine de modèle · de Modell-Stammverzeichnis · es raíz de modelo · ru корневая папка моделей · he שורש מודלים — note `sidebar.modelRoot` alone is the shorter 根目录 / 根目錄 / ルート / 루트 / Racine / Stammverzeichnis / Raíz / Корень / שורש |
|
||||||
|
| tree view / list view | zh-CN 树形视图 / 列表视图 · zh-TW 樹狀檢視 / 清單檢視 · ja ツリー表示 / リスト表示 · ko 트리 보기 / 목록 보기 · fr Vue arborescente / Vue liste · de Baumansicht / Listenansicht · es Vista de árbol / Vista de lista · ru Дерево / Список · he תצוגת עץ / תצוגת רשימה |
|
||||||
|
| sidebar | reuse each locale's `sidebar.hideOnThisPage` noun: zh-CN 侧边栏 · zh-TW 側邊欄 · ja サイドバー · ko 사이드바 · fr barre latérale · de Seitenleiste · es barra lateral · ru боковая панель · he סרגל צד |
|
||||||
|
|
||||||
|
Deleting a folder **never cascades over model files** — the backend refuses it and
|
||||||
|
`sidebar.deleteFolderModal.notEmptyMessage` states the rule in every locale, so keep that
|
||||||
|
clause (and its `—`) when the copy is edited. The `{name}` / `{count}` / `{message}` tokens in
|
||||||
|
`sidebar.createFolderResult.*`, `sidebar.deleteFolderResult.*` and `sidebar.renameFolderResult.*`
|
||||||
|
are verbatim §1-R2 placeholders; `successWithFiles` is the only key carrying `{count}`.
|
||||||
|
|
||||||
|
### Settings Organization tab
|
||||||
|
|
||||||
|
The settings modal's fourth nav tab groups everything about how files are arranged on
|
||||||
|
disk: download path templates, priority tags, and auto-organize exclusions. The label is
|
||||||
|
the **noun for arranging files**, matching each locale's existing
|
||||||
|
`settings.sections.autoOrganize` rendering minus the "auto":
|
||||||
|
|
||||||
|
| Locale | `settings.nav.organization` |
|
||||||
|
|---|---|
|
||||||
|
| fr | Organisation |
|
||||||
|
| zh-CN | 整理 |
|
||||||
|
| zh-TW | 整理 |
|
||||||
|
| ja | 整理 |
|
||||||
|
| ko | 정리 |
|
||||||
|
| de | Organisation |
|
||||||
|
| es | Organización |
|
||||||
|
| ru | Организация |
|
||||||
|
| he | ארגון |
|
||||||
|
|
||||||
|
zh-CN/zh-TW use 整理 ("tidying/arranging"), not 组织/組織 (an organization as a group).
|
||||||
|
|
||||||
|
### Filename Templates feature
|
||||||
|
|
||||||
|
Per-model-type templates that name downloaded model files; "Apply to Library Now"
|
||||||
|
bulk-renames existing files, and an **empty template restores the recorded original
|
||||||
|
filenames** (recorded in each model's metadata at its first rename). "Template" follows
|
||||||
|
each locale's existing download-path-template noun (zh-CN 模板 vs zh-TW 範本 — note the
|
||||||
|
split); progress strings mirror `loras.bulkOperations.autoOrganizeProgress` verbatim with
|
||||||
|
the locale's "moved" verb swapped for its "renamed" verb, and the toasts mirror the
|
||||||
|
`autoOrganize*` / `downloadTemplates*` toast shapes.
|
||||||
|
|
||||||
|
| Term | Rendering |
|
||||||
|
|---|---|
|
||||||
|
| filename template(s) | zh-CN 文件名模板 · zh-TW 檔案名稱範本 · ja ファイル名テンプレート · ko 파일명 템플릿 · fr modèle(s) de nom de fichier · de Dateinamen-Vorlage(n) · es plantilla(s) de nombres de archivo · ru шаблон(ы) имён файлов · he תבנית שם קובץ / תבניות שמות קבצים |
|
||||||
|
| Apply to Library Now (button) | zh-CN 立即应用到库 · zh-TW 立即套用至模型庫 · ja ライブラリに今すぐ適用 · ko 지금 라이브러리에 적용 · fr Appliquer à la bibliothèque maintenant · de Jetzt auf Bibliothek anwenden · es Aplicar a la biblioteca ahora · ru Применить к библиотеке сейчас · he החל על הספרייה כעת |
|
||||||
|
| Restore original filenames (modal title / button) | zh-CN 恢复原始文件名?/ 恢复原始文件名 · zh-TW 要還原原始檔案名稱嗎?/ 還原原始檔案名稱 · ja 元のファイル名を復元しますか?/ 元のファイル名を復元 · ko 원본 파일명을 복원하시겠습니까? / 원본 파일명 복원 · fr Restaurer les noms de fichier d'origine ? / Restaurer les noms de fichier d'origine · de Ursprüngliche Dateinamen wiederherstellen? / Ursprüngliche Dateinamen wiederherstellen · es ¿Restaurar los nombres de archivo originales? / Restaurar nombres de archivo originales · ru Восстановить исходные имена файлов? / Восстановить исходные имена файлов · he לשחזר שמות קבצים מקוריים? / שחזר שמות קבצים מקוריים |
|
||||||
|
| "renamed" (progress/toast counter) | zh-CN 已重命名 · zh-TW 已重新命名 · ja リネーム · ko 이름 변경 · fr renommés · de umbenannt · es renombrados · ru переименовано · he שונו שמותם |
|
||||||
|
|
||||||
|
### Chip reordering (model tags / trigger words)
|
||||||
|
|
||||||
|
Model tags and trigger-word chips share a single reorder affordance (drag the chip, or its
|
||||||
|
`⠿` grip where the chip body is click-to-edit), so the copy sits in `common.reorder.dragHandle`
|
||||||
|
instead of a feature namespace. It is used twice per editor: as the grip tooltip and as the
|
||||||
|
hint shown in the edit controls row. There is deliberately **no keyboard shortcut** — an
|
||||||
|
`Alt + Arrow` binding fought the browser's own Alt + Arrow handling and the modal's arrow-key
|
||||||
|
navigation, so reordering is pointer-only and the grip is a decorative, non-focusable
|
||||||
|
affordance. Do not reintroduce a shortcut or a "position X of Y" screen-reader string without
|
||||||
|
re-adding the corresponding keys.
|
||||||
|
|
||||||
|
`dragHandle` is a fragment, not a sentence: it labels both the grip and the hint, so keep it
|
||||||
|
short and imperative and do not append a keyboard hint in any locale.
|
||||||
|
|
||||||
|
| Term | Rendering |
|
||||||
|
|---|---|
|
||||||
|
| drag to reorder | zh-CN 拖拽以调整顺序 · zh-TW 拖曳以調整順序 · ja ドラッグして並べ替え · ko 드래그하여 순서 변경 · fr Glisser pour réordonner · de Zum Neuordnen ziehen · es Arrastra para reordenar · ru Перетащите, чтобы изменить порядок · he גרור כדי לשנות סדר |
|
||||||
|
|
||||||
|
The grip itself is an icon and is never translated.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Cross-cutting confusion hot-spots (must-fix list)
|
||||||
|
|
||||||
|
All items below were **resolved** in the 2026-08 sweep — treat them as a regression
|
||||||
|
watch-list: do not reintroduce these renderings.
|
||||||
|
|
||||||
|
1. **Checkpoint rendered as a literal security/road checkpoint** — fr, es, ru, he, zh-CN,
|
||||||
|
zh-TW all had 4–6 keys in the `statistics.*` domain reading as "control point"; reverted
|
||||||
|
to "Checkpoint".
|
||||||
|
2. **"recipe" variants that break the one-noun rule** — fr "recette" → "Recipe", zh
|
||||||
|
食谱/食譜 → 配方, de/ja/ru leftover English "Recipe" translated.
|
||||||
|
3. **ko `header.filter.tagLogicAny`** — was inverted ("모든 태그 일치 (OR)") and identical
|
||||||
|
to `tagLogicAll`; now "어느 하나의 태그와 일치 (OR)".
|
||||||
|
4. **ja `modals.model.versions.actions.viewLocalTooltip`** — was the stale "近日対応予定"
|
||||||
|
("coming soon"); all 9 locales now describe the actual action.
|
||||||
|
5. **Stale help texts** — `settings.downloadSkipBaseModels.help`,
|
||||||
|
`settings.aiProvider.apiBaseHelp`, `settings.hideEarlyAccessUpdates.help` retranslated
|
||||||
|
in all locales to the current `en.json` wording.
|
||||||
|
6. **en.json source bugs** (fixed in source, then mirrored):
|
||||||
|
- "Civitai" → "CivitAI" brand casing (values only; key names `relinkCivitai` etc. keep
|
||||||
|
their lowercase form and must not be renamed)
|
||||||
|
- `modals.relinkCivitai.helpText.format4` "CivitArchive" typo → "CivArchive"
|
||||||
|
- `zh-CN recipes.controls.import.downloadLocationPreview` invented `{path}` removed
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. Placeholder contract deviations (current)
|
||||||
|
|
||||||
|
`{...}` token sets must match `en.json` per key. All deviations found in the 2026-08 sweep
|
||||||
|
were fixed, with one *intentional* exception:
|
||||||
|
|
||||||
|
**`toast.settings.mappingsUpdated`** — the caller passes a hardcoded English inflection
|
||||||
|
(`plural: count !== 1 ? 's' : ''`). Languages that cannot build a plural by appending that
|
||||||
|
`s` (zh-CN/zh-TW, ja, ko, de, ru, he) **drop `{plural}`** and render a count-friendly form
|
||||||
|
(`({count})` or a measure word); fr and es keep it (`mappage{plural}`, `mapeo{plural}`).
|
||||||
|
|
||||||
|
```python
|
||||||
|
# keep a copy of this rule next to the key if it ever moves:
|
||||||
|
# fr/es: "... ({count} mappage{plural})"
|
||||||
|
# de/ru/he: "... ({count})"
|
||||||
|
# zh-CN: "({count} 条映射)" / zh-TW: "({count} 個對應)" / ja: "({count} マッピング)"
|
||||||
|
```
|
||||||
|
|
||||||
|
Do NOT add `{...}` tokens the source lacks (the caller will not supply them, and the literal
|
||||||
|
text renders in the UI), and do NOT rename source tokens (`{typePlural}` stays `{typePlural}`).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. One term, one rendering — offender matrix
|
||||||
|
|
||||||
|
Cross-locale summary of §2 inconsistencies. "✓" = already consistent. All ✗ cells were
|
||||||
|
resolved in the 2026-08 sweep; the row shows the single rendering now in force per locale.
|
||||||
|
|
||||||
|
| Term | fr | de | es | ru | he | ja | ko | zh-CN | zh-TW |
|
||||||
|
|---|---|---|---|---|---|---|---|---|---|
|
||||||
|
| recipe | Recipe | Rezept | receta | рецепт | מתכון | レシピ | 레시피 | 配方 | 配方 |
|
||||||
|
| Checkpoint | Checkpoint | Checkpoint | Checkpoint | Checkpoint | Checkpoint | Checkpoint | Checkpoint | Checkpoint | Checkpoint |
|
||||||
|
| Embedding | Embedding | Embedding | Embedding | Embedding | Embedding | Embedding | Embedding | Embedding | Embedding |
|
||||||
|
| prompt | Prompt | Prompt | prompt | промпт | פרומפט | プロンプト | 프롬프트 | 提示词 | 提示詞 |
|
||||||
|
| base model | modèle de base | Basismodell | modelo base | базовая модель | מודל בסיס | ベースモデル | 베이스 모델 | 基础模型 | 基礎模型 |
|
||||||
|
| preset | préréglage | Voreinstellung | preajuste | пресет | קביעה מראש | プリセット | 프리셋 | 预设 | 預設 |
|
||||||
|
| workflow | Workflow | Workflow | workflow | Workflow | workflow | ワークフロー | 워크플로 | 工作流 | 工作流 |
|
||||||
|
| hash | hash | Hash | hash | хеш | hash | ハッシュ | 해시 | 哈希 | 雜湊 |
|
||||||
|
| metadata | métadonnées | Metadaten | metadatos | метаданные | מטא-נתונים | メタデータ | 메타데이터 | 元数据 | 中繼資料 |
|
||||||
|
| tags | Tags | Tags | etiquetas | теги | תגיות | タグ | 태그 | 标签 | 標籤 |
|
||||||
|
| duplicates | en double | Duplikate | duplicados | дубликаты | כפילויות | 重複 | 중복 | 重复项 | 重複項 |
|
||||||
|
| bulk | groupé | Massen- | por lotes | пакетный | בכמות גדולה | 一括 | 일괄 | 批量 | 批量 |
|
||||||
|
|
||||||
|
Watch: ja/ko keep the model-type names **Checkpoint/Embedding** and `Diffusion Model` in
|
||||||
|
Latin (consistent with their model-type sections) — do not transliterate them as
|
||||||
|
チェックポイント/체크포인트.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. Untranslated English leftovers (status)
|
||||||
|
|
||||||
|
Values byte-identical to `en.json` that are actual UI sentences are bugs (brand names and
|
||||||
|
URL placeholders are the exception). As of the 2026-08 sweep, **all previously untranslated
|
||||||
|
blocks are translated** in every locale: `recipes.batchImport.*` + `toast.recipes.batchImport*`
|
||||||
|
(fr/de/es/ru/he/ja/ko), `banners.communitySupport.*`, `modals.model.license.*`,
|
||||||
|
`globalContextMenu.fetchMissingLicenses.*`, the `doctor.*` issue/action/label subset,
|
||||||
|
`toast.settings.libraryLoadFailed` / `libraryActivateFailed`, `toast.api.moveFailed`,
|
||||||
|
`settings.extraFolderPaths.restartRequired`, `toast.recipes.recipeSaved`,
|
||||||
|
`sidebar.dragDrop.moveUnsupported`, `checkpoints.modelTypes.diffusion_model`
|
||||||
|
(ja/ko keep the English loanword), `initialization.recipes.title`.
|
||||||
|
|
||||||
|
The only values that remain intentionally identical to `en.json` are non-translatable:
|
||||||
|
URL/path placeholders (`https://…`, `C:/…`), numeric presets (`5 (1080p), 6 (2K), 8 (4K)`),
|
||||||
|
example token lists (`character, concept, style(toon|toon_style)`), service/provider names
|
||||||
|
(`CivitAI → CivArchive → Archive DB`), model-type names (`settings.priorityTags.modelTypes.*`,
|
||||||
|
`settings.folderSettings.subTypeVae` … `subTypeControlnet` — see §2), and the external playlist
|
||||||
|
title (`help.updateVlogs.playlistTitle`, de: translated to "LoRA Manager-Update-Playlist").
|
||||||
|
|
||||||
|
Rule for `uiHelpers.workflow.noPromptTargets`: the second line (`Mark as → Send Prompt
|
||||||
|
Target`) quotes literal ComfyUI context-menu items — keep those menu labels in English in
|
||||||
|
every locale because that is what the user actually sees in ComfyUI.
|
||||||
|
|
||||||
|
License labels (`modals.model.license.*`): the restriction labels are now translated in all
|
||||||
|
locales (the sibling `creditRequired` has always been translated).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. Workflow for agents and translators
|
||||||
|
|
||||||
|
### Adding a new UI string
|
||||||
|
1. Add the key to `locales/en.json` only.
|
||||||
|
2. Run `python scripts/sync_translation_keys.py` — it inserts the key into the other 9
|
||||||
|
locales (as a `[TODO: Translate]` placeholder) preserving formatting.
|
||||||
|
3. **During feature development, stop here.** While the UI copy is still in flux, leave the
|
||||||
|
`[TODO: Translate]` placeholders as-is — translating churning strings into 9 locales is
|
||||||
|
wasted work. Placeholders are a normal intermediate state, not a bug.
|
||||||
|
4. Once the wording is final and the feature owner explicitly asks for translations,
|
||||||
|
translate **all** pending `[TODO: Translate]` keys in every locale (not just the latest
|
||||||
|
feature's), applying §1–§3 (placeholders verbatim, Recipe rule, term maps, register).
|
||||||
|
Find pending keys with: `grep -c "TODO: Translate" locales/*.json`
|
||||||
|
5. If the new string contains new terminology, extend §2 tables.
|
||||||
|
|
||||||
|
### Fixing a translation bug
|
||||||
|
1. Locate the key (dotted path) in the relevant locale file.
|
||||||
|
2. Check the corresponding `en.json` value and the actual caller (grep `static/js` or
|
||||||
|
`web/comfyui` for the key) to learn which placeholders are passed.
|
||||||
|
3. Fix trivially; for normalization sweeps (e.g. "recette" → "Recipe"), do it file-wide for
|
||||||
|
the offending keys only — do not touch unrelated lines.
|
||||||
|
4. If the bug is in `en.json` itself (R9), fix the source first, then re-sync and update all
|
||||||
|
locales.
|
||||||
|
|
||||||
|
### Verification
|
||||||
|
```bash
|
||||||
|
pytest tests/i18n/test_i18n.py # key parity + JSON validity + JS key references
|
||||||
|
python scripts/sync_translation_keys.py --dry-run # shows which keys would change; add --verbose for per-key detail
|
||||||
|
npm test # frontend tests incl. i18n helpers
|
||||||
|
```
|
||||||
|
|
||||||
|
`pytest tests/i18n` only checks structure. Quality conventions in this document are not
|
||||||
|
machine-enforced — a human/agent review pass is required.
|
||||||
|
|
||||||
|
### Anti-patterns checklist
|
||||||
|
- [ ] Placeholders `{x}` / `{{x}}` differ from `en.json`
|
||||||
|
- [ ] Same source term translated 2+ ways in the same file (see §5)
|
||||||
|
- [ ] "Checkpoint" became a literal checkpoint; "recipe" became menu/prescription/food-cookbook
|
||||||
|
- [ ] Brand names translated or transliterated (LoRA, CivitAI, ComfyUI, …)
|
||||||
|
- [ ] Latin locale using full-width `:()`; fr using `"` as apostrophe
|
||||||
|
- [ ] Mixed 你/您, du/Sie, tú/usted
|
||||||
|
- [ ] Full English sentences left behind (see §6)
|
||||||
|
- [ ] Register/typos/mojibake; source string is stale vs `en.json` (compare semantics, not
|
||||||
|
just words)
|
||||||
@@ -0,0 +1,107 @@
|
|||||||
|
# Plan: Filename Template Follow-ups
|
||||||
|
|
||||||
|
**Issue:** [#1071 — Lora Renaming](https://github.com/willmiao/ComfyUI-Lora-Manager/issues/1071)
|
||||||
|
**Status:** Core feature **implemented** (2026-09-19, commit `2bc9860b`,
|
||||||
|
preceded by the settings-tab split in `327da046`). Follow-ups 1 and 2 were
|
||||||
|
resolved together on 2026-09-19 by redefining the empty template as
|
||||||
|
"revert to recorded original filename" (see below). Follow-up 3 remains open.
|
||||||
|
|
||||||
|
## What shipped in `2bc9860b`
|
||||||
|
|
||||||
|
- Per-model-type `download_filename_templates` setting (empty = keep current
|
||||||
|
filename; opt-in). Placeholders: `{model_name}`, `{version_name}`,
|
||||||
|
`{base_model}`, `{author}`, `{first_tag}`, `{hash_short}`,
|
||||||
|
`{original_name}`.
|
||||||
|
- `calculate_filename_for_model()` in `py/utils/utils.py` renders the
|
||||||
|
template; templates containing path separators are rejected.
|
||||||
|
- Downloads apply the template post-download
|
||||||
|
(`DownloadManager._apply_download_filename_template`); rename conflicts
|
||||||
|
keep the original name and never fail the download.
|
||||||
|
- `ModelLifecycleService.rename_model` records `original_file_name` in the
|
||||||
|
`.metadata.json` sidecar (first rename wins via `setdefault`).
|
||||||
|
- Bulk apply: `GET|POST /api/lm/{prefix}/apply-filename-template`
|
||||||
|
(`FilenameTemplateUseCase`, shares the auto-organize lock, WS progress type
|
||||||
|
`filename_template_progress`).
|
||||||
|
- Settings UI: "Filename Templates" subsection in the new **Organization**
|
||||||
|
settings tab (`templates/components/modals/settings/organization.html`),
|
||||||
|
with validation, live preview, and per-type "Apply to Library Now".
|
||||||
|
|
||||||
|
Sandbox E2E verified: rename incl. companion files (previews, sidecars),
|
||||||
|
metadata pointer updates, `original_file_name` recording, idempotency,
|
||||||
|
conflict handling (failure counted, batch continues), empty-template no-op,
|
||||||
|
GET variant.
|
||||||
|
|
||||||
|
## Follow-ups 1 & 2 — RESOLVED: empty template = revert to recorded original
|
||||||
|
|
||||||
|
Follow-up 1 asked to reword the ambiguous "Valid (keep original filename)"
|
||||||
|
empty-template message; Follow-up 2 asked for a bulk revert to the recorded
|
||||||
|
`original_file_name`. Both were resolved by a single semantic change: **an
|
||||||
|
empty template now means "restore the recorded original filename"** instead of
|
||||||
|
"leave the current filename untouched".
|
||||||
|
|
||||||
|
Rationale: for never-renamed models a revert is a no-op (no recorded
|
||||||
|
original), for renamed models it restores the pre-rename name, and new
|
||||||
|
downloads with an empty template keep the download name as before — so the
|
||||||
|
two contexts (download path and bulk apply) share one coherent meaning, and
|
||||||
|
no separate revert feature or `{recorded_original}` placeholder is needed.
|
||||||
|
|
||||||
|
Implemented changes:
|
||||||
|
|
||||||
|
- `FilenameTemplateUseCase._process_model`: an empty template now resolves
|
||||||
|
the target name from the sidecar's `original_file_name` via the injected
|
||||||
|
`metadata_loader` (default `load_local_metadata`); models without a
|
||||||
|
recorded original or whose original matches the current name are skipped.
|
||||||
|
Cache entries do not project `original_file_name`, so the sidecar is read
|
||||||
|
per model.
|
||||||
|
- `SettingsManager.js`: removed the empty-template early return and the
|
||||||
|
apply-button disable (`updateFilenameTemplateApplyButton` deleted — the
|
||||||
|
button is now always enabled). The browser-native `confirm()` was replaced
|
||||||
|
with `filenameTemplateConfirmModal`
|
||||||
|
(`templates/components/modals/confirm_modals.html`), a **self-managed**
|
||||||
|
modal (like `DirectoryPickerModal`, NOT registered with ModalManager):
|
||||||
|
ModalManager's "close current modal on open" behavior would kill the
|
||||||
|
settings modal underneath. It stacks via `z-index: 10010`
|
||||||
|
(`delete-modal.css`), handles ESC in capture phase with
|
||||||
|
`stopPropagation`, and shows apply vs revert wording
|
||||||
|
(`modals.filenameTemplateConfirm.titleApply` / `titleRevert` /
|
||||||
|
`revertButton`; messages reuse `settings.filenameTemplates.confirmApply` /
|
||||||
|
`confirmRevert`).
|
||||||
|
- `locales/en.json`: reworded `help` / `applyHelp`, replaced
|
||||||
|
`validation.keepOriginal` with `validation.restoreOriginal`
|
||||||
|
("Valid (empty template restores original filenames)"), added
|
||||||
|
`confirmRevert`, removed the now-unused `emptyTemplateInfo`. Other locales
|
||||||
|
re-synced with `[TODO: Translate]` placeholders — retranslation waits for
|
||||||
|
the feature owner's request per `docs/i18n-translation-guidelines.md` §7.
|
||||||
|
- Tests: revert / no-record-skip / same-name-skip cases in
|
||||||
|
`tests/services/test_use_cases.py`; modal confirm-and-revert and
|
||||||
|
cancel paths in
|
||||||
|
`tests/frontend/managers/settingsManager.filenameTemplates.test.js`.
|
||||||
|
|
||||||
|
Sandbox E2E verified (standalone server, sandboxed settings + library under
|
||||||
|
`/tmp`, 2026-09-19): template apply renames and records
|
||||||
|
`original_file_name`; empty-template apply reverts to the recorded name;
|
||||||
|
revert target occupied by a newer file counts as failure and keeps the
|
||||||
|
current name; models without a recorded original are skipped;
|
||||||
|
apply → revert → re-apply cycles repeat cleanly.
|
||||||
|
|
||||||
|
Standing caveats (unchanged):
|
||||||
|
|
||||||
|
- The revert target may collide with an existing file — the existing conflict
|
||||||
|
handling (count as failure, keep current name) covers this.
|
||||||
|
- `original_file_name` only exists for models renamed after `2bc9860b`;
|
||||||
|
older renames have no recorded original and are skipped.
|
||||||
|
- `original_file_name` is kept (not cleared) after a revert, so
|
||||||
|
apply → revert → re-apply stays repeatable.
|
||||||
|
|
||||||
|
## Follow-up 3 — Cross-page refresh after bulk apply
|
||||||
|
|
||||||
|
**Problem:** the settings-modal "Apply to Library Now" button calls
|
||||||
|
`resetAndReload(true)`, which refreshes only the page type currently open.
|
||||||
|
Applying the checkpoint template while on the loras page leaves the loras
|
||||||
|
view refreshed but does not touch the checkpoints page state (same
|
||||||
|
limitation as the existing bulk auto-organize flow in
|
||||||
|
`static/js/managers/SettingsManager.js#applyFilenameTemplate`).
|
||||||
|
|
||||||
|
**Fix options:** broadcast a generic "library changed" event that every
|
||||||
|
page's state listens to, or accept the limitation (the other page reloads
|
||||||
|
its cache on next visit). Low priority.
|
||||||
@@ -0,0 +1,337 @@
|
|||||||
|
# Plan: Global Rate-Limit Abidance for Recipe Ingest & Metadata Fetching
|
||||||
|
|
||||||
|
**Issue:** [#1085 — Large Recipe Ingest Appears to not abide by vendor rate limits, possibly a few other errors?](https://github.com/willmiao/ComfyUI-Lora-Manager/issues/1085)
|
||||||
|
**Status:** v2 — reviewed; decisions recorded in §10. **Phase 1 implemented**
|
||||||
|
(2026-08-27, commit `c2a2048c`): coordinator + downloader gate + Fix C
|
||||||
|
failover semantics + helper double-wait fix + settings. **Phase 2
|
||||||
|
implemented** (2026-08-27): batch-import rate-limit failures map to
|
||||||
|
`SKIPPED` + `rate_limited` WebSocket flag + UI slowdown hint (toast + status
|
||||||
|
text, i18n keys synced); `download_to_memory` / `get_response_headers` /
|
||||||
|
`download_file` register 429 cooldowns. Changes vs v1: Fix C moved to
|
||||||
|
Phase 1, helper double-wait resolved in Phase 1, gate/guard ordering
|
||||||
|
specified.
|
||||||
|
**Scope:** HTTP API traffic to CivitAI (`civitai.red`) and CivArchive (`civarchive.com`) from metadata fetching (bulk refresh, metadata sync, recipe analysis/enrichment, usage-control lookups). Large binary downloads (model files / preview images via `download_file`) are out of scope for *pacing* (they are already single-connection transfers) but their 429 responses should still be *registered*.
|
||||||
|
|
||||||
|
> Context: a first batch of fixes for this issue was already committed as
|
||||||
|
> `ee233548` ("fix(recipes): enforce batch-import concurrency bound and harden
|
||||||
|
> ingest errors (#1085)"): the batch-import concurrency controller now shares a
|
||||||
|
> real semaphore (bounds 1–5 actually apply), the Comfy parser tolerates
|
||||||
|
> list/`None` `ckpt_name`, CivArchive treats empty error payloads as failures,
|
||||||
|
> and offline-cooldown short-circuits log at DEBUG. This plan covers the two
|
||||||
|
> remaining orchestration-level fixes:
|
||||||
|
> **Fix 2** — slow down globally when a vendor rate limit is hit (respect
|
||||||
|
> `Retry-After`, queue instead of hammering); **Fix 3** — stop immediately
|
||||||
|
> failing over to CivArchive when CivitAI is rate-limited.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. Problem Statement
|
||||||
|
|
||||||
|
During a large recipe ingest (e.g. importing the example-images directory,
|
||||||
|
which can be thousands of images), the manager fires one metadata request per
|
||||||
|
checkpoint + per LoRA per image through the fallback provider chain
|
||||||
|
(`civitai_api → civarchive_api → sqlite`). Consequences observed in #1085:
|
||||||
|
|
||||||
|
1. **CivitAI gets hammered** → 429s. The consumer then *immediately* tries
|
||||||
|
CivArchive for the same lookup, so **CivArchive gets hammered too** before
|
||||||
|
it was ever naturally needed (its only real job is recovering metadata for
|
||||||
|
models deleted from CivitAI).
|
||||||
|
2. Requests are retried per-call after `Retry-After`, but **each concurrent
|
||||||
|
call sleeps independently** → thundering herd: thousands of coroutines wake
|
||||||
|
at the same moment and re-flood the vendor.
|
||||||
|
3. While CivArchive is in the `ConnectivityGuard` cooldown, every batch item
|
||||||
|
short-circuits and is marked `FAILED` — the batch import's success/failure
|
||||||
|
accounting is polluted by a transient vendor state (log spam was fixed in
|
||||||
|
`ee233548`; the item-failure accounting is not).
|
||||||
|
4. `ConnectivityGuard` (`py/services/connectivity_guard.py`) only treats
|
||||||
|
transport-level unreachability as offline; **HTTP 429 is invisible to it**,
|
||||||
|
so nothing ever intentionally paces request rate.
|
||||||
|
|
||||||
|
User expectation from the issue: *"once a vendor rate limit time out is hit,
|
||||||
|
you should trigger a slow down with intentional reduction in request rate"*.
|
||||||
|
|
||||||
|
## 2. Current State (verified against code)
|
||||||
|
|
||||||
|
### 2.1 Where 429s are surfaced
|
||||||
|
|
||||||
|
- `Downloader.make_request` (`py/services/downloader.py:1120-1132`): HTTP 429 →
|
||||||
|
returns `RateLimitError(message, retry_after=…)` parsed from `Retry-After`
|
||||||
|
(missing header defaults to `None`).
|
||||||
|
- `CivitaiClient._make_request` (`py/services/civitai_client.py:97-100`):
|
||||||
|
converts `RateLimitError` to a raise immediately; no waiting. Transient
|
||||||
|
5xx/connection errors are retried 3× with 1s/2s/4s backoff.
|
||||||
|
- `CivArchiveClient._make_request` (`py/services/civarchive_client.py`):
|
||||||
|
raises `RateLimitError` with `provider="civarchive_api"` when not set.
|
||||||
|
- `_RateLimitRetryHelper` (`py/services/model_metadata_provider.py:45-102`):
|
||||||
|
per-call retry loop — sleeps `retry_after` (capped at 1800 s; `≥120 s` ⇒ no
|
||||||
|
retry), then re-raises. Because every concurrent call runs its own helper,
|
||||||
|
they sleep in parallel and re-fire in parallel.
|
||||||
|
- `FallbackMetadataProvider` (`py/services/model_metadata_provider.py:488-508,
|
||||||
|
564-584` etc.): on a final `RateLimitError` from one provider it logs
|
||||||
|
"skipping to next provider" and **continues to the next network provider** —
|
||||||
|
this is the direct cause of the CivArchive flood.
|
||||||
|
- `MetadataSyncService.fetch_and_update_model`
|
||||||
|
(`py/services/metadata_sync_service.py:248-333`): manually iterates
|
||||||
|
`provider_attempts`; on `RateLimitError` it `continue`s to the next provider
|
||||||
|
(same failover problem), then reports `"Rate limited"` when nothing
|
||||||
|
succeeded.
|
||||||
|
- `Downloader.make_request` has a per-destination scope already available:
|
||||||
|
`_guard_destination(url)` returns the hostname (`downloader.py:1194-1199`),
|
||||||
|
used by `ConnectivityGuard`.
|
||||||
|
|
||||||
|
### 2.2 What pacing exists today
|
||||||
|
|
||||||
|
- `ConnectivityGuard`: per-destination cooldown (30 s base, ×2 per extra
|
||||||
|
failure batch, 300 s cap) triggered only by transport errors
|
||||||
|
(`connectivity_guard.py:168-197`).
|
||||||
|
- `AdaptiveConcurrencyController` (batch import, fixed in `ee233548`): shared
|
||||||
|
semaphore enforces 1–5 concurrent items; *duration*-based adjustment only —
|
||||||
|
it never sees HTTP statuses, so it cannot distinguish "slow because rate
|
||||||
|
limited" from "slow because big image".
|
||||||
|
- No token bucket, no minimum inter-request interval, no shared
|
||||||
|
`Retry-After` gate anywhere (`grep` for throttle/token-bucket/rate-limiter:
|
||||||
|
0 hits).
|
||||||
|
|
||||||
|
## 3. Requirements & Constraints
|
||||||
|
|
||||||
|
R1. **Respect `Retry-After`.** After a 429, no further request to that
|
||||||
|
destination may be sent before the vendor's retry window elapses.
|
||||||
|
R2. **No thundering herd.** Concurrent waiters must share one wake-up (gate),
|
||||||
|
not sleep independently.
|
||||||
|
R3. **No double load.** A CivitAI 429 must not trigger a CivArchive request
|
||||||
|
for the same lookup. CivArchive should only be consulted when CivitAI
|
||||||
|
legitimately has no answer (404 / "not found"), or when CivitAI is
|
||||||
|
unreachable long-term.
|
||||||
|
R4. **No spurious item failures.** A rate-limited request must not turn a
|
||||||
|
batch-import item into `FAILED`; it should wait (bounded) and retry, or at
|
||||||
|
worst be `SKIPPED` with a clear "rate limited" reason (re-runnable import).
|
||||||
|
R5. **Never hang forever.** All waiting is bounded by a configurable cap; on
|
||||||
|
expiry the caller receives the `RateLimitError` and can decide.
|
||||||
|
R6. **Keep legitimate failover.** Deleted-model recovery via CivArchive/sqlite
|
||||||
|
must keep working (404 paths unchanged).
|
||||||
|
R7. **Single choke point.** The pacing gate should live where every API call
|
||||||
|
passes (the `Downloader`), so bulk refresh, metadata sync, recipe
|
||||||
|
analysis, and usage-control lookups all benefit without per-feature work.
|
||||||
|
|
||||||
|
## 4. Approach Comparison
|
||||||
|
|
||||||
|
### A. Reactive gate — shared `Retry-After` deadman clock (recommended core)
|
||||||
|
|
||||||
|
A process-wide, per-destination coordinator records the *next-allowed-send*
|
||||||
|
timestamp from each 429 (`now + max(retry_after, backoff)`). Every request
|
||||||
|
through `Downloader.make_request` consults the gate *before sending* and *when
|
||||||
|
a 429 arrives*; waiters block on a shared `asyncio.Event` that fires when the
|
||||||
|
cooldown expires.
|
||||||
|
|
||||||
|
- Pros: single choke point (R7); herd-free (R2); honors server guidance (R1);
|
||||||
|
no guessing at vendor limits; covers all providers automatically; reuses
|
||||||
|
existing per-destination scoping.
|
||||||
|
- Cons: still experiences 429s before slowing down (reactive); long
|
||||||
|
`Retry-After` windows (CivArchive has been observed at ~1500 s) need a sane
|
||||||
|
wait cap + skip/retry UX.
|
||||||
|
|
||||||
|
### B. Preemptive pacing — minimum inter-request interval (recommended companion)
|
||||||
|
|
||||||
|
Per-destination token bucket (simplest form: capacity 1 — at least `N` seconds
|
||||||
|
between consecutive API requests; `N` configurable, default ~0.75 s ≈ 80
|
||||||
|
r/min ceiling).
|
||||||
|
|
||||||
|
- Pros: prevents most 429s before they happen — exactly the "intentional
|
||||||
|
reduction in request rate" the issue asks for; trivial to implement on top
|
||||||
|
of A's coordinator.
|
||||||
|
- Cons: adds latency to bulk operations (thousands of models × `N`); the *exact*
|
||||||
|
vendor limits are unknown (CivitAI anonymous vs keyed vs `civitai.red`
|
||||||
|
mirror differ), so the default must be conservative-but-not-crippling and
|
||||||
|
settings-tunable.
|
||||||
|
|
||||||
|
### C. Fallback semantics change — stop network→network failover on 429 (must-do, low risk)
|
||||||
|
|
||||||
|
`FallbackMetadataProvider` (and `MetadataSyncService.fetch_and_update_model`'s
|
||||||
|
manual loop) must treat a final `RateLimitError` as a **terminal, non-failover
|
||||||
|
result** for network providers. Local-only providers (sqlite archive DB) may
|
||||||
|
stay as a last resort (no vendor cost).
|
||||||
|
|
||||||
|
- Pros: directly removes the CivArchive flood; small, surgical change.
|
||||||
|
- Cons: none significant; requires care to keep 404-failover intact (R6).
|
||||||
|
|
||||||
|
### Rejected / deferred
|
||||||
|
|
||||||
|
- **Per-feature retry queues** (batch import pauses & resumes whole batches):
|
||||||
|
richer UX but much larger change (batch state machine, WebSocket states);
|
||||||
|
unnecessary once A+B make requests wait at the choke point. Defer unless
|
||||||
|
review finds the bounded-wait UX insufficient.
|
||||||
|
- **Full token bucket with burst credit**: overkill; capacity-1 interval is
|
||||||
|
enough given the shared semaphore already caps concurrency at 5.
|
||||||
|
- **Retrying in `connectivity_guard`**: wrong layer — the guard is about
|
||||||
|
transport reachability, not vendor quota.
|
||||||
|
|
||||||
|
## 5. Recommended Architecture
|
||||||
|
|
||||||
|
New singleton **`RateLimitCoordinator`** (`py/services/rate_limit_coordinator.py`,
|
||||||
|
mirroring `ConnectivityGuard`'s singleton + per-destination patterns):
|
||||||
|
|
||||||
|
```
|
||||||
|
state per destination (hostname):
|
||||||
|
next_allowed_send: float (monotonic) # from 429 Retry-After + backoff
|
||||||
|
consecutive_429: int # for backoff growth
|
||||||
|
last_send_at: float # for min-interval pacing
|
||||||
|
waiters: list[Future] | asyncio.Event # shared wake-up per cooldown cycle
|
||||||
|
```
|
||||||
|
|
||||||
|
API:
|
||||||
|
|
||||||
|
- `async wait_for_slot(destination, request_started_within_window: bool)`
|
||||||
|
— called by `Downloader.make_request` *before* sending (blocks until
|
||||||
|
`min(now >= next_allowed_send)` and inter-request interval elapses) and
|
||||||
|
re-armable after a 429.
|
||||||
|
- `register_rate_limit(destination, retry_after: float | None)`
|
||||||
|
— called on 429: `next_allowed_send = max(now + retry_after_or_backoff, current)`;
|
||||||
|
`consecutive_429 += 1`; backoff = `retry_after` honored, else exponential
|
||||||
|
`30 · 2^(n-1)` capped at 1800 s; creates/re-arms the shared wake-up event.
|
||||||
|
- `register_success(destination)` — resets `consecutive_429` (called from the
|
||||||
|
existing 200 path in `make_request`).
|
||||||
|
- `remaining_seconds(destination)`, `in_cooldown(destination)` — for tests and
|
||||||
|
diagnostics.
|
||||||
|
|
||||||
|
Enforcement points:
|
||||||
|
|
||||||
|
1. **`Downloader.make_request`** (`downloader.py:1102-1132`): ordering inside
|
||||||
|
the method is **connectivity-guard fail-fast first** (offline short-circuit
|
||||||
|
costs nothing to check), **then** `await coordinator.wait_for_slot(destination)`
|
||||||
|
before `session.request`. On 429: `coordinator.register_rate_limit(...)`,
|
||||||
|
then *wait for the gate and re-send* (loop, bounded by
|
||||||
|
`rate_limit_max_wait_seconds`, default 300; `retry_after ≥ cap` ⇒ fail
|
||||||
|
immediately). After the loop, return the `RateLimitError` to the caller
|
||||||
|
(unchanged contract) **with `exc.gate_handled = True` set** so downstream
|
||||||
|
retry helpers know the wait already happened. 200 path calls
|
||||||
|
`register_success`.
|
||||||
|
2. **`Downloader.download_to_memory` / `get_response_headers`** (phase 2):
|
||||||
|
register 429s (so API calls queue); waiting only in `make_request`
|
||||||
|
initially.
|
||||||
|
3. **`FallbackMetadataProvider`** (`model_metadata_provider.py`): remove
|
||||||
|
network→network failover on `RateLimitError` — re-raise; only sqlite stays
|
||||||
|
as a local last resort (implementation: per-method `except RateLimitError`
|
||||||
|
handler that marks the chain rate-limited and stops iterating).
|
||||||
|
4. **`MetadataSyncService.fetch_and_update_model`**
|
||||||
|
(`metadata_sync_service.py:248-333`): on `RateLimitError` from the default
|
||||||
|
provider, stop appending further network providers (sqlite may remain);
|
||||||
|
the existing `any_rate_limited` merge already produces `"Rate limited"`.
|
||||||
|
5. **Batch import** (`batch_import_service.py`): no structural change needed —
|
||||||
|
items now wait inside `make_request`; optionally (phase 2) map residual
|
||||||
|
rate-limit failures (after the wait cap) to `SKIPPED` with
|
||||||
|
`"rate limited (retry_after=…s); re-run the import later"` instead of
|
||||||
|
`FAILED`, and surface a `rate_limited` flag in the WebSocket progress
|
||||||
|
broadcast.
|
||||||
|
6. **`_RateLimitRetryHelper` retries** (`model_metadata_provider.py`):
|
||||||
|
**Phase 1** — when the raised `RateLimitError` carries `gate_handled = True`
|
||||||
|
(set by the downloader after honoring the gate), the helper skips its own
|
||||||
|
`retry_after` sleep and re-raises immediately, eliminating the double wait.
|
||||||
|
The wiring stays so a `RateLimitError` still propagates cleanly; full
|
||||||
|
demotion/removal can follow once the gate proves out.
|
||||||
|
|
||||||
|
Settings (`settings.json`, schema extension in `SettingsManager`):
|
||||||
|
|
||||||
|
| key | default | meaning |
|
||||||
|
|---|---|---|
|
||||||
|
| `rate_limit_gate_enabled` | `true` | master switch for the coordinator |
|
||||||
|
| `rate_limit_max_wait_seconds` | `300` | how long `make_request` waits on a 429 gate before returning the error |
|
||||||
|
| `rate_limit_min_interval_seconds` | `0.75` | minimum seconds between API requests per destination (pacing, R6-friendly conservative default) |
|
||||||
|
|
||||||
|
## 6. Changes by File
|
||||||
|
|
||||||
|
| File | Change |
|
||||||
|
|---|---|
|
||||||
|
| `py/services/rate_limit_coordinator.py` (new) | coordinator singleton + per-destination state + tests seam |
|
||||||
|
| `py/services/downloader.py` | gate pre-check + 429 register/wait/retry loop + `register_success`; log the 429 notice at INFO once per cooldown, then DEBUG |
|
||||||
|
| `py/services/model_metadata_provider.py` | `FallbackMetadataProvider`: stop network failover on `RateLimitError`; helper skips its sleep when the error is marked `gate_handled` |
|
||||||
|
| `py/services/metadata_sync_service.py` | `fetch_and_update_model`: same failover semantics; keep sqlite last resort |
|
||||||
|
| `py/services/batch_import_service.py` | (phase 2) rate-limit failures → `SKIPPED` + `rate_limited` progress flag |
|
||||||
|
| `py/services/settings_manager.py` | new settings keys + defaults |
|
||||||
|
| `tests/services/test_rate_limit_coordinator.py` (new) | gate unit tests |
|
||||||
|
| `tests/services/test_civitai_client.py` / `test_civarchive_client.py` | provider-level 429 behavior |
|
||||||
|
| `tests/services/test_metadata_service.py` | failover-chain tests |
|
||||||
|
| `tests/services/test_batch_import_service.py` | SKIPPED-on-rate-limit |
|
||||||
|
|
||||||
|
## 7. Impact, Risks, Open Questions
|
||||||
|
|
||||||
|
- **Behavior change**: with the gate in `make_request`, any request can block
|
||||||
|
up to the wait cap — UI actions that call the API (e.g. a model-details
|
||||||
|
fetch) may take longer during cooldowns. Mitigation: bounded cap + INFO log
|
||||||
|
+ the existing async request handling already tolerates slow responses.
|
||||||
|
**Decided (§10): interactive requests take the same bounded wait** — one
|
||||||
|
behavior, no call-source plumbing; cooldowns are usually short.
|
||||||
|
- **Gate waits occupy batch slots**: with the 1–5 batch semaphore, all slots
|
||||||
|
can park on a gate simultaneously, freezing visible progress for up to one
|
||||||
|
wait cap per wave. Bounded and acceptable; the phase-2 `SKIPPED` mapping +
|
||||||
|
WebSocket `rate_limited` flag (both confirmed in scope, §10) make the stall
|
||||||
|
visible and recoverable.
|
||||||
|
- **Rate limit reality check**: CivitAI anonymous vs keyed limits, and whether
|
||||||
|
`civitai.red` differs, is unverified. Default pacing `0.75 s/req` is a
|
||||||
|
conservative guess (R6). Open question for maintainer: preferred default
|
||||||
|
and whether an API-keyed ceiling should be higher.
|
||||||
|
- **Long CivArchive windows**: `Retry-After ~1500 s` observed in code
|
||||||
|
comments. **Decided (§10): keep the 300 s default cap** — such lookups
|
||||||
|
fail/skip rather than park a request path for 25 minutes; batch import maps
|
||||||
|
them to `SKIPPED` (phase 2) so the user can re-run later.
|
||||||
|
- **Double waiting**: `_RateLimitRetryHelper` + gate could stack waits.
|
||||||
|
**Resolved in Phase 1**: the downloader marks gate-honored errors with
|
||||||
|
`gate_handled = True` and the helper skips its own sleep for those.
|
||||||
|
- **Downloads**: `download_file` 429s return an error to download managers
|
||||||
|
unchanged (already handled); only *registration* is proposed, so future
|
||||||
|
API calls queue behind a large `Retry-After` from a download burst.
|
||||||
|
|
||||||
|
## 8. Test Plan
|
||||||
|
|
||||||
|
1. **Coordinator unit tests** (new file):
|
||||||
|
- 429 with `retry_after` → `wait_for_slot` blocks ~that long, then passes.
|
||||||
|
- N concurrent waiters all wake together (herd test, wall-clock ≈ one
|
||||||
|
window, not N windows).
|
||||||
|
- Consecutive 429s grow backoff; `register_success` resets.
|
||||||
|
- Missing `Retry-After` → default backoff path.
|
||||||
|
- Wait cap: request fails after `rate_limit_max_wait_seconds` with
|
||||||
|
`RateLimitError`.
|
||||||
|
2. **Downloader tests** (mock aiohttp session): 429 then 200 → `make_request`
|
||||||
|
returns success after gate delay; two back-to-back calls to the same
|
||||||
|
destination are spaced ≥ `min_interval`; different destinations are not
|
||||||
|
spaced.
|
||||||
|
3. **Provider tests**: `FallbackMetadataProvider.get_model_version_info` —
|
||||||
|
Civitai raises `RateLimitError` → CivArchive mock **not called**; 404 still
|
||||||
|
falls through to CivArchive; sqlite still tried after network 429.
|
||||||
|
4. **Sync-service test**: `fetch_and_update_model` with a rate-limited default
|
||||||
|
provider → result error contains `"Rate limited"` and sqlite attempt state
|
||||||
|
unchanged.
|
||||||
|
5. **Batch-import test**: analysis provider 429s first, then succeeds →
|
||||||
|
item ends `SUCCESS` (wait path), and post-cap 429 → `SKIPPED` with
|
||||||
|
rate-limit reason (phase 2).
|
||||||
|
6. Full regression: `pytest tests/services tests/routes tests/standalone`
|
||||||
|
(currently 1582 passing).
|
||||||
|
|
||||||
|
## 9. Implementation Phases
|
||||||
|
|
||||||
|
- **Phase 1 (this plan, reviewed):** `RateLimitCoordinator` +
|
||||||
|
`Downloader.make_request` integration (guard fail-fast → gate pre-check
|
||||||
|
pacing → 429 register/wait/retry loop with cap → `gate_handled` marking) +
|
||||||
|
settings + **Fix C failover semantics** (`FallbackMetadataProvider`,
|
||||||
|
`fetch_and_update_model` — moved up from phase 2: smallest diff, kills the
|
||||||
|
CivArchive flood immediately, independent of coordinator correctness) +
|
||||||
|
`_RateLimitRetryHelper` double-wait fix + coordinator/downloader/provider/
|
||||||
|
sync tests.
|
||||||
|
- **Phase 2:** batch-import `SKIPPED`-on-rate-limit + `rate_limited` WebSocket
|
||||||
|
progress flag + slowdown hint (confirmed, §10),
|
||||||
|
`download_to_memory`/HEAD 429 registration, batch tests.
|
||||||
|
- **Phase 3:** full regression + docs + commit referencing `(#1085)`.
|
||||||
|
|
||||||
|
## 10. Review Checklist — Decisions (2026-08-27)
|
||||||
|
|
||||||
|
- [x] Default pacing interval `0.75 s` — **accepted** as conservative default;
|
||||||
|
tunable via `rate_limit_min_interval_seconds`. Revisit if CivitAI
|
||||||
|
publishes keyed/anonymous ceilings.
|
||||||
|
- [x] Wait cap `300 s` — **accepted**; long-window CivArchive lookups fail →
|
||||||
|
batch import marks them `SKIPPED` with a rate-limit reason (phase 2).
|
||||||
|
- [x] Interactive API calls also wait (bounded) — **yes**, same behavior for
|
||||||
|
all callers.
|
||||||
|
- [x] Keep sqlite as last resort behind a network rate limit — **yes**
|
||||||
|
(local-only, no vendor cost).
|
||||||
|
- [x] UI hint — **yes**: WebSocket `rate_limited` flag + "rate limited —
|
||||||
|
slowing down" hint in batch-import progress (phase 2); INFO logging
|
||||||
|
regardless.
|
||||||
@@ -0,0 +1,363 @@
|
|||||||
|
# Plan: "Other Models" Page — Unified Management for VAE / Upscaler / Text Encoder / etc.
|
||||||
|
|
||||||
|
**Status:** v2 — **Phase 1 implemented** (2026-09-12, commits `27da7b3c` backend + `fa7ce725` frontend; verified live against a running ComfyUI instance: scan/hash/sub_type-derivation/fetch/previews all green). **Phase 2 implemented** (2026-09-12, per §9 design; full pytest + vitest green). **Phase 3 implemented** (§11: opt-in management toggles; default off). **i18n done** (2026-09-13): all 36 new keys translated in the 9 non-English locales — the `[TODO: Translate]` placeholders left by the sync script during development are gone (see `docs/i18n-translation-guidelines.md` §2, "Other Models feature"). **Default set revised (pre-release):** only `vae` / `upscaler` / `text_encoder` are managed by default — `clip_vision` and `controlnet` are both opt-in (§2, §11.1.1).
|
||||||
|
**Scope (Phase 1):** scan + manage (list, search, filter, tags, folders, preview, rename, move, delete/exclude, CivitAI metadata fetch) for a new model type `other`, exposed as a new web page. **Phase 2 (§9):** one-click download from CivitAI for these types.
|
||||||
|
|
||||||
|
## 1. Goal
|
||||||
|
|
||||||
|
Today the manager supports three model types:
|
||||||
|
|
||||||
|
| page | model_type | sub_types |
|
||||||
|
|---|---|---|
|
||||||
|
| `/loras` | `lora` | `lora`, `locon`, `dora` |
|
||||||
|
| `/checkpoints` | `checkpoint` | `checkpoint`, `diffusion_model` |
|
||||||
|
| `/embeddings` | `embedding` | `embedding` |
|
||||||
|
|
||||||
|
Add a fourth page that manages "everything else" — VAE, upscalers, text encoders / CLIP, CLIP vision, optionally ControlNet — with a folder→sub_type mapping table so new ComfyUI folder categories can be added later by configuration, not code.
|
||||||
|
|
||||||
|
## 2. Locked Decisions
|
||||||
|
|
||||||
|
1. **Architecture: one scanner + one service + one page, sub_type derived by location.**
|
||||||
|
Replicates the checkpoint pattern (`CheckpointScanner` aggregates `checkpoints` + `unet` roots and derives `checkpoint` vs `diffusion_model` from the root containing the file, `py/services/checkpoint_scanner.py:384-415`). One `OtherScanner` aggregates all enabled folder roots; `resolve_sub_type_for_path()` maps each root to a sub_type. No per-category scanners.
|
||||||
|
|
||||||
|
2. **Naming: internal `model_type = "other"`, route prefix `/other`, page id `other`.**
|
||||||
|
- `misc` is rejected: `py/routes/misc_routes.py` already owns that name for system/settings routes (`/api/lm/settings`, `/api/lm/doctor/*`).
|
||||||
|
- `components` is rejected: `templates/components/` and `static/js/components/` directories would make `components.html` / `components.js` confusing neighbors.
|
||||||
|
- `other` matches CivitAI's `Other` fallback type semantics. The **display name** is an i18n string (`other.title`, e.g. "Other Models") and can be renamed later without touching code.
|
||||||
|
|
||||||
|
3. **sub_type values:** snake_case, aligned with CivitAI `ModelType` semantics:
|
||||||
|
|
||||||
|
| sub_type | ComfyUI `folder_paths` key(s) | CivitAI ModelType | enabled by default |
|
||||||
|
|---|---|---|---|
|
||||||
|
| `vae` | `vae` | `VAE` | yes |
|
||||||
|
| `upscaler` | `upscale_models` | `Upscaler` | yes |
|
||||||
|
| `text_encoder` | `text_encoders`, `clip` (legacy) | `TextEncoder` (CLIP is retired upstream) | yes |
|
||||||
|
| `clip_vision` | `clip_vision` | `CLIPVision` | no (mapping present, opt-in) |
|
||||||
|
| `controlnet` | `controlnet` | `Controlnet` | no (mapping present, opt-in) |
|
||||||
|
|
||||||
|
New folder categories = one line in the mapping table (see §4.1).
|
||||||
|
|
||||||
|
**Why only three are on by default** (revised in Phase 3, before release):
|
||||||
|
VAE, upscalers and text encoders are dependency-style assets every pipeline
|
||||||
|
needs, and "which one am I actually using" is the recurring problem they
|
||||||
|
solve. `clip_vision` and `controlnet` are workflow-driven instead
|
||||||
|
(IPAdapter/SVD image conditioning; per-workflow ControlNet variants), and
|
||||||
|
ControlNet libraries routinely run to dozens of files, so both are treated
|
||||||
|
symmetrically as opt-in. Enumerating all five as "the default set" was not
|
||||||
|
defensible on demand breadth alone.
|
||||||
|
|
||||||
|
4. **Phase 1 = scan/manage only.** Downloads from CivitAI (`download_manager.py` type mapping, default-root settings keys, download routing) are Phase 2 (§9). CivitAI **metadata fetch** for existing files IS in Phase 1 (hash-based lookup is type-agnostic; only the type-validation hook needs new values).
|
||||||
|
|
||||||
|
5. **Out of scope (default off, revisit later):** usage statistics buckets, recipe matching (`recipe_scanner.py` only merges lora+checkpoint scanners), statistics page, embeddings re-classification (stays its own page — merging would be a breaking change).
|
||||||
|
|
||||||
|
## 3. Why This Works With Minimal Churn
|
||||||
|
|
||||||
|
- `ModelScanner` (`py/services/model_scanner.py:93`) is specialized entirely via constructor params (`model_type`, `model_class`, `file_extensions`) + optional hooks (`adjust_metadata`, `adjust_cached_entry`, `resolve_sub_type_for_path`, `model_scanner.py:1429-1443`).
|
||||||
|
- `BaseModelService` subclasses can be one method (`EmbeddingService` implements only `format_response`, `py/services/embedding_service.py:12`).
|
||||||
|
- Routes: `ModelServiceFactory.register_model_type()` (`py/services/model_service_factory.py:120-136`) + `COMMON_ROUTE_DEFINITIONS` (`py/routes/model_route_registrar.py:23-149`) generate the full `/api/lm/{prefix}/*` surface (~50 endpoints) plus the `GET /{prefix}` page route.
|
||||||
|
- `PersistentModelCache` (`py/services/persistent_model_cache.py:526-606`) is a single `models` table keyed `(model_type, file_path)` with `model_type` as free text — **zero schema change**.
|
||||||
|
- Frontend `apiConfig.js` (`static/js/api/apiConfig.js:51`) generates all endpoints from the model-type string; `ModelCard.js:670-675` renders the sub_type badge from data; the checkpoints page already demonstrates the "one page, multiple sub_types" filter (`header.html:298`).
|
||||||
|
|
||||||
|
## 4. Backend Changes
|
||||||
|
|
||||||
|
### 4.1 New constants — `py/utils/constants.py`
|
||||||
|
|
||||||
|
```python
|
||||||
|
# folder_paths key -> sub_type; single source of truth for extensibility
|
||||||
|
OTHER_MODEL_FOLDER_SUBTYPES = {
|
||||||
|
"vae": "vae",
|
||||||
|
"upscale_models": "upscaler",
|
||||||
|
"text_encoders": "text_encoder",
|
||||||
|
"clip": "text_encoder", # legacy ComfyUI key
|
||||||
|
"clip_vision": "clip_vision",
|
||||||
|
"controlnet": "controlnet",
|
||||||
|
}
|
||||||
|
DEFAULT_OTHER_MODEL_FOLDERS = ("vae", "upscale_models", "text_encoders", "clip", "clip_vision")
|
||||||
|
VALID_OTHER_SUB_TYPES = ["vae", "upscaler", "text_encoder", "clip_vision", "controlnet"]
|
||||||
|
# CivitAI model.type values accepted for this page (fetch-metadata validation)
|
||||||
|
VALID_OTHER_CIVITAI_TYPES = {"vae", "upscaler", "textencoder", "clipvision", "controlnet", "other"}
|
||||||
|
```
|
||||||
|
|
||||||
|
Also extend `CIVITAI_USER_MODEL_TYPES` (`constants.py:90`) if user-model queries should include these types.
|
||||||
|
|
||||||
|
### 4.2 New files (mirror the embedding/checkpoint implementations)
|
||||||
|
|
||||||
|
1. **`py/utils/models.py`** — add `OtherModelMetadata(BaseModelMetadata)`: default `sub_type="vae"` placeholder overridden by scanner hook; `from_civitai_info` mapping CivitAI types → our sub_types (`TextEncoder`→`text_encoder`, `CLIPVision`→`clip_vision`, `Upscaler`→`upscaler`, `VAE`→`vae`, `Controlnet`→`controlnet`, else `other`-ish fallback to folder-derived sub_type).
|
||||||
|
2. **`py/services/other_scanner.py`** — `OtherScanner(ModelScanner)`:
|
||||||
|
- `model_type="other"`, extensions: reuse the checkpoint set (`safetensors/pt/pt2/bin/pth/pkl/sft/gguf`).
|
||||||
|
- `get_model_roots()`: iterate `OTHER_MODEL_FOLDER_SUBTYPES` ∩ enabled keys, pull each from `config` (§4.3); dedupe; build `root → sub_type` map (normalized abspaths; multiple keys may share a sub_type).
|
||||||
|
- Implement all three hooks like `CheckpointScanner` (`checkpoint_scanner.py:384-415`): `resolve_sub_type_for_path` by longest-prefix root match, `adjust_metadata`, `adjust_cached_entry` (sub_type is re-derived on cache load, never persisted).
|
||||||
|
- **Lazy hashing, checkpoint-style**: text encoders (T5-XXL ≈ 10 GB) make eager sha256 painful. Copy the `hash_status="pending"` + singleflight `calculate_hash_for_model` pattern from `CheckpointScanner`.
|
||||||
|
3. **`py/services/other_model_service.py`** — `OtherModelService(BaseModelService)`, `format_response` only (no usage_count, like `EmbeddingService`).
|
||||||
|
4. **`py/routes/other_routes.py`** — `OtherRoutes(BaseModelRoutes)`, `template_name="other.html"`, hooks:
|
||||||
|
- `_validate_civitai_model_type` → `VALID_OTHER_CIVITAI_TYPES`
|
||||||
|
- `_get_expected_model_types`, `_parse_specific_params` (no type-specific download params in Phase 1)
|
||||||
|
- `initialize_services()` on `app.on_startup` pulling `ServiceRegistry.get_other_scanner()`.
|
||||||
|
|
||||||
|
### 4.3 `py/config.py`
|
||||||
|
|
||||||
|
- New `other_roots` property: for each enabled key in `OTHER_MODEL_FOLDER_SUBTYPES`, `folder_paths.get_folder_paths(key)` (plugin mode) — standalone mode needs nothing new: `MockFolderPaths` (`standalone.py:66-105`) already serves arbitrary keys from `settings.json.folder_paths`.
|
||||||
|
- Follow the existing per-type recipe: an `_prepare_other_paths()` (dedupe + symlink registration; also **cross-scanner overlap detection** — warn if an `other` root is already covered by checkpoints/unet/embedding roots, mirroring the checkpoint/unet overlap check).
|
||||||
|
- Wire into: `_apply_library_paths`, `_symlink_roots()`, `_rebuild_preview_roots()` (hard requirement — preview images are served per registered root), `save_folder_paths_to_settings()`.
|
||||||
|
|
||||||
|
### 4.4 Existing-file edits (the "type string scatter" — each is a small branch/entry)
|
||||||
|
|
||||||
|
| file | change |
|
||||||
|
|---|---|
|
||||||
|
| `py/services/model_service_factory.py:120` | register `("other", OtherModelService, OtherRoutes)` in `register_default_model_types()` |
|
||||||
|
| `py/services/service_registry.py` | add `get_other_scanner()` (mirror `:297` `get_embedding_scanner`) |
|
||||||
|
| `py/services/model_scanner.py:67` | `PAGE_TYPE_MAP['other'] = 'other'` (WebSocket progress) |
|
||||||
|
| `py/services/base_model_service.py:896-906` | `get_model_types()` branch → `VALID_OTHER_SUB_TYPES` |
|
||||||
|
| `py/lora_manager.py` | `_initialize_services` scanner task list (`:219-242`), `_cleanup` cancel list (`:463`), `_cleanup_backup_files` roots (`:327-330`) |
|
||||||
|
| `py/routes/handlers/misc_handlers.py` | `scanner_getters` (`:657-661`) + `scanner_factories` (`:757-759`) so Doctor / init-status / refresh-all see the new scanner |
|
||||||
|
| `py/services/pending_delete_service.py` | `_PAGE_TYPE` map (`:57-61`) + scanner getter list (`:983-985`) |
|
||||||
|
| `py/metadata_ops/__init__.py:36-38` | `SCANNER_TYPE_MAP['other']` |
|
||||||
|
| `settings.json.example` | document optional `folder_paths` keys: `vae`, `upscale_models`, `text_encoders`, `clip_vision` |
|
||||||
|
|
||||||
|
**Explicitly NOT touched in Phase 1:** `py/services/download_manager.py`, `py/services/download_routing.py`, `py/services/settings_manager.py` default-root keys, `py/routes/stats_routes.py`, `py/utils/usage_stats.py`, `py/services/recipe_scanner.py`, `py/metadata_collector/`, `py/nodes/`.
|
||||||
|
|
||||||
|
**Zero-change confirmations (verified):** `PersistentModelCache`, `ModelUpdateService`, `DownloadedVersionHistoryService`, `MetadataSyncService` + provider chain (type-agnostic hash lookups), `ModelFileService` / `ModelMoveService` / `ModelLifecycleService` (scanner + model_type injected), `ModelCache` / `ModelHashIndex`, `AutoV3BackfillService`.
|
||||||
|
|
||||||
|
## 5. Frontend Changes
|
||||||
|
|
||||||
|
1. **`static/js/api/apiConfig.js`** — `MODEL_TYPES.OTHER = 'other'`; `MODEL_CONFIG.other` entry (displayName, singularName, `supportsMove`, `supportsBulkOperations`; no letter filter); endpoints come free from `getApiEndpoints()` (`:51`).
|
||||||
|
2. **`static/js/api/otherApi.js`** — thin `OtherApiClient extends BaseModelApiClient` (mirror `embeddingApi.js`); register in `modelApiFactory.js`.
|
||||||
|
3. **`static/js/other.js`** — page entry (mirror `embeddings.js`): `appCore.initialize()` + `createPageControls('other')` + `initializePageFeatures()` + `ModelDuplicatesManager` + `initActiveFiltersSync('other')`.
|
||||||
|
4. **Controls & context menu** — `OtherControls extends PageControls` and `OtherContextMenu` (start from the embedding variants — the smallest); add branches in the two factories (`components/controls/index.js:15`, `components/ContextMenu/index.js:15`). Context-menu template block lives in `templates/other.html` (`{% block additional_components %}`, the checkpoints/embeddings pattern — do NOT touch the shared `context_menu.html`).
|
||||||
|
5. **`templates/other.html`** — copy `embeddings.html`: same content blocks (controls + breadcrumb + duplicates banner + folder sidebar + `#modelGrid`), `data-page="other"`, main script `/loras_static/js/other.js`.
|
||||||
|
6. **`templates/components/header.html`** — nav entry (`:23-43`, active when `request.path.startswith('/other')`); enable the `modelTypes` sub_type filter panel for `other` (`:298-305` pattern from checkpoints); check search-options panel conditions (`:199-224`).
|
||||||
|
7. **`static/js/utils/constants.js`** — `MODEL_SUBTYPE_ABBREVIATIONS` (`:115`): `vae→VAE`, `upscaler→UPS`, `text_encoder→TE`, `clip_vision→CV`, `controlnet→CN`; matching `MODEL_SUBTYPE_DISPLAY_NAMES` (`:99`). (Unknown fallback already uppercases 4 chars, but explicit mappings read better.)
|
||||||
|
8. **`static/js/core.js:110` `getPageType()`** — verify `data-page="other"` flows through `state.pages` generically; add only if the page list is enumerated anywhere.
|
||||||
|
9. No change to `web/comfyui/top_menu_extension.js` (it opens `/loras`; page-to-page nav is the header bar).
|
||||||
|
|
||||||
|
## 6. i18n
|
||||||
|
|
||||||
|
- `locales/en.json`: add `other.title` (e.g. "Other Models") + minimal `other.contextMenu.*` / `other.modelTypes.*` keys; reuse `modelCard.*`, `loras.contextMenu.*`, `common.*` wherever possible (the established pattern — checkpoints/embeddings already reuse lora keys).
|
||||||
|
- Run `python scripts/sync_translation_keys.py`; leave `[TODO: Translate]` placeholders in other locales (per `docs/i18n-translation-guidelines.md` §7 — do not translate proactively).
|
||||||
|
|
||||||
|
## 7. Testing
|
||||||
|
|
||||||
|
Follow existing conventions (`pytest.ini`, `tests/frontend/` vitest):
|
||||||
|
|
||||||
|
1. **Backend (pytest, async where needed):**
|
||||||
|
- `OtherScanner` root aggregation + `resolve_sub_type_for_path` (file under `vae/` root → `vae`; `text_encoders` and legacy `clip` both → `text_encoder`; disabled `controlnet` root not scanned).
|
||||||
|
- Cache round-trip: sub_type re-derived via `adjust_cached_entry` (not persisted).
|
||||||
|
- Lazy hash: `hash_status="pending"` default; `calculate_hash_for_model` singleflight.
|
||||||
|
- `OtherRoutes` registration smoke test: `/api/lm/other/...` endpoints exist; `_validate_civitai_model_type` accepts `vae`/`upscaler`/`textencoder`, rejects `lora`.
|
||||||
|
- Config: `other_roots` in both modes (mock `folder_paths`, and standalone `settings.json.folder_paths`).
|
||||||
|
2. **Frontend (vitest + jsdom, `tests/frontend/`):**
|
||||||
|
- `apiConfig`: `getApiEndpoints('other')` URL shapes; `modelApiFactory` returns the Other client.
|
||||||
|
- `ModelCard` badge rendering for new sub_types.
|
||||||
|
- `createPageControls('other')` / `createPageContextMenu('other')` factories.
|
||||||
|
3. **Manual UI verification by the user** (per AGENTS.md — no sandbox/browser automation): page loads, scans a real library, sub_type filter + badges, context menu actions.
|
||||||
|
|
||||||
|
## 8. Execution Order
|
||||||
|
|
||||||
|
1. `constants.py` + `OtherModelMetadata` + `config.py` roots
|
||||||
|
2. `OtherScanner` (+ registry, factory, `PAGE_TYPE_MAP`) → scanner unit tests green
|
||||||
|
3. `OtherModelService` + `OtherRoutes` + handler/registrar wiring + `lora_manager.py` lifecycle → route tests green
|
||||||
|
4. Doctor/pending-delete/metadata-ops scatter entries
|
||||||
|
5. Template + header nav + frontend API/controls/context-menu/card badges → vitest green
|
||||||
|
6. i18n keys + sync script
|
||||||
|
7. `pytest` + `npm test` full runs; hand to user for manual UI check
|
||||||
|
|
||||||
|
## 9. Phase 2 Detailed Design — CivitAI Downloads for `other`
|
||||||
|
|
||||||
|
Designed 2026-09-12 against the Phase-1 code on this branch; decisions marked **[locked]** follow the same recommendations the feature owner approved for Phase 1.
|
||||||
|
|
||||||
|
### 9.1 Download pipeline touch points
|
||||||
|
|
||||||
|
Flow: `POST /api/lm/download-model` (`py/routes/model_route_registrar.py:104`; GET variant `:105` for the browser extension) → `ModelDownloadHandler.download_model` (`model_handlers.py:1740`) → `DownloadModelUseCase.execute` → `DownloadCoordinator.schedule_download` → `DownloadManager.download_from_civitai` (`download_manager.py:386`) → `_execute_original_download` (`:1415`). Inside, seven scatter points need an `other` branch:
|
||||||
|
|
||||||
|
1. **Type map** (`:1496-1507`): accept `model.type.lower() in VALID_OTHER_CIVITAI_TYPES` → `model_type = "other"` (reuses the Phase-1 set, incl. `"other"` itself).
|
||||||
|
2. **Early version-exists gate** (`:1436-1463`): add `other_scanner.check_model_version_exists`.
|
||||||
|
3. **File-level exists gate** (`:1640-1655` → `_find_local_file_entry` `:320-346` → `_get_scanner_for_model_type` `:230-236`): add explicit `other` branch. **Trap**: the function currently falls through to the lora scanner for unknown types — `"other"` would silently dedupe against loras. Also narrow the fall-through to `"lora"` only / raise on unknown.
|
||||||
|
4. **Version-level fallback gate** (`:1656-1688`): add `elif model_type == "other"`.
|
||||||
|
5. **Default-root selection** (`:1690-1727`): for `other`, first resolve sub_type (§9.2), then read `default_other_roots[sub_type]` (§9.3); if sub_type is undecidable or no default root configured → error guiding the user to pick a folder explicitly.
|
||||||
|
6. **Metadata class selection** (`:1909-1928`) + `_build_metadata_for_resume` (`:969-981`): add `OtherModelMetadata.from_civitai_info` branches.
|
||||||
|
7. **Post-download cache write** (`_execute_download_pipeline` `:2622-2679`): add `other` scanner branch; `adjust_metadata` re-derives sub_type from the on-disk root automatically. `_get_supported_extensions_for_type` (`:2720-2744`): `other` reuses the checkpoint extension set.
|
||||||
|
|
||||||
|
Hooks: `_record_downloaded_version_history` (model_type is free text — zero change); `_sync_downloaded_version` (`:1984` → scanner dispatch `:2130-2135`) add `other`; `py/utils/example_images_download_manager.py` scanner dispatch at `:411-421`, `:591-601`, `:1089+` — add `other` at all three (silent no-scanner otherwise).
|
||||||
|
|
||||||
|
Path templates: `get_download_path_template("other")` is unset, so `other` resolves to a **flat** layout (empty template) — downloads land directly under the resolved sub_type root. This is deliberate: other-model roots are already split per sub_type (`default_other_roots`), and `priority_tags` has no `other` entry, so `{first_tag}` would fall back to an arbitrary CivitAI tag and scatter files into unstable folders. Users who want nesting can still set `download_path_templates["other"]` in `settings.json`. See `DEFAULT_DOWNLOAD_PATH_TEMPLATES` (`py/utils/constants.py`) and `DEFAULT_PATH_TEMPLATES` (`static/js/utils/constants.js`).
|
||||||
|
|
||||||
|
### 9.2 File-level routing (model.type / file.type → sub_type) **[locked]**
|
||||||
|
|
||||||
|
Table-driven, mirroring Phase 1. New in `py/utils/constants.py`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
CIVITAI_FILE_TYPE_TO_OTHER_SUB_TYPE = {
|
||||||
|
"VAE": "vae", "Upscaler": "upscaler", "Text Encoder": "text_encoder",
|
||||||
|
"Vision Encoder": "clip_vision", "CLIPVision": "clip_vision",
|
||||||
|
"ControlNet": "controlnet",
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`download_routing.py` gains `resolve_other_download_sub_type(civitai_model_type, file_types, selected_file_type=None)` with fixed priority:
|
||||||
|
|
||||||
|
1. **Explicit user file pick** (`file_params` from #1058's `_resolve_target_file`) — if the picked file's type maps, it wins even when model.type is `Checkpoint`.
|
||||||
|
2. **model.type** via the existing `CIVITAI_TYPE_TO_OTHER_SUB_TYPE` (`constants.py:120-127`).
|
||||||
|
3. **file.type fallback** — only when model.type maps to nothing (e.g. model.type `Other` or retired `CLIP`). MUST NOT override a mapped model.type: checkpoint models routinely bundle VAE/Text Encoder component files, and unconditional file-type routing would misroute them.
|
||||||
|
4. Still undecidable → `None`; `use_default_paths` errors and the UI offers all other roots for manual selection.
|
||||||
|
|
||||||
|
HTTP: extend `DownloadRoutingHandler.get_download_routing` (`download_routing_handlers.py:23`) with an `other` branch returning `{root_kind: "other", sub_type: ...}`; add `GET /api/lm/other/roots_by_subtype` in `OtherRoutes.setup_specific_routes` (data from `config._prepare_other_paths`'s per-key roots, aggregating `text_encoders` + legacy `clip` under `text_encoder`).
|
||||||
|
|
||||||
|
### 9.3 Settings: single dict key `default_other_roots` **[locked]**
|
||||||
|
|
||||||
|
Rejected: four flat keys (`default_vae_root`…) — each flat key costs ~13 touch points in `settings_manager.py` (defaults `:82-85`, `_check_and_auto_set` `:890-895`, `set()` `:1621-1628`, `_update_active_library_entry` `:738-805`, upsert/create signatures `:1953-2132`, `_build_library_payload` `:552-612`, `_sync_active_library_to_root` `:519-547`, three library constructors, frontend `DEFAULT_SETTINGS_BASE`), repeated per future sub_type.
|
||||||
|
|
||||||
|
Chosen: one mapping key `default_other_roots: {sub_type: path}`, copying the `extra_folder_paths` precedent (generic Mapping handling at `:533-535`, `:573-578`, `:763-767`). `_check_and_auto_set` generalizes to per-sub_type candidates (union over that sub_type's folder keys — `text_encoder` → `text_encoders` + `clip`). `set()` validates keys against `VALID_OTHER_SUB_TYPES`.
|
||||||
|
|
||||||
|
Also fix the Phase-1 omission: add `"other_scanner"` to `_notify_library_change` (`:2150-2156`) and `_notify_model_name_display_change` (`:1795-1800`) — otherwise switching libraries leaves the other page stale.
|
||||||
|
|
||||||
|
### 9.4 Settings UI
|
||||||
|
|
||||||
|
- `templates/components/modals/settings/library.html:34-40`: sub_type selectors after the existing four `setting_select`s (Jinja loop; controlnet selector only when `enabled_other_folders` includes it). Dict-subkey save helper `saveOtherRootSetting(subType, value)` alongside the flat `saveSelectSetting`.
|
||||||
|
- `static/js/managers/SettingsManager.js:1547-1697`: `loadOtherRoots()` mirroring `loadUnetRoots()`, fed by `/api/lm/other/roots_by_subtype`; current values from `state.global.settings.default_other_roots`. `state/index.js:24` `DEFAULT_SETTINGS_BASE` += `default_other_roots: {}`.
|
||||||
|
- Optional: one `other` row in the download-path-template block (`library.html:153-211`).
|
||||||
|
- i18n: `settings.folderSettings.*` keys into `locales/en.json` + sync script; other locales keep `[TODO: Translate]`.
|
||||||
|
- Settings GET (`misc_handlers.py:1528-1536`) already returns all non-sensitive keys — new key reaches the frontend for free.
|
||||||
|
|
||||||
|
### 9.5 Frontend download entry
|
||||||
|
|
||||||
|
- `templates/components/controls.html:83`: drop the `page_id != 'other'` exclusion on the download button (keyboard shortcut D self-enables via `PageControls.js:196-198`).
|
||||||
|
- `OtherControls.js:22-55`: add `showDownloadModal: () => downloadManager.showDownloadModal()` (mirror `EmbeddingsControls.js:43-45`).
|
||||||
|
- `DownloadManager.js` `proceedToLocationContent` (`:955-1017`): add `_resolveOtherSubType()` (mirror `_resolveIsDiffusionModel` `:1026`): selected file type → `/api/lm/download/routing` → `otherApiClient.fetchModelRoots(subType)` (new); default-root preselect reads `default_other_roots[subType]` instead of `` `default_${singularType}_root` `` (`:974`). Undecidable → list all other roots (`/api/lm/other/roots`) for manual pick; an explicit save_dir skips backend default-root logic, so the two paths cannot disagree.
|
||||||
|
- `ModelVersionsTab` download buttons are modelType-generic and already work via `getModelApiClient('other')`; context menu has no CivitAI download entry — no change.
|
||||||
|
- Version-list type validation (`get_civitai_versions` → `_validate_civitai_model_type`) already accepts `VALID_OTHER_CIVITAI_TYPES` from Phase 1.
|
||||||
|
|
||||||
|
### 9.6 CivitAI type mapping decisions **[locked]**
|
||||||
|
|
||||||
|
- Download accepts exactly `VALID_OTHER_CIVITAI_TYPES` (`VAE, Upscaler, TextEncoder, CLIP, CLIPVision, Controlnet, Other`) — reuse the Phase-1 tables; do NOT create new ones.
|
||||||
|
- Extend `CIVITAI_USER_MODEL_TYPES` (`constants.py:133-137`) with the 7 aliases, and point them at the other scanner / `"other"` history bucket in `misc_handlers.py` (`type_scanner_map` `:2793-2797`, `downloaded_version_map` `:2821-2827`) — otherwise creator pages silently filter these models while downloads claim support.
|
||||||
|
- Fix (small Phase-1 bug): `OtherModelMetadata.from_civitai_info` (`py/utils/models.py:343`) reads `version_info.get("type")`, but the type lives at `version["model"]["type"]` — the mapping never fires and always degrades to the placeholder. Read `version_info.get("model", {}).get("type")` instead. (`CheckpointMetadata:290` has the same shape; leave it alone here.)
|
||||||
|
|
||||||
|
### 9.7 Tests
|
||||||
|
|
||||||
|
Existing base: `tests/services/test_download_manager_basic.py` (incl. `test_download_rejects_unsupported_model_type` `:1336`), `test_download_manager_error.py`, `test_download_manager_concurrent.py`, `tests/integration/test_download_flow.py`, `tests/services/test_settings_manager.py`; frontend `tests/frontend/managers/downloadManager.routing.test.js`, `settingsManager.library.test.js`.
|
||||||
|
|
||||||
|
Add: (1) `resolve_other_download_sub_type` unit tests — every priority tier, bundled-component anti-misrouting, undecidable → None, civarchive-shaped payload; (2) download_manager — six model.types accepted → other scanner (mock), unknown still rejected, no lora-scanner fall-through, per-sub_type default roots + unconfigured error, resume metadata, extension set; (3) settings_manager — `default_other_roots` defaults/auto-set (incl. text_encoder dual-key union)/library sync/upsert passthrough/illegal sub_type rejection; (4) routes — `/api/lm/download/routing` other branch, `roots_by_subtype` shape; (5) example-images dispatch accepts `other` (3 sites); (6) vitest — `_resolveOtherSubType` + root select + default preselect, `loadOtherRoots`; (7) user-models existsLocally for VAE.
|
||||||
|
|
||||||
|
### 9.8 Phase 2 file list
|
||||||
|
|
||||||
|
Backend: `py/utils/constants.py`, `py/services/download_routing.py`, `py/routes/handlers/download_routing_handlers.py`, `py/services/download_manager.py`, `py/utils/example_images_download_manager.py`, `py/services/settings_manager.py`, `py/utils/models.py`, `py/routes/other_routes.py`, `py/routes/handlers/misc_handlers.py`, `settings.json.example`.
|
||||||
|
Frontend/templates: `templates/components/controls.html`, `static/js/components/controls/OtherControls.js`, `static/js/managers/DownloadManager.js`, `static/js/api/otherApi.js`, `templates/components/modals/settings/library.html`, `static/js/managers/SettingsManager.js`, `static/js/state/index.js`, `locales/en.json` + sync.
|
||||||
|
|
||||||
|
## 10. Risks / Open Questions
|
||||||
|
|
||||||
|
- **Root overlap**: a user may point `text_encoders` at a directory already scanned as checkpoints/unet. Realpath dedup inside one scanner won't catch cross-scanner overlap → the `_prepare_other_paths` overlap warning (§4.3) is the mitigation; duplicate cards across pages are cosmetic, not corrupting (cache keyed by `(model_type, file_path)`).
|
||||||
|
- **Huge text encoders + lazy hash**: CivitAI fetch for a pending-hash model must trigger on-demand hash like checkpoints do — verify that flow (`calculate_hash_for_model`) is reachable from the `other` routes' fetch-metadata handler.
|
||||||
|
- **Retired CivitAI types**: `CLIP`/`CLIPVision` are retired upstream (grandfathered for existing models); metadata fetch must tolerate both retired and current types — `VALID_OTHER_CIVITAI_TYPES` includes them deliberately.
|
||||||
|
- **Standalone users** must add the new `folder_paths` keys to `settings.json` themselves; document in `settings.json.example` and the feature doc.
|
||||||
|
- **Page display name** is i18n-only; if "Other Models" tests poorly, rename `other.title` without code changes.
|
||||||
|
|
||||||
|
### Phase 2 risks
|
||||||
|
|
||||||
|
- **Bundled component files**: checkpoint models routinely ship VAE/Text Encoder component files — file.type routing must stay a fallback (or explicit user pick), never an override (§9.2 priority is load-bearing; test it).
|
||||||
|
- **`_get_scanner_for_model_type` lora fall-through** (`download_manager.py:236`): without an explicit `other` branch, dedupe checks run against the lora scanner — the most insidious trap in Phase 2.
|
||||||
|
- **text_encoder dual folder keys** (`text_encoders` + legacy `clip`): default-root candidates, `roots_by_subtype`, and auto-set must all merge both keys; miss one and the default-root dropdown comes up empty.
|
||||||
|
- **Undecidable sub_type** (model.type `Other` + unknown file types): must error and ask, never silently default to the vae folder.
|
||||||
|
- **Lazy hash after download**: downloads carry CivitAI SHA256 (no recompute needed) — ensure the post-download cache write doesn't leave `hash_status="pending"`, or the next metadata fetch re-hashes a 10 GB file.
|
||||||
|
- **CivArchive source**: same `_execute_original_download` path, same payload shape — cover it once in tests.
|
||||||
|
|
||||||
|
## 11. Phase 3 — Opt-in Management Toggles (implemented)
|
||||||
|
|
||||||
|
Designed 2026-09-13 against the Phase-1/2 code. Other Models is **opt-in**: after
|
||||||
|
Phase 3 the feature ships disabled, so no other-model folder is scanned and the
|
||||||
|
page shows an "enable" empty state until the user turns it on.
|
||||||
|
|
||||||
|
### 11.1 Settings (global, not per-library)
|
||||||
|
|
||||||
|
| key | type | default | meaning |
|
||||||
|
|---|---|---|---|
|
||||||
|
| `enable_other_models` | bool | `false` | master switch |
|
||||||
|
| `enabled_other_sub_types` | list[str] | `["vae","upscaler","text_encoder"]` | allow-list; `clip_vision` and `controlnet` are opt-in (see §2) |
|
||||||
|
|
||||||
|
`enabled_other_folders` (the unreleased, additive, no-UI backend key) was removed
|
||||||
|
and replaced by the sub_type-level allow-list; there is no migration because the
|
||||||
|
feature never shipped. `text_encoder` expands to `text_encoders` + legacy `clip`
|
||||||
|
via `OTHER_SUB_TYPE_FOLDER_KEYS`.
|
||||||
|
|
||||||
|
The default allow-list lives on five surfaces that must stay in sync:
|
||||||
|
`DEFAULT_ENABLED_OTHER_SUB_TYPES` (`py/utils/constants.py`), `DEFAULT_SETTINGS`
|
||||||
|
(`py/services/settings_manager.py`), the two `DEFAULT_SETTINGS_BASE` /
|
||||||
|
`createDefaultSettings` lists (`static/js/state/index.js`), the
|
||||||
|
`updateOtherModelsControls()` fallback (`static/js/managers/SettingsManager.js`)
|
||||||
|
and the server-rendered Jinja fallback
|
||||||
|
(`templates/components/modals/settings/library.html`).
|
||||||
|
|
||||||
|
### 11.1.1 Legacy key handling in `Config._init_other_paths`
|
||||||
|
|
||||||
|
ComfyUI's `folder_paths` rewrites legacy names before every access (`clip` →
|
||||||
|
`text_encoders`, `unet` → `diffusion_models`) and registers both legacy
|
||||||
|
directories under the canonical key, so `get_folder_paths("clip")` returns
|
||||||
|
exactly the same list as `get_folder_paths("text_encoders")`. Querying both keys
|
||||||
|
made the overlap guard fire twice with `please fix your path configuration` for a
|
||||||
|
configuration the user cannot fix. `Config._collapse_legacy_folder_keys()` now
|
||||||
|
drops a key when the host exposes `map_legacy` and resolves it to another queried
|
||||||
|
key, and `_prepare_other_paths()` downgrades a same-`sub_type` duplicate to
|
||||||
|
`debug` (a cross-`sub_type` collision still warns). In standalone mode
|
||||||
|
`MockFolderPaths` has no `map_legacy` and its keys are independent
|
||||||
|
`settings.json` entries, so every key is still queried there.
|
||||||
|
|
||||||
|
`settings.json.example` intentionally stays minimal (only `use_portable_settings`,
|
||||||
|
`civitai_api_key`, and the four core `folder_paths` keys: `loras`, `checkpoints`,
|
||||||
|
`unet`, `embeddings`). Optional keys — including the other-model folder paths and
|
||||||
|
`enable_other_models` — are NOT documented there; they live in `DEFAULT_SETTINGS`
|
||||||
|
and reach the user's `settings.json` on demand. This supersedes the Phase-1/Phase-2
|
||||||
|
notes that proposed adding the other-model folder keys to the example.
|
||||||
|
|
||||||
|
### 11.2 Behaviour matrix
|
||||||
|
|
||||||
|
| state | scan | nav / `/other` | other downloads | `default_other_roots` | Doctor / refresh-all |
|
||||||
|
|---|---|---|---|---|---|
|
||||||
|
| master off | nothing (`other_roots == []`) | nav entry hidden (`nav-item--hidden`); `/other` still renders the disabled empty state + Enable button; one-time dismissible announcement banner on first visit | rejected | preserved, never auto-set | scanner skipped |
|
||||||
|
| sub_type off | that sub_type's folder keys excluded | page keeps working, type disappears from data | auto-routing refused (manual folder still allowed) | preserved, not preselected | normal |
|
||||||
|
| all on (after enabling) | Phase-1/2 behaviour | normal | normal | normal | normal |
|
||||||
|
|
||||||
|
### 11.3 Backend touch points
|
||||||
|
|
||||||
|
- `py/utils/constants.py` — `DEFAULT_ENABLED_OTHER_SUB_TYPES`, `OTHER_SUB_TYPE_FOLDER_KEYS`, `normalize_other_sub_types`.
|
||||||
|
- `py/config.py` — `_get_enabled_other_folder_keys()` is the single scan gate (master switch + allow-list); new `refresh_other_roots()` rebuilds roots + preview roots on toggle.
|
||||||
|
- `py/services/settings_manager.py` — new defaults, `set()` normalization, `is_other_models_enabled()` / `get_enabled_other_sub_types()` / `is_other_sub_type_enabled()`, and `_apply_other_model_settings_change()` which reapplies config and calls `other_scanner.on_library_changed(reconcile=True)`.
|
||||||
|
- `py/services/model_scanner.py` — `_should_keep_cached_entry()` hydration hook (default keep) plus `on_library_changed(reconcile=...)` / `initialize_in_background(reconcile=...)`; the hook filters `raw_data` and the hash/autov3 index rows.
|
||||||
|
- `py/services/other_scanner.py` — drops persisted entries whose folder is no longer a managed root (sub_type is location-derived, so config is the source of truth).
|
||||||
|
- `py/routes/other_routes.py` — `_validate_civitai_model_type` rejects everything while off / mapped-but-disabled sub_types; `_get_page_context_provider()` injects `other_disabled` into the template.
|
||||||
|
- `py/routes/handlers/model_handlers.py` + `base_model_routes.py` — optional `page_context_provider` hook on `ModelPageView`.
|
||||||
|
- `py/routes/handlers/download_routing_handlers.py` — returns `{sub_type: None, disabled: true, reason}` instead of guessing.
|
||||||
|
- `py/services/download_manager.py` — rejects other-type downloads while off; disabled sub_type refuses default-path routing with a "pick a folder" error.
|
||||||
|
- `py/routes/handlers/misc_handlers.py` — Doctor / init-status / refresh-all skip the other scanner while off (`_active_scanner_factories` / `_active_scanner_getters`).
|
||||||
|
- `py/services/pending_delete_service.py` — deliberately untouched: the scanner stays registered so staged deletes still merge.
|
||||||
|
|
||||||
|
### 11.4 Frontend
|
||||||
|
|
||||||
|
Discoverability: the nav entry is hidden while the feature is off, and three
|
||||||
|
lightweight surfaces replace it — a one-time announcement banner, the download
|
||||||
|
toast, and the settings toggle itself.
|
||||||
|
|
||||||
|
- `templates/components/header.html` + `static/css/components/header.css` — `nav-item--hidden` class (server-rendered when off, client-toggled after enabling) and the `fa-shapes` icon.
|
||||||
|
- `templates/other.html` — `other_disabled` branch in `content` + `main_script`; page-scoped CSS for the empty state.
|
||||||
|
- `static/js/other_disabled.js` — boots `appCore` (shared header) and delegates to the shared enable helper.
|
||||||
|
- `static/js/utils/otherModels.js` — shared `enableOtherModels()` (POST settings + reload) and `openOtherModelsSettings()` (settings modal on the Library section); used by the disabled page, the banner and the download modal.
|
||||||
|
- `static/js/managers/BannerService.js` — `other-models-announcement` banner (only when off and not dismissed; `priority: 0`, dismissal persisted via `dismissed_banners`) with Enable / Open Settings actions; `removeOtherModelsAnnouncement()` drops it without persisting a dismissal.
|
||||||
|
- `templates/components/modals/settings/library.html` + `SettingsManager.updateOtherModelsControls()` / `saveEnabledOtherSubTypes()` / `updateOtherModelsNavVisibility()` — master toggle + five sub_type checkboxes; unchecked/disabled sub_types have their default-root select disabled.
|
||||||
|
- `static/js/managers/DownloadManager.js` — a disabled routing answer surfaces a `showActionToast` with an "Enable Other Models" action (opening settings) and falls back to manual selection.
|
||||||
|
- i18n: `settings.folderSettings.*`, `other.disabled.*` and `banners.otherModels.*` keys in `locales/en.json` + `scripts/sync_translation_keys.py` (other locales keep `[TODO: Translate]`).
|
||||||
|
|
||||||
|
### 11.5 Cache consistency
|
||||||
|
|
||||||
|
- Disabling purges rows from the in-memory view at hydration time (the
|
||||||
|
`_should_keep_cached_entry` hook) and from SQLite on the reconcile triggered by
|
||||||
|
the toggle; the `.metadata.json` sidecars survive, so re-enabling rescans
|
||||||
|
without recomputing hashes (critical for multi-GB text encoders).
|
||||||
|
- Enabling triggers a reconcile so newly managed roots are scanned immediately.
|
||||||
|
- Editing `settings.json` while the server is stopped is still covered by the
|
||||||
|
hydration hook, so disabled types never appear after a restart.
|
||||||
|
|
||||||
|
### 11.6 Tests
|
||||||
|
|
||||||
|
Backend: opt-in fixtures added to the other-related suites; new coverage for
|
||||||
|
"default off scans nothing", per-sub_type gating, routing/download rejection,
|
||||||
|
`_should_keep_cached_entry`, settings normalization and `other_disabled` page
|
||||||
|
context. Frontend: `updateOtherModelsControls` / `saveEnabledOtherSubTypes` and
|
||||||
|
the disabled-page enable flow.
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
# CivitAI image imports can end up with 0 LoRAs
|
||||||
|
|
||||||
|
## Symptom
|
||||||
|
|
||||||
|
Importing a CivitAI image URL can produce a recipe with **zero LoRA
|
||||||
|
entries**, even though the image page lists LoRAs in its resource panel.
|
||||||
|
|
||||||
|
Reported example: `https://civitai.red/images/140818889` was imported as a
|
||||||
|
local recipe with 0 LoRAs, while the page shows 3 LoRAs. Some images (e.g.
|
||||||
|
NSFW / higher browsing level) additionally require a login to view, so their
|
||||||
|
data is not publicly reachable at all.
|
||||||
|
|
||||||
|
## Root cause
|
||||||
|
|
||||||
|
URL imports use only two data sources:
|
||||||
|
|
||||||
|
1. **CivitAI REST image API** — `GET /api/v1/images?imageId=<id>&nsfw=X&withMeta=true` → `meta`
|
||||||
|
2. **Embedded image metadata** — EXIF/XMP read from the downloaded bytes
|
||||||
|
|
||||||
|
For the same image both sources can be empty, and the one source that does
|
||||||
|
contain the data is never queried. Verified for image 140818889:
|
||||||
|
|
||||||
|
| Source | What it returned |
|
||||||
|
|---|---|
|
||||||
|
| REST image API | `meta` holds only a prompt; `modelVersionIds: []`; no `resources`/`hashes`; `baseModel: null` |
|
||||||
|
| Downloaded image | PNG with **no EXIF/XMP** (the CDN URL ends in `.jpeg`, the body is PNG) |
|
||||||
|
| Image page HTML | `__NEXT_DATA__` embeds the trpc `image.getGenerationData` result → full `resources` list: 3 LoRAs, each with `modelId`, `modelVersionId`, `modelName`, `modelType`, `versionName`, `baseModel` |
|
||||||
|
|
||||||
|
Key points:
|
||||||
|
|
||||||
|
- The page's resource panel is fed by an **internal, non-public trpc
|
||||||
|
endpoint**, not by the public REST image API.
|
||||||
|
- That internal endpoint is **login-gated** for some content — the
|
||||||
|
"requires login" symptom.
|
||||||
|
- Even with the version IDs in hand, `/model-versions/{id}` for these
|
||||||
|
(Krea) versions returns **no `sha256`**, so an exact local-file hash match
|
||||||
|
is impossible; only model/version identity is recoverable.
|
||||||
|
|
||||||
|
## Conclusion / status
|
||||||
|
|
||||||
|
0-LoRA imports are a data-source gap: public REST meta and image EXIF are
|
||||||
|
both empty, while the only complete source (page generation data) is
|
||||||
|
internal, sometimes login-gated, and not used by the importer.
|
||||||
|
|
||||||
|
Such imports **cannot be reliably auto-repaired/completed** by the backend
|
||||||
|
alone. The old "Repair Metadata" feature only re-fetched the same incomplete
|
||||||
|
REST meta and could not fix them; it was deprecated and has been removed.
|
||||||
|
|
||||||
|
**Fixed via the companion browser extension.** When the extension is
|
||||||
|
installed with a valid license, it scrapes the image page's internal trpc
|
||||||
|
generation data with the user's session and calls the payload-capable
|
||||||
|
re-import endpoint (`POST /api/lm/recipe/{recipe_id}/reimport` with
|
||||||
|
`image_url`/`name`/`resources`/`gen_params`/`base_model`/`tags` query
|
||||||
|
params), which rebuilds the recipe from the caller-supplied metadata. The
|
||||||
|
web UI delegates re-import of CivitAI-image-sourced recipes to the extension
|
||||||
|
automatically (probe + `lm:reimport*` DOM events); without the extension,
|
||||||
|
re-import silently falls back to the native path, which remains limited by
|
||||||
|
the data-source gap documented above.
|
||||||
@@ -0,0 +1,92 @@
|
|||||||
|
# Reconcile 的 Windows 大小写回退分支 - 待验证清单
|
||||||
|
|
||||||
|
> **状态**: 待 Windows 环境验证 | **创建日期**: 2026-09-11
|
||||||
|
> **相关文件**: `py/services/model_scanner.py` (`ModelScanner._reconcile_cache`)
|
||||||
|
> **相关历史**: #871 (`76ee59cd`, 路径重叠去重)、#1108 (按文件夹扫描的需求)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 背景
|
||||||
|
|
||||||
|
Refresh 按钮走的是 `_reconcile_cache()`(快速增量对账)。2026-09-11 做了一轮性能优化,把两处"预防性"的
|
||||||
|
realpath 全量遍历改成按需触发(详见下方"已完成")。优化后,一次零变更 Refresh 在 5 万文件库上从
|
||||||
|
~1400 ms 降到 ~120 ms。
|
||||||
|
|
||||||
|
清理过程中发现**唯一一处遗留的可疑点**:Windows 专属的大小写不敏感回退分支。它无法在 Linux 上验证,
|
||||||
|
因此单独记录,留待 Windows 机器上确认。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 待验证分支(现状)
|
||||||
|
|
||||||
|
`py/services/model_scanner.py` 中 `_reconcile_cache()` 的 walk 循环内:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# Try case-insensitive match on Windows
|
||||||
|
if os.name == 'nt':
|
||||||
|
lower_path = file_path.lower()
|
||||||
|
matched = False
|
||||||
|
for cached_path in cached_paths: # 每个未命中文件都全量扫一遍缓存
|
||||||
|
if cached_path.lower() == lower_path:
|
||||||
|
found_paths.add(cached_path)
|
||||||
|
matched = True
|
||||||
|
break
|
||||||
|
if matched:
|
||||||
|
continue
|
||||||
|
```
|
||||||
|
|
||||||
|
它排在精确匹配(`file_path in cached_paths`)和 realpath 别名匹配之后,只有**未命中**的文件才会走到。
|
||||||
|
|
||||||
|
### 为什么可疑
|
||||||
|
|
||||||
|
1. **可能不可达**:Windows 上 `os.path.realpath()` 会返回磁盘上的真实大小写,因此"缓存路径大小写与磁盘
|
||||||
|
不一致"的情形,理论上已经被上一步的 realpath 别名匹配覆盖。若如此,这段就是纯冗余代码。
|
||||||
|
2. **一旦可达就是 O(N×M)**:每个未命中文件都要遍历全部 `cached_paths` 做小写比较。若某种路径写法让
|
||||||
|
整个库都变成"未命中"(例如缓存里的盘符/大小写形式与 walk 结果系统性不一致),一次 Refresh 会退化
|
||||||
|
成 文件数 × 缓存条目数 次字符串比较,比真实 IO 还贵。
|
||||||
|
3. **没有测试覆盖**:`tests/services/test_model_scanner.py` 没有任何针对该分支的用例(它在 Linux 上
|
||||||
|
被 `os.name == 'nt'` 短路,无法覆盖)。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 待办
|
||||||
|
|
||||||
|
- [ ] **验证可达性**:在 Windows 上构造"缓存路径与磁盘真实大小写不一致"的场景,确认 realpath 别名匹配
|
||||||
|
是否已经命中,即上面的 `if os.name == 'nt'` 分支是否还有进入的必要。
|
||||||
|
- [ ] **若不可达 / 冗余**:删除该分支,并在删除处留注释说明 realpath 已覆盖大小写归一(附验证记录)。
|
||||||
|
- [ ] **若可达**:保留语义但改成 O(1)——预先构建一次 `lower_path -> cached_path` 映射(与
|
||||||
|
`cached_real_paths` 同样按需、懒构建),把内层全量扫描换成一次字典查询。
|
||||||
|
- [ ] **补一个 Windows-only 的回归测试**(`pytest.mark.skipif(os.name != "nt", ...)`),锁定最终结论。
|
||||||
|
- [ ] 把验证结论回填到本文件,并同步更新状态行。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 验证方法(Windows)
|
||||||
|
|
||||||
|
1. **构造不一致的大小写**:让缓存里的 `file_path` 与磁盘实际路径大小写不同(例如改过盘符/目录大小写,
|
||||||
|
或从另一台机器迁移了 `settings.json` 与持久化缓存),然后在 UI 点 Refresh。
|
||||||
|
2. **看后端日志判据**:
|
||||||
|
- 若 realpath 已覆盖 → 日志应显示 `Cache reconciliation completed in X seconds. Added 0, removed 0 models.`,
|
||||||
|
且**没有** `Found N new files to process` / `Processing <path>`。
|
||||||
|
- 若回退分支在起作用 → 同样应该是 `Added 0, removed 0`(因为 `found_paths` 被补上),这是"分支可达"
|
||||||
|
的证据;反之若出现大量 `Processing ...` 并重新 hash,说明连回退分支也没命中,问题更严重
|
||||||
|
(缓存路径被当成了新文件 + 旧条目被删)。
|
||||||
|
3. **跑测试**:`python -m pytest tests/services/test_model_scanner.py -k reconcile`(该文件在 Windows 上会
|
||||||
|
真实执行 `os.name == 'nt'` 分支)。
|
||||||
|
4. **量化**:如果需要,可在 `_reconcile_cache` 里临时插桩统计该分支的进入次数与内层迭代次数,确认是否为 0。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 已完成(本轮优化,供对照)
|
||||||
|
|
||||||
|
同一次清理里已经落地并验证的部分(Linux,5 万文件库):
|
||||||
|
|
||||||
|
- `cached_real_paths` 别名映射改为**首次未命中时**懒构建(原来每次 Refresh 都对全部缓存条目算一次 realpath)。
|
||||||
|
- 每个文件的 `realpath` 移到精确命中检查**之后**(原来对每个文件都算,命中即丢弃)。
|
||||||
|
- `get_model_roots()` 在新增文件处理阶段只快照一次(原来每个新文件重读一次)。
|
||||||
|
- 全量去重 pass 加了 O(1) 前置判断(`cached_size_before != len(cached_paths) or total_added > 0`),
|
||||||
|
零变更且缓存干净时跳过;快照本身含重复路径时仍会自愈。
|
||||||
|
|
||||||
|
结果:零变更 Refresh 5 万文件 **~1400 ms → ~120 ms**;根目录顺序/符号链接别名翻转场景仍是
|
||||||
|
`re-processed=0`(不重新读 metadata、不重新 hash)。测试:`tests/services/test_model_scanner.py`
|
||||||
|
47 项、全量后端 2567 项全部通过。
|
||||||
+638
-241
File diff suppressed because it is too large
Load Diff
+509
-112
File diff suppressed because it is too large
Load Diff
+646
-249
File diff suppressed because it is too large
Load Diff
+656
-259
File diff suppressed because it is too large
Load Diff
+671
-274
File diff suppressed because it is too large
Load Diff
+604
-207
File diff suppressed because it is too large
Load Diff
+609
-212
File diff suppressed because it is too large
Load Diff
+626
-229
File diff suppressed because it is too large
Load Diff
+541
-144
File diff suppressed because it is too large
Load Diff
+555
-158
File diff suppressed because it is too large
Load Diff
+276
-1
@@ -17,6 +17,9 @@ import types as _types
|
|||||||
import time
|
import time
|
||||||
|
|
||||||
from .utils.cache_paths import CacheType, get_cache_file_path, get_legacy_cache_paths
|
from .utils.cache_paths import CacheType, get_cache_file_path, get_legacy_cache_paths
|
||||||
|
from .utils.constants import (
|
||||||
|
OTHER_MODEL_FOLDER_SUBTYPES,
|
||||||
|
)
|
||||||
from .utils.settings_paths import (
|
from .utils.settings_paths import (
|
||||||
ensure_settings_file,
|
ensure_settings_file,
|
||||||
get_settings_dir,
|
get_settings_dir,
|
||||||
@@ -172,6 +175,13 @@ class Config:
|
|||||||
self.embeddings_roots = None
|
self.embeddings_roots = None
|
||||||
self.base_models_roots = self._init_checkpoint_paths()
|
self.base_models_roots = self._init_checkpoint_paths()
|
||||||
self.embeddings_roots = self._init_embedding_paths()
|
self.embeddings_roots = self._init_embedding_paths()
|
||||||
|
# Other-model roots (VAE, upscalers, text encoders, ...): flat deduped
|
||||||
|
# list plus a normalized root -> sub_type map and per-folder_paths-key
|
||||||
|
# roots for settings persistence.
|
||||||
|
self.other_roots: Optional[List[str]] = None
|
||||||
|
self.other_root_subtypes: Dict[str, str] = {}
|
||||||
|
self.other_folder_roots: Dict[str, List[str]] = {}
|
||||||
|
self.other_roots = self._init_other_paths()
|
||||||
# Extra paths (only for LoRA Manager, not shared with ComfyUI)
|
# Extra paths (only for LoRA Manager, not shared with ComfyUI)
|
||||||
self.extra_loras_roots: List[str] = []
|
self.extra_loras_roots: List[str] = []
|
||||||
self.extra_checkpoints_roots: List[str] = []
|
self.extra_checkpoints_roots: List[str] = []
|
||||||
@@ -336,6 +346,10 @@ class Config:
|
|||||||
"unet": list(self.unet_roots or []),
|
"unet": list(self.unet_roots or []),
|
||||||
"embeddings": list(self.embeddings_roots or []),
|
"embeddings": list(self.embeddings_roots or []),
|
||||||
}
|
}
|
||||||
|
# Persist the other-model roots under their original folder_paths
|
||||||
|
# keys so library switching round-trips them.
|
||||||
|
for key, roots in (self.other_folder_roots or {}).items():
|
||||||
|
target_folder_paths[key] = list(roots)
|
||||||
|
|
||||||
normalized_target_paths = _normalize_folder_paths_for_comparison(
|
normalized_target_paths = _normalize_folder_paths_for_comparison(
|
||||||
target_folder_paths
|
target_folder_paths
|
||||||
@@ -522,6 +536,7 @@ class Config:
|
|||||||
roots.extend(self.loras_roots or [])
|
roots.extend(self.loras_roots or [])
|
||||||
roots.extend(self.base_models_roots or [])
|
roots.extend(self.base_models_roots or [])
|
||||||
roots.extend(self.embeddings_roots or [])
|
roots.extend(self.embeddings_roots or [])
|
||||||
|
roots.extend(self.other_roots or [])
|
||||||
# Include extra paths for scanning symlinks
|
# Include extra paths for scanning symlinks
|
||||||
roots.extend(self.extra_loras_roots or [])
|
roots.extend(self.extra_loras_roots or [])
|
||||||
roots.extend(self.extra_checkpoints_roots or [])
|
roots.extend(self.extra_checkpoints_roots or [])
|
||||||
@@ -862,6 +877,8 @@ class Config:
|
|||||||
preview_roots.update(self._expand_preview_root(root))
|
preview_roots.update(self._expand_preview_root(root))
|
||||||
for root in self.embeddings_roots or []:
|
for root in self.embeddings_roots or []:
|
||||||
preview_roots.update(self._expand_preview_root(root))
|
preview_roots.update(self._expand_preview_root(root))
|
||||||
|
for root in self.other_roots or []:
|
||||||
|
preview_roots.update(self._expand_preview_root(root))
|
||||||
# Include extra paths for preview access
|
# Include extra paths for preview access
|
||||||
for root in self.extra_loras_roots or []:
|
for root in self.extra_loras_roots or []:
|
||||||
preview_roots.update(self._expand_preview_root(root))
|
preview_roots.update(self._expand_preview_root(root))
|
||||||
@@ -882,7 +899,7 @@ class Config:
|
|||||||
path for path in preview_roots if path.is_absolute()
|
path for path in preview_roots if path.is_absolute()
|
||||||
}
|
}
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Preview roots rebuilt: %d paths from %d lora roots (%d extra), %d checkpoint roots (%d extra), %d embedding roots (%d extra), %d symlink mappings",
|
"Preview roots rebuilt: %d paths from %d lora roots (%d extra), %d checkpoint roots (%d extra), %d embedding roots (%d extra), %d other roots, %d symlink mappings",
|
||||||
len(self._preview_root_paths),
|
len(self._preview_root_paths),
|
||||||
len(self.loras_roots or []),
|
len(self.loras_roots or []),
|
||||||
len(self.extra_loras_roots or []),
|
len(self.extra_loras_roots or []),
|
||||||
@@ -890,6 +907,7 @@ class Config:
|
|||||||
len(self.extra_checkpoints_roots or []),
|
len(self.extra_checkpoints_roots or []),
|
||||||
len(self.embeddings_roots or []),
|
len(self.embeddings_roots or []),
|
||||||
len(self.extra_embeddings_roots or []),
|
len(self.extra_embeddings_roots or []),
|
||||||
|
len(self.other_roots or []),
|
||||||
len(self._path_mappings),
|
len(self._path_mappings),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1128,6 +1146,155 @@ class Config:
|
|||||||
|
|
||||||
return unique_paths
|
return unique_paths
|
||||||
|
|
||||||
|
def _get_enabled_other_folder_keys(self) -> List[str]:
|
||||||
|
"""Return the OTHER_MODEL_FOLDER_SUBTYPES keys that are enabled.
|
||||||
|
|
||||||
|
Other Models management is opt-in: while ``enable_other_models`` is
|
||||||
|
off (the default) no other-model folder is scanned at all. When it is
|
||||||
|
on, only the folder keys of the enabled sub_types are scanned
|
||||||
|
(text_encoder merges ``text_encoders`` with the legacy ``clip`` key).
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
from .services.settings_manager import get_settings_manager
|
||||||
|
|
||||||
|
enabled_sub_types = get_settings_manager().get_enabled_other_sub_types()
|
||||||
|
except Exception:
|
||||||
|
enabled_sub_types = []
|
||||||
|
if not enabled_sub_types:
|
||||||
|
return []
|
||||||
|
allowed = set(enabled_sub_types)
|
||||||
|
return [
|
||||||
|
key
|
||||||
|
for key, sub_type in OTHER_MODEL_FOLDER_SUBTYPES.items()
|
||||||
|
if sub_type in allowed
|
||||||
|
]
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _collapse_legacy_folder_keys(keys: List[str]) -> List[str]:
|
||||||
|
"""Drop folder keys the host already normalizes onto another queried key.
|
||||||
|
|
||||||
|
ComfyUI's ``folder_paths`` rewrites legacy names before every access
|
||||||
|
(``clip`` -> ``text_encoders``, ``unet`` -> ``diffusion_models``), and
|
||||||
|
registers both legacy directories under the canonical key, so
|
||||||
|
``get_folder_paths("clip")`` returns exactly the same list as
|
||||||
|
``get_folder_paths("text_encoders")``. Querying both therefore reports
|
||||||
|
every text-encoder folder twice and trips the overlap guard with a
|
||||||
|
conflict the user cannot fix.
|
||||||
|
|
||||||
|
When the host exposes ``map_legacy`` the alias is provably redundant and
|
||||||
|
is skipped (an empty canonical list implies an empty alias list).
|
||||||
|
Without it - the standalone mock, whose keys are independent
|
||||||
|
``settings.json`` entries - every key is kept, because a ``clip``-only
|
||||||
|
configuration is then genuinely distinct.
|
||||||
|
"""
|
||||||
|
map_legacy = getattr(folder_paths, "map_legacy", None)
|
||||||
|
if not callable(map_legacy):
|
||||||
|
return list(keys)
|
||||||
|
|
||||||
|
queried = set(keys)
|
||||||
|
collapsed: List[str] = []
|
||||||
|
for key in keys:
|
||||||
|
try:
|
||||||
|
canonical = map_legacy(key)
|
||||||
|
except Exception:
|
||||||
|
canonical = key
|
||||||
|
if canonical != key and canonical in queried:
|
||||||
|
logger.debug(
|
||||||
|
"Skipping legacy folder key '%s'; the host resolves it to "
|
||||||
|
"'%s', which is queried as well.",
|
||||||
|
key,
|
||||||
|
canonical,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
collapsed.append(key)
|
||||||
|
return collapsed
|
||||||
|
|
||||||
|
def _prepare_other_paths(
|
||||||
|
self, folder_path_map: Mapping[str, Iterable[str]]
|
||||||
|
) -> Tuple[List[str], Dict[str, str], Dict[str, List[str]]]:
|
||||||
|
"""Prepare other-model paths from a folder_paths-key -> raw paths map.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Tuple of (all_unique_roots, business_root -> sub_type map,
|
||||||
|
folder_paths key -> business roots). This method does NOT modify
|
||||||
|
instance variables - callers must set them.
|
||||||
|
"""
|
||||||
|
unique_paths: List[str] = []
|
||||||
|
sub_type_map: Dict[str, str] = {}
|
||||||
|
per_key_roots: Dict[str, List[str]] = {}
|
||||||
|
# real path -> (business path, sub_type) of the category that claimed it
|
||||||
|
seen_real_paths: Dict[str, Tuple[str, str]] = {}
|
||||||
|
|
||||||
|
# Cross-scanner overlap detection: warn when an "other" root is
|
||||||
|
# already covered by the checkpoints/unet or embeddings scanners.
|
||||||
|
# Kept (not dropped) on purpose - duplicate cards across pages are
|
||||||
|
# cosmetic, while dropping would silently unmanage the files.
|
||||||
|
covered_real_paths = {
|
||||||
|
os.path.normpath(os.path.realpath(path)).replace(os.sep, "/"): path
|
||||||
|
for path in [
|
||||||
|
*(self.base_models_roots or []),
|
||||||
|
*(self.embeddings_roots or []),
|
||||||
|
]
|
||||||
|
if isinstance(path, str) and path.strip() and os.path.exists(path)
|
||||||
|
}
|
||||||
|
|
||||||
|
for key, sub_type in OTHER_MODEL_FOLDER_SUBTYPES.items():
|
||||||
|
raw_paths = folder_path_map.get(key)
|
||||||
|
if not raw_paths:
|
||||||
|
continue
|
||||||
|
path_map = self._dedupe_existing_paths(raw_paths)
|
||||||
|
key_roots: List[str] = []
|
||||||
|
for real_path, business_path in sorted(
|
||||||
|
path_map.items(), key=lambda item: item[1].lower()
|
||||||
|
):
|
||||||
|
seen = seen_real_paths.get(real_path)
|
||||||
|
if seen is not None:
|
||||||
|
seen_business_path, seen_sub_type = seen
|
||||||
|
if seen_sub_type == sub_type:
|
||||||
|
# Same category reached through a second folder_paths
|
||||||
|
# key (legacy alias, or a sub_type spanning two keys).
|
||||||
|
# Expected, so never a "fix your configuration" warning.
|
||||||
|
logger.debug(
|
||||||
|
"Ignoring duplicate folder '%s' for category '%s' "
|
||||||
|
"(already covered by '%s').",
|
||||||
|
business_path,
|
||||||
|
sub_type,
|
||||||
|
seen_business_path,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Detected the same folder '%s' under multiple other-model "
|
||||||
|
"categories ('%s' is already mapped as '%s'). Keeping the "
|
||||||
|
"first category; please fix your path configuration.",
|
||||||
|
business_path,
|
||||||
|
seen_business_path,
|
||||||
|
seen_sub_type,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
seen_real_paths[real_path] = (business_path, sub_type)
|
||||||
|
unique_paths.append(business_path)
|
||||||
|
key_roots.append(business_path)
|
||||||
|
sub_type_map[business_path] = sub_type
|
||||||
|
|
||||||
|
if real_path != business_path:
|
||||||
|
self.add_path_mapping(business_path, real_path)
|
||||||
|
|
||||||
|
covered_by = covered_real_paths.get(real_path)
|
||||||
|
if covered_by:
|
||||||
|
logger.warning(
|
||||||
|
"Detected an other-model root ('%s', category '%s') that "
|
||||||
|
"overlaps an existing checkpoints/embeddings root ('%s'). "
|
||||||
|
"The same files will appear on both pages; please review "
|
||||||
|
"your path configuration.",
|
||||||
|
business_path,
|
||||||
|
key,
|
||||||
|
covered_by,
|
||||||
|
)
|
||||||
|
if key_roots:
|
||||||
|
per_key_roots[key] = key_roots
|
||||||
|
|
||||||
|
return unique_paths, sub_type_map, per_key_roots
|
||||||
|
|
||||||
def _apply_library_paths(
|
def _apply_library_paths(
|
||||||
self,
|
self,
|
||||||
folder_paths: Mapping[str, Any],
|
folder_paths: Mapping[str, Any],
|
||||||
@@ -1151,6 +1318,16 @@ class Config:
|
|||||||
) = self._prepare_checkpoint_paths(checkpoint_paths, unet_paths)
|
) = self._prepare_checkpoint_paths(checkpoint_paths, unet_paths)
|
||||||
self.embeddings_roots = self._prepare_embedding_paths(embedding_paths)
|
self.embeddings_roots = self._prepare_embedding_paths(embedding_paths)
|
||||||
|
|
||||||
|
other_path_map = {
|
||||||
|
key: folder_paths.get(key, []) or []
|
||||||
|
for key in self._get_enabled_other_folder_keys()
|
||||||
|
}
|
||||||
|
(
|
||||||
|
self.other_roots,
|
||||||
|
self.other_root_subtypes,
|
||||||
|
self.other_folder_roots,
|
||||||
|
) = self._prepare_other_paths(other_path_map)
|
||||||
|
|
||||||
# Process extra paths (only for LoRA Manager, not shared with ComfyUI)
|
# Process extra paths (only for LoRA Manager, not shared with ComfyUI)
|
||||||
extra_paths = extra_folder_paths or {}
|
extra_paths = extra_folder_paths or {}
|
||||||
extra_lora_paths = extra_paths.get("loras", []) or []
|
extra_lora_paths = extra_paths.get("loras", []) or []
|
||||||
@@ -1267,6 +1444,104 @@ class Config:
|
|||||||
logger.warning(f"Error initializing embedding paths: {e}")
|
logger.warning(f"Error initializing embedding paths: {e}")
|
||||||
return []
|
return []
|
||||||
|
|
||||||
|
def _init_other_paths(self) -> List[str]:
|
||||||
|
"""Initialize and validate other-model paths from ComfyUI settings.
|
||||||
|
|
||||||
|
Iterates the enabled OTHER_MODEL_FOLDER_SUBTYPES keys and pulls each
|
||||||
|
from ``folder_paths.get_folder_paths(key)`` (in standalone mode the
|
||||||
|
mock serves arbitrary keys from ``settings.json.folder_paths``).
|
||||||
|
Legacy aliases the host normalizes onto a canonical key (``clip`` ->
|
||||||
|
``text_encoders``) are collapsed first so the same folders are not
|
||||||
|
reported twice.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
folder_path_map: Dict[str, List[str]] = {}
|
||||||
|
for key in self._collapse_legacy_folder_keys(
|
||||||
|
self._get_enabled_other_folder_keys()
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
folder_path_map[key] = folder_paths.get_folder_paths(key)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("Error reading folder paths for '%s': %s", key, exc)
|
||||||
|
|
||||||
|
(
|
||||||
|
unique_paths,
|
||||||
|
self.other_root_subtypes,
|
||||||
|
self.other_folder_roots,
|
||||||
|
) = self._prepare_other_paths(folder_path_map)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Found other model roots:"
|
||||||
|
+ ("\n - " + "\n - ".join(unique_paths) if unique_paths else "[]")
|
||||||
|
)
|
||||||
|
|
||||||
|
if not unique_paths:
|
||||||
|
logger.info("No valid other-model folders found in configuration")
|
||||||
|
return []
|
||||||
|
|
||||||
|
return unique_paths
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"Error initializing other model paths: {e}")
|
||||||
|
return []
|
||||||
|
|
||||||
|
def refresh_other_roots(self) -> None:
|
||||||
|
"""Rebuild other-model roots after the management toggles changed.
|
||||||
|
|
||||||
|
Called when ``enable_other_models`` / ``enabled_other_sub_types`` are
|
||||||
|
updated so the scanner immediately reflects the new folder set without
|
||||||
|
a full application restart.
|
||||||
|
"""
|
||||||
|
self.other_roots = self._init_other_paths()
|
||||||
|
self._rebuild_preview_roots()
|
||||||
|
|
||||||
|
def get_other_models_availability(self) -> Dict[str, Any]:
|
||||||
|
"""Report the other-model folders the host can actually expose.
|
||||||
|
|
||||||
|
Independent of the opt-in ``enable_other_models`` toggle: this answers
|
||||||
|
"could Other Models management work here at all?". ComfyUI mode almost
|
||||||
|
always has these folder keys registered, while standalone mode only
|
||||||
|
knows the keys present in ``settings.json.folder_paths`` - so the UI
|
||||||
|
uses this to decide whether announcing the feature would be actionable.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
``{"available": bool, "sub_types": {sub_type: [existing roots]}}``.
|
||||||
|
A folder only counts when it exists on disk; an empty folder still
|
||||||
|
counts because CivitAI downloads can target it.
|
||||||
|
"""
|
||||||
|
sub_types: Dict[str, List[str]] = {}
|
||||||
|
try:
|
||||||
|
keys = self._collapse_legacy_folder_keys(
|
||||||
|
list(OTHER_MODEL_FOLDER_SUBTYPES.keys())
|
||||||
|
)
|
||||||
|
except Exception: # pragma: no cover - defensive
|
||||||
|
keys = list(OTHER_MODEL_FOLDER_SUBTYPES.keys())
|
||||||
|
|
||||||
|
for key in keys:
|
||||||
|
sub_type = OTHER_MODEL_FOLDER_SUBTYPES.get(key)
|
||||||
|
if not sub_type:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
raw_paths = folder_paths.get_folder_paths(key)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug("Error probing folder paths for '%s': %s", key, exc)
|
||||||
|
continue
|
||||||
|
|
||||||
|
bucket = sub_types.setdefault(sub_type, [])
|
||||||
|
for root in sorted(
|
||||||
|
self._dedupe_existing_paths(raw_paths or []).values(),
|
||||||
|
key=lambda path: path.lower(),
|
||||||
|
):
|
||||||
|
if root not in bucket:
|
||||||
|
bucket.append(root)
|
||||||
|
|
||||||
|
available_sub_types = {
|
||||||
|
sub_type: roots for sub_type, roots in sub_types.items() if roots
|
||||||
|
}
|
||||||
|
return {
|
||||||
|
"available": bool(available_sub_types),
|
||||||
|
"sub_types": available_sub_types,
|
||||||
|
}
|
||||||
|
|
||||||
def get_preview_static_url(self, preview_path: str) -> str:
|
def get_preview_static_url(self, preview_path: str) -> str:
|
||||||
if not preview_path:
|
if not preview_path:
|
||||||
return ""
|
return ""
|
||||||
|
|||||||
+7
-8
@@ -219,6 +219,7 @@ class LoraManager:
|
|||||||
lora_scanner = await ServiceRegistry.get_lora_scanner()
|
lora_scanner = await ServiceRegistry.get_lora_scanner()
|
||||||
checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner()
|
checkpoint_scanner = await ServiceRegistry.get_checkpoint_scanner()
|
||||||
embedding_scanner = await ServiceRegistry.get_embedding_scanner()
|
embedding_scanner = await ServiceRegistry.get_embedding_scanner()
|
||||||
|
other_scanner = await ServiceRegistry.get_other_scanner()
|
||||||
|
|
||||||
# Initialize recipe scanner if needed
|
# Initialize recipe scanner if needed
|
||||||
recipe_scanner = await ServiceRegistry.get_recipe_scanner()
|
recipe_scanner = await ServiceRegistry.get_recipe_scanner()
|
||||||
@@ -236,6 +237,10 @@ class LoraManager:
|
|||||||
embedding_scanner.initialize_in_background(),
|
embedding_scanner.initialize_in_background(),
|
||||||
name="embedding_cache_init",
|
name="embedding_cache_init",
|
||||||
),
|
),
|
||||||
|
asyncio.create_task(
|
||||||
|
other_scanner.initialize_in_background(),
|
||||||
|
name="other_cache_init",
|
||||||
|
),
|
||||||
asyncio.create_task(
|
asyncio.create_task(
|
||||||
recipe_scanner.initialize_in_background(), name="recipe_cache_init"
|
recipe_scanner.initialize_in_background(), name="recipe_cache_init"
|
||||||
),
|
),
|
||||||
@@ -328,6 +333,7 @@ class LoraManager:
|
|||||||
all_roots.update(config.loras_roots)
|
all_roots.update(config.loras_roots)
|
||||||
all_roots.update(config.base_models_roots or [])
|
all_roots.update(config.base_models_roots or [])
|
||||||
all_roots.update(config.embeddings_roots or [])
|
all_roots.update(config.embeddings_roots or [])
|
||||||
|
all_roots.update(config.other_roots or [])
|
||||||
|
|
||||||
total_deleted = 0
|
total_deleted = 0
|
||||||
total_size_freed = 0
|
total_size_freed = 0
|
||||||
@@ -460,18 +466,11 @@ class LoraManager:
|
|||||||
# Cancel any in-flight scanner initialization tasks so thread-pool
|
# Cancel any in-flight scanner initialization tasks so thread-pool
|
||||||
# workers (e.g. _initialize_cache_sync) can break out of their loops
|
# workers (e.g. _initialize_cache_sync) can break out of their loops
|
||||||
# when the server shuts down (e.g. Ctrl+C on WSL).
|
# when the server shuts down (e.g. Ctrl+C on WSL).
|
||||||
for name in ("lora_scanner", "checkpoint_scanner", "embedding_scanner"):
|
for name in ("lora_scanner", "checkpoint_scanner", "embedding_scanner", "other_scanner"):
|
||||||
scanner = ServiceRegistry.get_service_sync(name)
|
scanner = ServiceRegistry.get_service_sync(name)
|
||||||
if scanner is not None and hasattr(scanner, "cancel_task"):
|
if scanner is not None and hasattr(scanner, "cancel_task"):
|
||||||
scanner.cancel_task()
|
scanner.cancel_task()
|
||||||
logger.debug("LoRA Manager: Cancelled %s", name)
|
logger.debug("LoRA Manager: Cancelled %s", name)
|
||||||
|
|
||||||
# Close shared aiohttp sessions to avoid "Unclosed client session" warnings
|
|
||||||
try:
|
|
||||||
from py.routes.handlers.hf_handlers import close_hf_api_session
|
|
||||||
await close_hf_api_session()
|
|
||||||
except Exception as exc:
|
|
||||||
logger.debug("Error closing HF API session: %s", exc)
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error during cleanup: {e}", exc_info=True)
|
logger.error(f"Error during cleanup: {e}", exc_info=True)
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ SCANNER_TYPE_MAP: dict[str, str] = {
|
|||||||
"get_lora_scanner": "lora",
|
"get_lora_scanner": "lora",
|
||||||
"get_checkpoint_scanner": "checkpoint",
|
"get_checkpoint_scanner": "checkpoint",
|
||||||
"get_embedding_scanner": "embedding",
|
"get_embedding_scanner": "embedding",
|
||||||
|
"get_other_scanner": "other",
|
||||||
}
|
}
|
||||||
|
|
||||||
SCANNER_GETTER_NAMES = tuple(SCANNER_TYPE_MAP.keys())
|
SCANNER_GETTER_NAMES = tuple(SCANNER_TYPE_MAP.keys())
|
||||||
@@ -80,8 +81,8 @@ async def _find_scanner_for_model(
|
|||||||
|
|
||||||
|
|
||||||
async def identify_model_type(model_path: str) -> str:
|
async def identify_model_type(model_path: str) -> str:
|
||||||
"""Determine the model type (``\"lora\"``, ``\"checkpoint\"``, or
|
"""Determine the model type (``\"lora\"``, ``\"checkpoint\"``,
|
||||||
``\"embedding\"``) for *model_path*.
|
``\"embedding\"``, or ``\"other\"``) for *model_path*.
|
||||||
|
|
||||||
Falls back to ``\"lora\"`` when unknown.
|
Falls back to ``\"lora\"`` when unknown.
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -46,6 +46,16 @@ async def api_json_error(
|
|||||||
if request.path.startswith("/api/lm/previews") and exc.status == 404:
|
if request.path.startswith("/api/lm/previews") and exc.status == 404:
|
||||||
logger_method = logger.debug
|
logger_method = logger.debug
|
||||||
|
|
||||||
|
# Download-progress 404 is routine too: in-memory tracking is removed
|
||||||
|
# once a download finishes/fails, so the extension's final polls 404.
|
||||||
|
# The extension relies on the 404 status itself (failure detection),
|
||||||
|
# so only the log level is lowered.
|
||||||
|
if (
|
||||||
|
request.path.startswith("/api/lm/download-progress/")
|
||||||
|
and exc.status == 404
|
||||||
|
):
|
||||||
|
logger_method = logger.debug
|
||||||
|
|
||||||
logger_method(
|
logger_method(
|
||||||
"API %s %s returned HTTP %d: %s",
|
"API %s %s returned HTTP %d: %s",
|
||||||
request.method,
|
request.method,
|
||||||
|
|||||||
@@ -78,7 +78,7 @@ class CheckpointLoaderLM:
|
|||||||
|
|
||||||
# Filter only checkpoint type (not diffusion_model) and format names
|
# Filter only checkpoint type (not diffusion_model) and format names
|
||||||
names = []
|
names = []
|
||||||
for item in cache.raw_data:
|
for item in list(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", "")
|
||||||
# Only offer models that still exist on disk so ComfyUI
|
# Only offer models that still exist on disk so ComfyUI
|
||||||
@@ -126,7 +126,7 @@ class CheckpointLoaderLM:
|
|||||||
cache = await scanner.get_cached_data()
|
cache = await scanner.get_cached_data()
|
||||||
|
|
||||||
base_models = set()
|
base_models = set()
|
||||||
for item in cache.raw_data:
|
for item in list(cache.raw_data):
|
||||||
if item.get("sub_type") != "checkpoint":
|
if item.get("sub_type") != "checkpoint":
|
||||||
continue
|
continue
|
||||||
base_model = item.get("base_model")
|
base_model = item.get("base_model")
|
||||||
|
|||||||
@@ -1,214 +0,0 @@
|
|||||||
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,)
|
|
||||||
@@ -1,326 +0,0 @@
|
|||||||
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)}"
|
|
||||||
)
|
|
||||||
@@ -601,7 +601,7 @@ class SaveImageLM:
|
|||||||
os.path.basename(name),
|
os.path.basename(name),
|
||||||
os.path.splitext(os.path.basename(name))[0],
|
os.path.splitext(os.path.basename(name))[0],
|
||||||
]
|
]
|
||||||
for model in getattr(cache, "raw_data", []):
|
for model in list(getattr(cache, "raw_data", [])):
|
||||||
file_name = model.get("file_name")
|
file_name = model.get("file_name")
|
||||||
if file_name in candidates:
|
if file_name in candidates:
|
||||||
return model
|
return model
|
||||||
|
|||||||
@@ -93,7 +93,7 @@ class UNETLoaderLM:
|
|||||||
|
|
||||||
# Filter only diffusion_model type and format names
|
# Filter only diffusion_model type and format names
|
||||||
names = []
|
names = []
|
||||||
for item in cache.raw_data:
|
for item in list(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", "")
|
||||||
# Only offer models that still exist on disk so ComfyUI
|
# Only offer models that still exist on disk so ComfyUI
|
||||||
@@ -141,7 +141,7 @@ class UNETLoaderLM:
|
|||||||
cache = await scanner.get_cached_data()
|
cache = await scanner.get_cached_data()
|
||||||
|
|
||||||
base_models = set()
|
base_models = set()
|
||||||
for item in cache.raw_data:
|
for item in list(cache.raw_data):
|
||||||
if item.get("sub_type") != "diffusion_model":
|
if item.get("sub_type") != "diffusion_model":
|
||||||
continue
|
continue
|
||||||
base_model = item.get("base_model")
|
base_model = item.get("base_model")
|
||||||
|
|||||||
+1
-1
@@ -156,7 +156,7 @@ def _find_missing_loras(names: list[str]) -> list[str]:
|
|||||||
|
|
||||||
lookup = {}
|
lookup = {}
|
||||||
basename_candidates = {}
|
basename_candidates = {}
|
||||||
for item in cache.raw_data:
|
for item in list(cache.raw_data):
|
||||||
file_path = item.get("file_path")
|
file_path = item.get("file_path")
|
||||||
if not file_path:
|
if not file_path:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from typing import Dict, Any
|
|||||||
from ..base import RecipeMetadataParser
|
from ..base import RecipeMetadataParser
|
||||||
from ..constants import GEN_PARAM_KEYS
|
from ..constants import GEN_PARAM_KEYS
|
||||||
from ...services.metadata_service import get_default_metadata_provider
|
from ...services.metadata_service import get_default_metadata_provider
|
||||||
|
from ...utils.constants import is_empty_placeholder_hash
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -146,15 +147,13 @@ class AutomaticMetadataParser(RecipeMetadataParser):
|
|||||||
# Initialize hashes dict if it doesn't exist
|
# Initialize hashes dict if it doesn't exist
|
||||||
if "hashes" not in metadata:
|
if "hashes" not in metadata:
|
||||||
metadata["hashes"] = {}
|
metadata["hashes"] = {}
|
||||||
# Add as lora type in the same format as
|
# Lora hashes carries the 12-char AutoV3
|
||||||
# regular hashes. Only override an
|
# hash (resolvable on CivitAI and the local
|
||||||
# existing entry if its value is empty
|
# autov3 index); the Hashes JSON value is
|
||||||
# (Lora hashes is the more reliable
|
# only the 10-char AutoV2 prefix, so on
|
||||||
# source when Hashes JSON has blanks).
|
# conflict the Lora hashes value wins.
|
||||||
key = f"lora:{lora_name}"
|
key = f"lora:{lora_name}"
|
||||||
existing = metadata["hashes"].get(key, "")
|
metadata["hashes"][key] = lora_hash
|
||||||
if not existing:
|
|
||||||
metadata["hashes"][key] = lora_hash
|
|
||||||
|
|
||||||
# Remove lora hashes from params section
|
# Remove lora hashes from params section
|
||||||
params_section = params_section.replace(lora_hashes_match.group(0), '')
|
params_section = params_section.replace(lora_hashes_match.group(0), '')
|
||||||
@@ -526,6 +525,26 @@ class AutomaticMetadataParser(RecipeMetadataParser):
|
|||||||
weight = prompt_entries[0][1] if len(prompt_entries) == 1 else 1.0
|
weight = prompt_entries[0][1] if len(prompt_entries) == 1 else 1.0
|
||||||
lora_entry = make_lora_entry(lora_type, lora_name, weight, lora_hash)
|
lora_entry = make_lora_entry(lora_type, lora_name, weight, lora_hash)
|
||||||
|
|
||||||
|
if is_empty_placeholder_hash(lora_hash):
|
||||||
|
# The empty-hash placeholder (SHA256 of an empty byte
|
||||||
|
# string) is not a real hash: never look it up in the
|
||||||
|
# local hash index or on CivitAI. Match by filename;
|
||||||
|
# otherwise keep the item as unresolved (no hash, flagged
|
||||||
|
# hashInvalid so the UI shows the unresolvable-hash state
|
||||||
|
# and offers reconnect instead of download) rather than
|
||||||
|
# dropping it.
|
||||||
|
if recipe_scanner and lora_type == 'lora' and basename_key not in queried_local_basenames:
|
||||||
|
local_lora = await recipe_scanner.get_local_lora(lora_name, recipe_base_model)
|
||||||
|
if local_lora:
|
||||||
|
local_entry = self.populate_lora_from_local(lora_entry, local_lora)
|
||||||
|
merge_or_append_local(local_entry)
|
||||||
|
continue
|
||||||
|
lora_entry['hash'] = ''
|
||||||
|
lora_entry['hashInvalid'] = True
|
||||||
|
if not resource_lora_count:
|
||||||
|
loras.append(lora_entry)
|
||||||
|
continue
|
||||||
|
|
||||||
if lora_hash and recipe_scanner and lora_type == 'lora':
|
if lora_hash and recipe_scanner and lora_type == 'lora':
|
||||||
local_lora = await recipe_scanner.get_local_lora_by_hash(lora_hash)
|
local_lora = await recipe_scanner.get_local_lora_by_hash(lora_hash)
|
||||||
if local_lora:
|
if local_lora:
|
||||||
|
|||||||
@@ -115,6 +115,27 @@ class CivitaiApiMetadataParser(RecipeMetadataParser):
|
|||||||
):
|
):
|
||||||
metadata = inner_meta
|
metadata = inner_meta
|
||||||
|
|
||||||
|
# Civitai's image API meta parser mangles the A1111 "Lora hashes"
|
||||||
|
# text field into a quote-wrapped dict entry:
|
||||||
|
# '"Daphne Blake Cosplay_v1": "e67ebd5e315f"'
|
||||||
|
# The 12-char AutoV3 it carries is more reliable than the stale
|
||||||
|
# 10-char AutoV2 value in the "hashes" dict, so recover it and
|
||||||
|
# let it override the conflicting entry.
|
||||||
|
if isinstance(metadata, dict):
|
||||||
|
for key, hash_value in list(metadata.items()):
|
||||||
|
if (
|
||||||
|
isinstance(key, str)
|
||||||
|
and key.startswith('"')
|
||||||
|
and isinstance(hash_value, str)
|
||||||
|
and hash_value.endswith('"')
|
||||||
|
):
|
||||||
|
clean_name = key.strip('"').strip()
|
||||||
|
clean_hash = hash_value.strip('"').strip()
|
||||||
|
if clean_name and clean_hash:
|
||||||
|
hashes_dict = metadata.get("hashes")
|
||||||
|
if isinstance(hashes_dict, dict):
|
||||||
|
hashes_dict[f"lora:{clean_name}"] = clean_hash
|
||||||
|
|
||||||
# Initialize result structure
|
# Initialize result structure
|
||||||
result: Dict[str, Any] = {
|
result: Dict[str, Any] = {
|
||||||
"base_model": None,
|
"base_model": None,
|
||||||
|
|||||||
+28
-18
@@ -40,24 +40,34 @@ class ComfyMetadataParser(RecipeMetadataParser):
|
|||||||
checkpoint_node = next(iter(checkpoint_nodes.values()))
|
checkpoint_node = next(iter(checkpoint_nodes.values()))
|
||||||
if 'inputs' in checkpoint_node and 'ckpt_name' in checkpoint_node['inputs']:
|
if 'inputs' in checkpoint_node and 'ckpt_name' in checkpoint_node['inputs']:
|
||||||
checkpoint_name = checkpoint_node['inputs']['ckpt_name']
|
checkpoint_name = checkpoint_node['inputs']['ckpt_name']
|
||||||
checkpoint_match = re.search(r'civitai:(\d+)@(\d+)', checkpoint_name)
|
# Some ComfyUI workflows serialize ckpt_name as a
|
||||||
if checkpoint_match:
|
# single-element list (e.g. ["model.safetensors"]) or leave
|
||||||
checkpoint_id = checkpoint_match.group(1)
|
# the value unset (None). Neither is a string, so skip the
|
||||||
checkpoint_version_id = checkpoint_match.group(2)
|
# CivitAI-URN lookup instead of crashing re.search with a
|
||||||
checkpoint = {
|
# TypeError that fails the whole image import.
|
||||||
'id': checkpoint_version_id,
|
if isinstance(checkpoint_name, list):
|
||||||
'modelId': checkpoint_id,
|
checkpoint_name = (
|
||||||
'name': f"Checkpoint {checkpoint_id}",
|
checkpoint_name[0] if checkpoint_name else None
|
||||||
'version': '',
|
)
|
||||||
'type': 'checkpoint'
|
if isinstance(checkpoint_name, str):
|
||||||
}
|
checkpoint_match = re.search(r'civitai:(\d+)@(\d+)', checkpoint_name)
|
||||||
if metadata_provider:
|
if checkpoint_match:
|
||||||
try:
|
checkpoint_id = checkpoint_match.group(1)
|
||||||
civitai_info_tuple = await metadata_provider.get_model_version_info(checkpoint_version_id)
|
checkpoint_version_id = checkpoint_match.group(2)
|
||||||
civitai_info, _ = civitai_info_tuple if isinstance(civitai_info_tuple, tuple) else (civitai_info_tuple, None)
|
checkpoint = {
|
||||||
checkpoint = await self.populate_checkpoint_from_civitai(checkpoint, civitai_info)
|
'id': checkpoint_version_id,
|
||||||
except Exception as e:
|
'modelId': checkpoint_id,
|
||||||
logger.error(f"Error fetching Civitai info for checkpoint: {e}")
|
'name': f"Checkpoint {checkpoint_id}",
|
||||||
|
'version': '',
|
||||||
|
'type': 'checkpoint'
|
||||||
|
}
|
||||||
|
if metadata_provider:
|
||||||
|
try:
|
||||||
|
civitai_info_tuple = await metadata_provider.get_model_version_info(checkpoint_version_id)
|
||||||
|
civitai_info, _ = civitai_info_tuple if isinstance(civitai_info_tuple, tuple) else (civitai_info_tuple, None)
|
||||||
|
checkpoint = await self.populate_checkpoint_from_civitai(checkpoint, civitai_info)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error fetching Civitai info for checkpoint: {e}")
|
||||||
|
|
||||||
recipe_base_model = checkpoint.get('baseModel') if checkpoint else None
|
recipe_base_model = checkpoint.get('baseModel') if checkpoint else None
|
||||||
loras = []
|
loras = []
|
||||||
|
|||||||
@@ -196,7 +196,7 @@ class RecipeFormatParser(RecipeMetadataParser):
|
|||||||
filtered_gen_params[key] = value
|
filtered_gen_params[key] = value
|
||||||
|
|
||||||
return {
|
return {
|
||||||
'base_model': checkpoint['baseModel'] if checkpoint and checkpoint.get('baseModel') else recipe_metadata.get('base_model', ''),
|
'base_model': checkpoint['baseModel'] if checkpoint and checkpoint.get('baseModel') else (recipe_metadata.get('base_model') or None),
|
||||||
'loras': loras,
|
'loras': loras,
|
||||||
'gen_params': filtered_gen_params,
|
'gen_params': filtered_gen_params,
|
||||||
'tags': recipe_metadata.get('tags', []),
|
'tags': recipe_metadata.get('tags', []),
|
||||||
@@ -208,3 +208,24 @@ class RecipeFormatParser(RecipeMetadataParser):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error parsing recipe format metadata: {e}", exc_info=True)
|
logger.error(f"Error parsing recipe format metadata: {e}", exc_info=True)
|
||||||
return {"error": str(e), "loras": []}
|
return {"error": str(e), "loras": []}
|
||||||
|
|
||||||
|
|
||||||
|
def strip_recipe_metadata(metadata_text: str) -> str:
|
||||||
|
"""Strip the ``Recipe metadata: {...}`` block appended by LoRA Manager.
|
||||||
|
|
||||||
|
The saved recipe image carries the original generation metadata followed
|
||||||
|
by an appended recipe JSON block (see ``ExifUtils.append_recipe_metadata``).
|
||||||
|
Re-import wants to re-parse the original embedded metadata, so this returns
|
||||||
|
only the text before the appended marker. The input is returned unchanged
|
||||||
|
when no marker is present.
|
||||||
|
"""
|
||||||
|
if not metadata_text:
|
||||||
|
return metadata_text
|
||||||
|
match = re.search(
|
||||||
|
RecipeFormatParser.METADATA_MARKER,
|
||||||
|
metadata_text,
|
||||||
|
re.IGNORECASE | re.DOTALL,
|
||||||
|
)
|
||||||
|
if not match:
|
||||||
|
return metadata_text
|
||||||
|
return metadata_text[: match.start()].strip()
|
||||||
|
|||||||
@@ -24,9 +24,11 @@ from ..services.use_cases import (
|
|||||||
AutoOrganizeUseCase,
|
AutoOrganizeUseCase,
|
||||||
BulkMetadataRefreshUseCase,
|
BulkMetadataRefreshUseCase,
|
||||||
DownloadModelUseCase,
|
DownloadModelUseCase,
|
||||||
|
FilenameTemplateUseCase,
|
||||||
)
|
)
|
||||||
from ..services.websocket_progress_callback import (
|
from ..services.websocket_progress_callback import (
|
||||||
WebSocketBroadcastCallback,
|
WebSocketBroadcastCallback,
|
||||||
|
WebSocketFilenameTemplateProgressCallback,
|
||||||
WebSocketProgressCallback,
|
WebSocketProgressCallback,
|
||||||
)
|
)
|
||||||
from ..utils.exif_utils import ExifUtils
|
from ..utils.exif_utils import ExifUtils
|
||||||
@@ -37,6 +39,7 @@ from .handlers.model_handlers import (
|
|||||||
ModelAutoOrganizeHandler,
|
ModelAutoOrganizeHandler,
|
||||||
ModelCivitaiHandler,
|
ModelCivitaiHandler,
|
||||||
ModelDownloadHandler,
|
ModelDownloadHandler,
|
||||||
|
ModelFilenameTemplateHandler,
|
||||||
ModelHandlerSet,
|
ModelHandlerSet,
|
||||||
ModelListingHandler,
|
ModelListingHandler,
|
||||||
ModelManagementHandler,
|
ModelManagementHandler,
|
||||||
@@ -83,6 +86,9 @@ class BaseModelRoutes(ABC):
|
|||||||
self.model_lifecycle_service: ModelLifecycleService | None = None
|
self.model_lifecycle_service: ModelLifecycleService | None = None
|
||||||
self.websocket_progress_callback = WebSocketProgressCallback()
|
self.websocket_progress_callback = WebSocketProgressCallback()
|
||||||
self.metadata_progress_callback = WebSocketBroadcastCallback()
|
self.metadata_progress_callback = WebSocketBroadcastCallback()
|
||||||
|
self.filename_template_progress_callback = (
|
||||||
|
WebSocketFilenameTemplateProgressCallback()
|
||||||
|
)
|
||||||
|
|
||||||
self._handler_set: ModelHandlerSet | None = None
|
self._handler_set: ModelHandlerSet | None = None
|
||||||
self._handler_mapping: Dict[str, Callable[[web.Request], Awaitable[web.Response]]] | None = None
|
self._handler_mapping: Dict[str, Callable[[web.Request], Awaitable[web.Response]]] | None = None
|
||||||
@@ -149,6 +155,7 @@ class BaseModelRoutes(ABC):
|
|||||||
settings_service=self._settings,
|
settings_service=self._settings,
|
||||||
server_i18n=self._server_i18n,
|
server_i18n=self._server_i18n,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
|
page_context_provider=self._get_page_context_provider(),
|
||||||
)
|
)
|
||||||
listing = ModelListingHandler(
|
listing = ModelListingHandler(
|
||||||
service=service,
|
service=service,
|
||||||
@@ -201,6 +208,17 @@ class BaseModelRoutes(ABC):
|
|||||||
ws_manager=self._ws_manager,
|
ws_manager=self._ws_manager,
|
||||||
logger=logger,
|
logger=logger,
|
||||||
)
|
)
|
||||||
|
filename_template_use_case = FilenameTemplateUseCase(
|
||||||
|
scanner=service.scanner,
|
||||||
|
lifecycle_service=self._ensure_lifecycle_service(),
|
||||||
|
lock_provider=self._ws_manager,
|
||||||
|
model_type=service.model_type,
|
||||||
|
)
|
||||||
|
filename_template = ModelFilenameTemplateHandler(
|
||||||
|
use_case=filename_template_use_case,
|
||||||
|
progress_callback=self.filename_template_progress_callback,
|
||||||
|
logger=logger,
|
||||||
|
)
|
||||||
updates = ModelUpdateHandler(
|
updates = ModelUpdateHandler(
|
||||||
service=service,
|
service=service,
|
||||||
update_service=update_service,
|
update_service=update_service,
|
||||||
@@ -217,6 +235,7 @@ class BaseModelRoutes(ABC):
|
|||||||
civitai=civitai,
|
civitai=civitai,
|
||||||
move=move,
|
move=move,
|
||||||
auto_organize=auto_organize,
|
auto_organize=auto_organize,
|
||||||
|
filename_template=filename_template,
|
||||||
updates=updates,
|
updates=updates,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -250,6 +269,10 @@ class BaseModelRoutes(ABC):
|
|||||||
"""Get expected model types string for error messages - to be overridden by subclasses."""
|
"""Get expected model types string for error messages - to be overridden by subclasses."""
|
||||||
return "any model type"
|
return "any model type"
|
||||||
|
|
||||||
|
def _get_page_context_provider(self):
|
||||||
|
"""Optional hook returning extra template context for the page view."""
|
||||||
|
return None
|
||||||
|
|
||||||
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_WEIGHT_FILE_TYPES 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)
|
||||||
|
|||||||
@@ -47,15 +47,16 @@ class CheckpointRoutes(BaseModelRoutes):
|
|||||||
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/checkpoints_roots', prefix, self.get_checkpoints_roots)
|
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/checkpoints_roots', prefix, self.get_checkpoints_roots)
|
||||||
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/unet_roots', prefix, self.get_unet_roots)
|
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/unet_roots', prefix, self.get_unet_roots)
|
||||||
|
|
||||||
# Name/base_model pool for the Random Checkpoint/Unet Loader nodes
|
# Name/base_model pool for the Checkpoint/Unet Loader nodes' base_model filtering
|
||||||
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/loader-pool', prefix, self.get_loader_pool)
|
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/loader-pool', prefix, self.get_loader_pool)
|
||||||
|
|
||||||
async def get_loader_pool(self, request: web.Request) -> web.Response:
|
async def get_loader_pool(self, request: web.Request) -> web.Response:
|
||||||
"""Return ComfyUI-formatted model names with their base_model.
|
"""Return ComfyUI-formatted model names with their base_model.
|
||||||
|
|
||||||
Backing data for the Random Checkpoint/Unet Loader nodes: the front-end
|
Backing data for the Checkpoint/Unet Loader nodes'
|
||||||
filters the ckpt_name/unet_name combo options by base_model using this
|
control_after_generate feature: the front-end filters the
|
||||||
pool, so control_after_generate randomizes within the narrowed set.
|
ckpt_name/unet_name combo options by base_model using this pool, so
|
||||||
|
randomize mode picks within the narrowed set.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
sub_type = request.query.get("sub_type", "checkpoint")
|
sub_type = request.query.get("sub_type", "checkpoint")
|
||||||
|
|||||||
@@ -0,0 +1,110 @@
|
|||||||
|
"""HTTP handler for download target routing decisions."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
from ...services.download_routing import (
|
||||||
|
is_diffusion_model_download,
|
||||||
|
resolve_other_download_sub_type,
|
||||||
|
)
|
||||||
|
from ...utils.constants import VALID_OTHER_CIVITAI_TYPES
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class DownloadRoutingHandler:
|
||||||
|
"""Expose the download-time checkpoint/diffusion-model routing decision.
|
||||||
|
|
||||||
|
The web UI calls this when the user reaches the download location step
|
||||||
|
so the root dropdown offers the same root set (checkpoint vs unet) that
|
||||||
|
the download manager would pick for ``use_default_paths``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def get_download_routing(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
payload = await request.json()
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON payload"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
model_type = payload.get("model_type", "")
|
||||||
|
base_model = payload.get("base_model") or ""
|
||||||
|
file_types = payload.get("file_types") or []
|
||||||
|
selected_file_type = payload.get("selected_file_type")
|
||||||
|
|
||||||
|
if not isinstance(model_type, str) or not model_type:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "model_type is required"}, status=400
|
||||||
|
)
|
||||||
|
if not isinstance(base_model, str) or not isinstance(file_types, list):
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "base_model must be a string and file_types a list",
|
||||||
|
},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
if selected_file_type is not None and not isinstance(selected_file_type, str):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "selected_file_type must be a string"},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
if model_type.lower() in VALID_OTHER_CIVITAI_TYPES:
|
||||||
|
from ...services.settings_manager import get_settings_manager
|
||||||
|
|
||||||
|
settings = get_settings_manager()
|
||||||
|
if not settings.is_other_models_enabled():
|
||||||
|
# Opt-in feature is off: never auto-route, the UI falls back to
|
||||||
|
# manual folder selection and the download manager rejects it.
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": True,
|
||||||
|
"root_kind": "other",
|
||||||
|
"sub_type": None,
|
||||||
|
"disabled": True,
|
||||||
|
"reason": "other_models_disabled",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
sub_type = resolve_other_download_sub_type(
|
||||||
|
model_type,
|
||||||
|
file_types=(str(t) for t in file_types),
|
||||||
|
selected_file_type=selected_file_type,
|
||||||
|
)
|
||||||
|
if sub_type and not settings.is_other_sub_type_enabled(sub_type):
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": True,
|
||||||
|
"root_kind": "other",
|
||||||
|
"sub_type": None,
|
||||||
|
"disabled": True,
|
||||||
|
"reason": "other_sub_type_disabled",
|
||||||
|
"requested_sub_type": sub_type,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": True,
|
||||||
|
"root_kind": "other",
|
||||||
|
"sub_type": sub_type,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
is_diffusion = is_diffusion_model_download(
|
||||||
|
model_type,
|
||||||
|
file_types=(str(t) for t in file_types),
|
||||||
|
base_model=base_model,
|
||||||
|
)
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": True,
|
||||||
|
"is_diffusion_model": is_diffusion,
|
||||||
|
"root_kind": "unet" if is_diffusion else model_type,
|
||||||
|
}
|
||||||
|
)
|
||||||
@@ -1,508 +0,0 @@
|
|||||||
"""Handlers for Hugging Face model listing and download.
|
|
||||||
|
|
||||||
Minimal MVP implementation — uses direct HTTP to the HF API for file
|
|
||||||
listing and the project's existing aiohttp-based Downloader for
|
|
||||||
downloading. No huggingface_hub dependency required.
|
|
||||||
"""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import json
|
|
||||||
import logging
|
|
||||||
import os
|
|
||||||
import re
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
import aiohttp
|
|
||||||
from aiohttp import web
|
|
||||||
|
|
||||||
from ...config import config
|
|
||||||
from ...services.downloader import (
|
|
||||||
DownloadProgress,
|
|
||||||
get_downloader,
|
|
||||||
)
|
|
||||||
from ...services.aria2_downloader import Aria2Downloader
|
|
||||||
from ...services.settings_manager import get_settings_manager
|
|
||||||
from ...services.service_registry import ServiceRegistry
|
|
||||||
from ...services.websocket_manager import ws_manager
|
|
||||||
from ...utils.constants import MODEL_FILE_EXTENSIONS
|
|
||||||
from ...utils.metadata_manager import MetadataManager
|
|
||||||
from ...utils.models import LoraMetadata, CheckpointMetadata, EmbeddingMetadata
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
_DEFAULT_MODEL_CLASS = LoraMetadata
|
|
||||||
_DEFAULT_SCANNER_GETTER = "get_lora_scanner"
|
|
||||||
|
|
||||||
# Shared aiohttp session for HF API calls (created on first use)
|
|
||||||
_hf_api_session: aiohttp.ClientSession | None = None
|
|
||||||
|
|
||||||
|
|
||||||
async def _get_hf_api_session() -> aiohttp.ClientSession:
|
|
||||||
"""Get or create the shared aiohttp session for HF API calls."""
|
|
||||||
global _hf_api_session # needed because we reassign the module-level name
|
|
||||||
if _hf_api_session is None or _hf_api_session.closed:
|
|
||||||
_hf_api_session = aiohttp.ClientSession(
|
|
||||||
headers={"User-Agent": "ComfyUI-LoRA-Manager/1.0"},
|
|
||||||
timeout=aiohttp.ClientTimeout(total=30),
|
|
||||||
)
|
|
||||||
return _hf_api_session
|
|
||||||
|
|
||||||
|
|
||||||
async def close_hf_api_session() -> None:
|
|
||||||
"""Close the shared HF API session, if it was ever created."""
|
|
||||||
global _hf_api_session
|
|
||||||
if _hf_api_session is not None and not _hf_api_session.closed:
|
|
||||||
await _hf_api_session.close()
|
|
||||||
_hf_api_session = None
|
|
||||||
|
|
||||||
|
|
||||||
def _infer_model_type(model_root: str) -> tuple[Any, str]:
|
|
||||||
"""Determine model class and scanner by matching ``model_root`` against the
|
|
||||||
configured root paths for each model type (from ``Config``).
|
|
||||||
|
|
||||||
The ``model_root`` value comes from the frontend's model-root dropdown,
|
|
||||||
which is populated from the current page's scanner roots. By checking
|
|
||||||
which scanner's root list it belongs to, we avoid fragile heuristics
|
|
||||||
like substring-matching path names.
|
|
||||||
"""
|
|
||||||
norm = os.path.normpath(model_root).replace(os.sep, "/")
|
|
||||||
|
|
||||||
# LoRA roots
|
|
||||||
for p in (config.loras_roots or []) + (config.extra_loras_roots or []):
|
|
||||||
if os.path.normpath(p).replace(os.sep, "/") == norm:
|
|
||||||
return LoraMetadata, "get_lora_scanner"
|
|
||||||
|
|
||||||
# Checkpoint / UNet roots
|
|
||||||
for p in (
|
|
||||||
(config.checkpoints_roots or [])
|
|
||||||
+ (config.extra_checkpoints_roots or [])
|
|
||||||
+ (config.unet_roots or [])
|
|
||||||
+ (config.extra_unet_roots or [])
|
|
||||||
):
|
|
||||||
if os.path.normpath(p).replace(os.sep, "/") == norm:
|
|
||||||
return CheckpointMetadata, "get_checkpoint_scanner"
|
|
||||||
|
|
||||||
# Embedding roots
|
|
||||||
for p in (config.embeddings_roots or []) + (config.extra_embeddings_roots or []):
|
|
||||||
if os.path.normpath(p).replace(os.sep, "/") == norm:
|
|
||||||
return EmbeddingMetadata, "get_embedding_scanner"
|
|
||||||
|
|
||||||
# Fallback — should not happen in normal use
|
|
||||||
logger.warning(
|
|
||||||
"Could not determine model type for root '%s'; defaulting to LoRA",
|
|
||||||
model_root,
|
|
||||||
)
|
|
||||||
return _DEFAULT_MODEL_CLASS, _DEFAULT_SCANNER_GETTER
|
|
||||||
|
|
||||||
|
|
||||||
async def _save_hf_metadata(dest_path: str, repo: str, model_root: str) -> None:
|
|
||||||
"""Create a proper .metadata.json and add the model to the scanner cache.
|
|
||||||
|
|
||||||
Uses ``MetadataManager.create_default_metadata()`` which computes the
|
|
||||||
SHA256 hash, extracts safetensors header metadata (base_model), and
|
|
||||||
produces a fully-populated ``LoraMetadata`` (or ``CheckpointMetadata`` /
|
|
||||||
``EmbeddingMetadata``) object. We then overlay HF-specific fields and
|
|
||||||
register the model in the in-memory scanner cache so it appears
|
|
||||||
immediately without a full filesystem walk.
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
hf_url = f"https://huggingface.co/{repo}"
|
|
||||||
model_class, scanner_getter_name = _infer_model_type(model_root)
|
|
||||||
|
|
||||||
# 1. Create proper metadata (computes SHA256, reads safetensors headers)
|
|
||||||
metadata = await MetadataManager.create_default_metadata(
|
|
||||||
dest_path, model_class=model_class
|
|
||||||
)
|
|
||||||
if metadata is None:
|
|
||||||
logger.warning("create_default_metadata returned None for %s", dest_path)
|
|
||||||
return
|
|
||||||
|
|
||||||
# 2. Overlay HF-specific fields
|
|
||||||
metadata._unknown_fields["hf_url"] = hf_url
|
|
||||||
metadata.from_civitai = False # HF models are not from CivitAI
|
|
||||||
|
|
||||||
metadata_dict = metadata.to_dict()
|
|
||||||
if "trainedWords" in metadata_dict and not metadata_dict["trainedWords"]:
|
|
||||||
del metadata_dict["trainedWords"]
|
|
||||||
|
|
||||||
# 3. Save metadata atomically
|
|
||||||
await MetadataManager.save_metadata(dest_path, metadata_dict)
|
|
||||||
logger.info("Saved HF metadata (with hf_url) for %s", dest_path)
|
|
||||||
|
|
||||||
# 4. Determine relative folder path for cache
|
|
||||||
# model_root is an absolute path; dest_path is under it
|
|
||||||
folder = ""
|
|
||||||
if os.path.isabs(model_root) and dest_path.startswith(model_root):
|
|
||||||
rel = os.path.relpath(os.path.dirname(dest_path), model_root)
|
|
||||||
folder = rel.replace(os.sep, "/") if rel != "." else ""
|
|
||||||
|
|
||||||
# 5. Add to scanner cache (same as CivitAI's _execute_download does)
|
|
||||||
scanner_getter = getattr(ServiceRegistry, scanner_getter_name, None)
|
|
||||||
if scanner_getter is not None:
|
|
||||||
scanner = await scanner_getter()
|
|
||||||
if scanner is not None:
|
|
||||||
metadata_dict = metadata.to_dict()
|
|
||||||
metadata_dict["hf_url"] = hf_url
|
|
||||||
await scanner.add_model_to_cache(metadata_dict, folder)
|
|
||||||
logger.info("Added %s to scanner cache (folder=%s)", dest_path, folder)
|
|
||||||
|
|
||||||
except Exception as exc:
|
|
||||||
logger.warning("Failed to save HF metadata for %s: %s", dest_path, exc)
|
|
||||||
|
|
||||||
|
|
||||||
def _find_matching_root(dest_dir: str) -> str | None:
|
|
||||||
"""Walk up *dest_dir* to find which configured scanner root it belongs to."""
|
|
||||||
norm = os.path.normpath(dest_dir).replace(os.sep, "/")
|
|
||||||
all_roots = []
|
|
||||||
for root_list in (
|
|
||||||
config.loras_roots or [],
|
|
||||||
config.extra_loras_roots or [],
|
|
||||||
config.checkpoints_roots or [],
|
|
||||||
config.extra_checkpoints_roots or [],
|
|
||||||
config.unet_roots or [],
|
|
||||||
config.extra_unet_roots or [],
|
|
||||||
config.embeddings_roots or [],
|
|
||||||
config.extra_embeddings_roots or [],
|
|
||||||
):
|
|
||||||
all_roots.extend([os.path.normpath(p).replace(os.sep, "/") for p in root_list])
|
|
||||||
# Find the longest matching prefix
|
|
||||||
match: str | None = None
|
|
||||||
for root in all_roots:
|
|
||||||
if norm.startswith(root):
|
|
||||||
if match is None or len(root) > len(match):
|
|
||||||
match = root
|
|
||||||
return match
|
|
||||||
|
|
||||||
|
|
||||||
async def _add_to_scanner_cache(dest_path: str, metadata: dict[str, Any]) -> None:
|
|
||||||
model_dir = os.path.dirname(dest_path)
|
|
||||||
model_root = _find_matching_root(model_dir)
|
|
||||||
if not model_root:
|
|
||||||
raise ValueError(f"File path {dest_path} is not within any configured scanner root")
|
|
||||||
scanner_getter_name = _infer_model_type(model_root)[1]
|
|
||||||
scanner_getter = getattr(ServiceRegistry, scanner_getter_name, None)
|
|
||||||
if scanner_getter is None:
|
|
||||||
raise RuntimeError(f"Scanner getter '{scanner_getter_name}' not found in ServiceRegistry")
|
|
||||||
scanner = await scanner_getter()
|
|
||||||
if scanner is None:
|
|
||||||
raise RuntimeError(f"Scanner '{scanner_getter_name}' returned None")
|
|
||||||
await scanner.update_single_model_cache(dest_path, dest_path, metadata)
|
|
||||||
|
|
||||||
|
|
||||||
class HfHandler:
|
|
||||||
"""Handle Hugging Face model browsing and download."""
|
|
||||||
|
|
||||||
async def set_hf_url(self, request: web.Request) -> web.Response:
|
|
||||||
try:
|
|
||||||
payload: dict[str, Any] = await request.json()
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
return web.json_response({"success": False, "error": "Invalid JSON"}, status=400)
|
|
||||||
|
|
||||||
file_path = (payload.get("file_path") or "").strip()
|
|
||||||
hf_url = (payload.get("hf_url") or "").strip()
|
|
||||||
|
|
||||||
if not file_path or not hf_url:
|
|
||||||
return web.json_response(
|
|
||||||
{"success": False, "error": "Missing required fields: 'file_path' and 'hf_url'"},
|
|
||||||
status=400,
|
|
||||||
)
|
|
||||||
|
|
||||||
m = re.match(r"^https?://huggingface\.co/([^/]+/[^/]+)/?$", hf_url)
|
|
||||||
if not m:
|
|
||||||
return web.json_response(
|
|
||||||
{
|
|
||||||
"success": False,
|
|
||||||
"error": "Invalid HuggingFace URL. Expected format: https://huggingface.co/user/repo",
|
|
||||||
},
|
|
||||||
status=400,
|
|
||||||
)
|
|
||||||
|
|
||||||
if not os.path.isfile(file_path):
|
|
||||||
return web.json_response(
|
|
||||||
{"success": False, "error": f"File not found: {file_path}"},
|
|
||||||
status=404,
|
|
||||||
)
|
|
||||||
|
|
||||||
model_root = _find_matching_root(os.path.dirname(file_path))
|
|
||||||
if not model_root:
|
|
||||||
return web.json_response(
|
|
||||||
{
|
|
||||||
"success": False,
|
|
||||||
"error": "File is not within any configured model directory. Cannot link to HuggingFace.",
|
|
||||||
},
|
|
||||||
status=400,
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
|
||||||
existing = await MetadataManager.load_metadata_payload(file_path)
|
|
||||||
if existing.get("hf_url") == hf_url:
|
|
||||||
return web.json_response({
|
|
||||||
"success": True,
|
|
||||||
"message": "hf_url already set",
|
|
||||||
"hf_url": hf_url,
|
|
||||||
})
|
|
||||||
|
|
||||||
existing["hf_url"] = hf_url
|
|
||||||
existing["from_civitai"] = False
|
|
||||||
await MetadataManager.save_metadata(file_path, existing)
|
|
||||||
|
|
||||||
await _add_to_scanner_cache(file_path, existing)
|
|
||||||
|
|
||||||
logger.info("Set hf_url=%s for %s", hf_url, file_path)
|
|
||||||
return web.json_response({
|
|
||||||
"success": True,
|
|
||||||
"message": f"hf_url set to {hf_url}",
|
|
||||||
"hf_url": hf_url,
|
|
||||||
})
|
|
||||||
except Exception as exc:
|
|
||||||
logger.error("Failed to set hf_url for %s: %s", file_path, exc)
|
|
||||||
return web.json_response(
|
|
||||||
{"success": False, "error": str(exc)},
|
|
||||||
status=500,
|
|
||||||
)
|
|
||||||
|
|
||||||
async def get_hf_repo_files(self, request: web.Request) -> web.Response:
|
|
||||||
"""List model-weight files from a HF repo with real file sizes.
|
|
||||||
|
|
||||||
Uses the HF tree API endpoint which returns accurate file sizes
|
|
||||||
(including LFS-tracked files), unlike the model info endpoint.
|
|
||||||
"""
|
|
||||||
repo = request.query.get("repo", "").strip()
|
|
||||||
if not repo or "/" not in repo:
|
|
||||||
return web.json_response(
|
|
||||||
{"error": "Missing or invalid 'repo' parameter (expected user/repo)"},
|
|
||||||
status=400,
|
|
||||||
)
|
|
||||||
|
|
||||||
url = f"https://huggingface.co/api/models/{repo}/tree/main"
|
|
||||||
|
|
||||||
try:
|
|
||||||
session = await _get_hf_api_session()
|
|
||||||
async with session.get(url) as resp:
|
|
||||||
if resp.status == 404:
|
|
||||||
return web.json_response(
|
|
||||||
{"error": f"Repo '{repo}' not found"}, status=404
|
|
||||||
)
|
|
||||||
if resp.status != 200:
|
|
||||||
text = await resp.text()
|
|
||||||
return web.json_response(
|
|
||||||
{"error": f"HF API error {resp.status}: {text[:200]}"},
|
|
||||||
status=resp.status,
|
|
||||||
)
|
|
||||||
tree: list[dict[str, Any]] = await resp.json()
|
|
||||||
except Exception as exc:
|
|
||||||
logger.error("Failed to fetch HF repo files: %s", exc)
|
|
||||||
return web.json_response({"error": str(exc)}, status=502)
|
|
||||||
|
|
||||||
files: list[dict[str, Any]] = []
|
|
||||||
for entry in tree:
|
|
||||||
path: str = entry.get("path", "")
|
|
||||||
ext = os.path.splitext(path)[1].lower()
|
|
||||||
if ext not in MODEL_FILE_EXTENSIONS:
|
|
||||||
continue
|
|
||||||
size = entry.get("size", 0) or 0
|
|
||||||
if size == 0 and "lfs" in entry:
|
|
||||||
size = entry["lfs"].get("size", 0) or 0
|
|
||||||
files.append({
|
|
||||||
"filename": path,
|
|
||||||
"size": size,
|
|
||||||
})
|
|
||||||
|
|
||||||
files.sort(key=lambda f: f["size"], reverse=True)
|
|
||||||
return web.json_response(files)
|
|
||||||
|
|
||||||
async def download_hf_model(self, request: web.Request) -> web.Response:
|
|
||||||
"""Download a single file from Hugging Face into the model directory.
|
|
||||||
|
|
||||||
POST JSON body::
|
|
||||||
|
|
||||||
{
|
|
||||||
"repo": "dx8152/Flux2-Klein-9B-Consistency",
|
|
||||||
"filename": "Flux2-Klein-9B-consistency-V2.safetensors",
|
|
||||||
"revision": "main",
|
|
||||||
"model_root": "loras",
|
|
||||||
"relative_path": "",
|
|
||||||
"use_default_paths": false,
|
|
||||||
"download_id": "optional-batch-id"
|
|
||||||
}
|
|
||||||
|
|
||||||
If ``download_id`` is provided, real-time progress (bytes, speed,
|
|
||||||
percentage) is broadcast via the WebSocket progress system, matching
|
|
||||||
the CivitAI download experience.
|
|
||||||
|
|
||||||
Respects the ``download_backend`` setting (``aria2`` or ``default``).
|
|
||||||
"""
|
|
||||||
try:
|
|
||||||
payload: dict[str, Any] = await request.json()
|
|
||||||
except json.JSONDecodeError:
|
|
||||||
return web.json_response({"error": "Invalid JSON"}, status=400)
|
|
||||||
|
|
||||||
repo = (payload.get("repo") or "").strip()
|
|
||||||
filename = (payload.get("filename") or "").strip()
|
|
||||||
revision = (payload.get("revision") or "main").strip()
|
|
||||||
model_root = (payload.get("model_root") or "").strip()
|
|
||||||
relative_path = (payload.get("relative_path") or "").strip()
|
|
||||||
use_default_paths = bool(payload.get("use_default_paths", False))
|
|
||||||
download_id: str | None = payload.get("download_id")
|
|
||||||
|
|
||||||
logger.info(
|
|
||||||
"download_hf_model: repo=%s file=%s root=%s download_id=%s",
|
|
||||||
repo, filename, model_root, download_id,
|
|
||||||
)
|
|
||||||
|
|
||||||
if not repo or not filename:
|
|
||||||
return web.json_response(
|
|
||||||
{"error": "Missing required fields: 'repo' and 'filename'"}, status=400
|
|
||||||
)
|
|
||||||
|
|
||||||
# Validate repo format — must be user/repo_name
|
|
||||||
if repo.count("/") != 1 or not re.match(r"^[a-zA-Z0-9_.-]+/[a-zA-Z0-9_.-]+$", repo):
|
|
||||||
return web.json_response({"error": f"Invalid repo format: {repo}"}, status=400)
|
|
||||||
author, repo_name = repo.split("/", 1)
|
|
||||||
if ".." in (author, repo_name) or "." in (author, repo_name):
|
|
||||||
return web.json_response({"error": f"Invalid repo format: {repo}"}, status=400)
|
|
||||||
|
|
||||||
# Validate filename — must not contain path traversal
|
|
||||||
if ".." in filename:
|
|
||||||
return web.json_response({"error": "Invalid filename"}, status=400)
|
|
||||||
|
|
||||||
# Validate relative_path — must not be absolute or escape base directory
|
|
||||||
if relative_path:
|
|
||||||
if os.path.isabs(relative_path):
|
|
||||||
return web.json_response({"error": "relative_path must not be absolute"}, status=400)
|
|
||||||
if ".." in relative_path.split("/") or "\\" in relative_path:
|
|
||||||
return web.json_response({"error": "Invalid relative_path"}, status=400)
|
|
||||||
|
|
||||||
# Use model_root directly as the base directory — same approach as
|
|
||||||
# CivitAI's download path (download_manager.py). No realpath, no
|
|
||||||
# allowed-roots validation, no path-traversal check; those are
|
|
||||||
# unnecessary when the frontend sends the path from its own dropdown
|
|
||||||
# (populated from scanner roots). Using the "business path" directly
|
|
||||||
# keeps dest_path consistent with scanner roots so that later folder
|
|
||||||
# derivation (in _save_hf_metadata) works correctly.
|
|
||||||
if os.path.isabs(model_root):
|
|
||||||
base_dir = os.path.normpath(model_root)
|
|
||||||
else:
|
|
||||||
base_dir = os.path.normpath(os.path.join(os.getcwd(), "models", model_root))
|
|
||||||
|
|
||||||
if use_default_paths:
|
|
||||||
target_dir = os.path.join(base_dir, "huggingface", author, repo_name)
|
|
||||||
elif relative_path:
|
|
||||||
target_dir = os.path.join(base_dir, relative_path)
|
|
||||||
else:
|
|
||||||
target_dir = base_dir
|
|
||||||
|
|
||||||
# Strip HF repo subdirectory — "diffusion_models/xxx.safetensors"
|
|
||||||
# is an HF repo convention, not meaningful for local storage.
|
|
||||||
file_base = os.path.basename(filename)
|
|
||||||
|
|
||||||
os.makedirs(target_dir, exist_ok=True)
|
|
||||||
dest_path = os.path.join(target_dir, file_base)
|
|
||||||
|
|
||||||
# Check if already exists (simple skip)
|
|
||||||
if os.path.exists(dest_path) and os.path.getsize(dest_path) > 0:
|
|
||||||
logger.info("download_hf_model: file already exists, skipping — %s", dest_path)
|
|
||||||
return web.json_response({
|
|
||||||
"success": True,
|
|
||||||
"message": f"File already exists: {dest_path}",
|
|
||||||
"path": dest_path,
|
|
||||||
})
|
|
||||||
|
|
||||||
# Build HF resolve URL
|
|
||||||
resolve_url = (
|
|
||||||
f"https://huggingface.co/{repo}/resolve/{revision}/{filename}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Set up progress callback if download_id is provided
|
|
||||||
progress_callback = None
|
|
||||||
if download_id:
|
|
||||||
|
|
||||||
async def _progress_callback(
|
|
||||||
progress: float | DownloadProgress,
|
|
||||||
snapshot: DownloadProgress | None = None,
|
|
||||||
) -> None:
|
|
||||||
percent = 0.0
|
|
||||||
metrics = snapshot if isinstance(snapshot, DownloadProgress) else None
|
|
||||||
|
|
||||||
if isinstance(progress, DownloadProgress):
|
|
||||||
percent = progress.percent_complete
|
|
||||||
metrics = progress
|
|
||||||
elif isinstance(snapshot, DownloadProgress):
|
|
||||||
percent = snapshot.percent_complete
|
|
||||||
else:
|
|
||||||
percent = float(progress)
|
|
||||||
|
|
||||||
broadcast: dict[str, Any] = {
|
|
||||||
"status": "progress",
|
|
||||||
"progress": round(percent),
|
|
||||||
}
|
|
||||||
if metrics:
|
|
||||||
broadcast["bytes_downloaded"] = metrics.bytes_downloaded
|
|
||||||
broadcast["total_bytes"] = metrics.total_bytes
|
|
||||||
broadcast["bytes_per_second"] = metrics.bytes_per_second
|
|
||||||
|
|
||||||
await ws_manager.broadcast_download_progress(download_id, broadcast)
|
|
||||||
|
|
||||||
progress_callback = _progress_callback
|
|
||||||
|
|
||||||
# Respect download backend setting (aria2 vs default)
|
|
||||||
download_backend = (
|
|
||||||
get_settings_manager().get("download_backend", "default")
|
|
||||||
)
|
|
||||||
|
|
||||||
if download_backend == "aria2":
|
|
||||||
aria2 = await Aria2Downloader.get_instance()
|
|
||||||
aid = download_id or f"hf_{repo}_{filename}"
|
|
||||||
try:
|
|
||||||
hf_success, hf_result = await aria2.download_file(
|
|
||||||
url=resolve_url,
|
|
||||||
save_path=dest_path,
|
|
||||||
download_id=aid,
|
|
||||||
progress_callback=progress_callback,
|
|
||||||
)
|
|
||||||
if hf_success:
|
|
||||||
await _save_hf_metadata(dest_path, repo, model_root)
|
|
||||||
return web.json_response({
|
|
||||||
"success": True,
|
|
||||||
"message": f"Downloaded to {dest_path}",
|
|
||||||
"path": dest_path,
|
|
||||||
})
|
|
||||||
else:
|
|
||||||
return web.json_response(
|
|
||||||
{"success": False, "error": hf_result or "aria2 download failed"},
|
|
||||||
status=500,
|
|
||||||
)
|
|
||||||
except Exception as exc:
|
|
||||||
logger.error("HF download (aria2) failed: %s", exc)
|
|
||||||
return web.json_response(
|
|
||||||
{"success": False, "error": str(exc)}, status=500
|
|
||||||
)
|
|
||||||
|
|
||||||
# Default: use built-in aiohttp Downloader
|
|
||||||
downloader = await get_downloader()
|
|
||||||
try:
|
|
||||||
success, result = await downloader.download_file(
|
|
||||||
url=resolve_url,
|
|
||||||
save_path=dest_path,
|
|
||||||
use_auth=False,
|
|
||||||
allow_resume=True,
|
|
||||||
progress_callback=progress_callback,
|
|
||||||
)
|
|
||||||
if success:
|
|
||||||
await _save_hf_metadata(dest_path, repo, model_root)
|
|
||||||
return web.json_response({
|
|
||||||
"success": True,
|
|
||||||
"message": f"Downloaded to {result}",
|
|
||||||
"path": result,
|
|
||||||
})
|
|
||||||
else:
|
|
||||||
return web.json_response(
|
|
||||||
{"success": False, "error": result or "Download failed"},
|
|
||||||
status=500,
|
|
||||||
)
|
|
||||||
except Exception as exc:
|
|
||||||
logger.error("HF download failed: %s", exc)
|
|
||||||
return web.json_response(
|
|
||||||
{"success": False, "error": str(exc)}, status=500
|
|
||||||
)
|
|
||||||
@@ -53,11 +53,15 @@ from ...utils.constants import (
|
|||||||
PREVIEW_EXTENSIONS,
|
PREVIEW_EXTENSIONS,
|
||||||
SUPPORTED_MEDIA_EXTENSIONS,
|
SUPPORTED_MEDIA_EXTENSIONS,
|
||||||
VALID_LORA_TYPES,
|
VALID_LORA_TYPES,
|
||||||
|
VALID_OTHER_CIVITAI_TYPES,
|
||||||
|
folder_path_schema,
|
||||||
)
|
)
|
||||||
from .hf_handlers import HfHandler
|
from .model_source_handlers import ModelSourceHandler
|
||||||
from .agent_handlers import AgentHandler
|
from .agent_handlers import AgentHandler
|
||||||
|
from .download_routing_handlers import DownloadRoutingHandler
|
||||||
from .model_handlers import ModelCivitaiHandler
|
from .model_handlers import ModelCivitaiHandler
|
||||||
from ...utils.civitai_utils import rewrite_preview_url
|
from ...utils.civitai_utils import rewrite_preview_url
|
||||||
|
from ...utils.directory_browser import browse_directory
|
||||||
from ...utils.example_images_paths import (
|
from ...utils.example_images_paths import (
|
||||||
find_non_compliant_items_in_example_images_root,
|
find_non_compliant_items_in_example_images_root,
|
||||||
is_valid_example_images_root,
|
is_valid_example_images_root,
|
||||||
@@ -419,6 +423,11 @@ def _wsl_to_windows_path(wsl_path: str) -> str | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _has_gui_display() -> bool:
|
||||||
|
"""Check whether a GUI session is reachable for xdg-open."""
|
||||||
|
return bool(os.environ.get("DISPLAY") or os.environ.get("WAYLAND_DISPLAY"))
|
||||||
|
|
||||||
|
|
||||||
class PromptServerProtocol(Protocol):
|
class PromptServerProtocol(Protocol):
|
||||||
"""Subset of PromptServer used by the handlers."""
|
"""Subset of PromptServer used by the handlers."""
|
||||||
|
|
||||||
@@ -657,9 +666,21 @@ class HealthCheckHandler:
|
|||||||
"lora": ServiceRegistry.get_lora_scanner,
|
"lora": ServiceRegistry.get_lora_scanner,
|
||||||
"checkpoint": ServiceRegistry.get_checkpoint_scanner,
|
"checkpoint": ServiceRegistry.get_checkpoint_scanner,
|
||||||
"embedding": ServiceRegistry.get_embedding_scanner,
|
"embedding": ServiceRegistry.get_embedding_scanner,
|
||||||
|
"other": ServiceRegistry.get_other_scanner,
|
||||||
"recipe": ServiceRegistry.get_recipe_scanner,
|
"recipe": ServiceRegistry.get_recipe_scanner,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def _active_scanner_getters(
|
||||||
|
self,
|
||||||
|
) -> Mapping[str, Callable[[], Awaitable[Any]]]:
|
||||||
|
"""Drop the opt-in other scanner while Other Models is disabled."""
|
||||||
|
getters = self._scanner_getters
|
||||||
|
if "other" not in getters:
|
||||||
|
return getters
|
||||||
|
if get_settings_manager().is_other_models_enabled():
|
||||||
|
return getters
|
||||||
|
return {name: getter for name, getter in getters.items() if name != "other"}
|
||||||
|
|
||||||
async def health_check(self, request: web.Request) -> web.Response:
|
async def health_check(self, request: web.Request) -> web.Response:
|
||||||
return web.json_response({"status": "ok"})
|
return web.json_response({"status": "ok"})
|
||||||
|
|
||||||
@@ -671,7 +692,7 @@ class HealthCheckHandler:
|
|||||||
page accepts the update and only reloads once all scanners are done.
|
page accepts the update and only reloads once all scanners are done.
|
||||||
"""
|
"""
|
||||||
pending: list[str] = []
|
pending: list[str] = []
|
||||||
for name, getter in self._scanner_getters.items():
|
for name, getter in self._active_scanner_getters().items():
|
||||||
try:
|
try:
|
||||||
scanner = await getter()
|
scanner = await getter()
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -756,10 +777,19 @@ class DoctorHandler:
|
|||||||
("lora", "LoRAs", ServiceRegistry.get_lora_scanner),
|
("lora", "LoRAs", ServiceRegistry.get_lora_scanner),
|
||||||
("checkpoint", "Checkpoints", ServiceRegistry.get_checkpoint_scanner),
|
("checkpoint", "Checkpoints", ServiceRegistry.get_checkpoint_scanner),
|
||||||
("embedding", "Embeddings", ServiceRegistry.get_embedding_scanner),
|
("embedding", "Embeddings", ServiceRegistry.get_embedding_scanner),
|
||||||
|
("other", "Other Models", ServiceRegistry.get_other_scanner),
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
self._app_version_getter = app_version_getter
|
self._app_version_getter = app_version_getter
|
||||||
|
|
||||||
|
def _active_scanner_factories(
|
||||||
|
self,
|
||||||
|
) -> Sequence[tuple[str, str, Callable[[], Awaitable[Any]]]]:
|
||||||
|
"""Drop the opt-in other scanner while Other Models is disabled."""
|
||||||
|
if self._settings.is_other_models_enabled():
|
||||||
|
return self._scanner_factories
|
||||||
|
return tuple(entry for entry in self._scanner_factories if entry[0] != "other")
|
||||||
|
|
||||||
async def get_doctor_diagnostics(self, request: web.Request) -> web.Response:
|
async def get_doctor_diagnostics(self, request: web.Request) -> web.Response:
|
||||||
try:
|
try:
|
||||||
client_version = (request.query.get("clientVersion") or "").strip()
|
client_version = (request.query.get("clientVersion") or "").strip()
|
||||||
@@ -807,7 +837,7 @@ class DoctorHandler:
|
|||||||
repaired: list[dict[str, Any]] = []
|
repaired: list[dict[str, Any]] = []
|
||||||
failures: list[dict[str, str]] = []
|
failures: list[dict[str, str]] = []
|
||||||
|
|
||||||
for model_type, label, factory in self._scanner_factories:
|
for model_type, label, factory in self._active_scanner_factories():
|
||||||
try:
|
try:
|
||||||
scanner = await factory()
|
scanner = await factory()
|
||||||
await scanner.get_cached_data(force_refresh=True, rebuild_cache=True)
|
await scanner.get_cached_data(force_refresh=True, rebuild_cache=True)
|
||||||
@@ -839,7 +869,7 @@ class DoctorHandler:
|
|||||||
renamed: list[dict[str, Any]] = []
|
renamed: list[dict[str, Any]] = []
|
||||||
|
|
||||||
try:
|
try:
|
||||||
for model_type, label, factory in self._scanner_factories:
|
for model_type, label, factory in self._active_scanner_factories():
|
||||||
try:
|
try:
|
||||||
scanner = await factory()
|
scanner = await factory()
|
||||||
hash_index = getattr(scanner, "_hash_index", None)
|
hash_index = getattr(scanner, "_hash_index", None)
|
||||||
@@ -1071,7 +1101,7 @@ class DoctorHandler:
|
|||||||
overall_status = "ok"
|
overall_status = "ok"
|
||||||
summary = "All model caches look healthy."
|
summary = "All model caches look healthy."
|
||||||
|
|
||||||
for model_type, label, factory in self._scanner_factories:
|
for model_type, label, factory in self._active_scanner_factories():
|
||||||
try:
|
try:
|
||||||
scanner = await factory()
|
scanner = await factory()
|
||||||
persisted = None
|
persisted = None
|
||||||
@@ -1156,7 +1186,7 @@ class DoctorHandler:
|
|||||||
total_conflict_groups = 0
|
total_conflict_groups = 0
|
||||||
total_conflict_files = 0
|
total_conflict_files = 0
|
||||||
|
|
||||||
for model_type, label, factory in self._scanner_factories:
|
for model_type, label, factory in self._active_scanner_factories():
|
||||||
# Duplicate filename detection targets LoRAs which use basename-only
|
# Duplicate filename detection targets LoRAs which use basename-only
|
||||||
# syntax (<lora:name:strength>). Checkpoints/embeddings reference
|
# syntax (<lora:name:strength>). Checkpoints/embeddings reference
|
||||||
# models via relative paths with extensions, so conflicts there would
|
# models via relative paths with extensions, so conflicts there would
|
||||||
@@ -1536,6 +1566,46 @@ class SettingsHandler:
|
|||||||
response_data["civitai_api_key_set"] = bool(raw_key)
|
response_data["civitai_api_key_set"] = bool(raw_key)
|
||||||
raw_llm_key = self._settings.get("llm_api_key")
|
raw_llm_key = self._settings.get("llm_api_key")
|
||||||
response_data["llm_api_key_set"] = bool(raw_llm_key)
|
response_data["llm_api_key_set"] = bool(raw_llm_key)
|
||||||
|
# Derived capability flag (not persisted): whether the host exposes
|
||||||
|
# any other-model folder at all. Standalone installs only know the
|
||||||
|
# folder_paths keys present in settings.json, so the announcement
|
||||||
|
# banner uses this to avoid promising a page that cannot list
|
||||||
|
# anything.
|
||||||
|
try:
|
||||||
|
availability = config.get_other_models_availability()
|
||||||
|
response_data["other_models_paths_available"] = bool(
|
||||||
|
availability.get("available")
|
||||||
|
)
|
||||||
|
except Exception as availability_error: # pragma: no cover - defensive
|
||||||
|
logger.debug(
|
||||||
|
"Could not resolve Other Models availability: %s",
|
||||||
|
availability_error,
|
||||||
|
)
|
||||||
|
response_data["other_models_paths_available"] = None
|
||||||
|
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"
|
||||||
|
response_data["standalone_mode"] = standalone_mode
|
||||||
|
if standalone_mode:
|
||||||
|
# Standalone reads its model roots exclusively from
|
||||||
|
# settings.json, so the Model Paths settings UI needs the
|
||||||
|
# current values plus the editable-key schema. In plugin mode
|
||||||
|
# the paths come from the ComfyUI host and stay hidden.
|
||||||
|
folder_paths = self._settings.get("folder_paths") or {}
|
||||||
|
# A fresh install is seeded from settings.json.example, whose
|
||||||
|
# folder_paths are documentation placeholders — hide them so
|
||||||
|
# the UI starts with empty editors instead of fake paths.
|
||||||
|
get_placeholders = getattr(
|
||||||
|
self._settings, "get_template_folder_path_placeholders", None
|
||||||
|
)
|
||||||
|
placeholders = get_placeholders() if get_placeholders else set()
|
||||||
|
if placeholders:
|
||||||
|
folder_paths = {
|
||||||
|
key: [p for p in paths if p not in placeholders]
|
||||||
|
if isinstance(paths, list)
|
||||||
|
else paths
|
||||||
|
for key, paths in folder_paths.items()
|
||||||
|
}
|
||||||
|
response_data["folder_paths"] = folder_paths
|
||||||
|
response_data["folder_path_schema"] = folder_path_schema()
|
||||||
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
|
||||||
@@ -2065,6 +2135,7 @@ class ServiceRegistryAdapter:
|
|||||||
get_embedding_scanner: Callable[[], Awaitable[Any]]
|
get_embedding_scanner: Callable[[], Awaitable[Any]]
|
||||||
get_downloaded_version_history_service: Callable[[], Awaitable[Any]]
|
get_downloaded_version_history_service: Callable[[], Awaitable[Any]]
|
||||||
get_backup_service: Callable[[], Awaitable[Any]] = _noop_backup_service
|
get_backup_service: Callable[[], Awaitable[Any]] = _noop_backup_service
|
||||||
|
get_other_scanner: Callable[[], Awaitable[Any]] = ServiceRegistry.get_other_scanner
|
||||||
|
|
||||||
|
|
||||||
class ModelLibraryHandler:
|
class ModelLibraryHandler:
|
||||||
@@ -2089,6 +2160,8 @@ class ModelLibraryHandler:
|
|||||||
return "checkpoint"
|
return "checkpoint"
|
||||||
if normalized in {"embedding", "textualinversion"}:
|
if normalized in {"embedding", "textualinversion"}:
|
||||||
return "embedding"
|
return "embedding"
|
||||||
|
if normalized in VALID_OTHER_CIVITAI_TYPES:
|
||||||
|
return "other"
|
||||||
return None
|
return None
|
||||||
|
|
||||||
async def _get_scanner_for_type(self, model_type: str | None):
|
async def _get_scanner_for_type(self, model_type: str | None):
|
||||||
@@ -2099,6 +2172,13 @@ class ModelLibraryHandler:
|
|||||||
return normalized_type, await self._service_registry.get_checkpoint_scanner()
|
return normalized_type, await self._service_registry.get_checkpoint_scanner()
|
||||||
if normalized_type == "embedding":
|
if normalized_type == "embedding":
|
||||||
return normalized_type, await self._service_registry.get_embedding_scanner()
|
return normalized_type, await self._service_registry.get_embedding_scanner()
|
||||||
|
if normalized_type == "other":
|
||||||
|
# Opt-in feature: the other scanner only resolves while the master
|
||||||
|
# switch is on, so callers keep returning the legacy "required"
|
||||||
|
# error (400) when it is off.
|
||||||
|
if not get_settings_manager().is_other_models_enabled():
|
||||||
|
return None, None
|
||||||
|
return normalized_type, await self._service_registry.get_other_scanner()
|
||||||
return None, None
|
return None, None
|
||||||
|
|
||||||
async def _get_download_history_service(self):
|
async def _get_download_history_service(self):
|
||||||
@@ -2190,6 +2270,11 @@ class ModelLibraryHandler:
|
|||||||
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()
|
||||||
embedding_scanner = await self._service_registry.get_embedding_scanner()
|
embedding_scanner = await self._service_registry.get_embedding_scanner()
|
||||||
|
# Opt-in: probe the other scanner only while Other Models is enabled,
|
||||||
|
# so the disabled behaviour stays byte-identical to the legacy one.
|
||||||
|
other_scanner = None
|
||||||
|
if get_settings_manager().is_other_models_enabled():
|
||||||
|
other_scanner = await self._service_registry.get_other_scanner()
|
||||||
|
|
||||||
if model_version_id_str:
|
if model_version_id_str:
|
||||||
try:
|
try:
|
||||||
@@ -2228,6 +2313,13 @@ class ModelLibraryHandler:
|
|||||||
exists = True
|
exists = True
|
||||||
model_type = "embedding"
|
model_type = "embedding"
|
||||||
matched_scanner = embedding_scanner
|
matched_scanner = embedding_scanner
|
||||||
|
elif (
|
||||||
|
other_scanner
|
||||||
|
and await other_scanner.check_model_version_exists(model_version_id)
|
||||||
|
):
|
||||||
|
exists = True
|
||||||
|
model_type = "other"
|
||||||
|
matched_scanner = other_scanner
|
||||||
|
|
||||||
if exists:
|
if exists:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -2245,7 +2337,7 @@ class ModelLibraryHandler:
|
|||||||
history_service = await self._get_download_history_service()
|
history_service = await self._get_download_history_service()
|
||||||
has_been_downloaded = False
|
has_been_downloaded = False
|
||||||
history_type = None
|
history_type = None
|
||||||
for candidate_type in ("lora", "checkpoint", "embedding"):
|
for candidate_type in ("lora", "checkpoint", "embedding", "other"):
|
||||||
if await history_service.has_been_downloaded(
|
if await history_service.has_been_downloaded(
|
||||||
candidate_type,
|
candidate_type,
|
||||||
model_version_id,
|
model_version_id,
|
||||||
@@ -2267,6 +2359,7 @@ class ModelLibraryHandler:
|
|||||||
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
|
lora_versions = await lora_scanner.get_model_versions_by_id(model_id)
|
||||||
checkpoint_versions = []
|
checkpoint_versions = []
|
||||||
embedding_versions = []
|
embedding_versions = []
|
||||||
|
other_versions = []
|
||||||
if not lora_versions and checkpoint_scanner:
|
if not lora_versions and checkpoint_scanner:
|
||||||
checkpoint_versions = await checkpoint_scanner.get_model_versions_by_id(
|
checkpoint_versions = await checkpoint_scanner.get_model_versions_by_id(
|
||||||
model_id
|
model_id
|
||||||
@@ -2275,6 +2368,13 @@ class ModelLibraryHandler:
|
|||||||
embedding_versions = await embedding_scanner.get_model_versions_by_id(
|
embedding_versions = await embedding_scanner.get_model_versions_by_id(
|
||||||
model_id
|
model_id
|
||||||
)
|
)
|
||||||
|
if (
|
||||||
|
not lora_versions
|
||||||
|
and not checkpoint_versions
|
||||||
|
and not embedding_versions
|
||||||
|
and other_scanner
|
||||||
|
):
|
||||||
|
other_versions = await other_scanner.get_model_versions_by_id(model_id)
|
||||||
|
|
||||||
model_type = None
|
model_type = None
|
||||||
versions = []
|
versions = []
|
||||||
@@ -2306,9 +2406,18 @@ class ModelLibraryHandler:
|
|||||||
"downloadedVersionIds": [],
|
"downloadedVersionIds": [],
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
if other_versions:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": True,
|
||||||
|
"modelType": "other",
|
||||||
|
"versions": self._with_downloaded_flag(other_versions),
|
||||||
|
"downloadedVersionIds": [],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
history_service = await self._get_download_history_service()
|
history_service = await self._get_download_history_service()
|
||||||
for candidate_type in ("lora", "checkpoint", "embedding"):
|
for candidate_type in ("lora", "checkpoint", "embedding", "other"):
|
||||||
candidate_downloaded_version_ids = (
|
candidate_downloaded_version_ids = (
|
||||||
await history_service.get_downloaded_version_ids(
|
await history_service.get_downloaded_version_ids(
|
||||||
candidate_type,
|
candidate_type,
|
||||||
@@ -2363,6 +2472,11 @@ class ModelLibraryHandler:
|
|||||||
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()
|
||||||
embedding_scanner = await self._service_registry.get_embedding_scanner()
|
embedding_scanner = await self._service_registry.get_embedding_scanner()
|
||||||
|
# Opt-in: keep the other probe last so model cards for lora /
|
||||||
|
# checkpoint / embedding ids are unaffected by the extra scanner.
|
||||||
|
other_scanner = None
|
||||||
|
if get_settings_manager().is_other_models_enabled():
|
||||||
|
other_scanner = await self._service_registry.get_other_scanner()
|
||||||
|
|
||||||
results: list[dict[str, Any]] = []
|
results: list[dict[str, Any]] = []
|
||||||
for model_id in model_ids:
|
for model_id in model_ids:
|
||||||
@@ -2398,6 +2512,17 @@ class ModelLibraryHandler:
|
|||||||
})
|
})
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
if other_scanner:
|
||||||
|
other_versions = await other_scanner.get_model_versions_by_id(model_id)
|
||||||
|
if other_versions:
|
||||||
|
results.append({
|
||||||
|
"modelId": model_id,
|
||||||
|
"modelType": "other",
|
||||||
|
"versions": self._with_downloaded_flag(other_versions),
|
||||||
|
"downloadedVersionIds": [],
|
||||||
|
})
|
||||||
|
continue
|
||||||
|
|
||||||
results.append({
|
results.append({
|
||||||
"modelId": model_id,
|
"modelId": model_id,
|
||||||
"modelType": None,
|
"modelType": None,
|
||||||
@@ -2665,12 +2790,40 @@ class ModelLibraryHandler:
|
|||||||
|
|
||||||
normalized_type, scanner = await self._get_scanner_for_type(model_type)
|
normalized_type, scanner = await self._get_scanner_for_type(model_type)
|
||||||
if not normalized_type:
|
if not normalized_type:
|
||||||
|
# The lookup cannot be served as a fully interactive list. Two
|
||||||
|
# cases share this branch: a CivitAI type with no scanner at all
|
||||||
|
# (Wildcards, Workflows, Hypernetwork, Poses, AestheticGradient)
|
||||||
|
# and an Other-model type while the opt-in master switch is off.
|
||||||
|
# Answer 200 with the CivitAI list marked read-only plus a
|
||||||
|
# machine-readable reason, so clients can still show the
|
||||||
|
# versions and explain why the actions are missing. Legacy
|
||||||
|
# clients keep working: they only read `success`/`versions`.
|
||||||
|
reason = (
|
||||||
|
"other_models_disabled"
|
||||||
|
if self._normalize_model_type(model_type) == "other"
|
||||||
|
else "model_type_unsupported"
|
||||||
|
)
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
"success": False,
|
"success": True,
|
||||||
"error": f'Model type "{model_type}" is not supported',
|
"modelId": model_id,
|
||||||
},
|
"modelName": model_name,
|
||||||
status=400,
|
"modelType": model_type,
|
||||||
|
"supported": False,
|
||||||
|
"reason": reason,
|
||||||
|
"versions": [
|
||||||
|
{
|
||||||
|
"id": version.get("id"),
|
||||||
|
"name": version.get("name", ""),
|
||||||
|
"thumbnailUrl": version.get("images")[0]["url"]
|
||||||
|
if version.get("images")
|
||||||
|
else None,
|
||||||
|
"inLibrary": False,
|
||||||
|
"hasBeenDownloaded": False,
|
||||||
|
}
|
||||||
|
for version in versions
|
||||||
|
],
|
||||||
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
if not scanner:
|
if not scanner:
|
||||||
@@ -2712,6 +2865,7 @@ class ModelLibraryHandler:
|
|||||||
"modelId": model_id,
|
"modelId": model_id,
|
||||||
"modelName": model_name,
|
"modelName": model_name,
|
||||||
"modelType": model_type,
|
"modelType": model_type,
|
||||||
|
"supported": True,
|
||||||
"versions": enriched_versions,
|
"versions": enriched_versions,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
@@ -2786,12 +2940,32 @@ class ModelLibraryHandler:
|
|||||||
model_type.lower() for model_type in CIVITAI_USER_MODEL_TYPES
|
model_type.lower() for model_type in CIVITAI_USER_MODEL_TYPES
|
||||||
}
|
}
|
||||||
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}
|
||||||
|
other_type_aliases = {
|
||||||
|
model_type.lower() for model_type in VALID_OTHER_CIVITAI_TYPES
|
||||||
|
}
|
||||||
|
|
||||||
|
# Acquire the other scanner lazily so adapters without it only
|
||||||
|
# fail when the payload actually contains other-type models.
|
||||||
|
# While the opt-in feature is off the scanner still exists (its
|
||||||
|
# cache is empty), so other types simply report inLibrary=False.
|
||||||
|
needs_other_scanner = any(
|
||||||
|
isinstance(model, dict)
|
||||||
|
and str(model.get("type", "")).lower() in other_type_aliases
|
||||||
|
for model in models
|
||||||
|
)
|
||||||
|
other_scanner = None
|
||||||
|
if needs_other_scanner:
|
||||||
|
other_scanner = await self._service_registry.get_other_scanner()
|
||||||
|
|
||||||
type_scanner_map: Dict[str, Any] = {
|
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,
|
||||||
}
|
}
|
||||||
|
if other_scanner is not None:
|
||||||
|
type_scanner_map.update(
|
||||||
|
{alias: other_scanner for alias in other_type_aliases}
|
||||||
|
)
|
||||||
|
|
||||||
versions: list[dict[str, Any]] = []
|
versions: list[dict[str, Any]] = []
|
||||||
history_service = await self._get_download_history_service()
|
history_service = await self._get_download_history_service()
|
||||||
@@ -2815,12 +2989,17 @@ class ModelLibraryHandler:
|
|||||||
"embedding",
|
"embedding",
|
||||||
model_ids,
|
model_ids,
|
||||||
)
|
)
|
||||||
|
other_downloaded = await history_service.get_downloaded_version_ids_bulk(
|
||||||
|
"other",
|
||||||
|
model_ids,
|
||||||
|
)
|
||||||
downloaded_version_map: Dict[str, Dict[int, set[int]]] = {
|
downloaded_version_map: Dict[str, Dict[int, set[int]]] = {
|
||||||
"lora": lora_downloaded,
|
"lora": lora_downloaded,
|
||||||
"locon": lora_downloaded,
|
"locon": lora_downloaded,
|
||||||
"dora": lora_downloaded,
|
"dora": lora_downloaded,
|
||||||
"checkpoint": checkpoint_downloaded,
|
"checkpoint": checkpoint_downloaded,
|
||||||
"textualinversion": embedding_downloaded,
|
"textualinversion": embedding_downloaded,
|
||||||
|
**{alias: other_downloaded for alias in VALID_OTHER_CIVITAI_TYPES},
|
||||||
}
|
}
|
||||||
for model in models:
|
for model in models:
|
||||||
if not isinstance(model, dict):
|
if not isinstance(model, dict):
|
||||||
@@ -3274,6 +3453,18 @@ class FileSystemHandler:
|
|||||||
subprocess.Popen(["open", "-R", settings_file])
|
subprocess.Popen(["open", "-R", settings_file])
|
||||||
else:
|
else:
|
||||||
folder = os.path.dirname(settings_file)
|
folder = os.path.dirname(settings_file)
|
||||||
|
if not _has_gui_display():
|
||||||
|
# Headless/SSH session: xdg-open cannot open a file
|
||||||
|
# manager, so hand the path to the browser for copying
|
||||||
|
# instead of reporting a success that never happened.
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": True,
|
||||||
|
"message": "Headless session: path available for copying",
|
||||||
|
"path": settings_file,
|
||||||
|
"mode": "clipboard",
|
||||||
|
}
|
||||||
|
)
|
||||||
subprocess.Popen(["xdg-open", folder])
|
subprocess.Popen(["xdg-open", folder])
|
||||||
|
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -3307,6 +3498,76 @@ class FileSystemHandler:
|
|||||||
logger.error("Failed to open wildcards location: %s", exc, exc_info=True)
|
logger.error("Failed to open wildcards location: %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 browse_directory(self, request: web.Request) -> web.Response:
|
||||||
|
"""Browse a directory for the settings-UI directory picker."""
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
payload, status = browse_directory(data.get("path", ""))
|
||||||
|
return web.json_response(payload, status=status)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON"}, status=400
|
||||||
|
)
|
||||||
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
|
logger.error("Failed to browse directory: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def validate_path(self, request: web.Request) -> web.Response:
|
||||||
|
"""Validate a filesystem path for the settings UI.
|
||||||
|
|
||||||
|
A well-formed request always returns HTTP 200; invalid paths are
|
||||||
|
reported via ``error_code`` in the payload. HTTP 400 is reserved for
|
||||||
|
malformed requests (missing path, invalid JSON).
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
raw_path = data.get("path")
|
||||||
|
expect = data.get("expect", "directory")
|
||||||
|
|
||||||
|
if not raw_path or not isinstance(raw_path, str):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Missing path parameter"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
# Business path convention: abspath only, never realpath.
|
||||||
|
path = os.path.abspath(os.path.expanduser(raw_path))
|
||||||
|
|
||||||
|
exists = os.path.exists(path)
|
||||||
|
is_directory = os.path.isdir(path) if exists else False
|
||||||
|
readable = bool(exists and os.access(path, os.R_OK))
|
||||||
|
writable = bool(exists and os.access(path, os.W_OK))
|
||||||
|
|
||||||
|
error_code = None
|
||||||
|
if not exists:
|
||||||
|
error_code = "path_not_found"
|
||||||
|
elif expect == "directory" and not is_directory:
|
||||||
|
error_code = "not_a_directory"
|
||||||
|
elif expect == "file" and not os.path.isfile(path):
|
||||||
|
error_code = "not_a_file"
|
||||||
|
elif not readable:
|
||||||
|
error_code = "not_readable"
|
||||||
|
elif not writable:
|
||||||
|
error_code = "not_writable"
|
||||||
|
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": True,
|
||||||
|
"path": path,
|
||||||
|
"exists": exists,
|
||||||
|
"is_directory": is_directory,
|
||||||
|
"readable": readable,
|
||||||
|
"writable": writable,
|
||||||
|
"error_code": error_code,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON"}, status=400
|
||||||
|
)
|
||||||
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
|
logger.error("Failed to validate path: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
|
||||||
class CustomWordsHandler:
|
class CustomWordsHandler:
|
||||||
"""Handler for autocomplete via TagFTSIndex."""
|
"""Handler for autocomplete via TagFTSIndex."""
|
||||||
@@ -3882,8 +4143,9 @@ class MiscHandlerSet:
|
|||||||
doctor: DoctorHandler,
|
doctor: DoctorHandler,
|
||||||
example_workflows: ExampleWorkflowsHandler,
|
example_workflows: ExampleWorkflowsHandler,
|
||||||
base_model: BaseModelHandlerSet,
|
base_model: BaseModelHandlerSet,
|
||||||
hf_handler: Any = None,
|
model_source_handler: Any = None,
|
||||||
agent_handler: Any = None,
|
agent_handler: Any = None,
|
||||||
|
download_routing: Any = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.health = health
|
self.health = health
|
||||||
self.settings = settings
|
self.settings = settings
|
||||||
@@ -3902,8 +4164,9 @@ class MiscHandlerSet:
|
|||||||
self.doctor = doctor
|
self.doctor = doctor
|
||||||
self.example_workflows = example_workflows
|
self.example_workflows = example_workflows
|
||||||
self.base_model = base_model
|
self.base_model = base_model
|
||||||
self.hf_handler = hf_handler
|
self.model_source_handler = model_source_handler
|
||||||
self.agent_handler = agent_handler
|
self.agent_handler = agent_handler
|
||||||
|
self.download_routing = download_routing
|
||||||
|
|
||||||
def to_route_mapping(
|
def to_route_mapping(
|
||||||
self,
|
self,
|
||||||
@@ -3949,19 +4212,27 @@ class MiscHandlerSet:
|
|||||||
"open_settings_location": self.filesystem.open_settings_location,
|
"open_settings_location": self.filesystem.open_settings_location,
|
||||||
"open_backup_location": self.filesystem.open_backup_location,
|
"open_backup_location": self.filesystem.open_backup_location,
|
||||||
"open_wildcards_location": self.filesystem.open_wildcards_location,
|
"open_wildcards_location": self.filesystem.open_wildcards_location,
|
||||||
|
"browse_directory": self.filesystem.browse_directory,
|
||||||
|
"validate_path": self.filesystem.validate_path,
|
||||||
"search_custom_words": self.custom_words.search_custom_words,
|
"search_custom_words": self.custom_words.search_custom_words,
|
||||||
"search_wildcards": self.wildcards.search_wildcards,
|
"search_wildcards": self.wildcards.search_wildcards,
|
||||||
"get_supporters": self.supporters.get_supporters,
|
"get_supporters": self.supporters.get_supporters,
|
||||||
"get_example_workflows": self.example_workflows.get_example_workflows,
|
"get_example_workflows": self.example_workflows.get_example_workflows,
|
||||||
"get_example_workflow": self.example_workflows.get_example_workflow,
|
"get_example_workflow": self.example_workflows.get_example_workflow,
|
||||||
# Hugging Face handlers
|
# Hugging Face handlers
|
||||||
"get_hf_repo_files": self.hf_handler.get_hf_repo_files,
|
# External model sources (Hugging Face / ModelScope)
|
||||||
"download_hf_model": self.hf_handler.download_hf_model,
|
"list_model_source_files": self.model_source_handler.list_model_source_files,
|
||||||
"set_hf_url": self.hf_handler.set_hf_url,
|
"download_model_source": self.model_source_handler.download_model_source,
|
||||||
|
"get_hf_repo_files": self.model_source_handler.list_model_source_files,
|
||||||
|
"download_hf_model": self.model_source_handler.download_model_source,
|
||||||
|
"set_hf_url": self.model_source_handler.set_hf_url,
|
||||||
|
"get_model_sources": self.model_source_handler.get_model_sources,
|
||||||
# Agent skill handlers
|
# Agent skill handlers
|
||||||
"get_agent_skills": self.agent_handler.get_agent_skills,
|
"get_agent_skills": self.agent_handler.get_agent_skills,
|
||||||
"execute_agent_skill": self.agent_handler.execute_agent_skill,
|
"execute_agent_skill": self.agent_handler.execute_agent_skill,
|
||||||
"cancel_agent_skill": self.agent_handler.cancel_agent_skill,
|
"cancel_agent_skill": self.agent_handler.cancel_agent_skill,
|
||||||
|
# Download routing handler
|
||||||
|
"get_download_routing": self.download_routing.get_download_routing,
|
||||||
# Base model handlers
|
# Base model handlers
|
||||||
"get_base_models": self.base_model.get_base_models,
|
"get_base_models": self.base_model.get_base_models,
|
||||||
"refresh_base_models": self.base_model.refresh_base_models,
|
"refresh_base_models": self.base_model.refresh_base_models,
|
||||||
@@ -3975,6 +4246,7 @@ def build_service_registry_adapter() -> ServiceRegistryAdapter:
|
|||||||
get_lora_scanner=ServiceRegistry.get_lora_scanner,
|
get_lora_scanner=ServiceRegistry.get_lora_scanner,
|
||||||
get_checkpoint_scanner=ServiceRegistry.get_checkpoint_scanner,
|
get_checkpoint_scanner=ServiceRegistry.get_checkpoint_scanner,
|
||||||
get_embedding_scanner=ServiceRegistry.get_embedding_scanner,
|
get_embedding_scanner=ServiceRegistry.get_embedding_scanner,
|
||||||
|
get_other_scanner=ServiceRegistry.get_other_scanner,
|
||||||
get_downloaded_version_history_service=ServiceRegistry.get_downloaded_version_history_service,
|
get_downloaded_version_history_service=ServiceRegistry.get_downloaded_version_history_service,
|
||||||
get_backup_service=ServiceRegistry.get_backup_service,
|
get_backup_service=ServiceRegistry.get_backup_service,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -15,6 +15,10 @@ from aiohttp import web
|
|||||||
import jinja2
|
import jinja2
|
||||||
|
|
||||||
from ...config import config
|
from ...config import config
|
||||||
|
from ...services.active_filters_store import (
|
||||||
|
ActiveFiltersStore,
|
||||||
|
active_filters_to_query_kwargs,
|
||||||
|
)
|
||||||
from ...services.download_coordinator import DownloadCoordinator
|
from ...services.download_coordinator import DownloadCoordinator
|
||||||
from ...services.connectivity_guard import (
|
from ...services.connectivity_guard import (
|
||||||
OFFLINE_FRIENDLY_MESSAGE,
|
OFFLINE_FRIENDLY_MESSAGE,
|
||||||
@@ -33,10 +37,14 @@ from ...services.use_cases import (
|
|||||||
DownloadModelEarlyAccessError,
|
DownloadModelEarlyAccessError,
|
||||||
DownloadModelUseCase,
|
DownloadModelUseCase,
|
||||||
DownloadModelValidationError,
|
DownloadModelValidationError,
|
||||||
|
FilenameTemplateUseCase,
|
||||||
MetadataRefreshProgressReporter,
|
MetadataRefreshProgressReporter,
|
||||||
)
|
)
|
||||||
from ...services.websocket_manager import WebSocketManager
|
from ...services.websocket_manager import WebSocketManager
|
||||||
from ...services.websocket_progress_callback import WebSocketProgressCallback
|
from ...services.websocket_progress_callback import (
|
||||||
|
WebSocketFilenameTemplateProgressCallback,
|
||||||
|
WebSocketProgressCallback,
|
||||||
|
)
|
||||||
from ...services.download_queue_service import DownloadQueueService
|
from ...services.download_queue_service import DownloadQueueService
|
||||||
from ...services.errors import RateLimitError, ResourceNotFoundError
|
from ...services.errors import RateLimitError, ResourceNotFoundError
|
||||||
from ...utils.civitai_utils import resolve_license_payload
|
from ...utils.civitai_utils import resolve_license_payload
|
||||||
@@ -86,6 +94,7 @@ class ModelPageView:
|
|||||||
settings_service: SettingsManager,
|
settings_service: SettingsManager,
|
||||||
server_i18n,
|
server_i18n,
|
||||||
logger: logging.Logger,
|
logger: logging.Logger,
|
||||||
|
page_context_provider: Callable[[web.Request], Dict[str, Any]] | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._template_env = template_env
|
self._template_env = template_env
|
||||||
self._template_name = template_name
|
self._template_name = template_name
|
||||||
@@ -93,6 +102,7 @@ class ModelPageView:
|
|||||||
self._settings = settings_service
|
self._settings = settings_service
|
||||||
self._server_i18n = server_i18n
|
self._server_i18n = server_i18n
|
||||||
self._logger = logger
|
self._logger = logger
|
||||||
|
self._page_context_provider = page_context_provider
|
||||||
|
|
||||||
def _load_supporters(self) -> dict[str, Any]:
|
def _load_supporters(self) -> dict[str, Any]:
|
||||||
"""Load supporters data from JSON file."""
|
"""Load supporters data from JSON file."""
|
||||||
@@ -206,6 +216,16 @@ class ModelPageView:
|
|||||||
self._logger.error("Error loading cache data: %s", cache_error)
|
self._logger.error("Error loading cache data: %s", cache_error)
|
||||||
template_context["is_initializing"] = True
|
template_context["is_initializing"] = True
|
||||||
|
|
||||||
|
if self._page_context_provider is not None:
|
||||||
|
try:
|
||||||
|
extra_context = self._page_context_provider(request)
|
||||||
|
if isinstance(extra_context, dict):
|
||||||
|
template_context.update(extra_context)
|
||||||
|
except Exception as context_error: # pragma: no cover - logging path
|
||||||
|
self._logger.error(
|
||||||
|
"Error building page context: %s", context_error
|
||||||
|
)
|
||||||
|
|
||||||
rendered = self._template_env.get_template(self._template_name).render(
|
rendered = self._template_env.get_template(self._template_name).render(
|
||||||
**template_context
|
**template_context
|
||||||
)
|
)
|
||||||
@@ -634,6 +654,16 @@ class ModelManagementHandler:
|
|||||||
file_path = data.get("file_path")
|
file_path = data.get("file_path")
|
||||||
model_id = data.get("model_id")
|
model_id = data.get("model_id")
|
||||||
model_version_id = data.get("model_version_id")
|
model_version_id = data.get("model_version_id")
|
||||||
|
source = data.get("source")
|
||||||
|
|
||||||
|
if source not in (None, "", "civarchive"):
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": f"Unsupported relink source: {source}",
|
||||||
|
},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
if not file_path or model_id is None:
|
if not file_path or model_id is None:
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -649,20 +679,33 @@ class ModelManagementHandler:
|
|||||||
metadata_path
|
metadata_path
|
||||||
)
|
)
|
||||||
|
|
||||||
|
relink_kwargs = {
|
||||||
|
"file_path": file_path,
|
||||||
|
"metadata": local_metadata,
|
||||||
|
"model_id": int(model_id),
|
||||||
|
"model_version_id": int(model_version_id) if model_version_id else None,
|
||||||
|
}
|
||||||
|
if source == "civarchive":
|
||||||
|
relink_kwargs["provider_name"] = "civarchive_api"
|
||||||
|
|
||||||
updated_metadata = await self._metadata_sync.relink_metadata(
|
updated_metadata = await self._metadata_sync.relink_metadata(
|
||||||
file_path=file_path,
|
**relink_kwargs
|
||||||
metadata=local_metadata,
|
|
||||||
model_id=int(model_id),
|
|
||||||
model_version_id=int(model_version_id) if model_version_id else None,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
await self._service.scanner.update_single_model_cache(
|
await self._service.scanner.update_single_model_cache(
|
||||||
file_path, file_path, updated_metadata
|
file_path, file_path, updated_metadata
|
||||||
)
|
)
|
||||||
|
|
||||||
message = f"Model successfully re-linked to Civitai model {model_id}" + (
|
if source == "civarchive":
|
||||||
f" version {model_version_id}" if model_version_id else ""
|
message = (
|
||||||
)
|
f"Model successfully re-linked to CivArchive model {model_id}"
|
||||||
|
+ (f" version {model_version_id}" if model_version_id else "")
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
message = (
|
||||||
|
f"Model successfully re-linked to Civitai model {model_id}"
|
||||||
|
+ (f" version {model_version_id}" if model_version_id else "")
|
||||||
|
)
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{
|
{
|
||||||
"success": True,
|
"success": True,
|
||||||
@@ -670,6 +713,8 @@ class ModelManagementHandler:
|
|||||||
"hash": updated_metadata.get("sha256", ""),
|
"hash": updated_metadata.get("sha256", ""),
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
except ValueError as exc:
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=400)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
if is_expected_offline_error(str(exc)):
|
if is_expected_offline_error(str(exc)):
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
@@ -1570,12 +1615,50 @@ class ModelQueryHandler:
|
|||||||
allow_selling_generated_content.lower() not in ("false", "0", "")
|
allow_selling_generated_content.lower() not in ("false", "0", "")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# When requested, merge the manager page's active filters stored
|
||||||
|
# server-side. Explicit query parameters take precedence over the
|
||||||
|
# stored values.
|
||||||
|
use_active_filters = (
|
||||||
|
request.query.get("use_active_filters", "").lower() in ("1", "true")
|
||||||
|
)
|
||||||
|
if use_active_filters:
|
||||||
|
stored = ActiveFiltersStore.get_instance().get_filters(
|
||||||
|
self._service.model_type
|
||||||
|
)
|
||||||
|
injected = active_filters_to_query_kwargs(stored)
|
||||||
|
if folder is None and "folder" in injected:
|
||||||
|
folder = injected["folder"]
|
||||||
|
if "recursive" not in request.query and "recursive" in injected:
|
||||||
|
recursive = injected["recursive"]
|
||||||
|
if not base_models and injected.get("base_models"):
|
||||||
|
base_models = injected["base_models"]
|
||||||
|
if not model_types and injected.get("model_types"):
|
||||||
|
model_types = injected["model_types"]
|
||||||
|
if not tag_filters and injected.get("tags"):
|
||||||
|
tag_filters = injected["tags"]
|
||||||
|
if not auto_tag_filters and injected.get("auto_tags"):
|
||||||
|
auto_tag_filters = injected["auto_tags"]
|
||||||
|
if "tag_logic" not in request.query and injected.get("tag_logic"):
|
||||||
|
injected_logic = str(injected["tag_logic"]).lower()
|
||||||
|
if injected_logic in ("any", "all"):
|
||||||
|
tag_logic = injected_logic
|
||||||
|
if credit_required is None and "credit_required" in injected:
|
||||||
|
credit_required = injected["credit_required"]
|
||||||
|
if (
|
||||||
|
allow_selling_generated_content is None
|
||||||
|
and "allow_selling_generated_content" in injected
|
||||||
|
):
|
||||||
|
allow_selling_generated_content = injected[
|
||||||
|
"allow_selling_generated_content"
|
||||||
|
]
|
||||||
|
|
||||||
# The presence of the recursive param (always sent by the loras
|
# The presence of the recursive param (always sent by the loras
|
||||||
# widget when filter mode is on) signals that the filter pipeline
|
# widget when filter mode is on) signals that the filter pipeline
|
||||||
# must run even when no concrete filter is set, so global settings
|
# must run even when no concrete filter is set, so global settings
|
||||||
# like show_only_sfw stay consistent with the list endpoint.
|
# like show_only_sfw stay consistent with the list endpoint.
|
||||||
apply_filters = (
|
apply_filters = (
|
||||||
"recursive" in request.query
|
use_active_filters
|
||||||
|
or "recursive" in request.query
|
||||||
or folder is not None
|
or folder is not None
|
||||||
or bool(base_models)
|
or bool(base_models)
|
||||||
or bool(model_types)
|
or bool(model_types)
|
||||||
@@ -1609,6 +1692,50 @@ class ModelQueryHandler:
|
|||||||
)
|
)
|
||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def update_active_filters(self, request: web.Request) -> web.Response:
|
||||||
|
"""Store the manager page's active filters for this model type."""
|
||||||
|
try:
|
||||||
|
payload = await request.json()
|
||||||
|
except Exception:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Body must be a JSON object"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
ActiveFiltersStore.get_instance().set_filters(
|
||||||
|
self._service.model_type, payload
|
||||||
|
)
|
||||||
|
return web.json_response({"success": True})
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error(
|
||||||
|
"Error updating active filters for %s: %s",
|
||||||
|
self._service.model_type,
|
||||||
|
exc,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def get_active_filters(self, request: web.Request) -> web.Response:
|
||||||
|
"""Return the stored active filters for this model type."""
|
||||||
|
try:
|
||||||
|
filters = ActiveFiltersStore.get_instance().get_filters(
|
||||||
|
self._service.model_type
|
||||||
|
)
|
||||||
|
return web.json_response({"success": True, "filters": filters})
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error(
|
||||||
|
"Error getting active filters for %s: %s",
|
||||||
|
self._service.model_type,
|
||||||
|
exc,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
|
||||||
class ModelDownloadHandler:
|
class ModelDownloadHandler:
|
||||||
"""Coordinate downloads and progress reporting."""
|
"""Coordinate downloads and progress reporting."""
|
||||||
@@ -1787,6 +1914,11 @@ class ModelDownloadHandler:
|
|||||||
response_payload["status"] = status
|
response_payload["status"] = status
|
||||||
if "message" in progress_data:
|
if "message" in progress_data:
|
||||||
response_payload["message"] = progress_data["message"]
|
response_payload["message"] = progress_data["message"]
|
||||||
|
# Post-transfer stage (indexing / source metadata); polling
|
||||||
|
# consumers need it to tell "working" from "stuck".
|
||||||
|
for field in ("stage", "platform"):
|
||||||
|
if field in progress_data:
|
||||||
|
response_payload[field] = progress_data[field]
|
||||||
elif status is None and "message" in progress_data:
|
elif status is None and "message" in progress_data:
|
||||||
response_payload["message"] = progress_data["message"]
|
response_payload["message"] = progress_data["message"]
|
||||||
|
|
||||||
@@ -1904,8 +2036,18 @@ class ModelDownloadHandler:
|
|||||||
try:
|
try:
|
||||||
status_filter = request.query.get("status") or None
|
status_filter = request.query.get("status") or None
|
||||||
service = await DownloadQueueService.get_instance()
|
service = await DownloadQueueService.get_instance()
|
||||||
cleared = await service.clear_queue(status_filter=status_filter)
|
cleared_ids = await service.clear_queue(status_filter=status_filter)
|
||||||
return web.json_response({"success": True, "cleared": cleared})
|
# Clearing the queue rows alone would orphan any in-memory tasks
|
||||||
|
# and persisted aria2 state for those downloads, leaving them
|
||||||
|
# polling the daemon invisibly. Tear that tracking down too.
|
||||||
|
try:
|
||||||
|
await self._download_coordinator.discard_cleared_downloads(cleared_ids)
|
||||||
|
except Exception:
|
||||||
|
self._logger.warning(
|
||||||
|
"Failed to discard in-memory state for cleared downloads",
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
return web.json_response({"success": True, "cleared": len(cleared_ids)})
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
self._logger.error(
|
self._logger.error(
|
||||||
"Error clearing download queue: %s", exc, exc_info=True
|
"Error clearing download queue: %s", exc, exc_info=True
|
||||||
@@ -1988,9 +2130,11 @@ class ModelDownloadHandler:
|
|||||||
item_id=item_id, download_id=download_id
|
item_id=item_id, download_id=download_id
|
||||||
)
|
)
|
||||||
if item is None:
|
if item is None:
|
||||||
|
# Missing or non-retryable history entry is a business
|
||||||
|
# outcome, not a routing error: 200 lets the extension's
|
||||||
|
# apiFetch 404-fallback and error middleware stay quiet.
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "History item not found or not retryable"},
|
{"success": False, "error": "History item not found or not retryable"}
|
||||||
status=404,
|
|
||||||
)
|
)
|
||||||
return web.json_response({"success": True, "item": item})
|
return web.json_response({"success": True, "item": item})
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -2041,8 +2185,12 @@ class ModelDownloadHandler:
|
|||||||
completed_at=completed_at,
|
completed_at=completed_at,
|
||||||
)
|
)
|
||||||
if item is None:
|
if item is None:
|
||||||
|
# A missing queue item (already completed, or never queued) is
|
||||||
|
# a normal business outcome, not a routing error. Return 200
|
||||||
|
# so the browser extension's apiFetch 404-fallback and the
|
||||||
|
# error middleware stay quiet.
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "Download not found in queue"}, status=404
|
{"success": False, "error": "Download not found in queue"}
|
||||||
)
|
)
|
||||||
return web.json_response({"success": True, "item": item})
|
return web.json_response({"success": True, "item": item})
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -2084,9 +2232,10 @@ class ModelDownloadHandler:
|
|||||||
service = await DownloadQueueService.get_instance()
|
service = await DownloadQueueService.get_instance()
|
||||||
updated = await service.update_status(download_id, status)
|
updated = await service.update_status(download_id, status)
|
||||||
if not updated:
|
if not updated:
|
||||||
|
# Same rationale as complete_download_in_queue: a missing
|
||||||
|
# queue item is a business outcome, not a routing error.
|
||||||
return web.json_response(
|
return web.json_response(
|
||||||
{"success": False, "error": "Download not found in queue"},
|
{"success": False, "error": "Download not found in queue"}
|
||||||
status=404,
|
|
||||||
)
|
)
|
||||||
return web.json_response({"success": True})
|
return web.json_response({"success": True})
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
@@ -2339,6 +2488,90 @@ class ModelMoveHandler:
|
|||||||
self._move_service = move_service
|
self._move_service = move_service
|
||||||
self._logger = logger
|
self._logger = logger
|
||||||
|
|
||||||
|
async def create_folder(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
except Exception:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
folder_path = data.get("folder_path")
|
||||||
|
if not folder_path:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Folder path is required"}, status=400
|
||||||
|
)
|
||||||
|
result = await self._move_service.create_folder(folder_path)
|
||||||
|
status = 200 if result.get("success") else 400
|
||||||
|
return web.json_response(result, status=status)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error creating folder: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def delete_folder(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
except Exception:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
folder_path = data.get("folder_path")
|
||||||
|
if not folder_path:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Folder path is required"}, status=400
|
||||||
|
)
|
||||||
|
dry_run = bool(data.get("dry_run"))
|
||||||
|
result = await self._move_service.delete_folder(
|
||||||
|
folder_path, dry_run=dry_run
|
||||||
|
)
|
||||||
|
if result.get("success"):
|
||||||
|
if not dry_run:
|
||||||
|
_broadcast_models_changed()
|
||||||
|
return web.json_response(result, status=200)
|
||||||
|
|
||||||
|
# "not_empty" / "busy" are conflicts between the tree the client
|
||||||
|
# rendered and the on-disk truth; everything else is a bad request.
|
||||||
|
code = result.get("code")
|
||||||
|
status = 409 if code in ("not_empty", "busy") else 400
|
||||||
|
return web.json_response(result, status=status)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error deleting folder: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
async def rename_folder(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
except Exception:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Invalid JSON body"}, status=400
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
folder_path = data.get("folder_path")
|
||||||
|
new_name = data.get("new_name")
|
||||||
|
if not folder_path:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "Folder path is required"}, status=400
|
||||||
|
)
|
||||||
|
if not new_name:
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": "New folder name is required"}, status=400
|
||||||
|
)
|
||||||
|
result = await self._move_service.rename_folder(folder_path, new_name)
|
||||||
|
if result.get("success"):
|
||||||
|
if result.get("renamed"):
|
||||||
|
_broadcast_models_changed()
|
||||||
|
return web.json_response(result, status=200)
|
||||||
|
|
||||||
|
# A name collision or a staged delete inside the subtree is a
|
||||||
|
# conflict with the state the client rendered, not a bad request.
|
||||||
|
code = result.get("code")
|
||||||
|
status = 409 if code in ("target_exists", "busy") else 400
|
||||||
|
return web.json_response(result, status=status)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error("Error renaming folder: %s", exc, exc_info=True)
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
async def move_model(self, request: web.Request) -> web.Response:
|
async def move_model(self, request: web.Request) -> web.Response:
|
||||||
try:
|
try:
|
||||||
data = await request.json()
|
data = await request.json()
|
||||||
@@ -2463,6 +2696,71 @@ class ModelAutoOrganizeHandler:
|
|||||||
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelFilenameTemplateHandler:
|
||||||
|
"""Apply the configured filename template to existing library models."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
use_case: FilenameTemplateUseCase,
|
||||||
|
progress_callback: WebSocketFilenameTemplateProgressCallback,
|
||||||
|
logger: logging.Logger,
|
||||||
|
) -> None:
|
||||||
|
self._use_case = use_case
|
||||||
|
self._progress_callback = progress_callback
|
||||||
|
self._logger = logger
|
||||||
|
|
||||||
|
async def apply_filename_template(self, request: web.Request) -> web.Response:
|
||||||
|
try:
|
||||||
|
file_paths = None
|
||||||
|
if request.method == "POST":
|
||||||
|
try:
|
||||||
|
data = await request.json()
|
||||||
|
file_paths = data.get("file_paths")
|
||||||
|
except Exception: # pragma: no cover - permissive path
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
# GET variant (browser extension is GET-only): comma-separated
|
||||||
|
# file_paths query parameter.
|
||||||
|
raw_file_paths = request.query.get("file_paths")
|
||||||
|
if raw_file_paths:
|
||||||
|
file_paths = [
|
||||||
|
path.strip()
|
||||||
|
for path in raw_file_paths.split(",")
|
||||||
|
if path.strip()
|
||||||
|
]
|
||||||
|
|
||||||
|
result = await self._use_case.execute(
|
||||||
|
file_paths=file_paths,
|
||||||
|
progress_callback=self._progress_callback,
|
||||||
|
)
|
||||||
|
_broadcast_models_changed()
|
||||||
|
return web.json_response(result.to_dict())
|
||||||
|
except AutoOrganizeInProgressError:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "Another library operation is already running. Please wait for it to complete.",
|
||||||
|
},
|
||||||
|
status=409,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
self._logger.error(
|
||||||
|
"Error in apply_filename_template: %s", exc, exc_info=True
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
await self._progress_callback.on_progress(
|
||||||
|
{
|
||||||
|
"type": "filename_template_progress",
|
||||||
|
"status": "error",
|
||||||
|
"error": str(exc),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
except Exception: # pragma: no cover - defensive reporting
|
||||||
|
pass
|
||||||
|
return web.json_response({"success": False, "error": str(exc)}, status=500)
|
||||||
|
|
||||||
|
|
||||||
class ModelUpdateHandler:
|
class ModelUpdateHandler:
|
||||||
"""Handle update tracking requests."""
|
"""Handle update tracking requests."""
|
||||||
|
|
||||||
@@ -2637,10 +2935,20 @@ class ModelUpdateHandler:
|
|||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
same_base_scope = self._uses_same_base_update_scope()
|
||||||
|
|
||||||
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 not callable(has_update_fn):
|
||||||
|
continue
|
||||||
|
scoped_fn = (
|
||||||
|
getattr(record, "has_update_for_local_bases", None)
|
||||||
|
if same_base_scope
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
qualifies_fn = scoped_fn if callable(scoped_fn) else has_update_fn
|
||||||
|
if qualifies_fn(
|
||||||
hide_early_access=hide_early_access,
|
hide_early_access=hide_early_access,
|
||||||
hide_paid=hide_paid,
|
hide_paid=hide_paid,
|
||||||
):
|
):
|
||||||
@@ -2653,6 +2961,26 @@ class ModelUpdateHandler:
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _uses_same_base_update_scope(self) -> bool:
|
||||||
|
"""Return True when update reporting must honor same-base scoping.
|
||||||
|
|
||||||
|
Mirrors ``BaseModelService._annotate_update_flags``: the Updates filter
|
||||||
|
evaluates updates per local base model when ``version_grouping`` is
|
||||||
|
``same_base`` (its default). The refresh summary counts with the same
|
||||||
|
scope so the "Found N update(s)" toast matches what the filter
|
||||||
|
displays. See issue #1083.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if self._settings is None:
|
||||||
|
return True
|
||||||
|
try:
|
||||||
|
strategy_value = self._settings.get("version_grouping")
|
||||||
|
except Exception:
|
||||||
|
return True
|
||||||
|
if isinstance(strategy_value, str) and strategy_value.strip():
|
||||||
|
return strategy_value.strip().lower() == "same_base"
|
||||||
|
return True
|
||||||
|
|
||||||
async def set_model_update_ignore(self, request: web.Request) -> web.Response:
|
async def set_model_update_ignore(self, request: web.Request) -> web.Response:
|
||||||
payload = await self._read_json(request)
|
payload = await self._read_json(request)
|
||||||
model_id = self._normalize_model_id(payload.get("modelId"))
|
model_id = self._normalize_model_id(payload.get("modelId"))
|
||||||
@@ -3200,6 +3528,7 @@ class ModelHandlerSet:
|
|||||||
civitai: ModelCivitaiHandler
|
civitai: ModelCivitaiHandler
|
||||||
move: ModelMoveHandler
|
move: ModelMoveHandler
|
||||||
auto_organize: ModelAutoOrganizeHandler
|
auto_organize: ModelAutoOrganizeHandler
|
||||||
|
filename_template: ModelFilenameTemplateHandler
|
||||||
updates: ModelUpdateHandler
|
updates: ModelUpdateHandler
|
||||||
|
|
||||||
def to_route_mapping(
|
def to_route_mapping(
|
||||||
@@ -3259,14 +3588,20 @@ class ModelHandlerSet:
|
|||||||
"get_civitai_model_by_hash": self.civitai.get_civitai_model_by_hash,
|
"get_civitai_model_by_hash": self.civitai.get_civitai_model_by_hash,
|
||||||
"move_model": self.move.move_model,
|
"move_model": self.move.move_model,
|
||||||
"move_models_bulk": self.move.move_models_bulk,
|
"move_models_bulk": self.move.move_models_bulk,
|
||||||
|
"create_folder": self.move.create_folder,
|
||||||
|
"delete_folder": self.move.delete_folder,
|
||||||
|
"rename_folder": self.move.rename_folder,
|
||||||
"auto_organize_models": self.auto_organize.auto_organize_models,
|
"auto_organize_models": self.auto_organize.auto_organize_models,
|
||||||
"get_auto_organize_progress": self.auto_organize.get_auto_organize_progress,
|
"get_auto_organize_progress": self.auto_organize.get_auto_organize_progress,
|
||||||
|
"apply_filename_template": self.filename_template.apply_filename_template,
|
||||||
"get_model_notes": self.query.get_model_notes,
|
"get_model_notes": self.query.get_model_notes,
|
||||||
"get_model_preview_url": self.query.get_model_preview_url,
|
"get_model_preview_url": self.query.get_model_preview_url,
|
||||||
"get_model_civitai_url": self.query.get_model_civitai_url,
|
"get_model_civitai_url": self.query.get_model_civitai_url,
|
||||||
"get_model_metadata": self.query.get_model_metadata,
|
"get_model_metadata": self.query.get_model_metadata,
|
||||||
"get_model_description": self.query.get_model_description,
|
"get_model_description": self.query.get_model_description,
|
||||||
"get_relative_paths": self.query.get_relative_paths,
|
"get_relative_paths": self.query.get_relative_paths,
|
||||||
|
"update_active_filters": self.query.update_active_filters,
|
||||||
|
"get_active_filters": self.query.get_active_filters,
|
||||||
"refresh_model_updates": self.updates.refresh_model_updates,
|
"refresh_model_updates": self.updates.refresh_model_updates,
|
||||||
"fetch_missing_civitai_license_data": self.updates.fetch_missing_civitai_license_data,
|
"fetch_missing_civitai_license_data": self.updates.fetch_missing_civitai_license_data,
|
||||||
"set_model_update_ignore": self.updates.set_model_update_ignore,
|
"set_model_update_ignore": self.updates.set_model_update_ignore,
|
||||||
|
|||||||
@@ -0,0 +1,639 @@
|
|||||||
|
"""Handlers for external model sources: linking, file listing and downloads.
|
||||||
|
|
||||||
|
Covers every site registered in :mod:`py.services.model_sources`. The module
|
||||||
|
was Hugging Face only (``hf_handlers.py`` / ``HfHandler``) until ModelScope
|
||||||
|
downloads were added; the per-site differences now live in the providers, so
|
||||||
|
this file has no platform branches beyond the capability lookups.
|
||||||
|
|
||||||
|
The historical route paths (``/api/lm/set-hf-url``, ``/api/lm/hf-repo-files``,
|
||||||
|
``/api/lm/download-hf-model``) are still registered as aliases of the generic
|
||||||
|
handlers, so existing callers keep working.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
from ...config import config
|
||||||
|
from ...services.downloader import (
|
||||||
|
DownloadProgress,
|
||||||
|
get_downloader,
|
||||||
|
)
|
||||||
|
from ...services.aria2_downloader import Aria2Downloader
|
||||||
|
from ...services.model_sources import (
|
||||||
|
ModelSourceError,
|
||||||
|
SourceRef,
|
||||||
|
detect_source,
|
||||||
|
get_download_source,
|
||||||
|
hydrate_from_source,
|
||||||
|
is_valid_source_id,
|
||||||
|
list_sources,
|
||||||
|
normalize_metadata_source,
|
||||||
|
)
|
||||||
|
from ...services.settings_manager import get_settings_manager
|
||||||
|
from ...services.service_registry import ServiceRegistry
|
||||||
|
from ...services.websocket_manager import ws_manager
|
||||||
|
from ...utils.metadata_manager import MetadataManager
|
||||||
|
from ...utils.models import LoraMetadata, CheckpointMetadata, EmbeddingMetadata
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_DEFAULT_MODEL_CLASS = LoraMetadata
|
||||||
|
_DEFAULT_SCANNER_GETTER = "get_lora_scanner"
|
||||||
|
|
||||||
|
|
||||||
|
def _infer_model_type(model_root: str) -> tuple[Any, str]:
|
||||||
|
"""Determine model class and scanner by matching ``model_root`` against the
|
||||||
|
configured root paths for each model type (from ``Config``).
|
||||||
|
|
||||||
|
The ``model_root`` value comes from the frontend's model-root dropdown,
|
||||||
|
which is populated from the current page's scanner roots. By checking
|
||||||
|
which scanner's root list it belongs to, we avoid fragile heuristics
|
||||||
|
like substring-matching path names.
|
||||||
|
"""
|
||||||
|
norm = os.path.normpath(model_root).replace(os.sep, "/")
|
||||||
|
|
||||||
|
# LoRA roots
|
||||||
|
for p in (config.loras_roots or []) + (config.extra_loras_roots or []):
|
||||||
|
if os.path.normpath(p).replace(os.sep, "/") == norm:
|
||||||
|
return LoraMetadata, "get_lora_scanner"
|
||||||
|
|
||||||
|
# Checkpoint / UNet roots
|
||||||
|
for p in (
|
||||||
|
(config.checkpoints_roots or [])
|
||||||
|
+ (config.extra_checkpoints_roots or [])
|
||||||
|
+ (config.unet_roots or [])
|
||||||
|
+ (config.extra_unet_roots or [])
|
||||||
|
):
|
||||||
|
if os.path.normpath(p).replace(os.sep, "/") == norm:
|
||||||
|
return CheckpointMetadata, "get_checkpoint_scanner"
|
||||||
|
|
||||||
|
# Embedding roots
|
||||||
|
for p in (config.embeddings_roots or []) + (config.extra_embeddings_roots or []):
|
||||||
|
if os.path.normpath(p).replace(os.sep, "/") == norm:
|
||||||
|
return EmbeddingMetadata, "get_embedding_scanner"
|
||||||
|
|
||||||
|
# Fallback — should not happen in normal use
|
||||||
|
logger.warning(
|
||||||
|
"Could not determine model type for root '%s'; defaulting to LoRA",
|
||||||
|
model_root,
|
||||||
|
)
|
||||||
|
return _DEFAULT_MODEL_CLASS, _DEFAULT_SCANNER_GETTER
|
||||||
|
|
||||||
|
|
||||||
|
async def _report_phase(
|
||||||
|
download_id: str | None, stage: str, platform: str = ""
|
||||||
|
) -> None:
|
||||||
|
"""Tell the progress UI which post-transfer stage is running.
|
||||||
|
|
||||||
|
A download's byte counter stops the moment the last byte lands, but the
|
||||||
|
backend still has to index the file and read the model site's API. Without
|
||||||
|
this the bar sits at 100% reporting "0 B/s" and the download looks stuck for
|
||||||
|
several seconds. *stage* is machine-readable — the UI localises it — and
|
||||||
|
*platform* lets it name the site the metadata comes from.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not download_id:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
await ws_manager.broadcast_download_progress(
|
||||||
|
download_id,
|
||||||
|
{
|
||||||
|
"status": "metadata",
|
||||||
|
"stage": stage,
|
||||||
|
"platform": platform,
|
||||||
|
"progress": 100,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
except Exception as exc: # pragma: no cover - progress must never be fatal
|
||||||
|
logger.debug("Failed to report the '%s' phase: %s", stage, exc)
|
||||||
|
|
||||||
|
|
||||||
|
async def _save_source_metadata(
|
||||||
|
dest_path: str, ref: SourceRef, model_root: str, *, download_id: str | None = None
|
||||||
|
) -> None:
|
||||||
|
"""Create a proper .metadata.json and add the model to the scanner cache.
|
||||||
|
|
||||||
|
The metadata is created through the owning scanner rather than
|
||||||
|
``MetadataManager.create_default_metadata()``, because that is the only
|
||||||
|
factory that knows when hashing must be deferred: ``CheckpointScanner`` and
|
||||||
|
``OtherScanner`` deliberately record ``hash_status="pending"`` with an empty
|
||||||
|
``sha256`` for their multi-GB files, and the generic helper would read a
|
||||||
|
10 GB checkpoint end to end *inside the download request*. Scanners for the
|
||||||
|
small types delegate straight back to it, so nothing changes for them.
|
||||||
|
|
||||||
|
The external-source fields are then overlaid and the model is registered in
|
||||||
|
the in-memory scanner cache so it appears immediately without a full
|
||||||
|
filesystem walk.
|
||||||
|
|
||||||
|
Finally the site's own published metadata is applied (see
|
||||||
|
:func:`~py.services.model_sources.hydration.hydrate_from_source`), so a
|
||||||
|
ModelScope or Hugging Face download lands with the same populated model
|
||||||
|
card a CivitAI download produces instead of a bare filename and hash.
|
||||||
|
|
||||||
|
Both post-transfer stages are reported through *download_id* when the UI is
|
||||||
|
watching one, because neither advances the byte counter.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
model_class, scanner_getter_name = _infer_model_type(model_root)
|
||||||
|
|
||||||
|
scanner = None
|
||||||
|
scanner_getter = getattr(ServiceRegistry, scanner_getter_name, None)
|
||||||
|
if scanner_getter is not None:
|
||||||
|
scanner = await scanner_getter()
|
||||||
|
|
||||||
|
# 1. Create proper metadata (reads safetensors headers; hashes only for
|
||||||
|
# the model types whose scanner does not defer it)
|
||||||
|
await _report_phase(download_id, "indexing", ref.platform)
|
||||||
|
create_metadata = getattr(scanner, "_create_default_metadata", None)
|
||||||
|
if create_metadata is not None:
|
||||||
|
metadata = await create_metadata(dest_path)
|
||||||
|
else:
|
||||||
|
metadata = await MetadataManager.create_default_metadata(
|
||||||
|
dest_path, model_class=model_class
|
||||||
|
)
|
||||||
|
if metadata is None:
|
||||||
|
logger.warning("create_default_metadata returned None for %s", dest_path)
|
||||||
|
return
|
||||||
|
|
||||||
|
# 2. Overlay the external-source fields (`hf_url` is written by
|
||||||
|
# normalisation for Hugging Face only)
|
||||||
|
fields = metadata._unknown_fields
|
||||||
|
fields["source_url"] = ref.url
|
||||||
|
fields["source_platform"] = ref.platform
|
||||||
|
if ref.platform == "huggingface":
|
||||||
|
fields["hf_url"] = ref.url
|
||||||
|
metadata.from_civitai = False # externally-sourced models are not from CivitAI
|
||||||
|
|
||||||
|
# 3. Save metadata atomically
|
||||||
|
await MetadataManager.save_metadata(dest_path, metadata)
|
||||||
|
logger.info(
|
||||||
|
"Saved %s metadata (source=%s, hash_status=%s) for %s",
|
||||||
|
ref.platform, ref.url, getattr(metadata, "hash_status", "?"), dest_path,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 4. Determine relative folder path for cache
|
||||||
|
# model_root is an absolute path; dest_path is under it
|
||||||
|
folder = ""
|
||||||
|
if os.path.isabs(model_root) and dest_path.startswith(model_root):
|
||||||
|
rel = os.path.relpath(os.path.dirname(dest_path), model_root)
|
||||||
|
folder = rel.replace(os.sep, "/") if rel != "." else ""
|
||||||
|
|
||||||
|
# 5. Add to scanner cache (same as CivitAI's _execute_download does)
|
||||||
|
if scanner is not None:
|
||||||
|
metadata_dict = normalize_metadata_source(metadata.to_dict())
|
||||||
|
await scanner.add_model_to_cache(metadata_dict, folder)
|
||||||
|
logger.info("Added %s to scanner cache (folder=%s)", dest_path, folder)
|
||||||
|
|
||||||
|
# 6. Top up from the site's public API. Runs last so the scanner-cache
|
||||||
|
# refresh it performs lands on the entry created above. It never
|
||||||
|
# raises and never fails the download.
|
||||||
|
await _report_phase(download_id, "source", ref.platform)
|
||||||
|
await hydrate_from_source(dest_path, ref=ref)
|
||||||
|
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Failed to save source metadata for %s: %s", dest_path, exc)
|
||||||
|
|
||||||
|
|
||||||
|
def _find_matching_root(dest_dir: str) -> str | None:
|
||||||
|
"""Walk up *dest_dir* to find which configured scanner root it belongs to."""
|
||||||
|
norm = os.path.normpath(dest_dir).replace(os.sep, "/")
|
||||||
|
all_roots = []
|
||||||
|
for root_list in (
|
||||||
|
config.loras_roots or [],
|
||||||
|
config.extra_loras_roots or [],
|
||||||
|
config.checkpoints_roots or [],
|
||||||
|
config.extra_checkpoints_roots or [],
|
||||||
|
config.unet_roots or [],
|
||||||
|
config.extra_unet_roots or [],
|
||||||
|
config.embeddings_roots or [],
|
||||||
|
config.extra_embeddings_roots or [],
|
||||||
|
):
|
||||||
|
all_roots.extend([os.path.normpath(p).replace(os.sep, "/") for p in root_list])
|
||||||
|
# Find the longest matching prefix
|
||||||
|
match: str | None = None
|
||||||
|
for root in all_roots:
|
||||||
|
if norm.startswith(root):
|
||||||
|
if match is None or len(root) > len(match):
|
||||||
|
match = root
|
||||||
|
return match
|
||||||
|
|
||||||
|
|
||||||
|
async def _add_to_scanner_cache(dest_path: str, metadata: dict[str, Any]) -> None:
|
||||||
|
model_dir = os.path.dirname(dest_path)
|
||||||
|
model_root = _find_matching_root(model_dir)
|
||||||
|
if not model_root:
|
||||||
|
raise ValueError(f"File path {dest_path} is not within any configured scanner root")
|
||||||
|
scanner_getter_name = _infer_model_type(model_root)[1]
|
||||||
|
scanner_getter = getattr(ServiceRegistry, scanner_getter_name, None)
|
||||||
|
if scanner_getter is None:
|
||||||
|
raise RuntimeError(f"Scanner getter '{scanner_getter_name}' not found in ServiceRegistry")
|
||||||
|
scanner = await scanner_getter()
|
||||||
|
if scanner is None:
|
||||||
|
raise RuntimeError(f"Scanner '{scanner_getter_name}' returned None")
|
||||||
|
await scanner.update_single_model_cache(dest_path, dest_path, metadata)
|
||||||
|
|
||||||
|
|
||||||
|
def _unsupported_platform_error(platform: str) -> web.Response:
|
||||||
|
supported = ", ".join(source.label for source in list_sources() if source.supports_download)
|
||||||
|
return web.json_response(
|
||||||
|
{"error": f"'{platform}' does not support downloads. Supported: {supported}"},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelSourceHandler:
|
||||||
|
"""Handle external model browsing, linking and downloads."""
|
||||||
|
|
||||||
|
async def get_model_sources(self, request: web.Request) -> web.Response:
|
||||||
|
"""List the external model sites the UI can link a model to.
|
||||||
|
|
||||||
|
Used by the "Link Model" dialog to validate URLs client-side, to
|
||||||
|
explain which sites support AI metadata enrichment, and to pick the
|
||||||
|
right download endpoint/revision.
|
||||||
|
"""
|
||||||
|
|
||||||
|
return web.json_response([
|
||||||
|
{
|
||||||
|
"platform": source.platform,
|
||||||
|
"label": source.label,
|
||||||
|
"supports_enrichment": source.supports_enrichment,
|
||||||
|
"supports_download": source.supports_download,
|
||||||
|
"default_revision": source.default_revision,
|
||||||
|
"example_url": source.canonical_url(
|
||||||
|
"user/repo" if source.platform != "tensorart" else "827823520299086029"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
for source in list_sources()
|
||||||
|
])
|
||||||
|
|
||||||
|
async def set_hf_url(self, request: web.Request) -> web.Response:
|
||||||
|
"""Link a model file to its page on an external model site.
|
||||||
|
|
||||||
|
Accepts ``source_url`` (preferred) or the legacy ``hf_url`` / ``url``
|
||||||
|
payload key. Every registered site is recognised and the platform is
|
||||||
|
stored alongside the canonical URL. TensorArt models can be linked and
|
||||||
|
browsed, but not AI-enriched.
|
||||||
|
|
||||||
|
The route path keeps its historical ``set-hf-url`` name.
|
||||||
|
"""
|
||||||
|
|
||||||
|
try:
|
||||||
|
payload: dict[str, Any] = await request.json()
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return web.json_response({"success": False, "error": "Invalid JSON"}, status=400)
|
||||||
|
|
||||||
|
file_path = (payload.get("file_path") or "").strip()
|
||||||
|
raw_url = (
|
||||||
|
payload.get("source_url")
|
||||||
|
or payload.get("hf_url")
|
||||||
|
or payload.get("url")
|
||||||
|
or ""
|
||||||
|
)
|
||||||
|
source_url = raw_url.strip() if isinstance(raw_url, str) else ""
|
||||||
|
|
||||||
|
if not file_path or not source_url:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "Missing required fields: 'file_path' and 'source_url'",
|
||||||
|
},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
ref = detect_source(source_url, strict=True)
|
||||||
|
if ref is None:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": (
|
||||||
|
"Unsupported model URL. Supported formats: "
|
||||||
|
+ ", ".join(
|
||||||
|
f"{s.label} ({s.canonical_url('user/repo')})"
|
||||||
|
if s.platform != "tensorart"
|
||||||
|
else f"{s.label} (https://tensor.art/models/<id>)"
|
||||||
|
for s in list_sources()
|
||||||
|
)
|
||||||
|
),
|
||||||
|
},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not os.path.isfile(file_path):
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": f"File not found: {file_path}"},
|
||||||
|
status=404,
|
||||||
|
)
|
||||||
|
|
||||||
|
model_root = _find_matching_root(os.path.dirname(file_path))
|
||||||
|
if not model_root:
|
||||||
|
return web.json_response(
|
||||||
|
{
|
||||||
|
"success": False,
|
||||||
|
"error": "File is not within any configured model directory. Cannot link to a model source.",
|
||||||
|
},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
existing = await MetadataManager.load_metadata_payload(file_path)
|
||||||
|
|
||||||
|
already_linked = (
|
||||||
|
(existing.get("source_url") or "").strip() == ref.url
|
||||||
|
and (existing.get("source_platform") or "").strip().lower()
|
||||||
|
== ref.platform
|
||||||
|
) or (
|
||||||
|
not existing.get("source_url")
|
||||||
|
and ref.platform == "huggingface"
|
||||||
|
and (existing.get("hf_url") or "").strip() == ref.url
|
||||||
|
)
|
||||||
|
if already_linked:
|
||||||
|
return web.json_response({
|
||||||
|
"success": True,
|
||||||
|
"message": "source_url already set",
|
||||||
|
"source_url": ref.url,
|
||||||
|
"source_platform": ref.platform,
|
||||||
|
"hf_url": ref.url if ref.platform == "huggingface" else "",
|
||||||
|
})
|
||||||
|
|
||||||
|
existing["source_url"] = ref.url
|
||||||
|
existing["source_platform"] = ref.platform
|
||||||
|
if ref.platform == "huggingface":
|
||||||
|
existing["hf_url"] = ref.url
|
||||||
|
else:
|
||||||
|
existing.pop("hf_url", None)
|
||||||
|
normalize_metadata_source(existing)
|
||||||
|
|
||||||
|
# NOTE: deliberately do NOT touch `from_civitai` here. It records
|
||||||
|
# where the metadata came from, and the UI must show the CivitAI
|
||||||
|
# link whenever CivitAI data is present — linking an external
|
||||||
|
# source must not hide it (#1094). Source provenance is tracked
|
||||||
|
# via `source_platform` / `source_url`.
|
||||||
|
await MetadataManager.save_metadata(file_path, existing)
|
||||||
|
|
||||||
|
await _add_to_scanner_cache(file_path, existing)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Linked %s to %s source (%s)", file_path, ref.platform, ref.url
|
||||||
|
)
|
||||||
|
return web.json_response({
|
||||||
|
"success": True,
|
||||||
|
"message": f"Linked to {ref.url}",
|
||||||
|
"source_url": ref.url,
|
||||||
|
"source_platform": ref.platform,
|
||||||
|
"hf_url": existing.get("hf_url", ""),
|
||||||
|
})
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to link %s to a model source: %s", file_path, exc)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": str(exc)},
|
||||||
|
status=500,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def list_model_source_files(self, request: web.Request) -> web.Response:
|
||||||
|
"""List the downloadable weight files of an external repository.
|
||||||
|
|
||||||
|
Query params: ``platform``, ``repo`` (``owner/name``), ``revision``
|
||||||
|
(optional; each site has its own default branch).
|
||||||
|
|
||||||
|
Returns a JSON array of ``{"filename", "size"}``, largest first —
|
||||||
|
the same shape the Hugging Face endpoint has always returned.
|
||||||
|
"""
|
||||||
|
|
||||||
|
platform = (request.query.get("platform") or "").strip()
|
||||||
|
repo = (request.query.get("repo") or "").strip()
|
||||||
|
revision = (request.query.get("revision") or "").strip()
|
||||||
|
|
||||||
|
source = get_download_source(platform)
|
||||||
|
if source is None:
|
||||||
|
return _unsupported_platform_error(platform)
|
||||||
|
if not is_valid_source_id(repo):
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "Missing or invalid 'repo' parameter (expected owner/name)"},
|
||||||
|
status=400,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
files = await source.list_files(repo, revision)
|
||||||
|
except ModelSourceError as exc:
|
||||||
|
return web.json_response({"error": str(exc)}, status=exc.status)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to list %s files in %s: %s", platform, repo, exc)
|
||||||
|
return web.json_response({"error": str(exc)}, status=502)
|
||||||
|
|
||||||
|
return web.json_response(files)
|
||||||
|
|
||||||
|
async def download_model_source(self, request: web.Request) -> web.Response:
|
||||||
|
"""Download a single file from an external repository.
|
||||||
|
|
||||||
|
POST JSON body::
|
||||||
|
|
||||||
|
{
|
||||||
|
"platform": "modelscope",
|
||||||
|
"repo": "owner/name",
|
||||||
|
"filename": "subdir/model.safetensors",
|
||||||
|
"revision": "master",
|
||||||
|
"model_root": "loras",
|
||||||
|
"relative_path": "",
|
||||||
|
"use_default_paths": false,
|
||||||
|
"download_id": "optional-batch-id"
|
||||||
|
}
|
||||||
|
|
||||||
|
``platform`` defaults to ``huggingface`` when omitted, which keeps the
|
||||||
|
legacy ``/api/lm/download-hf-model`` payload working unchanged.
|
||||||
|
|
||||||
|
If ``download_id`` is provided, real-time progress (bytes, speed,
|
||||||
|
percentage) is broadcast via the WebSocket progress system.
|
||||||
|
|
||||||
|
Respects the ``download_backend`` setting (``aria2`` or ``default``).
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
payload: dict[str, Any] = await request.json()
|
||||||
|
except json.JSONDecodeError:
|
||||||
|
return web.json_response({"error": "Invalid JSON"}, status=400)
|
||||||
|
|
||||||
|
platform = (payload.get("platform") or "huggingface").strip()
|
||||||
|
repo = (payload.get("repo") or "").strip()
|
||||||
|
filename = (payload.get("filename") or "").strip()
|
||||||
|
revision = (payload.get("revision") or "").strip()
|
||||||
|
model_root = (payload.get("model_root") or "").strip()
|
||||||
|
relative_path = (payload.get("relative_path") or "").strip()
|
||||||
|
use_default_paths = bool(payload.get("use_default_paths", False))
|
||||||
|
download_id: str | None = payload.get("download_id")
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"download_model_source: platform=%s repo=%s file=%s root=%s download_id=%s",
|
||||||
|
platform, repo, filename, model_root, download_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
source = get_download_source(platform)
|
||||||
|
if source is None:
|
||||||
|
return _unsupported_platform_error(platform)
|
||||||
|
|
||||||
|
if not repo or not filename:
|
||||||
|
return web.json_response(
|
||||||
|
{"error": "Missing required fields: 'repo' and 'filename'"}, status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
# `owner/name` only; the components become path segments below.
|
||||||
|
if not is_valid_source_id(repo):
|
||||||
|
return web.json_response({"error": f"Invalid repo format: {repo}"}, status=400)
|
||||||
|
owner, repo_name = repo.split("/", 1)
|
||||||
|
|
||||||
|
# Validate filename — must not contain path traversal
|
||||||
|
if ".." in filename:
|
||||||
|
return web.json_response({"error": "Invalid filename"}, status=400)
|
||||||
|
|
||||||
|
# Validate relative_path — must not be absolute or escape base directory
|
||||||
|
if relative_path:
|
||||||
|
if os.path.isabs(relative_path):
|
||||||
|
return web.json_response({"error": "relative_path must not be absolute"}, status=400)
|
||||||
|
if ".." in relative_path.split("/") or "\\" in relative_path:
|
||||||
|
return web.json_response({"error": "Invalid relative_path"}, status=400)
|
||||||
|
|
||||||
|
# Use model_root directly as the base directory — same approach as
|
||||||
|
# CivitAI's download path (download_manager.py). No realpath, no
|
||||||
|
# allowed-roots validation, no path-traversal check; those are
|
||||||
|
# unnecessary when the frontend sends the path from its own dropdown
|
||||||
|
# (populated from scanner roots). Using the "business path" directly
|
||||||
|
# keeps dest_path consistent with scanner roots so that later folder
|
||||||
|
# derivation (in _save_source_metadata) works correctly.
|
||||||
|
if os.path.isabs(model_root):
|
||||||
|
base_dir = os.path.normpath(model_root)
|
||||||
|
else:
|
||||||
|
base_dir = os.path.normpath(os.path.join(os.getcwd(), "models", model_root))
|
||||||
|
|
||||||
|
if use_default_paths:
|
||||||
|
target_dir = os.path.join(base_dir, source.default_subdir, owner, repo_name)
|
||||||
|
elif relative_path:
|
||||||
|
target_dir = os.path.join(base_dir, relative_path)
|
||||||
|
else:
|
||||||
|
target_dir = base_dir
|
||||||
|
|
||||||
|
# Strip the repository sub-directory — "diffusion_models/xxx.safetensors"
|
||||||
|
# is a repository convention, not meaningful for local storage.
|
||||||
|
file_base = os.path.basename(filename)
|
||||||
|
|
||||||
|
os.makedirs(target_dir, exist_ok=True)
|
||||||
|
dest_path = os.path.join(target_dir, file_base)
|
||||||
|
|
||||||
|
# Built per request: sites that redirect to a CDN hand out a
|
||||||
|
# time-limited token in the redirect, so the URL must never be cached.
|
||||||
|
resolve_url = source.file_download_url(repo, filename, revision)
|
||||||
|
ref = SourceRef(
|
||||||
|
platform=source.platform, source_id=repo, url=source.canonical_url(repo)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check if already exists (simple skip)
|
||||||
|
if os.path.exists(dest_path) and os.path.getsize(dest_path) > 0:
|
||||||
|
logger.info("download_model_source: file already exists, skipping — %s", dest_path)
|
||||||
|
# The sidecar may predate the source metadata being fetched, or may
|
||||||
|
# have been deleted, so top it up instead of skipping past it.
|
||||||
|
# Hydration no-ops when there is no sidecar to update.
|
||||||
|
await _report_phase(download_id, "source", source.platform)
|
||||||
|
await hydrate_from_source(dest_path, ref=ref)
|
||||||
|
return web.json_response({
|
||||||
|
"success": True,
|
||||||
|
"message": f"File already exists: {dest_path}",
|
||||||
|
"path": dest_path,
|
||||||
|
})
|
||||||
|
|
||||||
|
# Set up progress callback if download_id is provided
|
||||||
|
progress_callback = None
|
||||||
|
if download_id:
|
||||||
|
|
||||||
|
async def _progress_callback(
|
||||||
|
progress: float | DownloadProgress,
|
||||||
|
snapshot: DownloadProgress | None = None,
|
||||||
|
) -> None:
|
||||||
|
percent = 0.0
|
||||||
|
metrics = snapshot if isinstance(snapshot, DownloadProgress) else None
|
||||||
|
|
||||||
|
if isinstance(progress, DownloadProgress):
|
||||||
|
percent = progress.percent_complete
|
||||||
|
metrics = progress
|
||||||
|
elif isinstance(snapshot, DownloadProgress):
|
||||||
|
percent = snapshot.percent_complete
|
||||||
|
else:
|
||||||
|
percent = float(progress)
|
||||||
|
|
||||||
|
broadcast: dict[str, Any] = {
|
||||||
|
"status": "progress",
|
||||||
|
"progress": round(percent),
|
||||||
|
}
|
||||||
|
if metrics:
|
||||||
|
broadcast["bytes_downloaded"] = metrics.bytes_downloaded
|
||||||
|
broadcast["total_bytes"] = metrics.total_bytes
|
||||||
|
broadcast["bytes_per_second"] = metrics.bytes_per_second
|
||||||
|
|
||||||
|
await ws_manager.broadcast_download_progress(download_id, broadcast)
|
||||||
|
|
||||||
|
progress_callback = _progress_callback
|
||||||
|
|
||||||
|
# Respect download backend setting (aria2 vs default)
|
||||||
|
download_backend = (
|
||||||
|
get_settings_manager().get("download_backend", "default")
|
||||||
|
)
|
||||||
|
|
||||||
|
if download_backend == "aria2":
|
||||||
|
aria2 = await Aria2Downloader.get_instance()
|
||||||
|
aid = download_id or f"{source.platform}_{repo}_{filename}"
|
||||||
|
try:
|
||||||
|
ok, result = await aria2.download_file(
|
||||||
|
url=resolve_url,
|
||||||
|
save_path=dest_path,
|
||||||
|
download_id=aid,
|
||||||
|
progress_callback=progress_callback,
|
||||||
|
)
|
||||||
|
if ok:
|
||||||
|
await _save_source_metadata(
|
||||||
|
dest_path, ref, model_root, download_id=download_id
|
||||||
|
)
|
||||||
|
return web.json_response({
|
||||||
|
"success": True,
|
||||||
|
"message": f"Downloaded to {dest_path}",
|
||||||
|
"path": dest_path,
|
||||||
|
})
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": result or "aria2 download failed"},
|
||||||
|
status=500,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("%s download (aria2) failed: %s", platform, exc)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": str(exc)}, status=500
|
||||||
|
)
|
||||||
|
|
||||||
|
# Default: use built-in aiohttp Downloader
|
||||||
|
downloader = await get_downloader()
|
||||||
|
try:
|
||||||
|
success, result = await downloader.download_file(
|
||||||
|
url=resolve_url,
|
||||||
|
save_path=dest_path,
|
||||||
|
use_auth=False,
|
||||||
|
allow_resume=True,
|
||||||
|
progress_callback=progress_callback,
|
||||||
|
)
|
||||||
|
if success:
|
||||||
|
await _save_source_metadata(
|
||||||
|
dest_path, ref, model_root, download_id=download_id
|
||||||
|
)
|
||||||
|
return web.json_response({
|
||||||
|
"success": True,
|
||||||
|
"message": f"Downloaded to {result}",
|
||||||
|
"path": result,
|
||||||
|
})
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": result or "Download failed"},
|
||||||
|
status=500,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("%s download failed: %s", platform, exc)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": str(exc)}, status=500
|
||||||
|
)
|
||||||
@@ -35,6 +35,7 @@ _MODEL_TYPE_GETTER_NAMES: Dict[str, str] = {
|
|||||||
"loras": "get_lora_scanner",
|
"loras": "get_lora_scanner",
|
||||||
"checkpoints": "get_checkpoint_scanner",
|
"checkpoints": "get_checkpoint_scanner",
|
||||||
"embeddings": "get_embedding_scanner",
|
"embeddings": "get_embedding_scanner",
|
||||||
|
"other": "get_other_scanner",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Staged batch ids are ``uuid.uuid4().hex`` (32 lowercase hex chars). The id is
|
# Staged batch ids are ``uuid.uuid4().hex`` (32 lowercase hex chars). The id is
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -37,6 +37,8 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
RouteDefinition("GET", "/api/lm/wildcards/search", "search_wildcards"),
|
RouteDefinition("GET", "/api/lm/wildcards/search", "search_wildcards"),
|
||||||
RouteDefinition("POST", "/api/lm/wildcards/open-location", "open_wildcards_location"),
|
RouteDefinition("POST", "/api/lm/wildcards/open-location", "open_wildcards_location"),
|
||||||
RouteDefinition("POST", "/api/lm/open-file-location", "open_file_location"),
|
RouteDefinition("POST", "/api/lm/open-file-location", "open_file_location"),
|
||||||
|
RouteDefinition("POST", "/api/lm/browse-directory", "browse_directory"),
|
||||||
|
RouteDefinition("POST", "/api/lm/validate-path", "validate_path"),
|
||||||
RouteDefinition("POST", "/api/lm/update-usage-stats", "update_usage_stats"),
|
RouteDefinition("POST", "/api/lm/update-usage-stats", "update_usage_stats"),
|
||||||
RouteDefinition("GET", "/api/lm/get-usage-stats", "get_usage_stats"),
|
RouteDefinition("GET", "/api/lm/get-usage-stats", "get_usage_stats"),
|
||||||
RouteDefinition("POST", "/api/lm/update-lora-code", "update_lora_code"),
|
RouteDefinition("POST", "/api/lm/update-lora-code", "update_lora_code"),
|
||||||
@@ -99,16 +101,31 @@ MISC_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"GET", "/api/lm/delete-model-version", "delete_model_version"
|
"GET", "/api/lm/delete-model-version", "delete_model_version"
|
||||||
),
|
),
|
||||||
# Hugging Face model endpoints
|
# External model source endpoints (Hugging Face / ModelScope).
|
||||||
|
# The hf-* paths are the historical names, kept as aliases.
|
||||||
|
RouteDefinition(
|
||||||
|
"GET", "/api/lm/model-source-files", "list_model_source_files"
|
||||||
|
),
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"GET", "/api/lm/hf-repo-files", "get_hf_repo_files"
|
"GET", "/api/lm/hf-repo-files", "get_hf_repo_files"
|
||||||
),
|
),
|
||||||
|
# Download target routing decision (checkpoint vs diffusion model roots)
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/download/routing", "get_download_routing"
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/download-model-source", "download_model_source"
|
||||||
|
),
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"POST", "/api/lm/download-hf-model", "download_hf_model"
|
"POST", "/api/lm/download-hf-model", "download_hf_model"
|
||||||
),
|
),
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"POST", "/api/lm/set-hf-url", "set_hf_url"
|
"POST", "/api/lm/set-hf-url", "set_hf_url"
|
||||||
),
|
),
|
||||||
|
# Supported external model sites (Hugging Face / ModelScope / TensorArt)
|
||||||
|
RouteDefinition(
|
||||||
|
"GET", "/api/lm/model-sources", "get_model_sources"
|
||||||
|
),
|
||||||
# Agent skill endpoints
|
# Agent skill endpoints
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"GET", "/api/lm/agent/skills", "get_agent_skills"
|
"GET", "/api/lm/agent/skills", "get_agent_skills"
|
||||||
|
|||||||
@@ -39,8 +39,9 @@ from .handlers.misc_handlers import (
|
|||||||
build_service_registry_adapter,
|
build_service_registry_adapter,
|
||||||
)
|
)
|
||||||
from .handlers.base_model_handlers import BaseModelHandlerSet
|
from .handlers.base_model_handlers import BaseModelHandlerSet
|
||||||
from .handlers.hf_handlers import HfHandler
|
from .handlers.model_source_handlers import ModelSourceHandler
|
||||||
from .handlers.agent_handlers import AgentHandler
|
from .handlers.agent_handlers import AgentHandler
|
||||||
|
from .handlers.download_routing_handlers import DownloadRoutingHandler
|
||||||
from .misc_route_registrar import MiscRouteRegistrar
|
from .misc_route_registrar import MiscRouteRegistrar
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -138,8 +139,9 @@ class MiscRoutes:
|
|||||||
doctor = DoctorHandler(settings_service=self._settings)
|
doctor = DoctorHandler(settings_service=self._settings)
|
||||||
example_workflows = ExampleWorkflowsHandler()
|
example_workflows = ExampleWorkflowsHandler()
|
||||||
base_model = BaseModelHandlerSet()
|
base_model = BaseModelHandlerSet()
|
||||||
hf_handler = HfHandler()
|
model_source_handler = ModelSourceHandler()
|
||||||
agent_handler = AgentHandler()
|
agent_handler = AgentHandler()
|
||||||
|
download_routing = DownloadRoutingHandler()
|
||||||
|
|
||||||
return self._handler_set_factory(
|
return self._handler_set_factory(
|
||||||
health=health,
|
health=health,
|
||||||
@@ -159,8 +161,9 @@ class MiscRoutes:
|
|||||||
doctor=doctor,
|
doctor=doctor,
|
||||||
example_workflows=example_workflows,
|
example_workflows=example_workflows,
|
||||||
base_model=base_model,
|
base_model=base_model,
|
||||||
hf_handler=hf_handler,
|
model_source_handler=model_source_handler,
|
||||||
agent_handler=agent_handler,
|
agent_handler=agent_handler,
|
||||||
|
download_routing=download_routing,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -40,11 +40,20 @@ COMMON_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
RouteDefinition("POST", "/api/lm/{prefix}/verify-duplicates", "verify_duplicates"),
|
RouteDefinition("POST", "/api/lm/{prefix}/verify-duplicates", "verify_duplicates"),
|
||||||
RouteDefinition("POST", "/api/lm/{prefix}/move_model", "move_model"),
|
RouteDefinition("POST", "/api/lm/{prefix}/move_model", "move_model"),
|
||||||
RouteDefinition("POST", "/api/lm/{prefix}/move_models_bulk", "move_models_bulk"),
|
RouteDefinition("POST", "/api/lm/{prefix}/move_models_bulk", "move_models_bulk"),
|
||||||
|
RouteDefinition("POST", "/api/lm/{prefix}/create-folder", "create_folder"),
|
||||||
|
RouteDefinition("POST", "/api/lm/{prefix}/delete-folder", "delete_folder"),
|
||||||
|
RouteDefinition("POST", "/api/lm/{prefix}/rename-folder", "rename_folder"),
|
||||||
RouteDefinition("GET", "/api/lm/{prefix}/auto-organize", "auto_organize_models"),
|
RouteDefinition("GET", "/api/lm/{prefix}/auto-organize", "auto_organize_models"),
|
||||||
RouteDefinition("POST", "/api/lm/{prefix}/auto-organize", "auto_organize_models"),
|
RouteDefinition("POST", "/api/lm/{prefix}/auto-organize", "auto_organize_models"),
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"GET", "/api/lm/{prefix}/auto-organize-progress", "get_auto_organize_progress"
|
"GET", "/api/lm/{prefix}/auto-organize-progress", "get_auto_organize_progress"
|
||||||
),
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"GET", "/api/lm/{prefix}/apply-filename-template", "apply_filename_template"
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/{prefix}/apply-filename-template", "apply_filename_template"
|
||||||
|
),
|
||||||
RouteDefinition("GET", "/api/lm/{prefix}/top-tags", "get_top_tags"),
|
RouteDefinition("GET", "/api/lm/{prefix}/top-tags", "get_top_tags"),
|
||||||
RouteDefinition("GET", "/api/lm/{prefix}/search-tags", "search_tags"),
|
RouteDefinition("GET", "/api/lm/{prefix}/search-tags", "search_tags"),
|
||||||
RouteDefinition("GET", "/api/lm/{prefix}/base-models", "get_base_models"),
|
RouteDefinition("GET", "/api/lm/{prefix}/base-models", "get_base_models"),
|
||||||
@@ -68,6 +77,8 @@ COMMON_ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
"GET", "/api/lm/{prefix}/model-description", "get_model_description"
|
"GET", "/api/lm/{prefix}/model-description", "get_model_description"
|
||||||
),
|
),
|
||||||
RouteDefinition("GET", "/api/lm/{prefix}/relative-paths", "get_relative_paths"),
|
RouteDefinition("GET", "/api/lm/{prefix}/relative-paths", "get_relative_paths"),
|
||||||
|
RouteDefinition("PUT", "/api/lm/{prefix}/active-filters", "update_active_filters"),
|
||||||
|
RouteDefinition("GET", "/api/lm/{prefix}/active-filters", "get_active_filters"),
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"GET", "/api/lm/{prefix}/civitai/versions/{model_id}", "get_civitai_versions"
|
"GET", "/api/lm/{prefix}/civitai/versions/{model_id}", "get_civitai_versions"
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -0,0 +1,144 @@
|
|||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from typing import Any, Dict, List
|
||||||
|
from aiohttp import web
|
||||||
|
|
||||||
|
from .base_model_routes import BaseModelRoutes
|
||||||
|
from .model_route_registrar import ModelRouteRegistrar
|
||||||
|
from ..config import config
|
||||||
|
from ..services.other_model_service import OtherModelService
|
||||||
|
from ..services.service_registry import ServiceRegistry
|
||||||
|
from ..utils.constants import (
|
||||||
|
CIVITAI_TYPE_TO_OTHER_SUB_TYPE,
|
||||||
|
OTHER_MODEL_FOLDER_SUBTYPES,
|
||||||
|
VALID_OTHER_CIVITAI_TYPES,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class OtherRoutes(BaseModelRoutes):
|
||||||
|
"""Other-model-specific route controller (VAE, upscaler, text encoder, ...)"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
"""Initialize Other-model routes with OtherModel service"""
|
||||||
|
super().__init__()
|
||||||
|
self.template_name = "other.html"
|
||||||
|
|
||||||
|
async def initialize_services(self):
|
||||||
|
"""Initialize services from ServiceRegistry"""
|
||||||
|
other_scanner = await ServiceRegistry.get_other_scanner()
|
||||||
|
update_service = await ServiceRegistry.get_model_update_service()
|
||||||
|
self.service = OtherModelService(other_scanner, update_service=update_service)
|
||||||
|
self.set_model_update_service(update_service)
|
||||||
|
|
||||||
|
# Attach service dependencies
|
||||||
|
self.attach_service(self.service)
|
||||||
|
|
||||||
|
def setup_routes(self, app: web.Application, prefix: str = "other"):
|
||||||
|
"""Setup Other-model routes"""
|
||||||
|
# Schedule service initialization on app startup
|
||||||
|
app.on_startup.append(lambda _: self.initialize_services())
|
||||||
|
|
||||||
|
# Setup common routes with 'other' prefix (includes page route)
|
||||||
|
super().setup_routes(app, prefix)
|
||||||
|
|
||||||
|
def setup_specific_routes(self, registrar: ModelRouteRegistrar, prefix: str):
|
||||||
|
"""Setup Other-model-specific routes"""
|
||||||
|
# Other-model info by name
|
||||||
|
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/info/{name}', prefix, self.get_other_model_info)
|
||||||
|
# Other-model roots grouped by sub_type (text_encoders + legacy clip
|
||||||
|
# are aggregated under text_encoder)
|
||||||
|
registrar.add_prefixed_route('GET', '/api/lm/{prefix}/roots_by_subtype', prefix, self.get_roots_by_subtype)
|
||||||
|
|
||||||
|
def _validate_civitai_model_type(self, model_type: str) -> bool:
|
||||||
|
"""Validate CivitAI model type for other models.
|
||||||
|
|
||||||
|
Accepts retired CivitAI types (CLIP, CLIPVision) as well — grandfathered
|
||||||
|
models on CivitAI still carry them. Types whose sub_type is currently
|
||||||
|
disabled (or every type while the opt-in feature is off) are rejected.
|
||||||
|
"""
|
||||||
|
normalized = (model_type or "").strip().lower()
|
||||||
|
if normalized not in VALID_OTHER_CIVITAI_TYPES:
|
||||||
|
return False
|
||||||
|
if not self._settings.is_other_models_enabled():
|
||||||
|
return False
|
||||||
|
|
||||||
|
sub_type = CIVITAI_TYPE_TO_OTHER_SUB_TYPE.get(normalized)
|
||||||
|
if sub_type is None:
|
||||||
|
# CivitAI "Other" has no sub_type of its own; it is only usable
|
||||||
|
# while at least one sub_type is enabled.
|
||||||
|
return bool(self._settings.get_enabled_other_sub_types())
|
||||||
|
return self._settings.is_other_sub_type_enabled(sub_type)
|
||||||
|
|
||||||
|
def _get_page_context_provider(self):
|
||||||
|
"""Expose the opt-in feature state to the Other Models page template."""
|
||||||
|
return self._page_context_for_other
|
||||||
|
|
||||||
|
def _page_context_for_other(self, request: web.Request) -> Dict[str, Any]:
|
||||||
|
if not self._settings.is_other_models_enabled():
|
||||||
|
return {"other_disabled": True, "other_no_paths": False}
|
||||||
|
|
||||||
|
# Enabled but nothing to scan: folder paths for the managed sub_types
|
||||||
|
# resolved to no existing folder. Render an actionable empty state
|
||||||
|
# instead of an apparently broken empty grid.
|
||||||
|
standalone_mode = os.environ.get("LORA_MANAGER_STANDALONE", "0") == "1"
|
||||||
|
context = {
|
||||||
|
"other_disabled": False,
|
||||||
|
"other_no_paths": not bool(config.other_roots),
|
||||||
|
"standalone_mode": standalone_mode,
|
||||||
|
}
|
||||||
|
if standalone_mode:
|
||||||
|
# The empty state points at the Model Paths settings section and
|
||||||
|
# shows the settings.json path as a fallback reference.
|
||||||
|
context["settings_file"] = getattr(self._settings, "settings_file", "") or ""
|
||||||
|
return context
|
||||||
|
|
||||||
|
def _get_expected_model_types(self) -> str:
|
||||||
|
"""Get expected model types string for error messages"""
|
||||||
|
return "VAE, Upscaler, TextEncoder, CLIPVision, Controlnet, or Other"
|
||||||
|
|
||||||
|
def _parse_specific_params(self, request: web.Request) -> Dict[str, Any]:
|
||||||
|
"""Parse other-model-specific parameters (none in Phase 1)."""
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async def get_roots_by_subtype(self, request: web.Request) -> web.Response:
|
||||||
|
"""Return other-model roots grouped by sub_type.
|
||||||
|
|
||||||
|
Aggregates the per-folder_paths-key roots from config
|
||||||
|
(``text_encoders`` and the legacy ``clip`` key both land under
|
||||||
|
``text_encoder``).
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
roots_by_subtype: Dict[str, List[str]] = {}
|
||||||
|
for key, roots in (config.other_folder_roots or {}).items():
|
||||||
|
sub_type = OTHER_MODEL_FOLDER_SUBTYPES.get(key)
|
||||||
|
if not sub_type:
|
||||||
|
continue
|
||||||
|
bucket = roots_by_subtype.setdefault(sub_type, [])
|
||||||
|
for root in roots:
|
||||||
|
if root and root not in bucket:
|
||||||
|
bucket.append(root)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": True, "roots_by_subtype": roots_by_subtype}
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error getting other roots by sub_type: {e}", exc_info=True)
|
||||||
|
return web.json_response(
|
||||||
|
{"success": False, "error": str(e)}, status=500
|
||||||
|
)
|
||||||
|
|
||||||
|
async def get_other_model_info(self, request: web.Request) -> web.Response:
|
||||||
|
"""Get detailed information for a specific other model by name"""
|
||||||
|
try:
|
||||||
|
name = request.match_info.get('name', '')
|
||||||
|
model_info = await self.service.get_model_info_by_name(name) # pyright: ignore[reportAttributeAccessIssue]
|
||||||
|
|
||||||
|
if model_info:
|
||||||
|
return web.json_response(model_info)
|
||||||
|
else:
|
||||||
|
return web.json_response({"error": "Model not found"}, status=404)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error in get_other_model_info: {e}", exc_info=True)
|
||||||
|
return web.json_response({"error": str(e)}, status=500)
|
||||||
@@ -49,6 +49,31 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
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"),
|
||||||
|
RouteDefinition("POST", "/api/lm/recipe/lora/restore", "restore_lora"),
|
||||||
|
RouteDefinition(
|
||||||
|
"GET",
|
||||||
|
"/api/lm/recipe/{recipe_id}/lora/{lora_index}/reconnect-suggestions",
|
||||||
|
"get_reconnect_suggestions",
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/recipe/lora/mark-hash-invalid", "mark_lora_hash_invalid"
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/recipe/checkpoint/reconnect", "reconnect_checkpoint"
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST", "/api/lm/recipe/checkpoint/restore", "restore_checkpoint"
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"GET",
|
||||||
|
"/api/lm/recipe/{recipe_id}/checkpoint/reconnect-suggestions",
|
||||||
|
"get_checkpoint_reconnect_suggestions",
|
||||||
|
),
|
||||||
|
RouteDefinition(
|
||||||
|
"POST",
|
||||||
|
"/api/lm/recipe/checkpoint/mark-hash-invalid",
|
||||||
|
"mark_checkpoint_hash_invalid",
|
||||||
|
),
|
||||||
RouteDefinition("GET", "/api/lm/recipes/find-duplicates", "find_duplicates"),
|
RouteDefinition("GET", "/api/lm/recipes/find-duplicates", "find_duplicates"),
|
||||||
RouteDefinition("POST", "/api/lm/recipes/bulk-delete", "bulk_delete"),
|
RouteDefinition("POST", "/api/lm/recipes/bulk-delete", "bulk_delete"),
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
@@ -59,11 +84,6 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
"GET", "/api/lm/recipes/for-checkpoint", "get_recipes_for_checkpoint"
|
"GET", "/api/lm/recipes/for-checkpoint", "get_recipes_for_checkpoint"
|
||||||
),
|
),
|
||||||
RouteDefinition("GET", "/api/lm/recipes/scan", "scan_recipes"),
|
RouteDefinition("GET", "/api/lm/recipes/scan", "scan_recipes"),
|
||||||
RouteDefinition("POST", "/api/lm/recipes/repair", "repair_recipes"),
|
|
||||||
RouteDefinition("POST", "/api/lm/recipes/cancel-repair", "cancel_repair"),
|
|
||||||
RouteDefinition("POST", "/api/lm/recipe/{recipe_id}/repair", "repair_recipe"),
|
|
||||||
RouteDefinition("POST", "/api/lm/recipes/repair-bulk", "repair_recipes_bulk"),
|
|
||||||
RouteDefinition("GET", "/api/lm/recipes/repair-progress", "get_repair_progress"),
|
|
||||||
RouteDefinition("POST", "/api/lm/recipes/rematch", "rematch_recipes"),
|
RouteDefinition("POST", "/api/lm/recipes/rematch", "rematch_recipes"),
|
||||||
RouteDefinition("POST", "/api/lm/recipes/rematch-bulk", "rematch_recipes_bulk"),
|
RouteDefinition("POST", "/api/lm/recipes/rematch-bulk", "rematch_recipes_bulk"),
|
||||||
RouteDefinition("POST", "/api/lm/recipe/{recipe_id}/rematch", "rematch_recipe"),
|
RouteDefinition("POST", "/api/lm/recipe/{recipe_id}/rematch", "rematch_recipe"),
|
||||||
@@ -90,6 +110,11 @@ ROUTE_DEFINITIONS: tuple[RouteDefinition, ...] = (
|
|||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"POST", "/api/lm/recipe/{recipe_id}/reimport", "reimport_recipe"
|
"POST", "/api/lm/recipe/{recipe_id}/reimport", "reimport_recipe"
|
||||||
),
|
),
|
||||||
|
# The companion browser extension only ever issues GET requests, so the
|
||||||
|
# payload-based re-import variant must also be reachable via GET.
|
||||||
|
RouteDefinition(
|
||||||
|
"GET", "/api/lm/recipe/{recipe_id}/reimport", "reimport_recipe"
|
||||||
|
),
|
||||||
RouteDefinition(
|
RouteDefinition(
|
||||||
"POST", "/api/lm/recipe/{recipe_id}/send-workflow", "send_recipe_workflow"
|
"POST", "/api/lm/recipe/{recipe_id}/send-workflow", "send_recipe_workflow"
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -21,7 +21,20 @@ NETWORK_EXCEPTIONS = (ClientError, OSError, asyncio.TimeoutError)
|
|||||||
# otherwise delete them because they are untracked and, in released tags,
|
# otherwise delete them because they are untracked and, in released tags,
|
||||||
# not listed in ``.gitignore``. ``-e`` excludes a path from cleaning
|
# not listed in ``.gitignore``. ``-e`` excludes a path from cleaning
|
||||||
# regardless of whether it is ignored.
|
# regardless of whether it is ignored.
|
||||||
_PRESERVE_DIRS = ('settings.json', 'civitai', 'wildcards', 'backups', 'stats', 'logs', 'cache', 'model_cache')
|
# ``cache`` covers the resolved cache tree (cache/model, cache/recipe,
|
||||||
|
# cache/fts, ...); the legacy ``recipe_cache`` / ``model_cache`` directories
|
||||||
|
# are listed too because a portable install can predate the cache/ move.
|
||||||
|
_PRESERVE_DIRS = (
|
||||||
|
'settings.json',
|
||||||
|
'civitai',
|
||||||
|
'wildcards',
|
||||||
|
'backups',
|
||||||
|
'stats',
|
||||||
|
'logs',
|
||||||
|
'cache',
|
||||||
|
'model_cache',
|
||||||
|
'recipe_cache',
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _clean_excludes() -> List[str]:
|
def _clean_excludes() -> List[str]:
|
||||||
|
|||||||
@@ -0,0 +1,135 @@
|
|||||||
|
"""In-memory store for the LoRA Manager page's active filters.
|
||||||
|
|
||||||
|
The manager page keeps its filter state in localStorage for its own
|
||||||
|
restoration, but the ComfyUI node autocomplete runs in a potentially
|
||||||
|
different browser/origin (or Electron shell) where that storage is not
|
||||||
|
shared. This store mirrors the active filters server-side so the
|
||||||
|
``/api/lm/{prefix}/relative-paths`` endpoint can inject them into
|
||||||
|
autocomplete searches regardless of which client set them.
|
||||||
|
|
||||||
|
State is process-local and intentionally not persisted; the manager page
|
||||||
|
re-pushes its restored state on load.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Keys copied from the manager page's persisted filter snapshot.
|
||||||
|
_FILTER_KEYS = (
|
||||||
|
"baseModel",
|
||||||
|
"tags",
|
||||||
|
"autoTags",
|
||||||
|
"modelTypes",
|
||||||
|
"tagLogic",
|
||||||
|
"license",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ActiveFiltersStore:
|
||||||
|
"""Process-local store of active filters, keyed by model type."""
|
||||||
|
|
||||||
|
_instance: Optional["ActiveFiltersStore"] = None
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._filters: Dict[str, Dict[str, Any]] = {}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_instance(cls) -> "ActiveFiltersStore":
|
||||||
|
if cls._instance is None:
|
||||||
|
cls._instance = cls()
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def reset_instance(cls) -> None:
|
||||||
|
"""Drop the singleton (test isolation)."""
|
||||||
|
cls._instance = None
|
||||||
|
|
||||||
|
def set_filters(self, model_type: str, payload: Dict[str, Any]) -> None:
|
||||||
|
"""Replace the stored active filters for a model type.
|
||||||
|
|
||||||
|
Only recognized keys are kept; everything else is discarded.
|
||||||
|
"""
|
||||||
|
filters = payload.get("filters")
|
||||||
|
sanitized: Dict[str, Any] = {
|
||||||
|
"activeFolder": payload.get("activeFolder"),
|
||||||
|
"recursiveSearch": bool(payload.get("recursiveSearch", True)),
|
||||||
|
"filters": (
|
||||||
|
{key: filters[key] for key in _FILTER_KEYS if key in filters}
|
||||||
|
if isinstance(filters, dict)
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
}
|
||||||
|
self._filters[model_type] = sanitized
|
||||||
|
|
||||||
|
def get_filters(self, model_type: str) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Return the stored payload for a model type, or None if unset."""
|
||||||
|
return self._filters.get(model_type)
|
||||||
|
|
||||||
|
def clear(self, model_type: str) -> None:
|
||||||
|
self._filters.pop(model_type, None)
|
||||||
|
|
||||||
|
|
||||||
|
def active_filters_to_query_kwargs(payload: Optional[Dict[str, Any]]) -> Dict[str, Any]:
|
||||||
|
"""Map a stored active-filters payload to ``search_relative_paths`` kwargs.
|
||||||
|
|
||||||
|
Mirrors the query-param mapping that the ComfyUI autocomplete used to
|
||||||
|
build client-side from localStorage (web/comfyui/autocomplete.js).
|
||||||
|
"""
|
||||||
|
kwargs: Dict[str, Any] = {}
|
||||||
|
if not payload:
|
||||||
|
return kwargs
|
||||||
|
|
||||||
|
active_folder = payload.get("activeFolder")
|
||||||
|
recursive = payload.get("recursiveSearch", True)
|
||||||
|
|
||||||
|
if active_folder and active_folder != "null":
|
||||||
|
kwargs["folder"] = active_folder
|
||||||
|
elif not recursive:
|
||||||
|
# Root folder with recursion disabled mirrors the page list,
|
||||||
|
# which matches only root-level files via folder=''.
|
||||||
|
kwargs["folder"] = ""
|
||||||
|
|
||||||
|
filters = payload.get("filters")
|
||||||
|
if isinstance(filters, dict):
|
||||||
|
base_models = filters.get("baseModel")
|
||||||
|
if isinstance(base_models, list):
|
||||||
|
kwargs["base_models"] = [m for m in base_models if m]
|
||||||
|
|
||||||
|
for source_key, target_key in (("tags", "tags"), ("autoTags", "auto_tags")):
|
||||||
|
states = filters.get(source_key)
|
||||||
|
if isinstance(states, dict):
|
||||||
|
mapped = {
|
||||||
|
tag: state
|
||||||
|
for tag, state in states.items()
|
||||||
|
if state in ("include", "exclude")
|
||||||
|
}
|
||||||
|
if mapped:
|
||||||
|
kwargs[target_key] = mapped
|
||||||
|
|
||||||
|
model_types = filters.get("modelTypes")
|
||||||
|
if isinstance(model_types, list):
|
||||||
|
kwargs["model_types"] = [t for t in model_types if t]
|
||||||
|
|
||||||
|
tag_logic = filters.get("tagLogic")
|
||||||
|
if tag_logic:
|
||||||
|
kwargs["tag_logic"] = tag_logic
|
||||||
|
|
||||||
|
license_filter = filters.get("license")
|
||||||
|
if isinstance(license_filter, dict):
|
||||||
|
no_credit = license_filter.get("noCredit")
|
||||||
|
if no_credit == "include":
|
||||||
|
kwargs["credit_required"] = False
|
||||||
|
elif no_credit == "exclude":
|
||||||
|
kwargs["credit_required"] = True
|
||||||
|
allow_selling = license_filter.get("allowSelling")
|
||||||
|
if allow_selling == "include":
|
||||||
|
kwargs["allow_selling_generated_content"] = True
|
||||||
|
elif allow_selling == "exclude":
|
||||||
|
kwargs["allow_selling_generated_content"] = False
|
||||||
|
|
||||||
|
kwargs["recursive"] = recursive
|
||||||
|
return kwargs
|
||||||
@@ -19,16 +19,21 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
import re
|
import re
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
import aiohttp
|
|
||||||
|
|
||||||
import os
|
|
||||||
|
|
||||||
from ...config import config
|
from ...config import config
|
||||||
from ..llm_service import LLMService
|
from ..llm_service import LLMService
|
||||||
|
from ..model_sources import (
|
||||||
|
ModelCardContext,
|
||||||
|
ModelSourceCache,
|
||||||
|
get_source,
|
||||||
|
resolve_source_ref,
|
||||||
|
source_label,
|
||||||
|
)
|
||||||
|
from ..model_sources.hydration import load_model_card, resolve_site_base_model
|
||||||
from ..websocket_manager import ws_manager
|
from ..websocket_manager import ws_manager
|
||||||
from .post_processor import PostProcessor
|
from .post_processor import PostProcessor
|
||||||
from .skill_registry import SkillRegistry
|
from .skill_registry import SkillRegistry
|
||||||
@@ -255,6 +260,11 @@ class AgentService:
|
|||||||
llm = await self._ensure_llm()
|
llm = await self._ensure_llm()
|
||||||
llm_configured = llm.is_configured() if skill.llm_required else True
|
llm_configured = llm.is_configured() if skill.llm_required else True
|
||||||
|
|
||||||
|
# A collection repository holds many model files under one source id;
|
||||||
|
# this memo keeps the README and the repository metadata from being
|
||||||
|
# re-fetched once per file. It lives for this run only.
|
||||||
|
source_cache = ModelSourceCache()
|
||||||
|
|
||||||
for model_path in model_paths:
|
for model_path in model_paths:
|
||||||
model_filename = os.path.basename(model_path)
|
model_filename = os.path.basename(model_path)
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -267,24 +277,50 @@ class AgentService:
|
|||||||
from ...metadata_ops import read_metadata
|
from ...metadata_ops import read_metadata
|
||||||
metadata = await read_metadata(model_path)
|
metadata = await read_metadata(model_path)
|
||||||
|
|
||||||
# Fast-fail: enrich_hf_metadata requires hf_url to have HF README context
|
# Fast-fail: enrich_hf_metadata needs an external model source
|
||||||
if skill_name == "enrich_hf_metadata" and not metadata.get("hf_url", ""):
|
# that exposes an accessible model card.
|
||||||
logger.info(
|
if skill_name == "enrich_hf_metadata":
|
||||||
"[%s] SKIP %s — no hf_url in metadata",
|
skip_reason = self._enrichment_skip_reason(metadata)
|
||||||
skill_name, model_filename,
|
if skip_reason:
|
||||||
)
|
logger.info(
|
||||||
skipped_count += 1
|
"[%s] SKIP %s — %s",
|
||||||
skip_model = True
|
skill_name, model_filename, skip_reason,
|
||||||
|
)
|
||||||
|
skipped_count += 1
|
||||||
|
skip_model = True
|
||||||
|
|
||||||
if not skip_model:
|
if not skip_model:
|
||||||
prompt_vars: Dict[str, Any] = {"model_path": model_path}
|
# The site's own data is deterministic and must land whether
|
||||||
if skill.llm_required and llm_configured:
|
# or not an LLM is available: a user without a key still gets
|
||||||
prompt_vars = await self._build_prompt_context(
|
# the author summary, the example images and the tags.
|
||||||
skill_name, model_path, metadata, registry, llm,
|
source_vars, source_context = await self._load_source_card(
|
||||||
|
model_path, metadata, cache=source_cache,
|
||||||
|
)
|
||||||
|
resolved_base_model = ""
|
||||||
|
if skill_name == "enrich_hf_metadata" and not (
|
||||||
|
metadata.get("base_model") or ""
|
||||||
|
).strip():
|
||||||
|
resolved_base_model = await self._resolve_site_base_model(
|
||||||
|
source_context,
|
||||||
)
|
)
|
||||||
|
|
||||||
llm_response: Optional[Dict[str, Any]] = None
|
llm_response: Optional[Dict[str, Any]] = None
|
||||||
if skill.llm_required and llm_configured:
|
if skill.llm_required and not llm_configured:
|
||||||
|
# Without a provider the deterministic model-source data
|
||||||
|
# still lands; the LLM-only fields simply stay untouched.
|
||||||
|
logger.info(
|
||||||
|
"[%s] No LLM configured for %s — applying %s data only",
|
||||||
|
skill_name, model_filename,
|
||||||
|
"model-source"
|
||||||
|
if not source_context.is_empty()
|
||||||
|
else "README",
|
||||||
|
)
|
||||||
|
elif skill.llm_required:
|
||||||
|
prompt_vars = await self._build_prompt_context(
|
||||||
|
skill_name, model_path, metadata, registry, llm,
|
||||||
|
source_vars=source_vars,
|
||||||
|
source_context=source_context,
|
||||||
|
)
|
||||||
prompt_template = registry.load_prompt(skill_name)
|
prompt_template = registry.load_prompt(skill_name)
|
||||||
rendered = _render_prompt(prompt_template, prompt_vars)
|
rendered = _render_prompt(prompt_template, prompt_vars)
|
||||||
llm_response = await llm.chat_completion_json(
|
llm_response = await llm.chat_completion_json(
|
||||||
@@ -307,7 +343,9 @@ class AgentService:
|
|||||||
model_path=model_path,
|
model_path=model_path,
|
||||||
llm_output=llm_response or {},
|
llm_output=llm_response or {},
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
readme_content=prompt_vars.get("readme_content_full", ""),
|
readme_content=source_vars.get("readme_content_full", ""),
|
||||||
|
source_context=source_context,
|
||||||
|
resolved_base_model=resolved_base_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
if model_result.get("success", True):
|
if model_result.get("success", True):
|
||||||
@@ -358,6 +396,28 @@ class AgentService:
|
|||||||
# Base model grouping (keeps the prompt compact)
|
# Base model grouping (keeps the prompt compact)
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _enrichment_skip_reason(metadata: Dict[str, Any]) -> str:
|
||||||
|
"""Return why ``enrich_hf_metadata`` cannot run, or ``""`` if it can.
|
||||||
|
|
||||||
|
Distinguishes the three cases the user can act on: no source linked,
|
||||||
|
a source we don't know, and a known source whose model card is not
|
||||||
|
reachable from the backend (TensorArt).
|
||||||
|
"""
|
||||||
|
|
||||||
|
ref = resolve_source_ref(metadata)
|
||||||
|
if ref is None:
|
||||||
|
return "no model source linked (source_url missing)"
|
||||||
|
source = get_source(ref.platform)
|
||||||
|
if source is None:
|
||||||
|
return f"unsupported model source platform '{ref.platform}'"
|
||||||
|
if not source.supports_enrichment:
|
||||||
|
return (
|
||||||
|
f"{source.label} does not expose a model card to the backend; "
|
||||||
|
"AI metadata enrichment is not available for this source"
|
||||||
|
)
|
||||||
|
return ""
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _format_base_models(models: List[str]) -> str:
|
def _format_base_models(models: List[str]) -> str:
|
||||||
"""Format the base model list as a flat, one-per-line list.
|
"""Format the base model list as a flat, one-per-line list.
|
||||||
@@ -368,6 +428,82 @@ class AgentService:
|
|||||||
"""
|
"""
|
||||||
return "\n".join(f"- {m}" for m in models)
|
return "\n".join(f"- {m}" for m in models)
|
||||||
|
|
||||||
|
async def _load_source_card(
|
||||||
|
self,
|
||||||
|
model_path: str,
|
||||||
|
metadata: Dict[str, Any],
|
||||||
|
*,
|
||||||
|
cache: Optional[ModelSourceCache] = None,
|
||||||
|
) -> tuple[Dict[str, Any], ModelCardContext]:
|
||||||
|
"""Fetch the model card and site-published extras for one model.
|
||||||
|
|
||||||
|
Runs for every source-backed enrichment regardless of LLM
|
||||||
|
availability, because everything it returns is deterministic data that
|
||||||
|
should be applied even without a configured provider.
|
||||||
|
|
||||||
|
*cache* is the per-run memo created by :meth:`execute_skill`. The
|
||||||
|
README is repository-wide, so it is fetched once per source id; only
|
||||||
|
successful reads are memoised, leaving a transient failure to be
|
||||||
|
retried for the next file.
|
||||||
|
"""
|
||||||
|
|
||||||
|
variables: Dict[str, Any] = {
|
||||||
|
"asset_base_url": "",
|
||||||
|
"source_description": "",
|
||||||
|
"source_base_model": "",
|
||||||
|
"source_official_tags": "",
|
||||||
|
"source_example_images": "",
|
||||||
|
"source_trigger_words": "",
|
||||||
|
"readme_content": "(README not available)",
|
||||||
|
"readme_content_full": "",
|
||||||
|
}
|
||||||
|
|
||||||
|
ref = resolve_source_ref(metadata)
|
||||||
|
source = get_source(ref.platform) if ref is not None else None
|
||||||
|
if ref is None or source is None or not source.supports_enrichment:
|
||||||
|
return variables, ModelCardContext()
|
||||||
|
|
||||||
|
raw_basename = os.path.splitext(os.path.basename(model_path))[0]
|
||||||
|
variables["asset_base_url"] = source.asset_base_url(ref.source_id)
|
||||||
|
|
||||||
|
readme = await load_model_card(source, ref.source_id, cache)
|
||||||
|
|
||||||
|
# Sites such as ModelScope keep part of the model card outside the
|
||||||
|
# README (author summary, curated tags, per-file example images). The
|
||||||
|
# recorded hash identifies the file even after the user renames it.
|
||||||
|
card_context = await source.fetch_model_card_context(
|
||||||
|
ref.source_id,
|
||||||
|
os.path.basename(model_path),
|
||||||
|
sha256=(metadata.get("sha256") or "").strip(),
|
||||||
|
cache=cache,
|
||||||
|
)
|
||||||
|
variables["source_description"] = card_context.description
|
||||||
|
variables["source_base_model"] = card_context.base_model
|
||||||
|
variables["source_official_tags"] = "\n".join(
|
||||||
|
f"- {tag}" for tag in card_context.official_tags
|
||||||
|
)
|
||||||
|
variables["source_example_images"] = "\n".join(
|
||||||
|
f"- {url}" for url in card_context.example_images
|
||||||
|
)
|
||||||
|
variables["source_trigger_words"] = ", ".join(card_context.trigger_words)
|
||||||
|
|
||||||
|
# Trim README to the section relevant to this model file
|
||||||
|
# (collection repos often have multiple models in one README).
|
||||||
|
if readme and raw_basename:
|
||||||
|
trimmed = extract_relevant_section(readme, raw_basename)
|
||||||
|
cleaned = clean_readme_for_llm(trimmed) if trimmed else ""
|
||||||
|
else:
|
||||||
|
cleaned = clean_readme_for_llm(readme) if readme else ""
|
||||||
|
variables["readme_content"] = cleaned if cleaned else "(README not available)"
|
||||||
|
variables["readme_content_full"] = readme or ""
|
||||||
|
|
||||||
|
return variables, card_context
|
||||||
|
|
||||||
|
async def _resolve_site_base_model(self, source_context: ModelCardContext) -> str:
|
||||||
|
"""Resolve the site's base-model hints to a canonical name, or ``""``."""
|
||||||
|
|
||||||
|
return await resolve_site_base_model(source_context)
|
||||||
|
|
||||||
async def _build_prompt_context(
|
async def _build_prompt_context(
|
||||||
self,
|
self,
|
||||||
skill_name: str,
|
skill_name: str,
|
||||||
@@ -375,19 +511,45 @@ class AgentService:
|
|||||||
metadata: Dict[str, Any],
|
metadata: Dict[str, Any],
|
||||||
registry: SkillRegistry,
|
registry: SkillRegistry,
|
||||||
llm: Any,
|
llm: Any,
|
||||||
|
*,
|
||||||
|
source_vars: Optional[Dict[str, Any]] = None,
|
||||||
|
source_context: Optional[ModelCardContext] = None,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""Gather variables for the skill's prompt template.
|
"""Gather variables for the skill's prompt template.
|
||||||
|
|
||||||
Reads metadata, fetches the HF README (if applicable), lists available
|
Reads metadata, fetches the model card (unless a pre-fetched
|
||||||
|
*source_vars* / *source_context* pair is supplied), lists available
|
||||||
base models, loads user priority tags, and returns a dict that maps to
|
base models, loads user priority tags, and returns a dict that maps to
|
||||||
``{{variable}}`` placeholders in ``prompt.md``.
|
``{{variable}}`` placeholders in ``prompt.md``.
|
||||||
"""
|
"""
|
||||||
from ...metadata_ops import identify_model_type, list_base_models
|
from ...metadata_ops import identify_model_type, list_base_models
|
||||||
from ..settings_manager import SettingsManager
|
from ..settings_manager import SettingsManager
|
||||||
|
|
||||||
|
if source_vars is None or source_context is None:
|
||||||
|
source_vars, source_context = await self._load_source_card(
|
||||||
|
model_path, metadata,
|
||||||
|
)
|
||||||
|
|
||||||
context: Dict[str, Any] = {
|
context: Dict[str, Any] = {
|
||||||
"model_path": model_path,
|
"model_path": model_path,
|
||||||
"model_basename": "",
|
"model_basename": "",
|
||||||
|
# Canonical external-source variables
|
||||||
|
"source_url": "",
|
||||||
|
"source_id": "",
|
||||||
|
"source_platform": "",
|
||||||
|
"source_label": "",
|
||||||
|
"asset_base_url": "",
|
||||||
|
# Site-provided card extras (see ModelSource.fetch_model_card_context)
|
||||||
|
"source_description": "",
|
||||||
|
"source_base_model": "",
|
||||||
|
"source_official_tags": "",
|
||||||
|
"source_example_images": "",
|
||||||
|
"source_trigger_words": "",
|
||||||
|
# Carrier for the structured context handed to the post-processor;
|
||||||
|
# never rendered into the prompt.
|
||||||
|
"source_context": ModelCardContext(),
|
||||||
|
# Legacy Hugging Face aliases (kept so older prompt templates and
|
||||||
|
# third-party skills keep rendering)
|
||||||
"hf_url": "",
|
"hf_url": "",
|
||||||
"repo": "",
|
"repo": "",
|
||||||
"readme_content": "",
|
"readme_content": "",
|
||||||
@@ -407,26 +569,33 @@ class AgentService:
|
|||||||
"base_model": metadata.get("base_model", ""),
|
"base_model": metadata.get("base_model", ""),
|
||||||
"tags": metadata.get("tags", []),
|
"tags": metadata.get("tags", []),
|
||||||
"modelDescription": metadata.get("modelDescription", ""),
|
"modelDescription": metadata.get("modelDescription", ""),
|
||||||
"trainedWords": metadata.get("trainedWords", []),
|
|
||||||
"sha256": (metadata.get("sha256") or "")[:16] + "..." if metadata.get("sha256") else "",
|
"sha256": (metadata.get("sha256") or "")[:16] + "..." if metadata.get("sha256") else "",
|
||||||
"size": metadata.get("size", 0),
|
"size": metadata.get("size", 0),
|
||||||
}
|
}
|
||||||
|
|
||||||
hf_url = metadata.get("hf_url", "")
|
ref = resolve_source_ref(metadata)
|
||||||
context["hf_url"] = hf_url
|
if ref is not None:
|
||||||
repo = self._extract_repo_from_url(hf_url) if hf_url else ""
|
context["source_url"] = ref.url
|
||||||
context["repo"] = repo or ""
|
context["source_id"] = ref.source_id
|
||||||
if repo:
|
context["source_platform"] = ref.platform
|
||||||
readme = await self._fetch_readme(repo)
|
context["source_label"] = source_label(ref.platform, ref.platform)
|
||||||
# Trim README to the section relevant to this model file
|
if ref.platform == "huggingface":
|
||||||
# (collection repos often have multiple models in one README).
|
context["hf_url"] = ref.url
|
||||||
if readme and raw_basename:
|
context["repo"] = ref.source_id
|
||||||
trimmed = extract_relevant_section(readme, raw_basename)
|
|
||||||
cleaned = clean_readme_for_llm(trimmed) if trimmed else ""
|
source = get_source(ref.platform) if ref is not None else None
|
||||||
else:
|
if ref is not None and source is not None and source.supports_enrichment:
|
||||||
cleaned = clean_readme_for_llm(readme) if readme else ""
|
# Values fetched once by _load_source_card and shared with the
|
||||||
context["readme_content"] = cleaned if cleaned else "(README not available)"
|
# post-processor, so the network is not hit twice per model.
|
||||||
context["readme_content_full"] = readme or ""
|
context["asset_base_url"] = source_vars["asset_base_url"]
|
||||||
|
context["source_context"] = source_context
|
||||||
|
context["source_description"] = source_vars["source_description"]
|
||||||
|
context["source_base_model"] = source_vars["source_base_model"]
|
||||||
|
context["source_official_tags"] = source_vars["source_official_tags"]
|
||||||
|
context["source_example_images"] = source_vars["source_example_images"]
|
||||||
|
context["source_trigger_words"] = source_vars["source_trigger_words"]
|
||||||
|
context["readme_content"] = source_vars["readme_content"]
|
||||||
|
context["readme_content_full"] = source_vars["readme_content_full"]
|
||||||
|
|
||||||
try:
|
try:
|
||||||
raw_models = await list_base_models()
|
raw_models = await list_base_models()
|
||||||
@@ -459,20 +628,14 @@ class AgentService:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def _fetch_readme(repo: str) -> str:
|
async def _fetch_readme(repo: str) -> str:
|
||||||
"""Fetch README.md from HuggingFace (tries ``main``, then ``master``)."""
|
"""Fetch a Hugging Face README (tries ``main``, then ``master``).
|
||||||
async with aiohttp.ClientSession(
|
|
||||||
headers={"User-Agent": "ComfyUI-LoRA-Manager/1.0"},
|
Kept for backward compatibility; new code should go through the
|
||||||
timeout=aiohttp.ClientTimeout(total=30),
|
model-source registry so every supported site works.
|
||||||
) as session:
|
"""
|
||||||
for branch in ("main", "master"):
|
from ..model_sources import HuggingFaceSource
|
||||||
url = f"https://huggingface.co/{repo}/raw/{branch}/README.md"
|
|
||||||
try:
|
return await HuggingFaceSource().fetch_model_card(repo)
|
||||||
async with session.get(url) as resp:
|
|
||||||
if resp.status == 200:
|
|
||||||
return await resp.text()
|
|
||||||
except Exception as exc:
|
|
||||||
logger.debug("Failed to fetch README from %s: %s", url, exc)
|
|
||||||
return ""
|
|
||||||
|
|
||||||
async def _emit_progress(
|
async def _emit_progress(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -0,0 +1,94 @@
|
|||||||
|
"""Map a site-reported base model onto this system's canonical vocabulary.
|
||||||
|
|
||||||
|
Model sites name base models in their own terms: ModelScope publishes
|
||||||
|
``krea/Krea-2-Turbo`` and ``KREA_2_TURBO`` where this system expects the
|
||||||
|
canonical ``Krea 2``. Turning one into the other is normally the LLM's job;
|
||||||
|
this module resolves the cases that can be decided safely so the canonical
|
||||||
|
field is still populated when the LLM returns nothing usable for it.
|
||||||
|
|
||||||
|
The resolver is deliberately strict, because a wrong base model written with
|
||||||
|
apparent authority is worse than no value at all:
|
||||||
|
|
||||||
|
* it only ever returns a name that is already present in *known_names*;
|
||||||
|
* matching is on the normalised form (lowercased, non-alphanumerics removed),
|
||||||
|
so separators and casing are ignored but nothing is inferred;
|
||||||
|
* a bounded set of published variant suffixes may be stripped, and only when
|
||||||
|
the remainder still matches a known name exactly.
|
||||||
|
|
||||||
|
Anything it cannot decide returns ``""``, and the caller falls back to the LLM.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
from typing import Iterable, Sequence
|
||||||
|
|
||||||
|
#: Variant suffixes sites append to a base-model *family* name. Stripping one
|
||||||
|
#: is only attempted when the remainder matches a known name exactly, so an
|
||||||
|
#: unrecognised suffix can never produce a bogus match.
|
||||||
|
_VARIANT_SUFFIXES: tuple[str, ...] = (
|
||||||
|
"turbo",
|
||||||
|
"schnell",
|
||||||
|
"lightning",
|
||||||
|
"dev",
|
||||||
|
"beta",
|
||||||
|
"alpha",
|
||||||
|
)
|
||||||
|
|
||||||
|
_NON_ALNUM = re.compile(r"[^a-z0-9]+")
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize(value: str) -> str:
|
||||||
|
"""Return the comparison form of *value*.
|
||||||
|
|
||||||
|
Lowercases and drops every non-alphanumeric character, so ``KREA_2``,
|
||||||
|
``Krea 2``, ``krea-2`` and ``krea.2`` all collapse to ``krea2``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
return _NON_ALNUM.sub("", (value or "").lower())
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_base_model(
|
||||||
|
hints: Iterable[str], known_names: Sequence[str]
|
||||||
|
) -> str:
|
||||||
|
"""Return the canonical base model that *hints* refers to, or ``""``.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
hints: Site-reported names, best first (e.g. an architecture enum
|
||||||
|
before a link-style repository id).
|
||||||
|
known_names: The canonical vocabulary; only these are ever returned.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
One of *known_names*, or ``""`` when nothing matches exactly.
|
||||||
|
"""
|
||||||
|
|
||||||
|
normalized: dict[str, str] = {}
|
||||||
|
for name in known_names:
|
||||||
|
key = _normalize(name)
|
||||||
|
if key and key not in normalized:
|
||||||
|
normalized[key] = name
|
||||||
|
if not normalized:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
ordered = [hint for hint in hints if hint]
|
||||||
|
|
||||||
|
# 1. Exact normalised match — the unambiguous case.
|
||||||
|
for hint in ordered:
|
||||||
|
candidate = _normalize(hint)
|
||||||
|
if candidate in normalized:
|
||||||
|
return normalized[candidate]
|
||||||
|
|
||||||
|
# 2. Drop one published variant suffix and retry exactly.
|
||||||
|
for hint in ordered:
|
||||||
|
candidate = _normalize(hint)
|
||||||
|
for suffix in _VARIANT_SUFFIXES:
|
||||||
|
if not candidate.endswith(suffix) or candidate == suffix:
|
||||||
|
continue
|
||||||
|
stem = candidate[: -len(suffix)]
|
||||||
|
if stem in normalized:
|
||||||
|
return normalized[stem]
|
||||||
|
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["resolve_base_model"]
|
||||||
@@ -10,12 +10,16 @@ refresh cache). All actual I/O is delegated to :mod:`~py.metadata_ops`.
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import html
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||||
|
|
||||||
|
if TYPE_CHECKING: # pragma: no cover - typing only
|
||||||
|
from ..model_sources import ModelCardContext
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -42,6 +46,9 @@ class PostProcessor:
|
|||||||
llm_output: Dict[str, Any],
|
llm_output: Dict[str, Any],
|
||||||
metadata: Dict[str, Any],
|
metadata: Dict[str, Any],
|
||||||
readme_content: str = "",
|
readme_content: str = "",
|
||||||
|
source_context: Optional["ModelCardContext"] = None,
|
||||||
|
resolved_base_model: str = "",
|
||||||
|
metadata_source: str = "agent:enrich_hf_metadata",
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""Route *llm_output* to the correct skill post-processor.
|
"""Route *llm_output* to the correct skill post-processor.
|
||||||
|
|
||||||
@@ -49,12 +56,26 @@ class PostProcessor:
|
|||||||
that is converted to HTML and stored as ``modelDescription`` for
|
that is converted to HTML and stored as ``modelDescription`` for
|
||||||
the description tab.
|
the description tab.
|
||||||
|
|
||||||
|
*source_context* carries the extras the model site publishes outside
|
||||||
|
the README (author description, per-file example images, trigger
|
||||||
|
words). It is ``None`` for callers that have none.
|
||||||
|
|
||||||
|
*resolved_base_model* is the canonical base-model name the site's own
|
||||||
|
hints resolve to, used when the LLM did not supply one (which is the
|
||||||
|
normal case when the LLM was skipped).
|
||||||
|
|
||||||
|
*metadata_source* records who produced the metadata. The AI skill
|
||||||
|
keeps its historical value; the deterministic download-time hydration
|
||||||
|
passes its own so the two remain distinguishable. ``llm_enriched_at``
|
||||||
|
is only stamped when *llm_output* actually carries a provider answer.
|
||||||
|
|
||||||
Returns a dict with keys ``success`` (bool), ``updated_fields`` (list),
|
Returns a dict with keys ``success`` (bool), ``updated_fields`` (list),
|
||||||
``preview_downloaded`` (bool), and ``errors`` (list).
|
``preview_downloaded`` (bool), and ``errors`` (list).
|
||||||
"""
|
"""
|
||||||
if skill_name == "enrich_hf_metadata":
|
if skill_name == "enrich_hf_metadata":
|
||||||
return await self._process_enrich_hf_metadata(
|
return await self._process_enrich_hf_metadata(
|
||||||
model_path, llm_output, metadata, readme_content,
|
model_path, llm_output, metadata, readme_content, source_context,
|
||||||
|
resolved_base_model, metadata_source,
|
||||||
)
|
)
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
@@ -72,12 +93,16 @@ class PostProcessor:
|
|||||||
llm_output: Dict[str, Any],
|
llm_output: Dict[str, Any],
|
||||||
metadata: Dict[str, Any],
|
metadata: Dict[str, Any],
|
||||||
readme_content: str = "",
|
readme_content: str = "",
|
||||||
|
source_context: Optional["ModelCardContext"] = None,
|
||||||
|
resolved_base_model: str = "",
|
||||||
|
metadata_source: str = "agent:enrich_hf_metadata",
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
from ...metadata_ops import (
|
from ...metadata_ops import (
|
||||||
apply_metadata_updates,
|
apply_metadata_updates,
|
||||||
download_preview,
|
download_preview,
|
||||||
refresh_cache,
|
refresh_cache,
|
||||||
)
|
)
|
||||||
|
from ..model_sources import get_source, has_external_source, resolve_source_ref
|
||||||
from .skills.enrich_hf_metadata.readme_processor import (
|
from .skills.enrich_hf_metadata.readme_processor import (
|
||||||
convert_readme_to_html,
|
convert_readme_to_html,
|
||||||
extract_gallery_images,
|
extract_gallery_images,
|
||||||
@@ -85,24 +110,49 @@ class PostProcessor:
|
|||||||
extract_relevant_section,
|
extract_relevant_section,
|
||||||
extract_simple_markdown_images,
|
extract_simple_markdown_images,
|
||||||
extract_html_img_tags,
|
extract_html_img_tags,
|
||||||
extract_repo_from_hf_url,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
updated_fields: List[str] = []
|
updated_fields: List[str] = []
|
||||||
preview_downloaded = False
|
preview_downloaded = False
|
||||||
|
|
||||||
# -- Determine whether this is an HF-sourced model -----------------
|
# -- Determine whether this is an externally-sourced model ---------
|
||||||
is_hf_model = not metadata.get("from_civitai", True)
|
# Key off the source fields directly: `from_civitai` records provenance
|
||||||
|
# and can be true for a model that is also linked to an external site
|
||||||
|
# (both sources coexist, see #1094), so it must not gate enrichment.
|
||||||
|
is_source_model = has_external_source(metadata)
|
||||||
|
|
||||||
|
source_ref = resolve_source_ref(metadata)
|
||||||
|
source = get_source(source_ref.platform) if source_ref else None
|
||||||
|
source_id = source_ref.source_id if source_ref else ""
|
||||||
|
asset_base_url = (
|
||||||
|
source.asset_base_url(source_id)
|
||||||
|
if source is not None and source_id
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
# -- Collect updates -----------------------------------------------
|
# -- Collect updates -----------------------------------------------
|
||||||
updates: Dict[str, Any] = {}
|
updates: Dict[str, Any] = {}
|
||||||
|
|
||||||
# base_model
|
# base_model — the LLM's mapping wins; when it returned nothing usable,
|
||||||
|
# fall back to the canonical name the site's own hints resolve to.
|
||||||
new_base = (llm_output.get("base_model") or "").strip()
|
new_base = (llm_output.get("base_model") or "").strip()
|
||||||
|
if not new_base:
|
||||||
|
new_base = (resolved_base_model or "").strip()
|
||||||
current_base = metadata.get("base_model", "") or ""
|
current_base = metadata.get("base_model", "") or ""
|
||||||
if new_base and self._should_overwrite(current_base, is_hf_model):
|
if new_base and self._should_overwrite(current_base, is_source_model):
|
||||||
updates["base_model"] = new_base
|
updates["base_model"] = new_base
|
||||||
|
|
||||||
|
# model_name — the site's own display name, so a source download never
|
||||||
|
# shows up under its local filename. Written only while the name is
|
||||||
|
# still the untouched file stem: once a user renames a model that
|
||||||
|
# choice is theirs to keep.
|
||||||
|
site_name = ((source_context.model_name if source_context else "") or "").strip()
|
||||||
|
if is_source_model and site_name:
|
||||||
|
current_name = (metadata.get("model_name") or "").strip()
|
||||||
|
file_stem = (metadata.get("file_name") or "").strip()
|
||||||
|
if not current_name or current_name == file_stem:
|
||||||
|
updates["model_name"] = site_name
|
||||||
|
|
||||||
# trigger words → civitai.trainedWords
|
# trigger words → civitai.trainedWords
|
||||||
new_triggers = llm_output.get("trigger_words", [])
|
new_triggers = llm_output.get("trigger_words", [])
|
||||||
trigger_words_empty = True
|
trigger_words_empty = True
|
||||||
@@ -110,45 +160,71 @@ class PostProcessor:
|
|||||||
cleaned = [t.strip() for t in new_triggers if t.strip()]
|
cleaned = [t.strip() for t in new_triggers if t.strip()]
|
||||||
cleaned = [t for t in cleaned if t.lower() not in ("none", "null", "n/a")]
|
cleaned = [t for t in cleaned if t.lower() not in ("none", "null", "n/a")]
|
||||||
trigger_words_empty = not cleaned
|
trigger_words_empty = not cleaned
|
||||||
current_civitai = metadata.get("civitai") or {}
|
current_triggers = (metadata.get("civitai") or {}).get("trainedWords") or []
|
||||||
current_triggers = current_civitai.get("trainedWords") or []
|
if self._should_overwrite_list(current_triggers, is_source_model):
|
||||||
if self._should_overwrite_list(current_triggers, is_hf_model):
|
self._merge_civitai(updates, metadata, trainedWords=cleaned)
|
||||||
trig_civitai = dict(current_civitai)
|
|
||||||
if "civitai" in updates and isinstance(updates["civitai"], dict):
|
|
||||||
trig_civitai.update(updates["civitai"])
|
|
||||||
trig_civitai["trainedWords"] = cleaned
|
|
||||||
updates["civitai"] = trig_civitai
|
|
||||||
|
|
||||||
# modelDescription — from raw README content (converted to HTML)
|
# modelDescription — the author's own summary (when the site keeps one
|
||||||
if readme_content and is_hf_model:
|
# outside the README, e.g. ModelScope's ``Description``) followed by the
|
||||||
converted = convert_readme_to_html(readme_content)
|
# README converted to HTML.
|
||||||
if converted:
|
site_description = (
|
||||||
updates["modelDescription"] = converted
|
(source_context.description if source_context else "") or ""
|
||||||
|
).strip()
|
||||||
|
if is_source_model and (site_description or readme_content):
|
||||||
|
parts: List[str] = []
|
||||||
|
if site_description:
|
||||||
|
parts.append(f"<p>{html.escape(site_description)}</p>")
|
||||||
|
if readme_content:
|
||||||
|
converted = convert_readme_to_html(readme_content)
|
||||||
|
if converted:
|
||||||
|
parts.append(converted)
|
||||||
|
if parts:
|
||||||
|
updates["modelDescription"] = "\n".join(parts)
|
||||||
|
|
||||||
# short_description → civitai.description (for "About this version")
|
# short_description → civitai.description (for "About this version").
|
||||||
|
# Falls back to the site's author summary, which for ModelScope AIGC
|
||||||
|
# models is frequently the only human-written text available.
|
||||||
short_desc = (llm_output.get("short_description") or "").strip()
|
short_desc = (llm_output.get("short_description") or "").strip()
|
||||||
if short_desc and is_hf_model:
|
if not short_desc:
|
||||||
current_civitai = metadata.get("civitai") or {}
|
short_desc = site_description
|
||||||
desc_civitai = dict(current_civitai)
|
if short_desc and is_source_model:
|
||||||
if "civitai" in updates and isinstance(updates["civitai"], dict):
|
self._merge_civitai(updates, metadata, description=short_desc)
|
||||||
desc_civitai.update(updates["civitai"])
|
|
||||||
desc_civitai["description"] = short_desc
|
# The version label completes the card the way a CivitAI download does:
|
||||||
updates["civitai"] = desc_civitai
|
# the UI renders `civitai.name` as the version chip. It is per file,
|
||||||
|
# so a collection repository shows that checkpoint's own label.
|
||||||
|
site_version = (
|
||||||
|
(source_context.version_name if source_context else "") or ""
|
||||||
|
).strip()
|
||||||
|
if is_source_model and site_version:
|
||||||
|
self._merge_civitai(updates, metadata, name=site_version)
|
||||||
|
|
||||||
|
# gallery images → civitai.images (site example images, YAML frontmatter
|
||||||
|
# widget entries, and Sample Gallery markdown tables in the README body)
|
||||||
|
rec_width = llm_output.get("recommended_width") or 0
|
||||||
|
rec_height = llm_output.get("recommended_height") or 0
|
||||||
|
|
||||||
|
# Example images the site publishes for *this* file. They are matched
|
||||||
|
# by filename, so they are the most precise preview source available
|
||||||
|
# and the only one for repositories whose README carries no images.
|
||||||
|
site_images: List[Dict[str, Any]] = []
|
||||||
|
if is_source_model and source_context is not None:
|
||||||
|
site_images = [
|
||||||
|
_example_image(url, rec_width, rec_height)
|
||||||
|
for url in source_context.example_images
|
||||||
|
if url
|
||||||
|
]
|
||||||
|
|
||||||
# gallery images → civitai.images (from YAML frontmatter widget entries
|
|
||||||
# and Sample Gallery markdown tables in the README body)
|
|
||||||
gallery_images: List[Dict[str, Any]] = []
|
gallery_images: List[Dict[str, Any]] = []
|
||||||
if readme_content and is_hf_model:
|
if (readme_content or site_images) and is_source_model:
|
||||||
hf_url = metadata.get("hf_url", "") or ""
|
repo = source_id
|
||||||
repo = extract_repo_from_hf_url(hf_url)
|
readme_images: List[Dict[str, Any]] = []
|
||||||
if repo:
|
if readme_content and repo:
|
||||||
rec_w = llm_output.get("recommended_width") or 0
|
|
||||||
rec_h = llm_output.get("recommended_height") or 0
|
|
||||||
|
|
||||||
# 1. Widget images (YAML frontmatter)
|
# 1. Widget images (YAML frontmatter)
|
||||||
gallery = extract_gallery_images(
|
gallery = extract_gallery_images(
|
||||||
readme_content, repo,
|
readme_content, repo,
|
||||||
default_width=rec_w, default_height=rec_h,
|
default_width=rec_width, default_height=rec_height,
|
||||||
|
base_url=asset_base_url,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Sample Gallery table images (markdown body), deduplicated
|
# 2. Sample Gallery table images (markdown body), deduplicated
|
||||||
@@ -156,7 +232,8 @@ class PostProcessor:
|
|||||||
table_images = extract_gallery_table_images(
|
table_images = extract_gallery_table_images(
|
||||||
readme_content, repo,
|
readme_content, repo,
|
||||||
existing_urls=existing_urls,
|
existing_urls=existing_urls,
|
||||||
default_width=rec_w, default_height=rec_h,
|
default_width=rec_width, default_height=rec_height,
|
||||||
|
base_url=asset_base_url,
|
||||||
)
|
)
|
||||||
existing_urls.update(img["url"] for img in table_images if img.get("url"))
|
existing_urls.update(img["url"] for img in table_images if img.get("url"))
|
||||||
|
|
||||||
@@ -164,7 +241,8 @@ class PostProcessor:
|
|||||||
simple_images = extract_simple_markdown_images(
|
simple_images = extract_simple_markdown_images(
|
||||||
readme_content, repo,
|
readme_content, repo,
|
||||||
existing_urls=existing_urls,
|
existing_urls=existing_urls,
|
||||||
default_width=rec_w, default_height=rec_h,
|
default_width=rec_width, default_height=rec_height,
|
||||||
|
base_url=asset_base_url,
|
||||||
)
|
)
|
||||||
existing_urls.update(img["url"] for img in simple_images if img.get("url"))
|
existing_urls.update(img["url"] for img in simple_images if img.get("url"))
|
||||||
|
|
||||||
@@ -172,54 +250,71 @@ class PostProcessor:
|
|||||||
html_images = extract_html_img_tags(
|
html_images = extract_html_img_tags(
|
||||||
readme_content, repo,
|
readme_content, repo,
|
||||||
existing_urls=existing_urls,
|
existing_urls=existing_urls,
|
||||||
default_width=rec_w, default_height=rec_h,
|
default_width=rec_width, default_height=rec_height,
|
||||||
|
base_url=asset_base_url,
|
||||||
)
|
)
|
||||||
|
|
||||||
all_images = gallery + table_images + simple_images + html_images
|
readme_images = gallery + table_images + simple_images + html_images
|
||||||
if all_images:
|
|
||||||
gallery_images = all_images
|
|
||||||
current_civitai = metadata.get("civitai") or {}
|
|
||||||
gallery_civitai = dict(current_civitai)
|
|
||||||
if "civitai" in updates and isinstance(updates["civitai"], dict):
|
|
||||||
gallery_civitai.update(updates["civitai"])
|
|
||||||
gallery_civitai["images"] = all_images
|
|
||||||
updates["civitai"] = gallery_civitai
|
|
||||||
|
|
||||||
# tags
|
# Site images come first so the preview fallback below prefers an
|
||||||
|
# image that is known to belong to this exact file.
|
||||||
|
all_images = _dedupe_images(site_images + readme_images)
|
||||||
|
if all_images:
|
||||||
|
gallery_images = all_images
|
||||||
|
self._merge_civitai(updates, metadata, images=all_images)
|
||||||
|
|
||||||
|
# tags — the site's curated tags are authoritative content vocabulary, so
|
||||||
|
# they are kept alongside whatever the LLM proposed (the LLM is skipped
|
||||||
|
# entirely when the site data is complete, which is why this cannot rely
|
||||||
|
# on ``llm_output`` alone).
|
||||||
new_tags = llm_output.get("tags", [])
|
new_tags = llm_output.get("tags", [])
|
||||||
if isinstance(new_tags, list) and new_tags:
|
candidate_tags: List[str] = []
|
||||||
|
if is_source_model and source_context is not None:
|
||||||
|
candidate_tags.extend(source_context.official_tags)
|
||||||
|
if isinstance(new_tags, list):
|
||||||
|
candidate_tags.extend(
|
||||||
|
tag for tag in new_tags if tag not in candidate_tags
|
||||||
|
)
|
||||||
|
if candidate_tags:
|
||||||
existing_tags = metadata.get("tags") or []
|
existing_tags = metadata.get("tags") or []
|
||||||
merged = self._merge_tags(existing_tags, new_tags)
|
merged = self._merge_tags(existing_tags, candidate_tags)
|
||||||
if len(merged) > len(existing_tags) or is_hf_model:
|
if len(merged) > len(existing_tags) or is_source_model:
|
||||||
updates["tags"] = merged
|
updates["tags"] = merged
|
||||||
|
|
||||||
# metadata_source & llm_enriched_at (always set)
|
# metadata_source is recorded for provenance; llm_enriched_at only means
|
||||||
updates["metadata_source"] = "agent:enrich_hf_metadata"
|
# something when a provider actually answered, so the deterministic
|
||||||
updates["llm_enriched_at"] = datetime.now(timezone.utc).isoformat()
|
# download-time hydration does not claim an enrichment that never ran.
|
||||||
|
updates["metadata_source"] = metadata_source
|
||||||
|
if llm_output:
|
||||||
|
updates["llm_enriched_at"] = datetime.now(timezone.utc).isoformat()
|
||||||
|
|
||||||
# Store LLM confidence in metadata so it's accessible for evaluation
|
# LLM confidence, stored for the enrichment evaluation harness. The key
|
||||||
|
# must NOT start with an underscore: `BaseModelMetadata.from_dict()`
|
||||||
|
# deliberately drops underscore-prefixed keys so they never round-trip,
|
||||||
|
# which silently erased this field on the next metadata write.
|
||||||
raw_confidence = (llm_output.get("confidence") or "").strip()
|
raw_confidence = (llm_output.get("confidence") or "").strip()
|
||||||
if raw_confidence:
|
if raw_confidence:
|
||||||
updates["_llm_confidence"] = raw_confidence
|
updates["llm_confidence"] = raw_confidence
|
||||||
|
|
||||||
# Fallback: extract instance_prompt from YAML frontmatter when the LLM
|
# Fallback: use the trigger words the site records for this exact file,
|
||||||
# returned empty trigger words but the README has instance_prompt.
|
# then the README's YAML `instance_prompt`, when the LLM returned none.
|
||||||
if trigger_words_empty:
|
if trigger_words_empty:
|
||||||
instance_prompt = _extract_yaml_instance_prompt(readme_content)
|
site_triggers = (
|
||||||
if instance_prompt:
|
list(source_context.trigger_words) if source_context else []
|
||||||
current_civitai = metadata.get("civitai") or {}
|
)
|
||||||
trig_civitai = dict(current_civitai)
|
if not site_triggers:
|
||||||
if "civitai" in updates and isinstance(updates["civitai"], dict):
|
instance_prompt = _extract_yaml_instance_prompt(readme_content)
|
||||||
trig_civitai.update(updates["civitai"])
|
if instance_prompt:
|
||||||
trig_civitai["trainedWords"] = [instance_prompt]
|
site_triggers = [instance_prompt]
|
||||||
updates["civitai"] = trig_civitai
|
if site_triggers:
|
||||||
|
self._merge_civitai(updates, metadata, trainedWords=site_triggers)
|
||||||
|
|
||||||
preview_remote_url = (llm_output.get("preview_url") or "").strip()
|
preview_remote_url = (llm_output.get("preview_url") or "").strip()
|
||||||
# Fallback: if the LLM couldn't find a preview image in the cleaned
|
# Fallback: if the LLM couldn't find a preview image in the cleaned
|
||||||
# README, find the first gallery image from the *model-specific
|
# README, find the first gallery image from the *model-specific
|
||||||
# section* of the README (not the repo-wide first image, which
|
# section* of the README (not the repo-wide first image, which
|
||||||
# belongs to a different model in collection repos).
|
# belongs to a different model in collection repos).
|
||||||
if not preview_remote_url and readme_content and is_hf_model:
|
if not preview_remote_url and readme_content and is_source_model:
|
||||||
model_basename = os.path.splitext(os.path.basename(model_path))[0]
|
model_basename = os.path.splitext(os.path.basename(model_path))[0]
|
||||||
relevant_section = extract_relevant_section(
|
relevant_section = extract_relevant_section(
|
||||||
readme_content, model_basename,
|
readme_content, model_basename,
|
||||||
@@ -245,8 +340,12 @@ class PostProcessor:
|
|||||||
if new_notes:
|
if new_notes:
|
||||||
updates["notes"] = new_notes
|
updates["notes"] = new_notes
|
||||||
|
|
||||||
# usage_tips — JSON string (e.g. {"strength_min":0.85,"strength_max":1.4})
|
# usage_tips — JSON string (e.g. {"strength_min":0.85,"strength_max":1.4}).
|
||||||
|
# When the LLM returned nothing, recover an explicitly stated strength
|
||||||
|
# range from the author summary so the value is not lost.
|
||||||
raw_tips = (llm_output.get("usage_tips") or "").strip()
|
raw_tips = (llm_output.get("usage_tips") or "").strip()
|
||||||
|
if not raw_tips or raw_tips == "{}":
|
||||||
|
raw_tips = _extract_usage_tips(site_description)
|
||||||
if raw_tips and raw_tips != "{}":
|
if raw_tips and raw_tips != "{}":
|
||||||
try:
|
try:
|
||||||
json.loads(raw_tips)
|
json.loads(raw_tips)
|
||||||
@@ -276,16 +375,35 @@ class PostProcessor:
|
|||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _should_overwrite(current_value: str, is_hf_model: bool) -> bool:
|
def _should_overwrite(current_value: str, is_source_model: bool) -> bool:
|
||||||
"""Return ``True`` when a scalar field should be overwritten."""
|
"""Return ``True`` when a scalar field should be overwritten."""
|
||||||
return is_hf_model or not current_value or current_value.lower() in (
|
return is_source_model or not current_value or current_value.lower() in (
|
||||||
"", "unknown",
|
"", "unknown",
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _should_overwrite_list(current_list: List[str], is_hf_model: bool) -> bool:
|
def _merge_civitai(
|
||||||
|
updates: Dict[str, Any], metadata: Dict[str, Any], **fields: Any
|
||||||
|
) -> None:
|
||||||
|
"""Layer *fields* onto the ``civitai`` block being assembled.
|
||||||
|
|
||||||
|
Description, version label, trigger words and gallery images all live
|
||||||
|
in the same dict and are contributed by separate branches, so each one
|
||||||
|
starts from what is already on disk and then applies whatever an
|
||||||
|
earlier branch queued in *updates*.
|
||||||
|
"""
|
||||||
|
|
||||||
|
merged = dict(metadata.get("civitai") or {})
|
||||||
|
queued = updates.get("civitai")
|
||||||
|
if isinstance(queued, dict):
|
||||||
|
merged.update(queued)
|
||||||
|
merged.update(fields)
|
||||||
|
updates["civitai"] = merged
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _should_overwrite_list(current_list: List[str], is_source_model: bool) -> bool:
|
||||||
"""Return ``True`` when a list field should be overwritten."""
|
"""Return ``True`` when a list field should be overwritten."""
|
||||||
return is_hf_model or not current_list
|
return is_source_model or not current_list
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _merge_tags(existing: List[str], new: List[str]) -> List[str]:
|
def _merge_tags(existing: List[str], new: List[str]) -> List[str]:
|
||||||
@@ -309,6 +427,129 @@ class PostProcessor:
|
|||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
#: Separator between a label and its value. Published model cards routinely
|
||||||
|
#: wrap the numbers in markdown emphasis or quotes (``strength: **0.85 - 1.4**``,
|
||||||
|
#: ``CLIP 强度「0.5」``), so those are absorbed rather than treated as a break.
|
||||||
|
_EMPHASIS = "[\"'\u201c\u201d\u300c\u300d*_`\\s]*"
|
||||||
|
|
||||||
|
#: An explicitly stated strength/weight range, e.g. ``权重0.5-1.2``,
|
||||||
|
#: ``强度 0.8 ~ 1.2``, ``strength: **0.85 - 1.4**``.
|
||||||
|
_RANGE_DASH = "(?:-|\u2010|\u2011|\u2012|\u2013|\u2014|\uff0d|~|\uff5e|\u81f3|\u5230|to)"
|
||||||
|
|
||||||
|
_STRENGTH_RANGE_RE = re.compile(
|
||||||
|
"(?:\u6743\u91cd|\u5f3a\u5ea6|strength|weight)" + _EMPHASIS + "[:\uff1a]?" + _EMPHASIS
|
||||||
|
+ r"(\d+(?:\.\d+)?)" + _EMPHASIS + _RANGE_DASH + _EMPHASIS
|
||||||
|
+ r"(\d+(?:\.\d+)?)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
#: A single strength/weight value, e.g. ``strength: 0.6``, ``权重 0.8``.
|
||||||
|
_STRENGTH_VALUE_RE = re.compile(
|
||||||
|
"(?:\u6743\u91cd|\u5f3a\u5ea6|strength|weight)" + _EMPHASIS + "[:\uff1a]?" + _EMPHASIS
|
||||||
|
+ r"(\d+(?:\.\d+)?)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
#: ``clip strength: 0.5`` / ``CLIP 强度 0.5``.
|
||||||
|
_CLIP_STRENGTH_RE = re.compile(
|
||||||
|
"clip" + _EMPHASIS + "(?:\u5f3a\u5ea6|strength)" + _EMPHASIS + "[:\uff1a]?" + _EMPHASIS
|
||||||
|
+ r"(\d+(?:\.\d+)?)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
#: ``clip skip: 2`` / ``CLIP 跳过 2``.
|
||||||
|
_CLIP_SKIP_RE = re.compile(
|
||||||
|
"clip" + _EMPHASIS + "(?:skip|\u8df3\u8fc7)" + _EMPHASIS + "[:\uff1a]?" + _EMPHASIS
|
||||||
|
+ r"(\d+)",
|
||||||
|
re.IGNORECASE,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_usage_tips(text: str) -> str:
|
||||||
|
"""Extract stated strength/CLIP recommendations from prose.
|
||||||
|
|
||||||
|
This is the deterministic counterpart to the LLM's ``usage_tips`` output,
|
||||||
|
used when the LLM was skipped. It only recognises explicitly written
|
||||||
|
values — it never infers a range — and returns ``""`` when it finds none.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A JSON string matching the skill's ``usage_tips`` schema, or ``""``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not text:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
tips: Dict[str, Any] = {}
|
||||||
|
|
||||||
|
# CLIP strength is resolved first and then blanked out, so the generic
|
||||||
|
# strength patterns cannot mistake `CLIP 强度 0.5` for the LoRA strength.
|
||||||
|
text_for_strength = text
|
||||||
|
clip_strength = _CLIP_STRENGTH_RE.search(text_for_strength)
|
||||||
|
if clip_strength:
|
||||||
|
tips["clip_strength"] = float(clip_strength.group(1))
|
||||||
|
text_for_strength = (
|
||||||
|
text_for_strength[: clip_strength.start()]
|
||||||
|
+ " "
|
||||||
|
+ text_for_strength[clip_strength.end() :]
|
||||||
|
)
|
||||||
|
|
||||||
|
range_match = _STRENGTH_RANGE_RE.search(text_for_strength)
|
||||||
|
if range_match:
|
||||||
|
low = float(range_match.group(1))
|
||||||
|
high = float(range_match.group(2))
|
||||||
|
if low > high:
|
||||||
|
low, high = high, low
|
||||||
|
tips["strength_min"] = low
|
||||||
|
tips["strength_max"] = high
|
||||||
|
tips["strength_range"] = f"{low:g}-{high:g}"
|
||||||
|
else:
|
||||||
|
value_match = _STRENGTH_VALUE_RE.search(text_for_strength)
|
||||||
|
if value_match:
|
||||||
|
tips["strength"] = float(value_match.group(1))
|
||||||
|
|
||||||
|
clip_skip = _CLIP_SKIP_RE.search(text)
|
||||||
|
if clip_skip:
|
||||||
|
tips["clip_skip"] = int(clip_skip.group(1))
|
||||||
|
|
||||||
|
if not tips:
|
||||||
|
return ""
|
||||||
|
return json.dumps(tips, ensure_ascii=False)
|
||||||
|
|
||||||
|
|
||||||
|
def _example_image(url: str, width: int, height: int) -> Dict[str, Any]:
|
||||||
|
"""Build a ``civitai.images`` entry for a site-provided example image.
|
||||||
|
|
||||||
|
The site publishes no prompt alongside these images, so the entry carries
|
||||||
|
empty prompt metadata and the LLM's recommended dimensions when it found
|
||||||
|
any (falling back to the same 512px placeholder the README extractors use).
|
||||||
|
"""
|
||||||
|
|
||||||
|
return {
|
||||||
|
"url": url,
|
||||||
|
"type": "image",
|
||||||
|
"nsfwLevel": 0,
|
||||||
|
"width": width or 512,
|
||||||
|
"height": height or 512,
|
||||||
|
"meta": {"prompt": "", "negativePrompt": ""},
|
||||||
|
"hasMeta": False,
|
||||||
|
"hasPositivePrompt": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _dedupe_images(images: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||||
|
"""Drop later entries that repeat an earlier image URL, keeping order."""
|
||||||
|
|
||||||
|
seen: set[str] = set()
|
||||||
|
unique: List[Dict[str, Any]] = []
|
||||||
|
for image in images:
|
||||||
|
url = image.get("url") or ""
|
||||||
|
if not url or url in seen:
|
||||||
|
continue
|
||||||
|
seen.add(url)
|
||||||
|
unique.append(image)
|
||||||
|
return unique
|
||||||
|
|
||||||
|
|
||||||
def _extract_yaml_instance_prompt(readme_content: str) -> str:
|
def _extract_yaml_instance_prompt(readme_content: str) -> str:
|
||||||
"""Extract ``instance_prompt`` from the YAML frontmatter of a HF README.
|
"""Extract ``instance_prompt`` from the YAML frontmatter of a HF README.
|
||||||
|
|
||||||
|
|||||||
@@ -1,20 +1,23 @@
|
|||||||
---
|
---
|
||||||
name: enrich_hf_metadata
|
name: enrich_hf_metadata
|
||||||
title: "Enrich Metadata from HuggingFace"
|
title: "Enrich Metadata from Model Card"
|
||||||
description: >
|
description: >
|
||||||
Parse the HuggingFace model card via LLM to extract description, trigger
|
Parse the model card (README) from HuggingFace, ModelScope, or any other
|
||||||
words, base model, tags, and preview image URL.
|
supported model site via LLM to extract description, trigger words, base
|
||||||
|
model, tags, and preview image URL.
|
||||||
llm_required: true
|
llm_required: true
|
||||||
---
|
---
|
||||||
|
|
||||||
You are an expert assistant for AI image generation models. Your task is to extract structured metadata from a HuggingFace model card (README.md).
|
You are an expert assistant for AI image generation models. Your task is to extract structured metadata from a model card (README).
|
||||||
|
|
||||||
## Model Information
|
## Model Information
|
||||||
|
|
||||||
- **Repository**: {{hf_url}}
|
- **Source site**: {{source_label}} ({{source_platform}})
|
||||||
|
- **Model page**: {{source_url}}
|
||||||
- **Model file path**: {{model_path}}
|
- **Model file path**: {{model_path}}
|
||||||
- **Model filename**: {{model_basename}}
|
- **Model filename**: {{model_basename}}
|
||||||
- **Repository ID**: {{repo}}
|
- **Repository ID**: {{source_id}}
|
||||||
|
- **Repository raw-file base URL**: {{asset_base_url}}
|
||||||
|
|
||||||
## Current Metadata (may be incomplete)
|
## Current Metadata (may be incomplete)
|
||||||
|
|
||||||
@@ -22,6 +25,34 @@ You are an expert assistant for AI image generation models. Your task is to extr
|
|||||||
{{current_metadata}}
|
{{current_metadata}}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Site-Provided Metadata (any field may be empty)
|
||||||
|
|
||||||
|
The model site publishes the following **alongside** the README. It is
|
||||||
|
first-hand information recorded by the site itself, so it outranks anything
|
||||||
|
you would otherwise guess:
|
||||||
|
|
||||||
|
- **Author description**: {{source_description}}
|
||||||
|
- **Base model reported by the site**: {{source_base_model}}
|
||||||
|
- **Trigger words recorded for this file**: {{source_trigger_words}}
|
||||||
|
- **Site-curated tags**:
|
||||||
|
{{source_official_tags}}
|
||||||
|
- **Example image URLs for this file**:
|
||||||
|
{{source_example_images}}
|
||||||
|
|
||||||
|
Use it as follows:
|
||||||
|
|
||||||
|
- A weight or strength range stated in the **author description** belongs in
|
||||||
|
``usage_tips`` (and in ``notes``); do not leave ``usage_tips`` empty when the
|
||||||
|
description states one.
|
||||||
|
- When the author description exists, base ``short_description`` on it rather
|
||||||
|
than on the README, which on some sites is auto-generated boilerplate.
|
||||||
|
- Treat the **site-curated tags** as strong signals for ``tags``: they are
|
||||||
|
already a curated content vocabulary, so prefer them over invented words.
|
||||||
|
- Treat the **base model reported by the site** as a strong hint for
|
||||||
|
``base_model``, but still map it to the EXACT canonical name from the
|
||||||
|
available base-model list.
|
||||||
|
- Use the **example image URLs** when the README contains no usable image.
|
||||||
|
|
||||||
## User Priority Tags Reference
|
## User Priority Tags Reference
|
||||||
|
|
||||||
The user has configured the following list of **meaningful tag categories** for this model type (`{{model_type}}`):
|
The user has configured the following list of **meaningful tag categories** for this model type (`{{model_type}}`):
|
||||||
@@ -39,7 +70,7 @@ name listed — do not invent aliases or modify variant suffixes.
|
|||||||
|
|
||||||
{{base_models}}
|
{{base_models}}
|
||||||
|
|
||||||
## HuggingFace README Content
|
## Model Card Content
|
||||||
|
|
||||||
```
|
```
|
||||||
{{readme_content}}
|
{{readme_content}}
|
||||||
@@ -52,10 +83,11 @@ Extract the following information from the README content above:
|
|||||||
### base_model
|
### base_model
|
||||||
The base model this model was trained on. Use EXACTLY one of the names from the **Available Base Models** list above. Do not invent new names or use aliases.
|
The base model this model was trained on. Use EXACTLY one of the names from the **Available Base Models** list above. Do not invent new names or use aliases.
|
||||||
|
|
||||||
Check the YAML frontmatter for ``base_model:`` first. If the frontmatter has no ``base_model:``, look at the **model filename** (``{{model_basename}}``), YAML ``tags:``, README title and first paragraph for clues — the base model family is often embedded in the name
|
Check the **base model reported by the site** (above) and the YAML frontmatter ``base_model:`` first. If neither yields a match, look at the **model filename** (``{{model_basename}}``), YAML ``tags:``, README title and first paragraph for clues — the base model family is often embedded in the name
|
||||||
|
|
||||||
### trigger_words
|
### trigger_words
|
||||||
The trigger words or activation prompts needed to use this LoRA. Look for:
|
The trigger words or activation prompts needed to use this LoRA. Look for:
|
||||||
|
- The **trigger words recorded for this file** in the site-provided metadata (most authoritative)
|
||||||
- `instance_prompt:` in the YAML frontmatter
|
- `instance_prompt:` in the YAML frontmatter
|
||||||
- Phrases like "trigger word:", "trigger:", "use this prompt:", "activation prompt:"
|
- Phrases like "trigger word:", "trigger:", "use this prompt:", "activation prompt:"
|
||||||
- In collection repos: the trigger section **specific to this model file** (look near matching download links or anchor IDs)
|
- In collection repos: the trigger section **specific to this model file** (look near matching download links or anchor IDs)
|
||||||
@@ -63,12 +95,13 @@ The trigger words or activation prompts needed to use this LoRA. Look for:
|
|||||||
Return as an array of strings. If none found, return an empty array `[]`. **Never** return `["None"]` or any placeholder value — a truly empty list means no trigger words exist.
|
Return as an array of strings. If none found, return an empty array `[]`. **Never** return `["None"]` or any placeholder value — a truly empty list means no trigger words exist.
|
||||||
|
|
||||||
### short_description
|
### short_description
|
||||||
A concise 1-2 sentence summary of what this model does. Extract from the "Model description" section or the first paragraph. For collection repos, focus on the **specific model version** matching `{{model_basename}}`, not the repo as a whole. Return empty string if the README is too minimal.
|
A concise 1-2 sentence summary of what this model does. For collection repos, focus on the **specific model version** matching `{{model_basename}}`, not the repo as a whole. Prefer the **author description** from the site-provided metadata when it is present; otherwise extract from the "Model description" section or the first paragraph. Return empty string if the available content is too minimal.
|
||||||
|
|
||||||
### tags
|
### tags
|
||||||
3-8 relevant tags for categorizing this model. **Quality over quantity.**
|
3-8 relevant tags for categorizing this model. **Quality over quantity.**
|
||||||
|
|
||||||
Sources to consider:
|
Sources to consider:
|
||||||
|
- The **site-curated tags** from the site-provided metadata (these are already filtered content tags — prefer them)
|
||||||
- The YAML frontmatter `tags:` list (filter out technical ones — see below)
|
- The YAML frontmatter `tags:` list (filter out technical ones — see below)
|
||||||
- The subject, style, character, or concept the model represents
|
- The subject, style, character, or concept the model represents
|
||||||
- The model filename itself may give clues (e.g. "pokemon", "anime", "pixelart")
|
- The model filename itself may give clues (e.g. "pokemon", "anime", "pixelart")
|
||||||
@@ -79,7 +112,9 @@ Sources to consider:
|
|||||||
|
|
||||||
2. **Cross-reference against the priority_tags reference.** Only include a tag if it meaningfully describes what the model actually creates (subject, style, character type) and is semantically close to one of the priority_tags. If none of the README's tags match meaningful categories, prefer returning a smaller set or an empty array over including low-value tags.
|
2. **Cross-reference against the priority_tags reference.** Only include a tag if it meaningfully describes what the model actually creates (subject, style, character type) and is semantically close to one of the priority_tags. If none of the README's tags match meaningful categories, prefer returning a smaller set or an empty array over including low-value tags.
|
||||||
|
|
||||||
3. **All lowercase, no spaces, no hyphens** (use single words like `"photorealistic"`, `"anime"`, `"character"`).
|
3. **All lowercase, and keep each tag's own wording.** Prefer the spelling already used by the site, the frontmatter, or the author — including hyphenated and multi-word tags such as `"sci-fi"`, `"semi-realistic"`, `"character-enhancement"` or `"art style"`. Do **not** strip separators or invent a single-word variant of a tag you are already including (e.g. do not emit both `"character-enhancement"` and `"character"`). When a tag is written in another script (e.g. Chinese), likewise keep it verbatim instead of translating it.
|
||||||
|
|
||||||
|
4. **Never invent a tag** that neither the site-provided metadata, the YAML frontmatter, nor the README text supports.
|
||||||
|
|
||||||
Return empty array if no meaningful content tags remain after filtering.
|
Return empty array if no meaningful content tags remain after filtering.
|
||||||
|
|
||||||
@@ -92,13 +127,13 @@ The URL of the most suitable preview image from the README. Look for:
|
|||||||
- The YAML frontmatter `widget:` section (which often has `output.url` fields)
|
- The YAML frontmatter `widget:` section (which often has `output.url` fields)
|
||||||
- In collection repos: the sample images listed **under the section** for this specific model version
|
- In collection repos: the sample images listed **under the section** for this specific model version
|
||||||
- Generic `` in the body
|
- Generic `` in the body
|
||||||
Choose the first image that appears to be a generation example (not a logo or diagram). Construct the absolute URL as `https://huggingface.co/{{repo}}/resolve/main/{filename}`. If no suitable image is found, return an empty string.
|
Choose the first image that appears to be a generation example (not a logo or diagram). Construct the absolute URL from the repository raw-file base URL (`{{asset_base_url}}`) plus the relative path. If the README has no suitable image, fall back to the site-provided **example image URLs** for this file. If nothing is available, return an empty string.
|
||||||
|
|
||||||
### notes
|
### notes
|
||||||
A plain-text summary of the model card's key practical usage information. Combine trigger words, style modifiers, recommended parameters (steps, CFG, resolution, sampler), and any setup tips into a readable paragraph. For collection repos, focus on the **specific model version** matching `{{model_basename}}`. Return empty string if the README has no useful usage info.
|
A plain-text summary of the model card's key practical usage information. Combine trigger words, style modifiers, recommended parameters (steps, CFG, resolution, sampler), and any setup tips into a readable paragraph. For collection repos, focus on the **specific model version** matching `{{model_basename}}`. Include the **author description** from the site-provided metadata when it is present. Return empty string if there is no useful usage info.
|
||||||
|
|
||||||
### usage_tips
|
### usage_tips
|
||||||
A JSON string with structured usage recommendations. Extract from the README any explicit ranges or recommended values (e.g. "Set LoRA strength: **0.85 - 1.4**", "CLIP strength: 0.5"). Possible fields (include only those you can determine):
|
A JSON string with structured usage recommendations. Extract from the **author description** (site-provided metadata) and the README any explicit ranges or recommended values (e.g. "Set LoRA strength: **0.85 - 1.4**", "CLIP strength: 0.5", "权重0.5-1.2"). Possible fields (include only those you can determine):
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
@@ -121,7 +156,7 @@ Your confidence level in the extracted data:
|
|||||||
|
|
||||||
## Important: Handling Collection Repos (multiple model files)
|
## Important: Handling Collection Repos (multiple model files)
|
||||||
|
|
||||||
Many HuggingFace repos contain **multiple model files** in a single repository
|
Many model repositories contain **multiple model files** in a single repository
|
||||||
(e.g. a "LoRA collection" with different styles/characters in separate files).
|
(e.g. a "LoRA collection" with different styles/characters in separate files).
|
||||||
|
|
||||||
The model file currently being enriched is: **`{{model_basename}}`**
|
The model file currently being enriched is: **`{{model_basename}}`**
|
||||||
|
|||||||
@@ -1,8 +1,15 @@
|
|||||||
"""HF README processing for the ``enrich_hf_metadata`` skill.
|
"""Model card (README) processing for the ``enrich_hf_metadata`` skill.
|
||||||
|
|
||||||
Provides README cleaning for LLM injection, gallery/image extraction from
|
Provides README cleaning for LLM injection, gallery/image extraction from
|
||||||
multiple formats (YAML widget, markdown, HTML ``<img>``, gallery tables),
|
multiple formats (YAML widget, markdown, HTML ``<img>``, gallery tables),
|
||||||
and section-based README trimming for collection repos.
|
and section-based README trimming for collection repos.
|
||||||
|
|
||||||
|
The extractors default to Hugging Face asset URLs, but every one of them
|
||||||
|
accepts an explicit ``base_url`` so the same parsing works for any model
|
||||||
|
source (ModelScope, ...). See :mod:`py.services.model_sources`.
|
||||||
|
|
||||||
|
This module deliberately has no package-relative imports: it is also loaded
|
||||||
|
standalone by the README-processing test harness.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -15,12 +22,25 @@ from typing import Any, List, Tuple
|
|||||||
_REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
|
_REPO_URL_PATTERN = re.compile(r"https?://huggingface\.co/([^/]+/[^/]+)")
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_asset_base_url(repo: str, base_url: str | None = None) -> str:
|
||||||
|
"""Return the base URL used to resolve repository-relative assets.
|
||||||
|
|
||||||
|
Falls back to the historical Hugging Face layout when *base_url* is not
|
||||||
|
supplied, so existing callers keep their behaviour.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if base_url:
|
||||||
|
return base_url.rstrip("/")
|
||||||
|
return f"https://huggingface.co/{repo}/resolve/main"
|
||||||
|
|
||||||
|
|
||||||
def extract_simple_markdown_images(
|
def extract_simple_markdown_images(
|
||||||
markdown_text: str,
|
markdown_text: str,
|
||||||
repo: str,
|
repo: str,
|
||||||
existing_urls: set[str] | None = None,
|
existing_urls: set[str] | None = None,
|
||||||
default_width: int = 512,
|
default_width: int = 512,
|
||||||
default_height: int = 512,
|
default_height: int = 512,
|
||||||
|
base_url: str | None = None,
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
"""Extract standalone markdown images from the README body.
|
"""Extract standalone markdown images from the README body.
|
||||||
|
|
||||||
@@ -32,10 +52,10 @@ def extract_simple_markdown_images(
|
|||||||
Returns a list of dicts in the same ``civitai.images`` format as
|
Returns a list of dicts in the same ``civitai.images`` format as
|
||||||
:func:`extract_gallery_images`.
|
:func:`extract_gallery_images`.
|
||||||
"""
|
"""
|
||||||
if not markdown_text or not repo:
|
if not markdown_text or not (repo or base_url):
|
||||||
return []
|
return []
|
||||||
|
|
||||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
base_url = resolve_asset_base_url(repo, base_url)
|
||||||
images: list[dict[str, Any]] = []
|
images: list[dict[str, Any]] = []
|
||||||
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
||||||
|
|
||||||
@@ -89,20 +109,21 @@ def extract_html_img_tags(
|
|||||||
existing_urls: set[str] | None = None,
|
existing_urls: set[str] | None = None,
|
||||||
default_width: int = 512,
|
default_width: int = 512,
|
||||||
default_height: int = 512,
|
default_height: int = 512,
|
||||||
|
base_url: str | None = None,
|
||||||
) -> list[dict[str, Any]]:
|
) -> 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
|
||||||
``<img>`` tags exclusively for their sample images, with no markdown
|
``<img>`` tags exclusively for their sample images, with no markdown
|
||||||
``![]()`` equivalents. This function finds those tags and constructs
|
``![]()`` equivalents. This function finds those tags and constructs
|
||||||
resolvable HF URLs.
|
resolvable URLs.
|
||||||
|
|
||||||
Returns a list of dicts in the ``civitai.images`` format.
|
Returns a list of dicts in the ``civitai.images`` format.
|
||||||
"""
|
"""
|
||||||
if not markdown_text or not repo:
|
if not markdown_text or not (repo or base_url):
|
||||||
return []
|
return []
|
||||||
|
|
||||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
base_url = resolve_asset_base_url(repo, base_url)
|
||||||
images: list[dict[str, Any]] = []
|
images: list[dict[str, Any]] = []
|
||||||
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
seen_urls: set[str] = set(existing_urls) if existing_urls else set()
|
||||||
|
|
||||||
@@ -166,7 +187,7 @@ def extract_html_img_tags(
|
|||||||
|
|
||||||
def extract_repo_from_hf_url(hf_url: str) -> str:
|
def extract_repo_from_hf_url(hf_url: str) -> str:
|
||||||
"""Extract ``user/repo`` from a HuggingFace URL."""
|
"""Extract ``user/repo`` from a HuggingFace URL."""
|
||||||
m = _REPO_URL_PATTERN.match(hf_url)
|
m = _REPO_URL_PATTERN.match(hf_url or "")
|
||||||
return m.group(1) if m else ""
|
return m.group(1) if m else ""
|
||||||
|
|
||||||
|
|
||||||
@@ -175,21 +196,23 @@ def extract_gallery_images(
|
|||||||
repo: str,
|
repo: str,
|
||||||
default_width: int = 512,
|
default_width: int = 512,
|
||||||
default_height: int = 512,
|
default_height: int = 512,
|
||||||
|
base_url: str | None = None,
|
||||||
) -> List[dict[str, Any]]:
|
) -> 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 README.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
markdown_text: Raw README content.
|
markdown_text: Raw README content.
|
||||||
repo: HF repo identifier (``user/repo``).
|
repo: Repository identifier (``user/repo``).
|
||||||
default_width: Fallback width when the README provides no dimension.
|
default_width: Fallback width when the README provides no dimension.
|
||||||
default_height: Fallback height when the README provides no dimension.
|
default_height: Fallback height when the README provides no dimension.
|
||||||
|
base_url: Overrides the asset base URL (defaults to Hugging Face).
|
||||||
|
|
||||||
Returns a list of dicts compatible with the ``civitai.images`` metadata
|
Returns a list of dicts compatible with the ``civitai.images`` metadata
|
||||||
format, each containing ``url`` (absolute HF URL), ``meta.prompt``,
|
format, each containing ``url`` (absolute), ``meta.prompt``,
|
||||||
``width``, ``height``, and ``type``. Returns an empty list when no
|
``width``, ``height``, and ``type``. Returns an empty list when no
|
||||||
widget entries are found or when *repo* is empty.
|
widget entries are found or when *repo* is empty.
|
||||||
"""
|
"""
|
||||||
if not markdown_text or not repo:
|
if not markdown_text or not (repo or base_url):
|
||||||
return []
|
return []
|
||||||
|
|
||||||
frontmatter = _extract_frontmatter(markdown_text)
|
frontmatter = _extract_frontmatter(markdown_text)
|
||||||
@@ -197,7 +220,7 @@ def extract_gallery_images(
|
|||||||
return []
|
return []
|
||||||
|
|
||||||
images: List[dict[str, Any]] = []
|
images: List[dict[str, Any]] = []
|
||||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
base_url = resolve_asset_base_url(repo, base_url)
|
||||||
w = default_width or 512
|
w = default_width or 512
|
||||||
h = default_height or 512
|
h = default_height or 512
|
||||||
|
|
||||||
@@ -279,10 +302,11 @@ def extract_gallery_table_images(
|
|||||||
existing_urls: set[str] | None = None,
|
existing_urls: set[str] | None = None,
|
||||||
default_width: int = 512,
|
default_width: int = 512,
|
||||||
default_height: int = 512,
|
default_height: int = 512,
|
||||||
|
base_url: str | None = None,
|
||||||
) -> list[dict[str, Any]]:
|
) -> 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 READMEs include a sample-gallery table in the body (outside
|
||||||
the YAML frontmatter) that shows generation examples with their
|
the YAML frontmatter) that shows generation examples with their
|
||||||
prompts. This function parses those tables and merges results with
|
prompts. This function parses those tables and merges results with
|
||||||
the widget-sourced images from :func:`extract_gallery_images`.
|
the widget-sourced images from :func:`extract_gallery_images`.
|
||||||
@@ -291,10 +315,10 @@ def extract_gallery_table_images(
|
|||||||
:func:`extract_gallery_images`. Already-seen URLs (from *existing_urls*)
|
:func:`extract_gallery_images`. Already-seen URLs (from *existing_urls*)
|
||||||
are skipped.
|
are skipped.
|
||||||
"""
|
"""
|
||||||
if not markdown_text or not repo:
|
if not markdown_text or not (repo or base_url):
|
||||||
return []
|
return []
|
||||||
|
|
||||||
base_url = f"https://huggingface.co/{repo}/resolve/main"
|
base_url = resolve_asset_base_url(repo, base_url)
|
||||||
images: list[dict[str, Any]] = []
|
images: list[dict[str, Any]] = []
|
||||||
seen_urls: set[str] = 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")
|
||||||
@@ -368,12 +392,18 @@ def _extract_frontmatter(text: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def convert_readme_to_html(markdown_text: str | None) -> str:
|
def convert_readme_to_html(markdown_text: str | None) -> str:
|
||||||
"""Convert HF README markdown to sanitised HTML."""
|
"""Convert HF README markdown to sanitised HTML.
|
||||||
|
|
||||||
|
Site-generated placeholder notices are dropped here too, so a repository
|
||||||
|
whose author wrote nothing does not store the download instructions as its
|
||||||
|
model description; the result is an empty string in that case.
|
||||||
|
"""
|
||||||
if not markdown_text:
|
if not markdown_text:
|
||||||
return ""
|
return ""
|
||||||
|
|
||||||
text = markdown_text
|
text = markdown_text
|
||||||
text = _strip_frontmatter(text)
|
text = _strip_frontmatter(text)
|
||||||
|
text = _strip_generated_card_boilerplate(text)
|
||||||
text = _strip_gallery(text)
|
text = _strip_gallery(text)
|
||||||
text = _strip_badge_images(text)
|
text = _strip_badge_images(text)
|
||||||
text = _strip_html_comments(text)
|
text = _strip_html_comments(text)
|
||||||
@@ -420,6 +450,59 @@ _MASSIVE_LIST_LINE_MIN_LEN = 150
|
|||||||
#: Minimum consecutive enumeration lines to trigger massive-list stripping.
|
#: Minimum consecutive enumeration lines to trigger massive-list stripping.
|
||||||
_MASSIVE_LIST_THRESHOLD = 8
|
_MASSIVE_LIST_THRESHOLD = 8
|
||||||
|
|
||||||
|
#: Substrings identifying text a *site* generated to fill a model card whose
|
||||||
|
#: author wrote nothing, as opposed to the author's own content. ModelScope
|
||||||
|
#: renders such a card as a placeholder notice, a block of SDK/git download
|
||||||
|
#: instructions, and a closing invitation to improve the card.
|
||||||
|
#:
|
||||||
|
#: Matched as substrings rather than whole headings because the notices are
|
||||||
|
#: prose, and because non-Latin scripts are not space-delimited — the notice
|
||||||
|
#: continues with a full-width period, so the ``title == kw`` style matching
|
||||||
|
#: used for :data:`_BOILERPLATE_HEADERS` would never fire.
|
||||||
|
_GENERATED_CARD_MARKERS: tuple[str, ...] = (
|
||||||
|
"当前模型的贡献者未提供更加详细的模型介绍",
|
||||||
|
"您可以通过如下",
|
||||||
|
"如果您是本模型的贡献者",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _strip_generated_card_boilerplate(text: str) -> str:
|
||||||
|
"""Remove the notices a site generates to fill an empty model card.
|
||||||
|
|
||||||
|
A repository whose uploader wrote no README still gets a card: ModelScope
|
||||||
|
answers with "the contributor provided no further description", the SDK
|
||||||
|
and git download commands, and an invitation to complete the card. None
|
||||||
|
of it describes the model, yet it was landing in both the LLM prompt and
|
||||||
|
the stored description.
|
||||||
|
|
||||||
|
A notice that is a heading takes its whole section with it, so the
|
||||||
|
download block goes too; a stand-alone notice line is dropped on its own.
|
||||||
|
Content the author added later — under a heading of equal or higher
|
||||||
|
level — is kept, so an improved card is not thrown away.
|
||||||
|
"""
|
||||||
|
|
||||||
|
lines = text.split("\n")
|
||||||
|
out: list[str] = []
|
||||||
|
skip_until_level: int | None = None
|
||||||
|
|
||||||
|
for line in lines:
|
||||||
|
level = _heading_level(line)
|
||||||
|
|
||||||
|
if any(marker in line for marker in _GENERATED_CARD_MARKERS):
|
||||||
|
if level > 0:
|
||||||
|
skip_until_level = level
|
||||||
|
continue
|
||||||
|
|
||||||
|
if skip_until_level is not None:
|
||||||
|
if level > 0 and level <= skip_until_level:
|
||||||
|
skip_until_level = None
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
|
||||||
|
out.append(line)
|
||||||
|
|
||||||
|
return "\n".join(out)
|
||||||
|
|
||||||
|
|
||||||
def clean_readme_for_llm(markdown_text: str | None, max_length: int = 6000) -> str:
|
def clean_readme_for_llm(markdown_text: str | None, max_length: int = 6000) -> str:
|
||||||
"""Clean a HF README for injection into an LLM metadata-extraction prompt.
|
"""Clean a HF README for injection into an LLM metadata-extraction prompt.
|
||||||
@@ -429,6 +512,8 @@ def clean_readme_for_llm(markdown_text: str | None, max_length: int = 6000) -> s
|
|||||||
|
|
||||||
* ``widget:`` YAML block (example prompts + output URLs)
|
* ``widget:`` YAML block (example prompts + output URLs)
|
||||||
* ``<Gallery />`` tags and wrappers
|
* ``<Gallery />`` tags and wrappers
|
||||||
|
* Site-generated placeholder notices for a card the author never wrote
|
||||||
|
(see :func:`_strip_generated_card_boilerplate`)
|
||||||
* Fenced code blocks (Python / bash / bibtex / yaml)
|
* Fenced code blocks (Python / bash / bibtex / yaml)
|
||||||
* Standalone ```` image lines and ``<img>`` tags
|
* Standalone ```` image lines and ``<img>`` tags
|
||||||
* Training-parameter tables
|
* Training-parameter tables
|
||||||
@@ -454,6 +539,7 @@ def clean_readme_for_llm(markdown_text: str | None, max_length: int = 6000) -> s
|
|||||||
# Order matters — broader strips first, then finer ones.
|
# Order matters — broader strips first, then finer ones.
|
||||||
text = _strip_gallery(text)
|
text = _strip_gallery(text)
|
||||||
text = _strip_widget_section(text)
|
text = _strip_widget_section(text)
|
||||||
|
text = _strip_generated_card_boilerplate(text)
|
||||||
text = _strip_fenced_code_blocks(text)
|
text = _strip_fenced_code_blocks(text)
|
||||||
text = _strip_standalone_images(text)
|
text = _strip_standalone_images(text)
|
||||||
text = _strip_training_tables(text)
|
text = _strip_training_tables(text)
|
||||||
|
|||||||
+170
-23
@@ -82,6 +82,17 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_no_uri_available_error(message: str) -> bool:
|
||||||
|
"""Return True for aria2's "No URI available" transfer failure.
|
||||||
|
|
||||||
|
aria2 reports this when every URI for the transfer has become unusable.
|
||||||
|
For CivitAI downloads this typically means the temporary signed URL
|
||||||
|
expired mid-download; the transfer can be recovered by resolving a fresh
|
||||||
|
signed URL and re-scheduling with ``continue=true``.
|
||||||
|
"""
|
||||||
|
return "no uri available" in message.lower()
|
||||||
|
|
||||||
|
|
||||||
class Aria2Error(RuntimeError):
|
class Aria2Error(RuntimeError):
|
||||||
"""Raised when aria2 integration fails."""
|
"""Raised when aria2 integration fails."""
|
||||||
|
|
||||||
@@ -145,8 +156,16 @@ class Aria2Downloader:
|
|||||||
disappears (e.g. another download restarted the daemon and
|
disappears (e.g. another download restarted the daemon and
|
||||||
``close()`` cleared ``_transfers``) or the RPC becomes unreachable,
|
``close()`` cleared ``_transfers``) or the RPC becomes unreachable,
|
||||||
the transfer is re-scheduled with ``continue=true`` so the download
|
the transfer is re-scheduled with ``continue=true`` so the download
|
||||||
resumes from the on-disk ``.aria2`` control file. Recovery is bounded
|
resumes from the on-disk ``.aria2`` control file. The same
|
||||||
by ``MAX_TRANSFER_RECOVERY_ATTEMPTS``.
|
re-scheduling happens when aria2 fails with "No URI available"
|
||||||
|
(typically an expired CivitAI signed URL): a fresh URL is resolved
|
||||||
|
and the partial download continues. Recovery is bounded by
|
||||||
|
``MAX_TRANSFER_RECOVERY_ATTEMPTS``.
|
||||||
|
|
||||||
|
Cancellation never leaks daemon transfers: the gid is tracked in
|
||||||
|
``_transfers`` before any post-``addUri`` await, and a gid accepted
|
||||||
|
by the daemon while the caller is being cancelled is removed again
|
||||||
|
before the ``CancelledError`` propagates.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
await self._ensure_process()
|
await self._ensure_process()
|
||||||
@@ -201,14 +220,47 @@ class Aria2Downloader:
|
|||||||
completed_path = self._resolve_completed_path(status, save_path)
|
completed_path = self._resolve_completed_path(status, save_path)
|
||||||
return True, completed_path
|
return True, completed_path
|
||||||
if state == "error":
|
if state == "error":
|
||||||
return False, status.get("errorMessage") or "aria2 download failed"
|
error_message = status.get("errorMessage") or "aria2 download failed"
|
||||||
|
if (
|
||||||
|
_is_no_uri_available_error(error_message)
|
||||||
|
and recovery_attempts < MAX_TRANSFER_RECOVERY_ATTEMPTS
|
||||||
|
):
|
||||||
|
# The signed URL (e.g. CivitAI's) expired before the
|
||||||
|
# transfer finished. Re-registering resolves a fresh
|
||||||
|
# URL and resumes from the on-disk partial payload and
|
||||||
|
# .aria2 control file via ``continue=true``.
|
||||||
|
recovery_attempts += 1
|
||||||
|
logger.warning(
|
||||||
|
"aria2 transfer %s failed with %r; refreshing the "
|
||||||
|
"URL and resuming the partial download "
|
||||||
|
"(attempt %d/%d)",
|
||||||
|
download_id,
|
||||||
|
error_message,
|
||||||
|
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
|
||||||
|
return False, error_message
|
||||||
if state == "removed":
|
if state == "removed":
|
||||||
return False, "Download was cancelled"
|
return False, "Download was cancelled"
|
||||||
|
|
||||||
await asyncio.sleep(self._poll_interval)
|
await asyncio.sleep(self._poll_interval)
|
||||||
finally:
|
finally:
|
||||||
current = self._transfers.get(download_id)
|
current = self._transfers.get(download_id)
|
||||||
if current is not None and current.gid == transfer.gid:
|
if (
|
||||||
|
transfer is not None
|
||||||
|
and current is not None
|
||||||
|
and current.gid == transfer.gid
|
||||||
|
):
|
||||||
self._transfers.pop(download_id, None)
|
self._transfers.pop(download_id, None)
|
||||||
|
|
||||||
async def _get_status_with_retry(
|
async def _get_status_with_retry(
|
||||||
@@ -217,8 +269,9 @@ class Aria2Downloader:
|
|||||||
"""Call get_status with retry for transient RPC failures.
|
"""Call get_status with retry for transient RPC failures.
|
||||||
|
|
||||||
Only retries on :exc:`Aria2Error` (RPC-level failure). Returns
|
Only retries on :exc:`Aria2Error` (RPC-level failure). Returns
|
||||||
``None`` immediately when the download_id is not tracked (a missing
|
``None`` immediately when the transfer is not tracked or its GID is
|
||||||
transfer is not a transient condition, so retrying is pointless).
|
gone from the daemon (a missing transfer is not a transient
|
||||||
|
condition, so retrying is pointless).
|
||||||
|
|
||||||
A single failed RPC call should not immediately fail the download,
|
A single failed RPC call should not immediately fail the download,
|
||||||
because aria2 may be temporarily busy (e.g. finalizing multiple
|
because aria2 may be temporarily busy (e.g. finalizing multiple
|
||||||
@@ -295,21 +348,43 @@ class Aria2Downloader:
|
|||||||
resolved_url != url,
|
resolved_url != url,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Shield the addUri RPC from cancellation: the daemon may accept the
|
||||||
|
# download even when the caller is cancelled while the request is in
|
||||||
|
# flight. On cancellation, wait for the RPC result so the freshly
|
||||||
|
# created gid can be removed instead of leaking an untracked
|
||||||
|
# download that keeps running in the daemon.
|
||||||
|
add_task = asyncio.ensure_future(
|
||||||
|
self._rpc_call("aria2.addUri", [[resolved_url], options])
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
gid = await self._rpc_call("aria2.addUri", [[resolved_url], options])
|
gid = await asyncio.shield(add_task)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
leaked_gid: Any = None
|
||||||
|
try:
|
||||||
|
leaked_gid = await add_task
|
||||||
|
except Exception:
|
||||||
|
leaked_gid = None
|
||||||
|
if isinstance(leaked_gid, str) and leaked_gid:
|
||||||
|
logger.info(
|
||||||
|
"Removing aria2 gid %s accepted while download %s was "
|
||||||
|
"being cancelled",
|
||||||
|
leaked_gid,
|
||||||
|
download_id,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
await self._rpc_call("aria2.forceRemove", [leaked_gid])
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to remove leaked aria2 gid %s for download %s: %s",
|
||||||
|
leaked_gid,
|
||||||
|
download_id,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
raise
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
raise Aria2Error(f"Failed to schedule aria2 download: {exc}") from exc
|
raise Aria2Error(f"Failed to schedule aria2 download: {exc}") from exc
|
||||||
|
|
||||||
logger.debug("aria2 accepted download %s with gid %s", download_id, gid)
|
logger.debug("aria2 accepted download %s with gid %s", download_id, gid)
|
||||||
await self._state_store.upsert(
|
|
||||||
download_id,
|
|
||||||
{
|
|
||||||
"gid": gid,
|
|
||||||
"save_path": save_path,
|
|
||||||
"status": "downloading",
|
|
||||||
"url": url,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
return gid
|
return gid
|
||||||
|
|
||||||
async def _register_transfer(
|
async def _register_transfer(
|
||||||
@@ -328,11 +403,56 @@ class Aria2Downloader:
|
|||||||
headers=headers,
|
headers=headers,
|
||||||
)
|
)
|
||||||
transfer = Aria2Transfer(gid=gid, save_path=os.path.abspath(save_path))
|
transfer = Aria2Transfer(gid=gid, save_path=os.path.abspath(save_path))
|
||||||
|
# Register the transfer before any further await: once the daemon
|
||||||
|
# holds the gid, cancel_download() must be able to find it. An await
|
||||||
|
# in between would open a window where a concurrent cancel reports
|
||||||
|
# "Download task not found" and the daemon keeps downloading
|
||||||
|
# untracked.
|
||||||
self._transfers[download_id] = transfer
|
self._transfers[download_id] = transfer
|
||||||
|
try:
|
||||||
|
await self._state_store.upsert(
|
||||||
|
download_id,
|
||||||
|
{
|
||||||
|
"gid": gid,
|
||||||
|
"save_path": transfer.save_path,
|
||||||
|
"status": "downloading",
|
||||||
|
"url": url,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
# The task was cancelled while persisting state and the
|
||||||
|
# coordinator's cancel ran before the transfer was registered
|
||||||
|
# above. Remove the daemon transfer unless it was deliberately
|
||||||
|
# paused (skip_download preserves paused transfers for resume).
|
||||||
|
status = None
|
||||||
|
try:
|
||||||
|
status = await self.get_status(download_id)
|
||||||
|
except Exception:
|
||||||
|
status = None
|
||||||
|
if status is not None and status.get("status") != "paused":
|
||||||
|
try:
|
||||||
|
await self._rpc_call("aria2.forceRemove", [gid])
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to remove aria2 gid %s for cancelled download %s: %s",
|
||||||
|
gid,
|
||||||
|
download_id,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
current = self._transfers.get(download_id)
|
||||||
|
if current is not None and current.gid == gid:
|
||||||
|
self._transfers.pop(download_id, None)
|
||||||
|
raise
|
||||||
return 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.
|
||||||
|
|
||||||
|
Returns ``None`` when the download_id is not tracked or the daemon no
|
||||||
|
longer knows the transfer's GID (daemon restart / forceRemove). A
|
||||||
|
forgotten GID is permanent, not transient, so the caller's recovery
|
||||||
|
path handles it instead of burning retry attempts on a dead GID.
|
||||||
|
"""
|
||||||
|
|
||||||
transfer = self._transfers.get(download_id)
|
transfer = self._transfers.get(download_id)
|
||||||
if transfer is None:
|
if transfer is None:
|
||||||
@@ -348,8 +468,17 @@ class Aria2Downloader:
|
|||||||
"files",
|
"files",
|
||||||
]
|
]
|
||||||
try:
|
try:
|
||||||
status = await self._rpc_call("aria2.tellStatus", [transfer.gid, keys])
|
status = await self._rpc_call(
|
||||||
|
"aria2.tellStatus", [transfer.gid, keys], log_errors=False
|
||||||
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
|
if "not found" in str(exc).lower():
|
||||||
|
logger.debug(
|
||||||
|
"aria2 GID %s for download %s is gone; treating as lost transfer",
|
||||||
|
transfer.gid,
|
||||||
|
download_id,
|
||||||
|
)
|
||||||
|
return None
|
||||||
raise Aria2Error(f"Failed to query aria2 download status: {exc}") from exc
|
raise Aria2Error(f"Failed to query aria2 download status: {exc}") from exc
|
||||||
|
|
||||||
if isinstance(status, dict):
|
if isinstance(status, dict):
|
||||||
@@ -367,7 +496,9 @@ class Aria2Downloader:
|
|||||||
"files",
|
"files",
|
||||||
]
|
]
|
||||||
try:
|
try:
|
||||||
status = await self._rpc_call("aria2.tellStatus", [gid, keys])
|
status = await self._rpc_call(
|
||||||
|
"aria2.tellStatus", [gid, keys], log_errors=False
|
||||||
|
)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
message = str(exc)
|
message = str(exc)
|
||||||
if "cannot be found" in message.lower() or "not found" in message.lower():
|
if "cannot be found" in message.lower() or "not found" in message.lower():
|
||||||
@@ -434,8 +565,19 @@ class Aria2Downloader:
|
|||||||
try:
|
try:
|
||||||
await self._rpc_call("aria2.forceRemove", [transfer.gid])
|
await self._rpc_call("aria2.forceRemove", [transfer.gid])
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
return {"success": False, "error": str(exc)}
|
if "not found" not in str(exc).lower():
|
||||||
|
return {"success": False, "error": str(exc)}
|
||||||
|
# The daemon already forgot this GID (restart / prior removal),
|
||||||
|
# so the transfer is effectively cancelled.
|
||||||
|
logger.debug(
|
||||||
|
"aria2 GID %s for download %s already gone during cancel",
|
||||||
|
transfer.gid,
|
||||||
|
download_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Drop the in-memory entry as well so a concurrent poll loop does
|
||||||
|
# not mistake the removal for a lost transfer and re-register it.
|
||||||
|
self._transfers.pop(download_id, None)
|
||||||
await self._state_store.remove(download_id)
|
await self._state_store.remove(download_id)
|
||||||
return {"success": True, "message": "Download cancelled successfully"}
|
return {"success": True, "message": "Download cancelled successfully"}
|
||||||
|
|
||||||
@@ -725,7 +867,9 @@ class Aria2Downloader:
|
|||||||
|
|
||||||
return isinstance(result, dict)
|
return isinstance(result, dict)
|
||||||
|
|
||||||
async def _rpc_call(self, method: str, params: list[Any]) -> Any:
|
async def _rpc_call(
|
||||||
|
self, method: str, params: list[Any], *, log_errors: bool = True
|
||||||
|
) -> Any:
|
||||||
if not self._rpc_url:
|
if not self._rpc_url:
|
||||||
raise Aria2Error("aria2 RPC endpoint is not initialized")
|
raise Aria2Error("aria2 RPC endpoint is not initialized")
|
||||||
|
|
||||||
@@ -756,7 +900,10 @@ class Aria2Downloader:
|
|||||||
error = body["error"] or {}
|
error = body["error"] or {}
|
||||||
code = error.get("code") if isinstance(error, dict) else None
|
code = error.get("code") if isinstance(error, dict) else None
|
||||||
message = error.get("message") if isinstance(error, dict) else str(error)
|
message = error.get("message") if isinstance(error, dict) else str(error)
|
||||||
logger.error(
|
# Probing calls (e.g. tellStatus for a GID the daemon may have
|
||||||
|
# forgotten) pass log_errors=False: an expected "not found" must
|
||||||
|
# not spam the log at ERROR level.
|
||||||
|
(logger.error if log_errors else logger.debug)(
|
||||||
"aria2 RPC %s failed with HTTP %s, code=%s, message=%s",
|
"aria2 RPC %s failed with HTTP %s, code=%s, message=%s",
|
||||||
method,
|
method,
|
||||||
response.status,
|
response.status,
|
||||||
@@ -771,7 +918,7 @@ class Aria2Downloader:
|
|||||||
raise Aria2Error(status_message or "Unknown aria2 RPC error")
|
raise Aria2Error(status_message or "Unknown aria2 RPC error")
|
||||||
|
|
||||||
if response.status != 200:
|
if response.status != 200:
|
||||||
logger.error(
|
(logger.error if log_errors else logger.debug)(
|
||||||
"aria2 RPC %s returned unexpected HTTP status %s without error payload: %s",
|
"aria2 RPC %s returned unexpected HTTP status %s without error payload: %s",
|
||||||
method,
|
method,
|
||||||
response.status,
|
response.status,
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
|
|
||||||
from ..utils.constants import VALID_LORA_SUB_TYPES, VALID_CHECKPOINT_SUB_TYPES
|
from ..utils.constants import VALID_LORA_SUB_TYPES, VALID_CHECKPOINT_SUB_TYPES, VALID_OTHER_SUB_TYPES
|
||||||
from ..utils.models import BaseModelMetadata
|
from ..utils.models import BaseModelMetadata
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from ..utils.usage_stats import UsageStats
|
from ..utils.usage_stats import UsageStats
|
||||||
@@ -21,6 +21,7 @@ from .model_query import (
|
|||||||
resolve_sub_type,
|
resolve_sub_type,
|
||||||
)
|
)
|
||||||
from .settings_manager import get_settings_manager
|
from .settings_manager import get_settings_manager
|
||||||
|
from .model_sources import source_group_key
|
||||||
from ..utils.civitai_utils import build_civitai_model_page_url
|
from ..utils.civitai_utils import build_civitai_model_page_url
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -742,29 +743,32 @@ class BaseModelService(ABC):
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def _extract_hf_group_key(item: Dict[str, Any]) -> Optional[str]:
|
def _extract_hf_group_key(item: Dict[str, Any]) -> Optional[str]:
|
||||||
"""Extract `hf:{owner}/{repo}` from item's ``hf_url``, or None."""
|
"""Extract `hf:{owner}/{repo}` from item's ``hf_url``, or None."""
|
||||||
hf_url = item.get("hf_url") if isinstance(item, dict) else None
|
key = BaseModelService._extract_source_group_key(item)
|
||||||
if not hf_url or not isinstance(hf_url, str):
|
return key if key and key.startswith("hf:") else None
|
||||||
return None
|
|
||||||
m = re.match(
|
@staticmethod
|
||||||
r"https?://huggingface\.co/([^/]+/[^/]+)", hf_url.strip()
|
def _extract_source_group_key(item: Dict[str, Any]) -> Optional[str]:
|
||||||
)
|
"""Return the external-source group key for *item*, or None.
|
||||||
if not m:
|
|
||||||
return None
|
Hugging Face keeps the historical ``hf:{owner}/{repo}`` shape; other
|
||||||
return f"hf:{m.group(1)}"
|
platforms use their own short prefix (``ms:`` / ``ta:``).
|
||||||
|
"""
|
||||||
|
return source_group_key(item)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _extract_group_key(item: Dict[str, Any]) -> Union[int, str, None]:
|
def _extract_group_key(item: Dict[str, Any]) -> Union[int, str, None]:
|
||||||
"""Return the group identity key: CivitAI modelId (int) or HF repo (str).
|
"""Return the group identity key.
|
||||||
|
|
||||||
Preference order:
|
Preference order:
|
||||||
1. CivitAI ``modelId`` (int)
|
1. CivitAI ``modelId`` (int)
|
||||||
2. HF repo identity ``hf:{owner}/{repo}`` (str)
|
2. External model source identity, e.g. ``hf:{owner}/{repo}``,
|
||||||
|
``ms:{owner}/{repo}``, ``ta:{model_id}`` (str)
|
||||||
3. ``None`` (no known grouping source)
|
3. ``None`` (no known grouping source)
|
||||||
"""
|
"""
|
||||||
mid = BaseModelService._extract_model_id(item)
|
mid = BaseModelService._extract_model_id(item)
|
||||||
if mid is not None:
|
if mid is not None:
|
||||||
return mid
|
return mid
|
||||||
return BaseModelService._extract_hf_group_key(item)
|
return BaseModelService._extract_source_group_key(item)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _extract_model_id(item: Dict[str, Any]) -> Optional[int]:
|
def _extract_model_id(item: Dict[str, Any]) -> Optional[int]:
|
||||||
@@ -904,6 +908,11 @@ class BaseModelService(ABC):
|
|||||||
and normalized_type not in VALID_CHECKPOINT_SUB_TYPES
|
and normalized_type not in VALID_CHECKPOINT_SUB_TYPES
|
||||||
):
|
):
|
||||||
continue
|
continue
|
||||||
|
if (
|
||||||
|
self.model_type == "other"
|
||||||
|
and normalized_type not in VALID_OTHER_SUB_TYPES
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
|
||||||
type_counts[normalized_type] = type_counts.get(normalized_type, 0) + 1
|
type_counts[normalized_type] = type_counts.get(normalized_type, 0) + 1
|
||||||
|
|
||||||
@@ -1295,6 +1304,27 @@ class BaseModelService(ABC):
|
|||||||
path_for_sorting,
|
path_for_sorting,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _relative_path_folder_group_sort_key(
|
||||||
|
relative_path: str, include_terms: List[str]
|
||||||
|
) -> tuple:
|
||||||
|
"""Group paths by folder, then sort by relevance within each group.
|
||||||
|
|
||||||
|
Folders are ordered alphabetically (case-insensitive) by their full
|
||||||
|
folder path, with root-level files (empty folder) first. Within a
|
||||||
|
folder, paths keep the relevance ordering of
|
||||||
|
``_relative_path_sort_key``. This keeps same-folder entries together
|
||||||
|
in the autocomplete dropdown instead of interleaving them by filename.
|
||||||
|
"""
|
||||||
|
path_for_sorting = BaseModelService._remove_model_extension(
|
||||||
|
relative_path.lower()
|
||||||
|
)
|
||||||
|
folder = path_for_sorting.rpartition(os.sep)[0]
|
||||||
|
|
||||||
|
return (folder,) + BaseModelService._relative_path_sort_key(
|
||||||
|
relative_path, include_terms
|
||||||
|
)
|
||||||
|
|
||||||
async def search_relative_paths(
|
async def search_relative_paths(
|
||||||
self,
|
self,
|
||||||
search_term: str,
|
search_term: str,
|
||||||
@@ -1404,9 +1434,13 @@ class BaseModelService(ABC):
|
|||||||
):
|
):
|
||||||
matching_paths.append(relative_path)
|
matching_paths.append(relative_path)
|
||||||
|
|
||||||
# Sort by relevance (prefix and earliest hits first, then by length and alphabetically)
|
# Group by folder (root first, then alphabetically) and sort by
|
||||||
|
# relevance (prefix and earliest hits, then length and alphabetically)
|
||||||
|
# within each folder group.
|
||||||
matching_paths.sort(
|
matching_paths.sort(
|
||||||
key=lambda relative: self._relative_path_sort_key(relative, include_terms)
|
key=lambda relative: self._relative_path_folder_group_sort_key(
|
||||||
|
relative, include_terms
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
# Apply offset and limit
|
# Apply offset and limit
|
||||||
|
|||||||
@@ -20,6 +20,11 @@ from .recipes import (
|
|||||||
RecipeDownloadError,
|
RecipeDownloadError,
|
||||||
RecipeNotFoundError,
|
RecipeNotFoundError,
|
||||||
)
|
)
|
||||||
|
from .recipes.import_info import (
|
||||||
|
CHANNEL_BATCH_IMPORT_LOCAL,
|
||||||
|
CHANNEL_BATCH_IMPORT_URL,
|
||||||
|
build_import_info,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ImportItemType(Enum):
|
class ImportItemType(Enum):
|
||||||
@@ -71,6 +76,9 @@ class BatchImportProgress:
|
|||||||
tags: List[str] = field(default_factory=list)
|
tags: List[str] = field(default_factory=list)
|
||||||
skip_no_metadata: bool = False
|
skip_no_metadata: bool = False
|
||||||
skip_duplicates: bool = False
|
skip_duplicates: bool = False
|
||||||
|
# Set once any item is skipped due to vendor rate limiting (#1085); lets
|
||||||
|
# the UI surface a "slowing down / try again later" hint.
|
||||||
|
rate_limited: bool = False
|
||||||
|
|
||||||
def to_dict(self) -> Dict[str, Any]:
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
return {
|
return {
|
||||||
@@ -82,6 +90,7 @@ class BatchImportProgress:
|
|||||||
"skipped": self.skipped,
|
"skipped": self.skipped,
|
||||||
"current_item": self.current_item,
|
"current_item": self.current_item,
|
||||||
"status": self.status,
|
"status": self.status,
|
||||||
|
"rate_limited": self.rate_limited,
|
||||||
"started_at": self.started_at,
|
"started_at": self.started_at,
|
||||||
"finished_at": self.finished_at,
|
"finished_at": self.finished_at,
|
||||||
"progress_percent": round((self.completed / self.total) * 100, 1)
|
"progress_percent": round((self.completed / self.total) * 100, 1)
|
||||||
@@ -118,6 +127,10 @@ class AdaptiveConcurrencyController:
|
|||||||
self._task_durations: List[float] = []
|
self._task_durations: List[float] = []
|
||||||
self._recent_errors = 0
|
self._recent_errors = 0
|
||||||
self._recent_successes = 0
|
self._recent_successes = 0
|
||||||
|
# Batch-wide shared semaphore; created lazily on first use so the
|
||||||
|
# controller can also be constructed outside a running event loop.
|
||||||
|
self._semaphore: Optional[asyncio.Semaphore] = None
|
||||||
|
self._semaphore_capacity = initial_concurrency
|
||||||
|
|
||||||
def record_result(self, duration: float, success: bool) -> None:
|
def record_result(self, duration: float, success: bool) -> None:
|
||||||
self._task_durations.append(duration)
|
self._task_durations.append(duration)
|
||||||
@@ -146,7 +159,37 @@ class AdaptiveConcurrencyController:
|
|||||||
self._recent_successes = 0
|
self._recent_successes = 0
|
||||||
|
|
||||||
def get_semaphore(self) -> asyncio.Semaphore:
|
def get_semaphore(self) -> asyncio.Semaphore:
|
||||||
return asyncio.Semaphore(self.current_concurrency)
|
"""Return the batch-wide shared semaphore.
|
||||||
|
|
||||||
|
The same semaphore instance is returned for every item of a batch so
|
||||||
|
the configured concurrency bounds are actually enforced. Previously a
|
||||||
|
fresh semaphore was created per call, letting every item run
|
||||||
|
concurrently and hammering remote metadata providers without any
|
||||||
|
limit.
|
||||||
|
"""
|
||||||
|
if self._semaphore is None:
|
||||||
|
self._semaphore = asyncio.Semaphore(self.current_concurrency)
|
||||||
|
self._semaphore_capacity = self.current_concurrency
|
||||||
|
return self._semaphore
|
||||||
|
|
||||||
|
async def apply_concurrency(self) -> None:
|
||||||
|
"""Synchronize the shared semaphore capacity with ``current_concurrency``.
|
||||||
|
|
||||||
|
Call after ``record_result`` (once per completed item). Growing the
|
||||||
|
capacity is immediate (release). Shrinking requires acquiring a permit
|
||||||
|
and holding it, which is best-effort while other tasks are still
|
||||||
|
running — the capacity converges on subsequent calls.
|
||||||
|
"""
|
||||||
|
semaphore = self.get_semaphore()
|
||||||
|
while self._semaphore_capacity < self.current_concurrency:
|
||||||
|
semaphore.release()
|
||||||
|
self._semaphore_capacity += 1
|
||||||
|
while self._semaphore_capacity > self.current_concurrency:
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(semaphore.acquire(), timeout=0.01)
|
||||||
|
except (asyncio.TimeoutError, asyncio.CancelledError):
|
||||||
|
break
|
||||||
|
self._semaphore_capacity -= 1
|
||||||
|
|
||||||
|
|
||||||
class BatchImportService:
|
class BatchImportService:
|
||||||
@@ -184,6 +227,7 @@ class BatchImportService:
|
|||||||
def cancel_import(self, operation_id: str) -> bool:
|
def cancel_import(self, operation_id: str) -> bool:
|
||||||
if operation_id in self._active_operations:
|
if operation_id in self._active_operations:
|
||||||
self._cancellation_flags[operation_id] = True
|
self._cancellation_flags[operation_id] = True
|
||||||
|
self._logger.info("Cancel requested for batch import operation %s", operation_id)
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -273,6 +317,14 @@ class BatchImportService:
|
|||||||
self._active_operations[operation_id] = progress
|
self._active_operations[operation_id] = progress
|
||||||
self._cancellation_flags[operation_id] = False
|
self._cancellation_flags[operation_id] = False
|
||||||
|
|
||||||
|
self._logger.info(
|
||||||
|
"Starting batch import operation %s: %d item(s) (%d URL(s), %d local path(s))",
|
||||||
|
operation_id,
|
||||||
|
len(import_items),
|
||||||
|
sum(1 for it in import_items if it.item_type == ImportItemType.URL),
|
||||||
|
sum(1 for it in import_items if it.item_type == ImportItemType.LOCAL_PATH),
|
||||||
|
)
|
||||||
|
|
||||||
asyncio.create_task(
|
asyncio.create_task(
|
||||||
self._run_batch_import(
|
self._run_batch_import(
|
||||||
operation_id=operation_id,
|
operation_id=operation_id,
|
||||||
@@ -295,6 +347,12 @@ class BatchImportService:
|
|||||||
skip_duplicates: bool = False,
|
skip_duplicates: bool = False,
|
||||||
) -> str:
|
) -> str:
|
||||||
image_paths = await self._discover_images(directory, recursive)
|
image_paths = await self._discover_images(directory, recursive)
|
||||||
|
self._logger.info(
|
||||||
|
"Batch import directory scan: %d image(s) discovered in %s (recursive=%s)",
|
||||||
|
len(image_paths),
|
||||||
|
directory,
|
||||||
|
recursive,
|
||||||
|
)
|
||||||
|
|
||||||
items = [{"source": path, "type": "local_path"} for path in image_paths]
|
items = [{"source": path, "type": "local_path"} for path in image_paths]
|
||||||
|
|
||||||
@@ -334,6 +392,13 @@ class BatchImportService:
|
|||||||
ext = os.path.splitext(filename)[1].lower()
|
ext = os.path.splitext(filename)[1].lower()
|
||||||
return ext in self.SUPPORTED_EXTENSIONS
|
return ext in self.SUPPORTED_EXTENSIONS
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_rate_limit_error(error: Optional[str]) -> bool:
|
||||||
|
"""Return True when an error payload represents vendor rate limiting."""
|
||||||
|
if not error:
|
||||||
|
return False
|
||||||
|
return "rate limit" in error.lower()
|
||||||
|
|
||||||
async def _run_batch_import(
|
async def _run_batch_import(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -379,6 +444,9 @@ class BatchImportService:
|
|||||||
self._concurrency_controller.record_result(
|
self._concurrency_controller.record_result(
|
||||||
duration, result.get("success", False)
|
duration, result.get("success", False)
|
||||||
)
|
)
|
||||||
|
# Keep the shared batch semaphore in sync with the adaptively
|
||||||
|
# adjusted concurrency so the bounds actually take effect.
|
||||||
|
await self._concurrency_controller.apply_concurrency()
|
||||||
|
|
||||||
if result.get("success"):
|
if result.get("success"):
|
||||||
item.status = ImportStatus.SUCCESS
|
item.status = ImportStatus.SUCCESS
|
||||||
@@ -389,6 +457,17 @@ class BatchImportService:
|
|||||||
item.status = ImportStatus.SKIPPED
|
item.status = ImportStatus.SKIPPED
|
||||||
item.error_message = result.get("error")
|
item.error_message = result.get("error")
|
||||||
progress.skipped += 1
|
progress.skipped += 1
|
||||||
|
elif self._is_rate_limit_error(result.get("error")):
|
||||||
|
# Vendor rate limit is a transient, external condition —
|
||||||
|
# do not pollute the failure count with it (#1085). The
|
||||||
|
# import can simply be re-run later.
|
||||||
|
item.status = ImportStatus.SKIPPED
|
||||||
|
item.error_message = (
|
||||||
|
f"Rate limited by metadata provider; "
|
||||||
|
f"re-run the import later ({result.get('error')})"
|
||||||
|
)
|
||||||
|
progress.skipped += 1
|
||||||
|
progress.rate_limited = True
|
||||||
else:
|
else:
|
||||||
item.status = ImportStatus.FAILED
|
item.status = ImportStatus.FAILED
|
||||||
item.error_message = result.get("error")
|
item.error_message = result.get("error")
|
||||||
@@ -396,13 +475,36 @@ class BatchImportService:
|
|||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self._logger.error(f"Error importing {item.source}: {e}")
|
self._logger.error(f"Error importing {item.source}: {e}")
|
||||||
item.status = ImportStatus.FAILED
|
|
||||||
item.error_message = str(e)
|
|
||||||
item.duration = time.time() - start_time
|
item.duration = time.time() - start_time
|
||||||
progress.failed += 1
|
if self._is_rate_limit_error(str(e)):
|
||||||
|
item.status = ImportStatus.SKIPPED
|
||||||
|
item.error_message = (
|
||||||
|
f"Rate limited by metadata provider; "
|
||||||
|
f"re-run the import later ({e})"
|
||||||
|
)
|
||||||
|
progress.skipped += 1
|
||||||
|
progress.rate_limited = True
|
||||||
|
else:
|
||||||
|
item.status = ImportStatus.FAILED
|
||||||
|
item.error_message = str(e)
|
||||||
|
progress.failed += 1
|
||||||
self._concurrency_controller.record_result(item.duration, False)
|
self._concurrency_controller.record_result(item.duration, False)
|
||||||
|
await self._concurrency_controller.apply_concurrency()
|
||||||
|
|
||||||
progress.completed += 1
|
progress.completed += 1
|
||||||
|
self._logger.info(
|
||||||
|
"Batch import %s: item %d/%d status=%s source=%s%s",
|
||||||
|
operation_id,
|
||||||
|
progress.completed,
|
||||||
|
progress.total,
|
||||||
|
item.status.value,
|
||||||
|
(
|
||||||
|
os.path.basename(item.source)
|
||||||
|
if item.item_type == ImportItemType.LOCAL_PATH
|
||||||
|
else item.source[:50]
|
||||||
|
),
|
||||||
|
(f" error={item.error_message}" if item.error_message else ""),
|
||||||
|
)
|
||||||
await self._broadcast_progress(progress)
|
await self._broadcast_progress(progress)
|
||||||
|
|
||||||
tasks = [process_item(item) for item in progress.items]
|
tasks = [process_item(item) for item in progress.items]
|
||||||
@@ -415,6 +517,15 @@ class BatchImportService:
|
|||||||
|
|
||||||
progress.finished_at = time.time()
|
progress.finished_at = time.time()
|
||||||
progress.current_item = ""
|
progress.current_item = ""
|
||||||
|
self._logger.info(
|
||||||
|
"Batch import %s finished: status=%s total=%d success=%d failed=%d skipped=%d",
|
||||||
|
operation_id,
|
||||||
|
progress.status,
|
||||||
|
progress.total,
|
||||||
|
progress.success,
|
||||||
|
progress.failed,
|
||||||
|
progress.skipped,
|
||||||
|
)
|
||||||
await self._broadcast_progress(progress)
|
await self._broadcast_progress(progress)
|
||||||
|
|
||||||
await asyncio.sleep(5)
|
await asyncio.sleep(5)
|
||||||
@@ -518,6 +629,17 @@ class BatchImportService:
|
|||||||
"loras": loras,
|
"loras": loras,
|
||||||
"gen_params": payload.get("gen_params", {}),
|
"gen_params": payload.get("gen_params", {}),
|
||||||
"source_path": item.source,
|
"source_path": item.source,
|
||||||
|
# Record why this import ended up with no LoRAs so the
|
||||||
|
# recipe modal can explain it (collapsed by default).
|
||||||
|
"import_info": build_import_info(
|
||||||
|
(
|
||||||
|
CHANNEL_BATCH_IMPORT_URL
|
||||||
|
if item.item_type == ImportItemType.URL
|
||||||
|
else CHANNEL_BATCH_IMPORT_LOCAL
|
||||||
|
),
|
||||||
|
payload.get("diagnostics"),
|
||||||
|
loras,
|
||||||
|
),
|
||||||
}
|
}
|
||||||
|
|
||||||
if payload.get("checkpoint"):
|
if payload.get("checkpoint"):
|
||||||
@@ -595,3 +717,6 @@ class BatchImportService:
|
|||||||
def _cleanup_operation(self, operation_id: str) -> None:
|
def _cleanup_operation(self, operation_id: str) -> None:
|
||||||
if operation_id in self._cancellation_flags:
|
if operation_id in self._cancellation_flags:
|
||||||
del self._cancellation_flags[operation_id]
|
del self._cancellation_flags[operation_id]
|
||||||
|
if operation_id in self._active_operations:
|
||||||
|
del self._active_operations[operation_id]
|
||||||
|
self._logger.info("Batch import operation %s cleaned up", operation_id)
|
||||||
|
|||||||
@@ -410,6 +410,10 @@ class CheckpointScanner(ModelScanner):
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def resolve_sub_type_for_path(self, file_path: Optional[str]) -> Optional[str]:
|
||||||
|
"""Resolve sub_type from the configured root that contains the file."""
|
||||||
|
return self._resolve_sub_type(self._find_root_for_file(file_path))
|
||||||
|
|
||||||
def adjust_metadata(self, metadata, file_path, root_path):
|
def adjust_metadata(self, metadata, file_path, root_path):
|
||||||
"""Adjust metadata during scanning to set sub_type."""
|
"""Adjust metadata during scanning to set sub_type."""
|
||||||
sub_type = self._resolve_sub_type(root_path)
|
sub_type = self._resolve_sub_type(root_path)
|
||||||
@@ -419,9 +423,7 @@ class CheckpointScanner(ModelScanner):
|
|||||||
|
|
||||||
def adjust_cached_entry(self, entry: Dict[str, Any]) -> Dict[str, Any]:
|
def adjust_cached_entry(self, entry: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
"""Adjust entries loaded from the persisted cache to ensure sub_type is set."""
|
"""Adjust entries loaded from the persisted cache to ensure sub_type is set."""
|
||||||
sub_type = self._resolve_sub_type(
|
sub_type = self.resolve_sub_type_for_path(entry.get("file_path"))
|
||||||
self._find_root_for_file(entry.get("file_path"))
|
|
||||||
)
|
|
||||||
if sub_type:
|
if sub_type:
|
||||||
entry["sub_type"] = sub_type
|
entry["sub_type"] = sub_type
|
||||||
return entry
|
return entry
|
||||||
|
|||||||
@@ -67,6 +67,8 @@ class CheckpointService(BaseModelService):
|
|||||||
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
|
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
|
||||||
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
||||||
"version_count": model_data.get("version_count"),
|
"version_count": model_data.get("version_count"),
|
||||||
|
"source_platform": model_data.get("source_platform", ""),
|
||||||
|
"source_url": model_data.get("source_url", ""),
|
||||||
"hf_url": model_data.get("hf_url", ""),
|
"hf_url": model_data.get("hf_url", ""),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import logging
|
|||||||
import asyncio
|
import asyncio
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from typing import Any, Optional, Dict, Tuple, List, cast
|
from typing import Any, Optional, Dict, Tuple, List, cast
|
||||||
|
from .connectivity_guard import is_expected_offline_error
|
||||||
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
|
||||||
@@ -46,7 +47,11 @@ class CivArchiveClient:
|
|||||||
"""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:
|
||||||
error = payload if isinstance(payload, str) else "Request failed"
|
# Normalize empty-string failure payloads (e.g. a throttled
|
||||||
|
# connection dropped without a message) so callers never see a
|
||||||
|
# falsy error alongside a None payload — that combination used to
|
||||||
|
# crash downstream None.get() calls.
|
||||||
|
error = payload if isinstance(payload, str) and payload else "Request failed"
|
||||||
return None, error
|
return None, error
|
||||||
if not isinstance(payload, dict):
|
if not isinstance(payload, dict):
|
||||||
return None, "Invalid response structure"
|
return None, "Invalid response structure"
|
||||||
@@ -298,6 +303,8 @@ class CivArchiveClient:
|
|||||||
|
|
||||||
async def _resolve_version_from_files(self, payload: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
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"""
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
return None
|
||||||
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 []
|
||||||
if not isinstance(files, list):
|
if not isinstance(files, list):
|
||||||
@@ -332,10 +339,13 @@ class CivArchiveClient:
|
|||||||
"""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()}")
|
||||||
if error:
|
# Treat a missing payload as an error even when the error string is
|
||||||
if "not found" in error.lower():
|
# falsy; passing None into the split/transform helpers below used to
|
||||||
|
# crash with "'NoneType' object has no attribute 'get'".
|
||||||
|
if error is not None or payload is None:
|
||||||
|
if error and "not found" in error.lower():
|
||||||
return None, "Model not found"
|
return None, "Model not found"
|
||||||
return None, error
|
return None, error or "Request failed"
|
||||||
|
|
||||||
context, version_data, fallback_files = self._split_context(cast(Dict[str, Any], 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)
|
||||||
@@ -352,7 +362,14 @@ class CivArchiveClient:
|
|||||||
except RateLimitError:
|
except RateLimitError:
|
||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error fetching CivArchive model by hash {model_hash[:10]}: {e}")
|
if is_expected_offline_error(str(e)):
|
||||||
|
logger.debug(
|
||||||
|
"Skipping CivArchive model by hash %s while offline: %s",
|
||||||
|
model_hash[:10],
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
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[str, Any]]:
|
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
|
||||||
@@ -362,7 +379,14 @@ class CivArchiveClient:
|
|||||||
if error or payload is None:
|
if error or payload is None:
|
||||||
if error and "not found" in error.lower():
|
if error and "not found" in error.lower():
|
||||||
return None
|
return None
|
||||||
logger.error(f"Error fetching CivArchive model versions for {model_id}: {error}")
|
if is_expected_offline_error(error):
|
||||||
|
logger.debug(
|
||||||
|
"Skipping CivArchive model versions fetch for %s while offline: %s",
|
||||||
|
model_id,
|
||||||
|
error,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.error(f"Error fetching CivArchive model versions for {model_id}: {error}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
data = self._normalize_payload(payload)
|
data = self._normalize_payload(payload)
|
||||||
@@ -426,7 +450,19 @@ class CivArchiveClient:
|
|||||||
if error or payload is None:
|
if error or payload is None:
|
||||||
if error and "not found" in error.lower():
|
if error and "not found" in error.lower():
|
||||||
return None
|
return None
|
||||||
logger.error(f"Error fetching CivArchive model version via API {model_id}/{version_id}: {error}")
|
# The connectivity guard short-circuits requests during its
|
||||||
|
# offline cooldown; that is an expected, transient state, so
|
||||||
|
# log it as DEBUG instead of spamming one ERROR per request
|
||||||
|
# (batch imports can hit this thousands of times).
|
||||||
|
if is_expected_offline_error(error):
|
||||||
|
logger.debug(
|
||||||
|
"Skipping CivArchive model version fetch %s/%s while offline: %s",
|
||||||
|
model_id,
|
||||||
|
version_id,
|
||||||
|
error,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.error(f"Error fetching CivArchive model version via API {model_id}/{version_id}: {error}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
context, version_data, fallback_files = self._split_context(payload)
|
context, version_data, fallback_files = self._split_context(payload)
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ 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
|
from ..utils.constants import MODEL_WEIGHT_FILE_TYPES, is_empty_placeholder_hash
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -180,6 +180,11 @@ 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[str, Any]], Optional[str]]:
|
) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||||
|
if is_empty_placeholder_hash(model_hash):
|
||||||
|
# The empty-hash placeholder (SHA256 of an empty byte string)
|
||||||
|
# matches no real file; CivitAI's by-hash index can contain
|
||||||
|
# polluted entries for it, so never resolve it.
|
||||||
|
return None, "Model not found"
|
||||||
try:
|
try:
|
||||||
success, version = await self._make_request(
|
success, version = await self._make_request(
|
||||||
"GET",
|
"GET",
|
||||||
@@ -500,9 +505,55 @@ 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 get_version_file_mini(
|
||||||
|
self, version_id: int, file_id: int
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Fetch raw stored file info via the model-versions/mini endpoint.
|
||||||
|
|
||||||
|
The public REST API rewrites ``files[].name`` to
|
||||||
|
``"{model}_{version}"`` for non-LoRA model types, so every
|
||||||
|
precision variant of a multi-file version shares one name (#1100).
|
||||||
|
The mini endpoint returns the raw ``ModelFile.name`` in
|
||||||
|
``fileName``. ``file_id`` is mandatory: without it mini picks a
|
||||||
|
file via its own primary-file logic, which can disagree with the
|
||||||
|
REST ``primary`` flag.
|
||||||
|
|
||||||
|
Returns the mini payload dict on success, None on any failure.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
success, data = await self._make_request(
|
||||||
|
"GET",
|
||||||
|
f"{self.base_url}/model-versions/mini/{version_id}",
|
||||||
|
params={"modelFileId": file_id},
|
||||||
|
use_auth=True,
|
||||||
|
)
|
||||||
|
if success and isinstance(data, dict):
|
||||||
|
return data
|
||||||
|
if is_expected_offline_error(data):
|
||||||
|
return None
|
||||||
|
logger.debug(
|
||||||
|
"Mini endpoint lookup failed for version %s file %s: %s",
|
||||||
|
version_id,
|
||||||
|
file_id,
|
||||||
|
data,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
except RateLimitError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug(
|
||||||
|
"Error fetching mini info for version %s file %s: %s",
|
||||||
|
version_id,
|
||||||
|
file_id,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
async def _fetch_version_by_hash(self, model_hash: Optional[str]) -> Optional[Dict[str, Any]]:
|
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
|
||||||
|
if is_empty_placeholder_hash(model_hash):
|
||||||
|
return None
|
||||||
|
|
||||||
success, version = await self._make_request(
|
success, version = await self._make_request(
|
||||||
"GET",
|
"GET",
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import Any, Awaitable, Callable, Dict, Optional
|
from typing import Any, Awaitable, Callable, Dict, Iterable, Optional
|
||||||
|
|
||||||
from .downloader import DownloadProgress
|
from .downloader import DownloadProgress
|
||||||
|
|
||||||
@@ -186,6 +186,14 @@ class DownloadCoordinator:
|
|||||||
download_manager = await self._download_manager_factory()
|
download_manager = await self._download_manager_factory()
|
||||||
return await download_manager.get_active_downloads()
|
return await download_manager.get_active_downloads()
|
||||||
|
|
||||||
|
async def discard_cleared_downloads(self, download_ids: Iterable[str]) -> int:
|
||||||
|
"""Tear down in-memory/aria2 tracking for queue-cleared downloads."""
|
||||||
|
|
||||||
|
if not download_ids:
|
||||||
|
return 0
|
||||||
|
download_manager = await self._download_manager_factory()
|
||||||
|
return await download_manager.discard_cleared_downloads(download_ids)
|
||||||
|
|
||||||
def _parse_optional_int(self, value: Any, field: str) -> Optional[int]:
|
def _parse_optional_int(self, value: Any, field: str) -> Optional[int]:
|
||||||
"""Parse an optional integer from user input."""
|
"""Parse an optional integer from user input."""
|
||||||
|
|
||||||
|
|||||||
+453
-42
@@ -2,6 +2,7 @@
|
|||||||
# Lazy (function-local) imports still count as static edges in basedpyright's
|
# Lazy (function-local) imports still count as static edges in basedpyright's
|
||||||
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
# reportImportCycles, so the ServiceRegistry singleton pattern necessarily forms
|
||||||
# import cycles. Breaking them would require an architectural refactor.
|
# import cycles. Breaking them would require an architectural refactor.
|
||||||
|
import contextlib
|
||||||
import copy
|
import copy
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
@@ -12,30 +13,39 @@ import shutil
|
|||||||
import zipfile
|
import zipfile
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
|
from dataclasses import dataclass, field
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Any, Dict, List, Optional, Set, Tuple, cast
|
from typing import Any, Dict, Iterable, 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,
|
||||||
|
OtherModelMetadata,
|
||||||
|
)
|
||||||
from ..utils.constants import (
|
from ..utils.constants import (
|
||||||
CARD_PREVIEW_WIDTH,
|
CARD_PREVIEW_WIDTH,
|
||||||
DIFFUSION_MODEL_BASE_MODELS,
|
|
||||||
MODEL_WEIGHT_FILE_TYPES,
|
MODEL_WEIGHT_FILE_TYPES,
|
||||||
SUPPORTED_DOWNLOAD_SKIP_BASE_MODELS,
|
SUPPORTED_DOWNLOAD_SKIP_BASE_MODELS,
|
||||||
VALID_LORA_TYPES,
|
VALID_LORA_TYPES,
|
||||||
|
VALID_OTHER_CIVITAI_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, calculate_autov3
|
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 calculate_filename_for_model, sanitize_folder_name
|
||||||
from ..utils.exif_utils import ExifUtils
|
from ..utils.exif_utils import ExifUtils
|
||||||
from ..utils.metadata_manager import MetadataManager
|
from ..utils.metadata_manager import MetadataManager
|
||||||
from .service_registry import ServiceRegistry
|
from .service_registry import ServiceRegistry
|
||||||
|
from .download_routing import is_diffusion_model_download, resolve_other_download_sub_type
|
||||||
from .settings_manager import get_settings_manager
|
from .settings_manager import get_settings_manager
|
||||||
from .metadata_service import get_default_metadata_provider, get_metadata_provider
|
from .metadata_service import get_default_metadata_provider, get_metadata_provider
|
||||||
from .downloader import get_downloader, DownloadProgress, DownloadStreamControl
|
from .downloader import get_downloader, DownloadProgress, DownloadStreamControl
|
||||||
|
from .errors import RateLimitError
|
||||||
from .aria2_downloader import Aria2Error, get_aria2_downloader
|
from .aria2_downloader import Aria2Error, get_aria2_downloader
|
||||||
from .aria2_transfer_state import Aria2TransferStateStore
|
from .aria2_transfer_state import Aria2TransferStateStore
|
||||||
from .download_queue_service import DownloadQueueService
|
from .download_queue_service import DownloadQueueService
|
||||||
|
from .model_lifecycle_service import ModelLifecycleService, load_local_metadata
|
||||||
|
|
||||||
# Download to temporary file first
|
# Download to temporary file first
|
||||||
import tempfile
|
import tempfile
|
||||||
@@ -53,6 +63,12 @@ CIVITAI_DOWNLOAD_URL_PREFIXES = (
|
|||||||
NON_DOWNLOADABLE_PRIMARY_TYPES = ("Config", "Archive", "Workflow", "Training Data")
|
NON_DOWNLOADABLE_PRIMARY_TYPES = ("Config", "Archive", "Workflow", "Training Data")
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _PathSlot:
|
||||||
|
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||||
|
refs: int = 0
|
||||||
|
|
||||||
|
|
||||||
class DownloadManager:
|
class DownloadManager:
|
||||||
_instance = None
|
_instance = None
|
||||||
_lock = asyncio.Lock()
|
_lock = asyncio.Lock()
|
||||||
@@ -82,6 +98,11 @@ class DownloadManager:
|
|||||||
self._aria2_state_store = Aria2TransferStateStore()
|
self._aria2_state_store = Aria2TransferStateStore()
|
||||||
self._restored_persisted_downloads = False
|
self._restored_persisted_downloads = False
|
||||||
self._restore_lock = asyncio.Lock()
|
self._restore_lock = asyncio.Lock()
|
||||||
|
# Refcounted per-target-path locks: two downloads resolving to the
|
||||||
|
# same save_path (e.g. model versions sharing one filename) must not
|
||||||
|
# overlap, or one task's failure cleanup can delete the other's file.
|
||||||
|
self._path_slot_guard: asyncio.Lock = asyncio.Lock()
|
||||||
|
self._path_slots: dict[str, _PathSlot] = {}
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _get_model_download_backend() -> str:
|
def _get_model_download_backend() -> str:
|
||||||
@@ -214,12 +235,21 @@ class DownloadManager:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
async def _get_scanner_for_model_type(self, model_type: str):
|
async def _get_scanner_for_model_type(self, model_type: str):
|
||||||
"""Return the scanner responsible for the given model type."""
|
"""Return the scanner responsible for the given model type.
|
||||||
|
|
||||||
|
Every supported type resolves explicitly — an unknown type must never
|
||||||
|
fall through to the lora scanner (an "other" download would silently
|
||||||
|
dedupe against the lora library).
|
||||||
|
"""
|
||||||
if model_type == "checkpoint":
|
if model_type == "checkpoint":
|
||||||
return await self._get_checkpoint_scanner()
|
return await self._get_checkpoint_scanner()
|
||||||
if model_type == "embedding":
|
if model_type == "embedding":
|
||||||
return await ServiceRegistry.get_embedding_scanner()
|
return await ServiceRegistry.get_embedding_scanner()
|
||||||
return await self._get_lora_scanner()
|
if model_type == "other":
|
||||||
|
return await ServiceRegistry.get_other_scanner()
|
||||||
|
if model_type == "lora":
|
||||||
|
return await self._get_lora_scanner()
|
||||||
|
raise ValueError(f'Unknown model type "{model_type}"')
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _resolve_target_file(
|
def _resolve_target_file(
|
||||||
@@ -704,6 +734,47 @@ class DownloadManager:
|
|||||||
await asyncio.sleep(delay)
|
await asyncio.sleep(delay)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _reconcile_failed_aria2_partial(save_path: str) -> None:
|
||||||
|
"""Reconcile on-disk partial state after a failed aria2 transfer.
|
||||||
|
|
||||||
|
The payload and its ``.aria2`` control file form a resumable pair and
|
||||||
|
are preserved together so a retry (with a refreshed URL when needed)
|
||||||
|
can resume via aria2's ``continue=true``. A control file without its
|
||||||
|
payload cannot resume anything, so the orphan is reported and removed.
|
||||||
|
"""
|
||||||
|
control_path = f"{save_path}.aria2"
|
||||||
|
payload_exists = os.path.exists(save_path)
|
||||||
|
control_exists = os.path.exists(control_path)
|
||||||
|
|
||||||
|
if payload_exists and not control_exists:
|
||||||
|
# If the .aria2 control file is missing, aria2 considers the
|
||||||
|
# download complete. A transient RPC failure may have made us
|
||||||
|
# think the download failed even though the file is fully on disk.
|
||||||
|
# Keep the file so a retry can find it already complete.
|
||||||
|
logger.warning(
|
||||||
|
"aria2 download reported failure but .aria2 file is absent "
|
||||||
|
"for %s — the file is likely complete. Preserving it for retry.",
|
||||||
|
save_path,
|
||||||
|
)
|
||||||
|
elif payload_exists and control_exists:
|
||||||
|
logger.info(
|
||||||
|
"Preserving aria2 partial download for resume: %s", save_path
|
||||||
|
)
|
||||||
|
elif control_exists:
|
||||||
|
logger.warning(
|
||||||
|
"Orphaned aria2 control file without payload: %s — removing it",
|
||||||
|
control_path,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
os.remove(control_path)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to remove orphaned aria2 control file %s: %s",
|
||||||
|
control_path,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
async def _cleanup_cancelled_download_files(
|
async def _cleanup_cancelled_download_files(
|
||||||
self,
|
self,
|
||||||
download_id: str,
|
download_id: str,
|
||||||
@@ -875,6 +946,42 @@ class DownloadManager:
|
|||||||
|
|
||||||
return download_urls
|
return download_urls
|
||||||
|
|
||||||
|
async def _fetch_raw_file_name(
|
||||||
|
self,
|
||||||
|
metadata_provider,
|
||||||
|
version_id: Optional[int],
|
||||||
|
file_id: Any,
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Best-effort lookup of the raw stored filename via the CivitAI
|
||||||
|
model-versions/mini endpoint (#1100). Returns None on any failure so
|
||||||
|
the caller can fall back to the (possibly rewritten) REST name."""
|
||||||
|
if version_id is None or file_id is None:
|
||||||
|
return None
|
||||||
|
fetch = getattr(metadata_provider, "get_version_file_mini", None)
|
||||||
|
if fetch is None:
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
mini_info = await fetch(int(version_id), int(file_id))
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return None
|
||||||
|
except RateLimitError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
logger.debug(
|
||||||
|
"Mini endpoint lookup failed for version %s file %s: %s",
|
||||||
|
version_id,
|
||||||
|
file_id,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
if not isinstance(mini_info, dict):
|
||||||
|
return None
|
||||||
|
raw_name = mini_info.get("fileName")
|
||||||
|
if not isinstance(raw_name, str) or not raw_name.strip():
|
||||||
|
return None
|
||||||
|
# Defensive: never let a path component slip into the filename.
|
||||||
|
return os.path.basename(raw_name.strip()) or None
|
||||||
|
|
||||||
def _build_metadata_for_resume(
|
def _build_metadata_for_resume(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -887,6 +994,8 @@ class DownloadManager:
|
|||||||
return CheckpointMetadata.from_civitai_info(version_info, file_info, save_path)
|
return CheckpointMetadata.from_civitai_info(version_info, file_info, save_path)
|
||||||
if model_type == "embedding":
|
if model_type == "embedding":
|
||||||
return EmbeddingMetadata.from_civitai_info(version_info, file_info, save_path)
|
return EmbeddingMetadata.from_civitai_info(version_info, file_info, save_path)
|
||||||
|
if model_type == "other":
|
||||||
|
return OtherModelMetadata.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[str, Any]) -> Optional[str]:
|
def _resolve_save_path_from_persisted_record(self, record: Dict[str, Any]) -> Optional[str]:
|
||||||
@@ -1100,6 +1209,11 @@ class DownloadManager:
|
|||||||
|
|
||||||
save_path = self._resolve_save_path_from_persisted_record(record)
|
save_path = self._resolve_save_path_from_persisted_record(record)
|
||||||
if save_path is None:
|
if save_path is None:
|
||||||
|
# No resolvable target path (e.g. a queued download whose
|
||||||
|
# paths were never resolved before shutdown): the record
|
||||||
|
# can never be restored, so drop it instead of letting it
|
||||||
|
# accumulate in the state store forever.
|
||||||
|
await self._aria2_state_store.remove(download_id)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if (
|
if (
|
||||||
@@ -1208,6 +1322,24 @@ class DownloadManager:
|
|||||||
)
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
if not os.path.exists(save_path) and os.path.exists(control_path):
|
||||||
|
# A control file without its payload cannot resume
|
||||||
|
# anything; report it and clean up the orphan.
|
||||||
|
logger.warning(
|
||||||
|
"Orphaned aria2 control file without payload for %s: "
|
||||||
|
"%s — removing it",
|
||||||
|
download_id,
|
||||||
|
control_path,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
os.remove(control_path)
|
||||||
|
except OSError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to remove orphaned aria2 control file %s: %s",
|
||||||
|
control_path,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
await self._aria2_state_store.remove(download_id)
|
await self._aria2_state_store.remove(download_id)
|
||||||
|
|
||||||
self._restored_persisted_downloads = True
|
self._restored_persisted_downloads = True
|
||||||
@@ -1324,6 +1456,7 @@ class DownloadManager:
|
|||||||
lora_scanner = await self._get_lora_scanner()
|
lora_scanner = await self._get_lora_scanner()
|
||||||
checkpoint_scanner = await self._get_checkpoint_scanner()
|
checkpoint_scanner = await self._get_checkpoint_scanner()
|
||||||
embedding_scanner = await ServiceRegistry.get_embedding_scanner()
|
embedding_scanner = await ServiceRegistry.get_embedding_scanner()
|
||||||
|
other_scanner = await ServiceRegistry.get_other_scanner()
|
||||||
|
|
||||||
# Check lora scanner first
|
# Check lora scanner first
|
||||||
if await lora_scanner.check_model_version_exists(model_version_id):
|
if await lora_scanner.check_model_version_exists(model_version_id):
|
||||||
@@ -1348,6 +1481,13 @@ class DownloadManager:
|
|||||||
"error": "Model version already exists in embedding library",
|
"error": "Model version already exists in embedding library",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# Check other scanner
|
||||||
|
if await other_scanner.check_model_version_exists(model_version_id):
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": "Model version already exists in other library",
|
||||||
|
}
|
||||||
|
|
||||||
# Use CivArchive provider directly when source is 'civarchive'
|
# Use CivArchive provider directly when source is 'civarchive'
|
||||||
# This prioritizes CivArchive metadata (with mirror availability info) over Civitai
|
# This prioritizes CivArchive metadata (with mirror availability info) over Civitai
|
||||||
if source == "civarchive":
|
if source == "civarchive":
|
||||||
@@ -1386,6 +1526,20 @@ class DownloadManager:
|
|||||||
model_type = "lora"
|
model_type = "lora"
|
||||||
elif model_type_from_info == "textualinversion":
|
elif model_type_from_info == "textualinversion":
|
||||||
model_type = "embedding"
|
model_type = "embedding"
|
||||||
|
elif model_type_from_info in VALID_OTHER_CIVITAI_TYPES:
|
||||||
|
if not get_settings_manager().is_other_models_enabled():
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": (
|
||||||
|
"Other Models management is disabled. Enable it in "
|
||||||
|
"Settings > Library before downloading VAE, upscaler, "
|
||||||
|
"text encoder or CLIP files."
|
||||||
|
),
|
||||||
|
# Machine-readable failure code consumed by the companion
|
||||||
|
# browser extension (docs/other-models-support.md C4).
|
||||||
|
"reason": "other_models_disabled",
|
||||||
|
}
|
||||||
|
model_type = "other"
|
||||||
else:
|
else:
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
@@ -1507,27 +1661,13 @@ class DownloadManager:
|
|||||||
}
|
}
|
||||||
|
|
||||||
# Check if this checkpoint should be treated as a diffusion model
|
# Check if this checkpoint should be treated as a diffusion model
|
||||||
# Priority: (1) any file has type "UNet" or "Diffusion Model",
|
# (shared with the download routing endpoint so the UI location
|
||||||
# (2) baseModel is in DIFFUSION_MODEL_BASE_MODELS
|
# step and the actual download agree on the target roots).
|
||||||
is_diffusion_model = False
|
is_diffusion_model = is_diffusion_model_download(
|
||||||
if model_type == "checkpoint":
|
model_type,
|
||||||
# Check file types first (more direct signal from CivitAI)
|
file_types=(f.get("type", "") for f in version_info.get("files", [])),
|
||||||
version_files = version_info.get("files", [])
|
base_model=base_model_value,
|
||||||
for f in version_files:
|
)
|
||||||
f_type = f.get("type", "")
|
|
||||||
if f_type in ("UNet", "Diffusion Model"):
|
|
||||||
is_diffusion_model = True
|
|
||||||
logger.info(
|
|
||||||
f"File type '{f_type}' detected, routing checkpoint to unet folder"
|
|
||||||
)
|
|
||||||
break
|
|
||||||
|
|
||||||
# Fallback to baseModel name check
|
|
||||||
if not is_diffusion_model and base_model_value in DIFFUSION_MODEL_BASE_MODELS:
|
|
||||||
is_diffusion_model = True
|
|
||||||
logger.info(
|
|
||||||
f"baseModel '{base_model_value}' is a known diffusion model, routing to unet folder"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Existence check after the metadata fetch (#1058):
|
# Existence check after the metadata fetch (#1058):
|
||||||
# - An explicit file selection only blocks when THIS file is
|
# - An explicit file selection only blocks when THIS file is
|
||||||
@@ -1586,6 +1726,13 @@ class DownloadManager:
|
|||||||
"success": False,
|
"success": False,
|
||||||
"error": "Model version already exists in embedding library",
|
"error": "Model version already exists in embedding library",
|
||||||
}
|
}
|
||||||
|
elif model_type == "other":
|
||||||
|
other_scanner = await ServiceRegistry.get_other_scanner()
|
||||||
|
if await other_scanner.check_model_version_exists(version_id):
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": "Model version already exists in other library",
|
||||||
|
}
|
||||||
|
|
||||||
# Handle use_default_paths
|
# Handle use_default_paths
|
||||||
if use_default_paths:
|
if use_default_paths:
|
||||||
@@ -1625,6 +1772,60 @@ class DownloadManager:
|
|||||||
"error": "Default embedding root path not set in settings",
|
"error": "Default embedding root path not set in settings",
|
||||||
}
|
}
|
||||||
save_dir = default_path
|
save_dir = default_path
|
||||||
|
elif model_type == "other":
|
||||||
|
other_sub_type = resolve_other_download_sub_type(
|
||||||
|
model_type_from_info,
|
||||||
|
file_types=(
|
||||||
|
f.get("type", "")
|
||||||
|
for f in version_info.get("files", [])
|
||||||
|
if isinstance(f, dict)
|
||||||
|
),
|
||||||
|
selected_file_type=(
|
||||||
|
target_file.get("type") if explicit_file else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
default_other_roots = (
|
||||||
|
settings_manager.get("default_other_roots") or {}
|
||||||
|
)
|
||||||
|
if other_sub_type and not settings_manager.is_other_sub_type_enabled(
|
||||||
|
other_sub_type
|
||||||
|
):
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": (
|
||||||
|
f"Other-model sub-type '{other_sub_type}' is "
|
||||||
|
f"disabled in settings. Please pick a destination "
|
||||||
|
f"folder explicitly instead of using default paths."
|
||||||
|
),
|
||||||
|
"reason": "other_sub_type_disabled",
|
||||||
|
}
|
||||||
|
default_path = (
|
||||||
|
default_other_roots.get(other_sub_type)
|
||||||
|
if other_sub_type
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
if not isinstance(default_path, str) or not default_path:
|
||||||
|
if other_sub_type:
|
||||||
|
detail = (
|
||||||
|
f"No default root configured for other-model "
|
||||||
|
f"sub-type '{other_sub_type}'"
|
||||||
|
)
|
||||||
|
reason = "other_no_default_root"
|
||||||
|
else:
|
||||||
|
detail = (
|
||||||
|
"Could not determine the other-model sub-type "
|
||||||
|
"from the model metadata"
|
||||||
|
)
|
||||||
|
reason = "other_sub_type_undecidable"
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": (
|
||||||
|
f"{detail}. Please pick a destination folder "
|
||||||
|
f"explicitly instead of using default paths."
|
||||||
|
),
|
||||||
|
"reason": reason,
|
||||||
|
}
|
||||||
|
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)
|
||||||
@@ -1781,6 +1982,24 @@ class DownloadManager:
|
|||||||
if not download_urls:
|
if not download_urls:
|
||||||
return {"success": False, "error": "No mirror URL found"}
|
return {"success": False, "error": "No mirror URL found"}
|
||||||
|
|
||||||
|
# The public REST API rewrites files[].name to
|
||||||
|
# "{model}_{version}" for non-LoRA model types, so every
|
||||||
|
# precision variant of a multi-file version shares one name and
|
||||||
|
# lands on disk with a random short-hash suffix. The mini
|
||||||
|
# endpoint returns the raw stored filename (#1100). CivArchive
|
||||||
|
# already serves raw names.
|
||||||
|
if source != "civarchive":
|
||||||
|
raw_file_name = await self._fetch_raw_file_name(
|
||||||
|
metadata_provider, resolved_version_id, file_info.get("id")
|
||||||
|
)
|
||||||
|
if raw_file_name and raw_file_name != file_info.get("name"):
|
||||||
|
logger.info(
|
||||||
|
"[download] Using raw stored filename '%s' instead of REST name '%s'",
|
||||||
|
raw_file_name,
|
||||||
|
file_info.get("name"),
|
||||||
|
)
|
||||||
|
file_info = {**file_info, "name": raw_file_name}
|
||||||
|
|
||||||
# 3. Prepare download
|
# 3. Prepare download
|
||||||
file_name = file_info.get("name", "")
|
file_name = file_info.get("name", "")
|
||||||
if not file_name:
|
if not file_name:
|
||||||
@@ -1803,6 +2022,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}")
|
||||||
|
elif model_type == "other":
|
||||||
|
metadata = OtherModelMetadata.from_civitai_info(
|
||||||
|
version_info, file_info, save_path
|
||||||
|
)
|
||||||
|
logger.info(f"Creating OtherModelMetadata for {file_name}")
|
||||||
else:
|
else:
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
@@ -2015,6 +2239,8 @@ class DownloadManager:
|
|||||||
scanner = await self._get_checkpoint_scanner()
|
scanner = await self._get_checkpoint_scanner()
|
||||||
elif model_type == "embedding":
|
elif model_type == "embedding":
|
||||||
scanner = await ServiceRegistry.get_embedding_scanner()
|
scanner = await ServiceRegistry.get_embedding_scanner()
|
||||||
|
elif model_type == "other":
|
||||||
|
scanner = await ServiceRegistry.get_other_scanner()
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.debug("Failed to acquire scanner for %s models: %s", model_type, exc)
|
logger.debug("Failed to acquire scanner for %s models: %s", model_type, exc)
|
||||||
|
|
||||||
@@ -2127,6 +2353,28 @@ class DownloadManager:
|
|||||||
|
|
||||||
return formatted_path
|
return formatted_path
|
||||||
|
|
||||||
|
@contextlib.asynccontextmanager
|
||||||
|
async def _exclusive_target_slot(self, target_key: str):
|
||||||
|
async with self._path_slot_guard:
|
||||||
|
slot = self._path_slots.get(target_key)
|
||||||
|
if slot is None:
|
||||||
|
slot = _PathSlot()
|
||||||
|
self._path_slots[target_key] = slot
|
||||||
|
slot.refs += 1
|
||||||
|
try:
|
||||||
|
async with slot.lock:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
async with self._path_slot_guard:
|
||||||
|
slot.refs -= 1
|
||||||
|
if slot.refs <= 0:
|
||||||
|
_ = self._path_slots.pop(target_key, None)
|
||||||
|
|
||||||
|
def _target_slot_key(self, save_dir: str, metadata) -> str:
|
||||||
|
return os.path.abspath(
|
||||||
|
os.path.join(save_dir, os.path.basename(metadata.file_path))
|
||||||
|
)
|
||||||
|
|
||||||
async def _execute_download(
|
async def _execute_download(
|
||||||
self,
|
self,
|
||||||
download_urls: List[str],
|
download_urls: List[str],
|
||||||
@@ -2138,6 +2386,33 @@ class DownloadManager:
|
|||||||
model_type: str = "lora",
|
model_type: str = "lora",
|
||||||
download_id: str | None = None,
|
download_id: str | None = None,
|
||||||
transfer_backend: Optional[str] = None,
|
transfer_backend: Optional[str] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
"""Execute the download serialized against other downloads targeting the same path."""
|
||||||
|
target_key = self._target_slot_key(save_dir, metadata)
|
||||||
|
async with self._exclusive_target_slot(target_key):
|
||||||
|
return await self._execute_download_pipeline(
|
||||||
|
download_urls=download_urls,
|
||||||
|
save_dir=save_dir,
|
||||||
|
metadata=metadata,
|
||||||
|
version_info=version_info,
|
||||||
|
relative_path=relative_path,
|
||||||
|
progress_callback=progress_callback,
|
||||||
|
model_type=model_type,
|
||||||
|
download_id=download_id,
|
||||||
|
transfer_backend=transfer_backend,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _execute_download_pipeline(
|
||||||
|
self,
|
||||||
|
download_urls: List[str],
|
||||||
|
save_dir: str,
|
||||||
|
metadata,
|
||||||
|
version_info: Dict[str, Any],
|
||||||
|
relative_path: str,
|
||||||
|
progress_callback=None,
|
||||||
|
model_type: str = "lora",
|
||||||
|
download_id: str | None = None,
|
||||||
|
transfer_backend: Optional[str] = None,
|
||||||
) -> Dict[str, Any]:
|
) -> 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[Any] = []
|
metadata_entries: List[Any] = []
|
||||||
@@ -2356,20 +2631,8 @@ class DownloadManager:
|
|||||||
break
|
break
|
||||||
|
|
||||||
last_error = result
|
last_error = result
|
||||||
# For aria2: if the .aria2 control file is missing, aria2 considers
|
if transfer_backend == "aria2":
|
||||||
# the download complete. A transient RPC failure may have made us
|
self._reconcile_failed_aria2_partial(save_path)
|
||||||
# think the download failed even though the file is fully on disk.
|
|
||||||
# Keep the file so a retry can find it already complete.
|
|
||||||
if (
|
|
||||||
transfer_backend == "aria2"
|
|
||||||
and os.path.exists(save_path)
|
|
||||||
and not os.path.exists(f"{save_path}.aria2")
|
|
||||||
):
|
|
||||||
logger.warning(
|
|
||||||
"aria2 download reported failure but .aria2 file is absent "
|
|
||||||
"for %s — the file is likely complete. Preserving it for retry.",
|
|
||||||
save_path,
|
|
||||||
)
|
|
||||||
elif os.path.exists(save_path):
|
elif os.path.exists(save_path):
|
||||||
try:
|
try:
|
||||||
os.remove(save_path)
|
os.remove(save_path)
|
||||||
@@ -2474,6 +2737,9 @@ class DownloadManager:
|
|||||||
elif model_type == "embedding":
|
elif model_type == "embedding":
|
||||||
scanner = await ServiceRegistry.get_embedding_scanner()
|
scanner = await ServiceRegistry.get_embedding_scanner()
|
||||||
logger.info(f"Updating embedding cache for {actual_file_paths[0]}")
|
logger.info(f"Updating embedding cache for {actual_file_paths[0]}")
|
||||||
|
elif model_type == "other":
|
||||||
|
scanner = await ServiceRegistry.get_other_scanner()
|
||||||
|
logger.info(f"Updating other-model cache for {actual_file_paths[0]}")
|
||||||
|
|
||||||
adjust_cached_entry = (
|
adjust_cached_entry = (
|
||||||
getattr(scanner, "adjust_cached_entry", None)
|
getattr(scanner, "adjust_cached_entry", None)
|
||||||
@@ -2481,6 +2747,7 @@ class DownloadManager:
|
|||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
downloaded_metadata: List[Dict[str, Any]] = []
|
||||||
for index, entry in enumerate(metadata_entries):
|
for index, entry in enumerate(metadata_entries):
|
||||||
file_path_for_adjust = getattr(
|
file_path_for_adjust = getattr(
|
||||||
entry, "file_path", actual_file_paths[index]
|
entry, "file_path", actual_file_paths[index]
|
||||||
@@ -2523,6 +2790,15 @@ class DownloadManager:
|
|||||||
if scanner is not None:
|
if scanner is not None:
|
||||||
await scanner.add_model_to_cache(metadata_dict, relative_path)
|
await scanner.add_model_to_cache(metadata_dict, relative_path)
|
||||||
|
|
||||||
|
downloaded_metadata.append(metadata_dict)
|
||||||
|
|
||||||
|
await self._apply_download_filename_template(
|
||||||
|
scanner=scanner,
|
||||||
|
model_type=model_type,
|
||||||
|
downloaded_metadata=downloaded_metadata,
|
||||||
|
download_id=download_id,
|
||||||
|
)
|
||||||
|
|
||||||
if transfer_backend == "aria2" and download_id:
|
if transfer_backend == "aria2" and download_id:
|
||||||
await self._aria2_state_store.remove(download_id)
|
await self._aria2_state_store.remove(download_id)
|
||||||
|
|
||||||
@@ -2562,8 +2838,85 @@ class DownloadManager:
|
|||||||
|
|
||||||
return {"success": False, "error": str(e)}
|
return {"success": False, "error": str(e)}
|
||||||
|
|
||||||
|
async def _apply_download_filename_template(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
scanner,
|
||||||
|
model_type: str,
|
||||||
|
downloaded_metadata: List[Dict[str, Any]],
|
||||||
|
download_id: Optional[str],
|
||||||
|
) -> None:
|
||||||
|
"""Rename freshly downloaded models according to the filename template.
|
||||||
|
|
||||||
|
Best-effort post-download step: any failure (including name conflicts)
|
||||||
|
is logged and skipped so a successful download is never turned into a
|
||||||
|
failure by a rename problem.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
if scanner is None or not downloaded_metadata:
|
||||||
|
return
|
||||||
|
|
||||||
|
template = get_settings_manager().get_download_filename_template(
|
||||||
|
model_type
|
||||||
|
)
|
||||||
|
if not template:
|
||||||
|
return
|
||||||
|
|
||||||
|
lifecycle_service = ModelLifecycleService(
|
||||||
|
scanner=scanner,
|
||||||
|
metadata_manager=MetadataManager,
|
||||||
|
metadata_loader=load_local_metadata,
|
||||||
|
recipe_scanner_factory=ServiceRegistry.get_recipe_scanner,
|
||||||
|
)
|
||||||
|
|
||||||
|
for metadata_dict in downloaded_metadata:
|
||||||
|
file_path = metadata_dict.get("file_path")
|
||||||
|
if not isinstance(file_path, str) or not file_path:
|
||||||
|
continue
|
||||||
|
|
||||||
|
new_stem = calculate_filename_for_model(metadata_dict, model_type)
|
||||||
|
if not new_stem:
|
||||||
|
continue
|
||||||
|
|
||||||
|
current_stem = os.path.splitext(os.path.basename(file_path))[0]
|
||||||
|
if new_stem == current_stem or os.path.normcase(
|
||||||
|
new_stem
|
||||||
|
) == os.path.normcase(current_stem):
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = await lifecycle_service.rename_model(
|
||||||
|
file_path=file_path, new_file_name=new_stem
|
||||||
|
)
|
||||||
|
except ValueError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Keeping original filename for %s: %s", file_path, exc
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
new_file_path = result.get("new_file_path")
|
||||||
|
if download_id and isinstance(new_file_path, str):
|
||||||
|
info = self._active_downloads.get(download_id)
|
||||||
|
if info is None:
|
||||||
|
continue
|
||||||
|
if info.get("file_path") == file_path:
|
||||||
|
info["file_path"] = new_file_path
|
||||||
|
extracted = info.get("extracted_paths")
|
||||||
|
if isinstance(extracted, list):
|
||||||
|
info["extracted_paths"] = [
|
||||||
|
new_file_path if path == file_path else path
|
||||||
|
for path in extracted
|
||||||
|
]
|
||||||
|
except Exception as exc: # Rename phase must never fail the download
|
||||||
|
logger.warning(
|
||||||
|
"Filename template rename failed for %s download: %s",
|
||||||
|
model_type,
|
||||||
|
exc,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
|
||||||
def _get_supported_extensions_for_type(self, model_type: str) -> Set[str]:
|
def _get_supported_extensions_for_type(self, model_type: str) -> Set[str]:
|
||||||
if model_type == "checkpoint":
|
if model_type in ("checkpoint", "other"):
|
||||||
return {
|
return {
|
||||||
".ckpt",
|
".ckpt",
|
||||||
".pt",
|
".pt",
|
||||||
@@ -2897,6 +3250,64 @@ 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 discard_cleared_downloads(self, download_ids: Iterable[str]) -> int:
|
||||||
|
"""Stop in-memory tracking for downloads cleared from the queue.
|
||||||
|
|
||||||
|
Cancels asyncio tasks, removes live aria2 transfers and drops the
|
||||||
|
persisted aria2 state so cleared downloads cannot keep polling the
|
||||||
|
daemon or be resurrected as ghost entries on the next restart.
|
||||||
|
Partial files on disk are preserved; unlike ``cancel_download`` no
|
||||||
|
files are deleted.
|
||||||
|
|
||||||
|
Returns the number of downloads that had any in-memory or persisted
|
||||||
|
tracking removed.
|
||||||
|
"""
|
||||||
|
discarded = 0
|
||||||
|
aria2_downloader = None
|
||||||
|
|
||||||
|
for download_id in download_ids:
|
||||||
|
task = self._download_tasks.get(download_id)
|
||||||
|
info = self._active_downloads.get(download_id)
|
||||||
|
persisted = await self._aria2_state_store.get(download_id)
|
||||||
|
if task is None and info is None and persisted is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
discarded += 1
|
||||||
|
|
||||||
|
if task is not None:
|
||||||
|
task.cancel()
|
||||||
|
|
||||||
|
pause_control = self._pause_events.pop(download_id, None)
|
||||||
|
if pause_control is not None:
|
||||||
|
pause_control.resume()
|
||||||
|
|
||||||
|
if task is not None:
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(asyncio.shield(task), timeout=2.0)
|
||||||
|
except (asyncio.CancelledError, asyncio.TimeoutError):
|
||||||
|
pass
|
||||||
|
|
||||||
|
self._download_tasks.pop(download_id, None)
|
||||||
|
self._active_downloads.pop(download_id, None)
|
||||||
|
|
||||||
|
backend = (info or persisted or {}).get("transfer_backend") or "python"
|
||||||
|
if backend == "aria2":
|
||||||
|
if aria2_downloader is None:
|
||||||
|
aria2_downloader = await get_aria2_downloader()
|
||||||
|
if await aria2_downloader.has_transfer(download_id):
|
||||||
|
try:
|
||||||
|
await aria2_downloader.cancel_download(download_id)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Failed to remove aria2 transfer for cleared download %s: %s",
|
||||||
|
download_id,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
await self._aria2_state_store.remove(download_id)
|
||||||
|
|
||||||
|
return discarded
|
||||||
|
|
||||||
async def pause_download(self, download_id: str) -> Dict[str, Any]:
|
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."""
|
||||||
|
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ import logging
|
|||||||
import os
|
import os
|
||||||
import sqlite3
|
import sqlite3
|
||||||
import time
|
import time
|
||||||
from typing import Any, Optional
|
from typing import Any, List, Optional
|
||||||
|
|
||||||
from ..utils.cache_paths import get_cache_base_dir
|
from ..utils.cache_paths import get_cache_base_dir
|
||||||
|
|
||||||
@@ -390,23 +390,31 @@ class DownloadQueueService:
|
|||||||
conn.commit()
|
conn.commit()
|
||||||
return True
|
return True
|
||||||
|
|
||||||
async def clear_queue(self, status_filter: Optional[str] = None) -> int:
|
async def clear_queue(self, status_filter: Optional[str] = None) -> List[str]:
|
||||||
"""Remove items from the queue.
|
"""Remove items from the queue.
|
||||||
|
|
||||||
When *status_filter* is provided only items with that status are
|
When *status_filter* is provided only items with that status are
|
||||||
deleted. Returns the number of deleted rows.
|
deleted. Returns the ``download_id`` values of the deleted rows so
|
||||||
|
callers can also tear down any in-memory tracking for them.
|
||||||
"""
|
"""
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
conn = self._get_conn()
|
conn = self._get_conn()
|
||||||
if status_filter is not None:
|
if status_filter is not None:
|
||||||
cursor = conn.execute(
|
rows = conn.execute(
|
||||||
|
"SELECT download_id FROM download_queue WHERE status = ?",
|
||||||
|
(status_filter,),
|
||||||
|
).fetchall()
|
||||||
|
conn.execute(
|
||||||
"DELETE FROM download_queue WHERE status = ?",
|
"DELETE FROM download_queue WHERE status = ?",
|
||||||
(status_filter,),
|
(status_filter,),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
cursor = conn.execute("DELETE FROM download_queue")
|
rows = conn.execute(
|
||||||
|
"SELECT download_id FROM download_queue"
|
||||||
|
).fetchall()
|
||||||
|
conn.execute("DELETE FROM download_queue")
|
||||||
conn.commit()
|
conn.commit()
|
||||||
return cursor.rowcount
|
return [row["download_id"] for row in rows]
|
||||||
|
|
||||||
async def complete_download(
|
async def complete_download(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@@ -0,0 +1,103 @@
|
|||||||
|
"""Shared download routing logic.
|
||||||
|
|
||||||
|
Decides whether a download initiated from the checkpoint library should be
|
||||||
|
routed to the unet/diffusion-model roots instead of the checkpoint roots.
|
||||||
|
Used by both the download manager (at download time) and the download
|
||||||
|
routing HTTP endpoint (when the user picks a location in the UI), so the
|
||||||
|
two can never disagree.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Iterable, Optional
|
||||||
|
|
||||||
|
from ..utils.constants import (
|
||||||
|
CIVITAI_FILE_TYPE_TO_OTHER_SUB_TYPE,
|
||||||
|
CIVITAI_TYPE_TO_OTHER_SUB_TYPE,
|
||||||
|
DIFFUSION_MODEL_BASE_MODELS,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# File types reported by the CivitAI API that indicate a raw diffusion
|
||||||
|
# model (loaded via UNETLoader in ComfyUI) rather than a full checkpoint.
|
||||||
|
DIFFUSION_FILE_TYPES = frozenset({"UNet", "Diffusion Model"})
|
||||||
|
|
||||||
|
|
||||||
|
def is_diffusion_model_download(
|
||||||
|
model_type: str,
|
||||||
|
file_types: Iterable[str] = (),
|
||||||
|
base_model: str = "",
|
||||||
|
) -> bool:
|
||||||
|
"""Return True when a download should be routed to the unet roots.
|
||||||
|
|
||||||
|
Only applies to downloads initiated from the checkpoint library.
|
||||||
|
Priority: (1) any file has type "UNet" or "Diffusion Model" (the more
|
||||||
|
direct signal from CivitAI), (2) baseModel is a known diffusion model.
|
||||||
|
"""
|
||||||
|
if model_type != "checkpoint":
|
||||||
|
return False
|
||||||
|
|
||||||
|
for file_type in file_types:
|
||||||
|
if file_type in DIFFUSION_FILE_TYPES:
|
||||||
|
logger.info(
|
||||||
|
"File type '%s' detected, routing checkpoint to unet folder",
|
||||||
|
file_type,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
if base_model in DIFFUSION_MODEL_BASE_MODELS:
|
||||||
|
logger.info(
|
||||||
|
"baseModel '%s' is a known diffusion model, routing to unet folder",
|
||||||
|
base_model,
|
||||||
|
)
|
||||||
|
return True
|
||||||
|
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_other_download_sub_type(
|
||||||
|
civitai_model_type: str,
|
||||||
|
file_types: Iterable[str] = (),
|
||||||
|
selected_file_type: Optional[str] = None,
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Resolve the "other"-page sub_type for a download.
|
||||||
|
|
||||||
|
Fixed priority (locked design, docs/plans/other-models-page.md §9.2):
|
||||||
|
|
||||||
|
1. Explicit user file pick — when the picked file's type maps, it wins
|
||||||
|
even when model.type maps to something else.
|
||||||
|
2. model.type via CIVITAI_TYPE_TO_OTHER_SUB_TYPE.
|
||||||
|
3. file.type fallback — only when model.type maps to nothing. Must NOT
|
||||||
|
override a mapped model.type: checkpoint models routinely bundle
|
||||||
|
VAE/Text Encoder component files.
|
||||||
|
4. Still undecidable -> None (caller must ask the user for a folder).
|
||||||
|
"""
|
||||||
|
if selected_file_type:
|
||||||
|
mapped = CIVITAI_FILE_TYPE_TO_OTHER_SUB_TYPE.get(selected_file_type)
|
||||||
|
if mapped:
|
||||||
|
logger.info(
|
||||||
|
"Explicit file pick type '%s' routes other download to '%s'",
|
||||||
|
selected_file_type,
|
||||||
|
mapped,
|
||||||
|
)
|
||||||
|
return mapped
|
||||||
|
|
||||||
|
normalized_model_type = (civitai_model_type or "").strip().lower()
|
||||||
|
mapped = CIVITAI_TYPE_TO_OTHER_SUB_TYPE.get(normalized_model_type)
|
||||||
|
if mapped:
|
||||||
|
return mapped
|
||||||
|
|
||||||
|
for file_type in file_types:
|
||||||
|
mapped = CIVITAI_FILE_TYPE_TO_OTHER_SUB_TYPE.get(file_type)
|
||||||
|
if mapped:
|
||||||
|
logger.info(
|
||||||
|
"model.type '%s' unmapped; file type '%s' routes other download to '%s'",
|
||||||
|
civitai_model_type,
|
||||||
|
file_type,
|
||||||
|
mapped,
|
||||||
|
)
|
||||||
|
return mapped
|
||||||
|
|
||||||
|
return None
|
||||||
@@ -236,24 +236,33 @@ class DownloadedVersionHistoryService:
|
|||||||
return
|
return
|
||||||
|
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
conn = self._get_conn()
|
# The connection is created with check_same_thread=False and all
|
||||||
conn.executemany(
|
# access is serialized by self._lock, so the executemany upsert +
|
||||||
"""
|
# commit can run in the default executor without blocking the
|
||||||
INSERT INTO downloaded_model_versions (
|
# event loop on large hydration payloads.
|
||||||
model_type, version_id, model_id, first_seen_at, last_seen_at,
|
loop = asyncio.get_running_loop()
|
||||||
source, last_file_path, last_library_name, is_deleted_override
|
await loop.run_in_executor(None, self._mark_downloaded_bulk_sync, payload)
|
||||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0)
|
|
||||||
ON CONFLICT(model_type, version_id) DO UPDATE SET
|
def _mark_downloaded_bulk_sync(self, payload: Sequence[tuple[object, ...]]) -> None:
|
||||||
model_id = COALESCE(excluded.model_id, downloaded_model_versions.model_id),
|
"""Synchronous executemany upsert + commit; runs in a worker thread."""
|
||||||
last_seen_at = excluded.last_seen_at,
|
conn = self._get_conn()
|
||||||
source = excluded.source,
|
conn.executemany(
|
||||||
last_file_path = COALESCE(excluded.last_file_path, downloaded_model_versions.last_file_path),
|
"""
|
||||||
last_library_name = COALESCE(excluded.last_library_name, downloaded_model_versions.last_library_name),
|
INSERT INTO downloaded_model_versions (
|
||||||
is_deleted_override = 0
|
model_type, version_id, model_id, first_seen_at, last_seen_at,
|
||||||
""",
|
source, last_file_path, last_library_name, is_deleted_override
|
||||||
payload,
|
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, 0)
|
||||||
)
|
ON CONFLICT(model_type, version_id) DO UPDATE SET
|
||||||
conn.commit()
|
model_id = COALESCE(excluded.model_id, downloaded_model_versions.model_id),
|
||||||
|
last_seen_at = excluded.last_seen_at,
|
||||||
|
source = excluded.source,
|
||||||
|
last_file_path = COALESCE(excluded.last_file_path, downloaded_model_versions.last_file_path),
|
||||||
|
last_library_name = COALESCE(excluded.last_library_name, downloaded_model_versions.last_library_name),
|
||||||
|
is_deleted_override = 0
|
||||||
|
""",
|
||||||
|
payload,
|
||||||
|
)
|
||||||
|
conn.commit()
|
||||||
|
|
||||||
async def mark_as_deleted(self, model_type: str, version_id: int) -> None:
|
async def mark_as_deleted(self, model_type: str, version_id: int) -> None:
|
||||||
normalized_type = _normalize_model_type(model_type)
|
normalized_type = _normalize_model_type(model_type)
|
||||||
|
|||||||
+125
-57
@@ -32,6 +32,7 @@ from .connectivity_guard import (
|
|||||||
ConnectivityGuard,
|
ConnectivityGuard,
|
||||||
)
|
)
|
||||||
from .errors import RateLimitError
|
from .errors import RateLimitError
|
||||||
|
from .rate_limit_coordinator import RateLimitCoordinator
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -595,6 +596,21 @@ class Downloader:
|
|||||||
False,
|
False,
|
||||||
"File not found - the download link may be invalid or expired.",
|
"File not found - the download link may be invalid or expired.",
|
||||||
)
|
)
|
||||||
|
elif response.status == 429:
|
||||||
|
# Register the vendor's cooldown so API calls through
|
||||||
|
# make_request queue behind it (#1085). The download
|
||||||
|
# itself fails as before; retry policy stays with the
|
||||||
|
# caller (download manager).
|
||||||
|
retry_after = self._extract_retry_after(response.headers)
|
||||||
|
coordinator = await RateLimitCoordinator.get_instance()
|
||||||
|
if coordinator.enabled:
|
||||||
|
coordinator.register_rate_limit(
|
||||||
|
self._guard_destination(url), retry_after
|
||||||
|
)
|
||||||
|
logger.warning(
|
||||||
|
f"Rate limited (429) for {url}, retry_after={retry_after}"
|
||||||
|
)
|
||||||
|
return False, f"Download rate limited (429), retry after {retry_after}s"
|
||||||
else:
|
else:
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Download failed for {url} with status {response.status}"
|
f"Download failed for {url} with status {response.status}"
|
||||||
@@ -972,6 +988,11 @@ class Downloader:
|
|||||||
elif response.status == 429:
|
elif response.status == 429:
|
||||||
raw_retry_after = response.headers.get("Retry-After")
|
raw_retry_after = response.headers.get("Retry-After")
|
||||||
retry_after = _parse_retry_after(raw_retry_after or "")
|
retry_after = _parse_retry_after(raw_retry_after or "")
|
||||||
|
# Register the vendor's cooldown so API calls through
|
||||||
|
# make_request queue behind it (#1085).
|
||||||
|
coordinator = await RateLimitCoordinator.get_instance()
|
||||||
|
if coordinator.enabled:
|
||||||
|
coordinator.register_rate_limit(destination, retry_after)
|
||||||
if raw_retry_after:
|
if raw_retry_after:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Rate limited (429) for %s, Retry-After: %ss", url, retry_after
|
"Rate limited (429) for %s, Retry-After: %ss", url, retry_after
|
||||||
@@ -1041,6 +1062,14 @@ class Downloader:
|
|||||||
if response.status == 200:
|
if response.status == 200:
|
||||||
guard.register_success(destination)
|
guard.register_success(destination)
|
||||||
return True, dict(response.headers)
|
return True, dict(response.headers)
|
||||||
|
elif response.status == 429:
|
||||||
|
# Register the vendor's cooldown so API calls through
|
||||||
|
# make_request queue behind it (#1085).
|
||||||
|
retry_after = self._extract_retry_after(response.headers)
|
||||||
|
coordinator = await RateLimitCoordinator.get_instance()
|
||||||
|
if coordinator.enabled:
|
||||||
|
coordinator.register_rate_limit(destination, retry_after)
|
||||||
|
return False, f"Head request rate limited (429), retry after {retry_after}s"
|
||||||
else:
|
else:
|
||||||
return False, f"Head request failed with status {response.status}"
|
return False, f"Head request failed with status {response.status}"
|
||||||
|
|
||||||
@@ -1074,74 +1103,113 @@ class Downloader:
|
|||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple[bool, Union[Dict, str]]: (success, response data or error message)
|
Tuple[bool, Union[Dict, str]]: (success, response data or error message)
|
||||||
|
|
||||||
|
When the rate-limit gate is enabled (``rate_limit_gate_enabled``),
|
||||||
|
requests are paced per destination and 429 responses are honored by
|
||||||
|
waiting out the ``Retry-After`` window (bounded by
|
||||||
|
``rate_limit_max_wait_seconds``) before re-sending. A ``RateLimitError``
|
||||||
|
returned after gate involvement is marked with ``gate_handled = True``
|
||||||
|
so downstream retry helpers do not wait a second time.
|
||||||
"""
|
"""
|
||||||
guard = await ConnectivityGuard.get_instance()
|
guard = await ConnectivityGuard.get_instance()
|
||||||
destination = self._guard_destination(url)
|
destination = self._guard_destination(url)
|
||||||
|
# Fail fast on transport-level outages before pacing: there is no
|
||||||
|
# point waiting out a vendor cooldown while the network is down.
|
||||||
if guard.should_block_request(destination):
|
if guard.should_block_request(destination):
|
||||||
return False, OFFLINE_COOLDOWN_ERROR
|
return False, OFFLINE_COOLDOWN_ERROR
|
||||||
|
|
||||||
try:
|
coordinator = await RateLimitCoordinator.get_instance()
|
||||||
session = await self.session
|
gate_enabled = coordinator.enabled
|
||||||
# Debug log for proxy mode at request time
|
# Safety bound on the wait-and-resend loop; each 429 normally exits
|
||||||
if self.proxy_url:
|
# via the wait cap in wait_for_slot, this covers pathological 429s
|
||||||
logger.debug(f"[make_request] Using app-level proxy: {self.proxy_url}")
|
# with tiny Retry-After values.
|
||||||
else:
|
max_resend_attempts = 5
|
||||||
logger.debug(
|
attempt = 0
|
||||||
"[make_request] Using system-level proxy (trust_env) if configured."
|
|
||||||
)
|
|
||||||
|
|
||||||
# Prepare headers
|
while True:
|
||||||
headers = self._get_auth_headers(use_auth)
|
if gate_enabled:
|
||||||
if custom_headers:
|
try:
|
||||||
headers.update(custom_headers)
|
await coordinator.wait_for_slot(destination)
|
||||||
|
except RateLimitError as exc:
|
||||||
|
exc.gate_handled = True
|
||||||
|
return False, exc
|
||||||
|
|
||||||
# Add proxy to kwargs if not already present
|
try:
|
||||||
if "proxy" not in kwargs:
|
session = await self.session
|
||||||
kwargs["proxy"] = self.proxy_url
|
# Debug log for proxy mode at request time
|
||||||
|
if self.proxy_url:
|
||||||
async with session.request(
|
logger.debug(f"[make_request] Using app-level proxy: {self.proxy_url}")
|
||||||
method, url, headers=headers, **kwargs
|
|
||||||
) as response:
|
|
||||||
if response.status == 200:
|
|
||||||
guard.register_success(destination)
|
|
||||||
# Try to parse as JSON, fall back to text
|
|
||||||
try:
|
|
||||||
data = await response.json()
|
|
||||||
return True, data
|
|
||||||
except:
|
|
||||||
text = await response.text()
|
|
||||||
return True, text
|
|
||||||
elif response.status == 401:
|
|
||||||
return False, "Unauthorized access - invalid or missing API key"
|
|
||||||
elif response.status == 403:
|
|
||||||
return False, "Access forbidden"
|
|
||||||
elif response.status == 404:
|
|
||||||
return False, "Resource not found"
|
|
||||||
elif response.status == 429:
|
|
||||||
retry_after = self._extract_retry_after(response.headers)
|
|
||||||
error_msg = "Request rate limited"
|
|
||||||
logger.warning(
|
|
||||||
"Rate limit encountered for %s %s; retry_after=%s",
|
|
||||||
method,
|
|
||||||
url,
|
|
||||||
retry_after,
|
|
||||||
)
|
|
||||||
return False, RateLimitError(
|
|
||||||
error_msg,
|
|
||||||
retry_after=retry_after,
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
return False, f"Request failed with status {response.status}"
|
logger.debug(
|
||||||
|
"[make_request] Using system-level proxy (trust_env) if configured."
|
||||||
|
)
|
||||||
|
|
||||||
except Exception as e:
|
# Prepare headers
|
||||||
if guard.is_network_unreachable_error(e):
|
headers = self._get_auth_headers(use_auth)
|
||||||
guard.register_network_failure(e, destination)
|
if custom_headers:
|
||||||
if guard.should_block_request(destination):
|
headers.update(custom_headers)
|
||||||
return False, OFFLINE_COOLDOWN_ERROR
|
|
||||||
logger.debug("Network unavailable for %s %s: %s", method, url, e)
|
# Add proxy to kwargs if not already present
|
||||||
|
if "proxy" not in kwargs:
|
||||||
|
kwargs["proxy"] = self.proxy_url
|
||||||
|
|
||||||
|
async with session.request(
|
||||||
|
method, url, headers=headers, **kwargs
|
||||||
|
) as response:
|
||||||
|
if response.status == 200:
|
||||||
|
guard.register_success(destination)
|
||||||
|
if gate_enabled:
|
||||||
|
coordinator.register_success(destination)
|
||||||
|
# Try to parse as JSON, fall back to text
|
||||||
|
try:
|
||||||
|
data = await response.json()
|
||||||
|
return True, data
|
||||||
|
except:
|
||||||
|
text = await response.text()
|
||||||
|
return True, text
|
||||||
|
elif response.status == 401:
|
||||||
|
return False, "Unauthorized access - invalid or missing API key"
|
||||||
|
elif response.status == 403:
|
||||||
|
return False, "Access forbidden"
|
||||||
|
elif response.status == 404:
|
||||||
|
return False, "Resource not found"
|
||||||
|
elif response.status == 429:
|
||||||
|
retry_after = self._extract_retry_after(response.headers)
|
||||||
|
error_msg = "Request rate limited"
|
||||||
|
if not gate_enabled:
|
||||||
|
logger.warning(
|
||||||
|
"Rate limit encountered for %s %s; retry_after=%s",
|
||||||
|
method,
|
||||||
|
url,
|
||||||
|
retry_after,
|
||||||
|
)
|
||||||
|
return False, RateLimitError(
|
||||||
|
error_msg,
|
||||||
|
retry_after=retry_after,
|
||||||
|
)
|
||||||
|
# The coordinator logs the cooldown notice (INFO once
|
||||||
|
# per window, DEBUG on extension).
|
||||||
|
coordinator.register_rate_limit(destination, retry_after)
|
||||||
|
attempt += 1
|
||||||
|
if attempt >= max_resend_attempts:
|
||||||
|
error = RateLimitError(error_msg, retry_after=retry_after)
|
||||||
|
error.gate_handled = True
|
||||||
|
return False, error
|
||||||
|
# Loop back: wait_for_slot blocks until the cooldown
|
||||||
|
# elapses (or raises once the wait exceeds the cap).
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
return False, f"Request failed with status {response.status}"
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
if guard.is_network_unreachable_error(e):
|
||||||
|
guard.register_network_failure(e, destination)
|
||||||
|
if guard.should_block_request(destination):
|
||||||
|
return False, OFFLINE_COOLDOWN_ERROR
|
||||||
|
logger.debug("Network unavailable for %s %s: %s", method, url, e)
|
||||||
|
return False, str(e)
|
||||||
|
logger.error(f"Error making {method} request to {url}: {e}")
|
||||||
return False, str(e)
|
return False, str(e)
|
||||||
logger.error(f"Error making {method} request to {url}: {e}")
|
|
||||||
return False, str(e)
|
|
||||||
|
|
||||||
async def close(self):
|
async def close(self):
|
||||||
"""Close the HTTP session"""
|
"""Close the HTTP session"""
|
||||||
|
|||||||
@@ -67,6 +67,8 @@ class EmbeddingService(BaseModelService):
|
|||||||
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
|
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
|
||||||
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
||||||
"version_count": model_data.get("version_count"),
|
"version_count": model_data.get("version_count"),
|
||||||
|
"source_platform": model_data.get("source_platform", ""),
|
||||||
|
"source_url": model_data.get("source_url", ""),
|
||||||
"hf_url": model_data.get("hf_url", ""),
|
"hf_url": model_data.get("hf_url", ""),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+150
-83
@@ -11,6 +11,7 @@ from __future__ import annotations
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
@@ -32,8 +33,26 @@ _catalog_cache: Optional[Dict[str, List[str]]] = None
|
|||||||
# ``{provider_id: {model_id: max_output_tokens}}``.
|
# ``{provider_id: {model_id: max_output_tokens}}``.
|
||||||
_model_output_limits: Dict[str, Dict[str, int]] = {}
|
_model_output_limits: Dict[str, Dict[str, int]] = {}
|
||||||
|
|
||||||
|
# Monotonic timestamp of the last failed catalog fetch (None = no failure
|
||||||
|
# yet). Failed fetches are negatively cached: further calls return the
|
||||||
|
# empty fallback without hitting the network until the cooldown elapses,
|
||||||
|
# so users on broken networks don't stall on every settings-modal open.
|
||||||
|
_catalog_last_failure: Optional[float] = None
|
||||||
|
_CATALOG_FAILURE_COOLDOWN = 600.0 # seconds
|
||||||
|
|
||||||
|
# Serializes catalog fetches so concurrent callers don't duplicate requests.
|
||||||
|
_catalog_lock = asyncio.Lock()
|
||||||
|
|
||||||
_CATALOG_TIMEOUT = aiohttp.ClientTimeout(total=30)
|
_CATALOG_TIMEOUT = aiohttp.ClientTimeout(total=30)
|
||||||
|
|
||||||
|
# Cloudflare serves brotli when the client advertises it, and brotli is a
|
||||||
|
# required dependency here — a corrupted br stream can crash the native
|
||||||
|
# decoder with a Windows access violation (issue #1099). Request gzip
|
||||||
|
# instead; zlib decompression is not affected and corrupt gzip data only
|
||||||
|
# raises ContentEncodingError (an aiohttp.ClientError subclass), which the
|
||||||
|
# exception handlers below already catch.
|
||||||
|
_NO_BROTLI_HEADERS = {"Accept-Encoding": "gzip, deflate"}
|
||||||
|
|
||||||
|
|
||||||
async def _load_model_catalog() -> Dict[str, List[str]]:
|
async def _load_model_catalog() -> Dict[str, List[str]]:
|
||||||
"""Fetch and parse the model catalog.
|
"""Fetch and parse the model catalog.
|
||||||
@@ -46,61 +65,85 @@ async def _load_model_catalog() -> Dict[str, List[str]]:
|
|||||||
value has a ``models`` sub-dict keyed by model ID. The result is cached
|
value has a ``models`` sub-dict keyed by model ID. The result is cached
|
||||||
in memory after the first successful fetch.
|
in memory after the first successful fetch.
|
||||||
Subsequent calls return the cached data immediately.
|
Subsequent calls return the cached data immediately.
|
||||||
|
|
||||||
|
Failed fetches are negatively cached: further calls return an empty
|
||||||
|
dict without hitting the network until ``_CATALOG_FAILURE_COOLDOWN``
|
||||||
|
has elapsed, so a broken network does not stall every settings-modal
|
||||||
|
open. Concurrent callers are serialized behind :data:`_catalog_lock`
|
||||||
|
so only one request is ever in flight.
|
||||||
"""
|
"""
|
||||||
global _catalog_cache, _model_output_limits
|
global _catalog_cache, _model_output_limits, _catalog_last_failure
|
||||||
if _catalog_cache is not None:
|
if _catalog_cache is not None:
|
||||||
return _catalog_cache
|
return _catalog_cache
|
||||||
|
|
||||||
try:
|
async with _catalog_lock:
|
||||||
async with aiohttp.ClientSession(timeout=_CATALOG_TIMEOUT) as session:
|
# Re-check under the lock: another caller may have fetched (or
|
||||||
async with session.get(_MODEL_CATALOG_URL) as resp:
|
# failed) while we were waiting.
|
||||||
if resp.status != 200:
|
if _catalog_cache is not None:
|
||||||
logger.warning("Model catalog returned HTTP %s", resp.status)
|
return _catalog_cache
|
||||||
return _catalog_cache or {}
|
if (
|
||||||
data = await resp.json()
|
_catalog_last_failure is not None
|
||||||
except (aiohttp.ClientError, asyncio.TimeoutError, json.JSONDecodeError) as exc:
|
and time.monotonic() - _catalog_last_failure < _CATALOG_FAILURE_COOLDOWN
|
||||||
logger.warning("Failed to fetch model catalog: %s", exc)
|
):
|
||||||
return _catalog_cache or {}
|
logger.debug(
|
||||||
|
"Skipping model catalog fetch: last attempt failed %.0fs ago",
|
||||||
|
time.monotonic() - _catalog_last_failure,
|
||||||
|
)
|
||||||
|
return {}
|
||||||
|
|
||||||
if not isinstance(data, dict):
|
try:
|
||||||
logger.warning("Model catalog is not a dict, got %s", type(data).__name__)
|
async with aiohttp.ClientSession(timeout=_CATALOG_TIMEOUT) as session:
|
||||||
return _catalog_cache or {}
|
async with session.get(_MODEL_CATALOG_URL, headers=_NO_BROTLI_HEADERS) as resp:
|
||||||
|
if resp.status != 200:
|
||||||
|
logger.warning("Model catalog returned HTTP %s", resp.status)
|
||||||
|
_catalog_last_failure = time.monotonic()
|
||||||
|
return {}
|
||||||
|
data = await resp.json()
|
||||||
|
except (aiohttp.ClientError, asyncio.TimeoutError, json.JSONDecodeError, UnicodeDecodeError) as exc:
|
||||||
|
logger.warning("Failed to fetch model catalog: %s", exc)
|
||||||
|
_catalog_last_failure = time.monotonic()
|
||||||
|
return {}
|
||||||
|
|
||||||
result: Dict[str, List[str]] = {}
|
if not isinstance(data, dict):
|
||||||
output_limits: Dict[str, Dict[str, int]] = {}
|
logger.warning("Model catalog is not a dict, got %s", type(data).__name__)
|
||||||
for provider_id, provider_info in data.items():
|
_catalog_last_failure = time.monotonic()
|
||||||
if not isinstance(provider_info, dict):
|
return {}
|
||||||
continue
|
|
||||||
models_dict = provider_info.get("models")
|
result: Dict[str, List[str]] = {}
|
||||||
if not isinstance(models_dict, dict):
|
output_limits: Dict[str, Dict[str, int]] = {}
|
||||||
continue
|
for provider_id, provider_info in data.items():
|
||||||
model_ids: List[str] = []
|
if not isinstance(provider_info, dict):
|
||||||
provider_limits: Dict[str, int] = {}
|
|
||||||
for mid, model_info in models_dict.items():
|
|
||||||
if not isinstance(mid, str):
|
|
||||||
continue
|
continue
|
||||||
model_ids.append(mid)
|
models_dict = provider_info.get("models")
|
||||||
if isinstance(model_info, dict):
|
if not isinstance(models_dict, dict):
|
||||||
limit = model_info.get("limit")
|
continue
|
||||||
if isinstance(limit, dict):
|
model_ids: List[str] = []
|
||||||
output = limit.get("output")
|
provider_limits: Dict[str, int] = {}
|
||||||
if isinstance(output, (int, float)) and output > 0:
|
for mid, model_info in models_dict.items():
|
||||||
provider_limits[mid] = int(output)
|
if not isinstance(mid, str):
|
||||||
if model_ids:
|
continue
|
||||||
result[provider_id] = model_ids
|
model_ids.append(mid)
|
||||||
if provider_limits:
|
if isinstance(model_info, dict):
|
||||||
output_limits[provider_id] = provider_limits
|
limit = model_info.get("limit")
|
||||||
|
if isinstance(limit, dict):
|
||||||
|
output = limit.get("output")
|
||||||
|
if isinstance(output, (int, float)) and output > 0:
|
||||||
|
provider_limits[mid] = int(output)
|
||||||
|
if model_ids:
|
||||||
|
result[provider_id] = model_ids
|
||||||
|
if provider_limits:
|
||||||
|
output_limits[provider_id] = provider_limits
|
||||||
|
|
||||||
_catalog_cache = result
|
_catalog_cache = result
|
||||||
_model_output_limits = output_limits
|
_model_output_limits = output_limits
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Loaded model catalog: %d providers, %d total models "
|
"Loaded model catalog: %d providers, %d total models "
|
||||||
"(%d providers have output limits)",
|
"(%d providers have output limits)",
|
||||||
len(result),
|
len(result),
|
||||||
sum(len(m) for m in result.values()),
|
sum(len(m) for m in result.values()),
|
||||||
len(output_limits),
|
len(output_limits),
|
||||||
)
|
)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _get_model_max_output(provider: str, model: str) -> Optional[int]:
|
def _get_model_max_output(provider: str, model: str) -> Optional[int]:
|
||||||
@@ -126,12 +169,12 @@ async def fetch_ollama_models(api_base: str) -> List[str]:
|
|||||||
url = f"{api_base.rstrip('/')}/models"
|
url = f"{api_base.rstrip('/')}/models"
|
||||||
try:
|
try:
|
||||||
async with aiohttp.ClientSession(timeout=_OLLAMA_API_TIMEOUT) as session:
|
async with aiohttp.ClientSession(timeout=_OLLAMA_API_TIMEOUT) as session:
|
||||||
async with session.get(url) as resp:
|
async with session.get(url, headers=_NO_BROTLI_HEADERS) as resp:
|
||||||
if resp.status != 200:
|
if resp.status != 200:
|
||||||
logger.debug("Ollama API returned HTTP %s from %s", resp.status, api_base)
|
logger.debug("Ollama API returned HTTP %s from %s", resp.status, api_base)
|
||||||
return []
|
return []
|
||||||
data = await resp.json()
|
data = await resp.json()
|
||||||
except (aiohttp.ClientError, asyncio.TimeoutError, json.JSONDecodeError) as exc:
|
except (aiohttp.ClientError, asyncio.TimeoutError, json.JSONDecodeError, UnicodeDecodeError) as exc:
|
||||||
logger.debug("Ollama not reachable at %s: %s", api_base, exc)
|
logger.debug("Ollama not reachable at %s: %s", api_base, exc)
|
||||||
return []
|
return []
|
||||||
|
|
||||||
@@ -224,6 +267,16 @@ _PROVIDER_DEFAULTS: Dict[str, str] = {
|
|||||||
# Request timeout for LLM calls (seconds)
|
# Request timeout for LLM calls (seconds)
|
||||||
_LLM_TIMEOUT = aiohttp.ClientTimeout(total=120)
|
_LLM_TIMEOUT = aiohttp.ClientTimeout(total=120)
|
||||||
|
|
||||||
|
# Providers that do NOT implement ``response_format: {"type": "json_schema"}``
|
||||||
|
# and reject it with HTTP 400. For these the weaker, widely supported
|
||||||
|
# ``json_object`` mode is used instead (the prompt already specifies the
|
||||||
|
# expected JSON shape, and ``_try_salvage_json`` repairs imperfect output).
|
||||||
|
# DeepSeek answers a json_schema request with
|
||||||
|
# ``{"error":{"message":"This response_format type is unavailable now"}}``.
|
||||||
|
# LM Studio and some other local OpenAI-compatible servers reject
|
||||||
|
# ``json_object`` but accept ``json_schema``, so they are not listed here.
|
||||||
|
_JSON_OBJECT_ONLY_PROVIDERS = frozenset({"deepseek"})
|
||||||
|
|
||||||
|
|
||||||
class LLMService:
|
class LLMService:
|
||||||
"""Centralized LLM API client.
|
"""Centralized LLM API client.
|
||||||
@@ -571,47 +624,61 @@ class LLMService:
|
|||||||
if effective_max is None:
|
if effective_max is None:
|
||||||
effective_max = 4096
|
effective_max = 4096
|
||||||
|
|
||||||
# Use json_schema (not json_object) for broader provider compatibility:
|
# Structured-output format. ``json_schema`` is preferred because LM
|
||||||
# LM Studio and some other OpenAI-compatible servers reject
|
# Studio and other local OpenAI-compatible servers reject
|
||||||
# json_object but accept json_schema. {"type": "object"} is
|
# ``json_object`` but accept ``json_schema``; ``{"type": "object"}``
|
||||||
# functionally equivalent — it accepts any JSON object without
|
# accepts any JSON object without constraining specific fields, so the
|
||||||
# constraining specific fields.
|
# two modes are functionally equivalent here. Providers known to
|
||||||
response_format = {
|
# reject json_schema (see _JSON_OBJECT_ONLY_PROVIDERS) get
|
||||||
|
# ``json_object`` instead.
|
||||||
|
schema_format: Dict[str, Any] = {
|
||||||
"type": "json_schema",
|
"type": "json_schema",
|
||||||
"json_schema": {
|
"json_schema": {
|
||||||
"name": "metadata",
|
"name": "metadata",
|
||||||
"schema": {"type": "object"},
|
"schema": {"type": "object"},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
json_object_format: Dict[str, Any] = {"type": "json_object"}
|
||||||
|
|
||||||
try:
|
if self._get_config()["provider"] in _JSON_OBJECT_ONLY_PROVIDERS:
|
||||||
result = await self.chat_completion(
|
format_chain: List[Optional[Dict[str, Any]]] = [
|
||||||
messages=messages,
|
json_object_format,
|
||||||
model=model,
|
None,
|
||||||
temperature=temperature,
|
]
|
||||||
response_format=response_format,
|
else:
|
||||||
max_tokens=effective_max,
|
format_chain = [schema_format, json_object_format, None]
|
||||||
)
|
|
||||||
except LLMResponseError as e:
|
result: Optional[Dict[str, Any]] = None
|
||||||
# Only fall back when the provider rejects the response_format
|
for index, fmt in enumerate(format_chain):
|
||||||
# type value (e.g. "'response_format.type' must be..."). Avoid
|
try:
|
||||||
# catching unrelated 400 errors whose body happens to mention
|
result = await self.chat_completion(
|
||||||
# "response_format" (e.g. "model does not support
|
messages=messages,
|
||||||
# response_format restrictions on this endpoint").
|
model=model,
|
||||||
if "'response_format.type'" not in str(e).lower():
|
temperature=temperature,
|
||||||
raise
|
response_format=fmt,
|
||||||
logger.info(
|
max_tokens=effective_max,
|
||||||
"Provider rejected response_format, retrying without it. "
|
)
|
||||||
"Falling back to prompt-only JSON mode. Error: %s",
|
break
|
||||||
e,
|
except LLMResponseError as e:
|
||||||
)
|
message = str(e).lower()
|
||||||
result = await self.chat_completion(
|
if index + 1 >= len(format_chain):
|
||||||
messages=messages,
|
raise
|
||||||
model=model,
|
# Only downgrade when the failure is about ``response_format``.
|
||||||
temperature=temperature,
|
# Everything else (auth, unknown model, rate limits) must
|
||||||
response_format=None,
|
# surface unchanged. Matching on the bare parameter name also
|
||||||
max_tokens=effective_max,
|
# covers variants such as DeepSeek's "This response_format
|
||||||
)
|
# type is unavailable now" without swallowing unrelated 400s.
|
||||||
|
if "response_format" not in message:
|
||||||
|
raise
|
||||||
|
logger.info(
|
||||||
|
"Provider rejected response_format=%s, retrying with %s. "
|
||||||
|
"Error: %s",
|
||||||
|
(fmt or {}).get("type", "none"),
|
||||||
|
(format_chain[index + 1] or {}).get("type", "none"),
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result is not None # non-empty chain always sets or raises
|
||||||
|
|
||||||
content = result.get("content", "") or ""
|
content = result.get("content", "") or ""
|
||||||
if not content:
|
if not content:
|
||||||
|
|||||||
@@ -79,6 +79,8 @@ class LoraService(BaseModelService):
|
|||||||
),
|
),
|
||||||
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
||||||
"version_count": model_data.get("version_count"),
|
"version_count": model_data.get("version_count"),
|
||||||
|
"source_platform": model_data.get("source_platform", ""),
|
||||||
|
"source_url": model_data.get("source_url", ""),
|
||||||
"hf_url": model_data.get("hf_url", ""),
|
"hf_url": model_data.get("hf_url", ""),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -712,12 +714,18 @@ class LoraService(BaseModelService):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Return minimal data needed for cycling
|
# Return minimal data needed for cycling. usage_tips is only included
|
||||||
return [
|
# when non-empty so widget consumers (recommended strength range cues)
|
||||||
{
|
# can build their lookup without inflating the payload.
|
||||||
|
result = []
|
||||||
|
for lora in available_loras:
|
||||||
|
entry = {
|
||||||
"file_name": f"{lora['folder']}/{lora['file_name']}" if lora.get("folder") else lora["file_name"],
|
"file_name": f"{lora['folder']}/{lora['file_name']}" if lora.get("folder") else lora["file_name"],
|
||||||
"model_name": lora.get("model_name", lora["file_name"]),
|
"model_name": lora.get("model_name", lora["file_name"]),
|
||||||
"folder": lora.get("folder", ""),
|
"folder": lora.get("folder", ""),
|
||||||
}
|
}
|
||||||
for lora in available_loras
|
usage_tips = lora.get("usage_tips")
|
||||||
]
|
if usage_tips:
|
||||||
|
entry["usage_tips"] = usage_tips
|
||||||
|
result.append(entry)
|
||||||
|
return result
|
||||||
|
|||||||
@@ -14,10 +14,33 @@ from ..utils.model_utils import determine_base_model
|
|||||||
from ..utils.models import autov3_from_civitai_files
|
from ..utils.models import autov3_from_civitai_files
|
||||||
from .connectivity_guard import OFFLINE_FRIENDLY_MESSAGE, is_expected_offline_error
|
from .connectivity_guard import OFFLINE_FRIENDLY_MESSAGE, is_expected_offline_error
|
||||||
from .errors import RateLimitError
|
from .errors import RateLimitError
|
||||||
|
from .model_sources import has_external_source
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _merge_ordered_unique(existing: Iterable[str], new: Iterable[str]) -> list[str]:
|
||||||
|
"""Concatenate two word lists, dropping duplicates without reordering.
|
||||||
|
|
||||||
|
Trigger word order is meaningful: the sequence stored in
|
||||||
|
``civitai.trainedWords`` is the order used when building prompts, and users
|
||||||
|
can reorder it in the UI. A plain ``set`` union used to shuffle that order on
|
||||||
|
every metadata refresh, so existing words are kept first (in their saved
|
||||||
|
order) and newly discovered ones are appended.
|
||||||
|
"""
|
||||||
|
|
||||||
|
merged: list[str] = []
|
||||||
|
seen: set[str] = set()
|
||||||
|
|
||||||
|
for word in list(existing) + list(new):
|
||||||
|
if word in seen:
|
||||||
|
continue
|
||||||
|
seen.add(word)
|
||||||
|
merged.append(word)
|
||||||
|
|
||||||
|
return merged
|
||||||
|
|
||||||
|
|
||||||
class MetadataProviderProtocol(Protocol):
|
class MetadataProviderProtocol(Protocol):
|
||||||
"""Subset of metadata provider interface consumed by the sync service."""
|
"""Subset of metadata provider interface consumed by the sync service."""
|
||||||
|
|
||||||
@@ -114,9 +137,10 @@ class MetadataSyncService:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if "trainedWords" in existing_civitai:
|
if "trainedWords" in existing_civitai:
|
||||||
existing_trained = existing_civitai.get("trainedWords", [])
|
existing_trained = existing_civitai.get("trainedWords", []) or []
|
||||||
new_trained = civitai_metadata.get("trainedWords", [])
|
new_trained = civitai_metadata.get("trainedWords", []) or []
|
||||||
merged_trained = list(set(existing_trained + new_trained))
|
# Order preserving merge: the saved order drives prompt order.
|
||||||
|
merged_trained = _merge_ordered_unique(existing_trained, new_trained)
|
||||||
merged_civitai["trainedWords"] = merged_trained
|
merged_civitai["trainedWords"] = merged_trained
|
||||||
|
|
||||||
local_metadata["civitai"] = merged_civitai
|
local_metadata["civitai"] = merged_civitai
|
||||||
@@ -222,9 +246,10 @@ class MetadataSyncService:
|
|||||||
error_msg = "CivitAI model is deleted and no archive provider is available"
|
error_msg = "CivitAI model is deleted and no archive provider is available"
|
||||||
return False, error_msg
|
return False, error_msg
|
||||||
else:
|
else:
|
||||||
is_hf_source = bool(model_data.get("hf_url"))
|
is_hf_source = has_external_source(model_data)
|
||||||
if is_hf_source:
|
if is_hf_source:
|
||||||
# HF-sourced model: only check CivitAI API directly.
|
# External-source model (Hugging Face / ModelScope /
|
||||||
|
# TensorArt): only check CivitAI API directly.
|
||||||
# CivArchive is almost guaranteed to have no record, and
|
# CivArchive is almost guaranteed to have no record, and
|
||||||
# hitting it wastes rate-limit budget.
|
# hitting it wastes rate-limit budget.
|
||||||
# Use a distinct provider name ("civitai_api" not None) so
|
# Use a distinct provider name ("civitai_api" not None) so
|
||||||
@@ -245,16 +270,23 @@ class MetadataSyncService:
|
|||||||
civitai_api_not_found = False
|
civitai_api_not_found = False
|
||||||
any_rate_limited = False
|
any_rate_limited = False
|
||||||
|
|
||||||
|
skip_network_providers = False
|
||||||
for provider_name, provider in provider_attempts:
|
for provider_name, provider in provider_attempts:
|
||||||
|
if skip_network_providers and provider_name != "sqlite":
|
||||||
|
# A network provider was already rate-limited; failing
|
||||||
|
# over to another network provider just spreads the flood
|
||||||
|
# (#1085). The local sqlite archive stays as last resort.
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
civitai_metadata_candidate, error = await provider.get_model_by_hash(sha256)
|
civitai_metadata_candidate, error = await provider.get_model_by_hash(sha256)
|
||||||
except RateLimitError as exc:
|
except RateLimitError as exc:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Provider %s is rate-limited (retry_after=%.0fs); skipping to next provider",
|
"Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers",
|
||||||
provider_name or provider.__class__.__name__,
|
provider_name or provider.__class__.__name__,
|
||||||
exc.retry_after or 0,
|
exc.retry_after or 0,
|
||||||
)
|
)
|
||||||
any_rate_limited = True
|
any_rate_limited = True
|
||||||
|
skip_network_providers = True
|
||||||
continue
|
continue
|
||||||
except Exception as exc: # pragma: no cover - defensive logging
|
except Exception as exc: # pragma: no cover - defensive logging
|
||||||
logger.error("Provider %s failed for hash %s: %s", provider_name, sha256, exc)
|
logger.error("Provider %s failed for hash %s: %s", provider_name, sha256, exc)
|
||||||
@@ -419,14 +451,37 @@ class MetadataSyncService:
|
|||||||
metadata: Dict[str, Any],
|
metadata: Dict[str, Any],
|
||||||
model_id: int,
|
model_id: int,
|
||||||
model_version_id: Optional[int],
|
model_version_id: Optional[int],
|
||||||
|
provider_name: Optional[str] = None,
|
||||||
) -> Dict[str, Any]:
|
) -> Dict[str, Any]:
|
||||||
"""Relink a local metadata record to a specific CivitAI model version."""
|
"""Relink a local metadata record to a specific CivitAI model version.
|
||||||
|
|
||||||
|
When ``provider_name`` is given, the named provider is resolved via the
|
||||||
|
metadata provider selector instead of the default fallback chain. A
|
||||||
|
missing/disabled provider surfaces a user-friendly error instead of the
|
||||||
|
raw selector exception.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if provider_name:
|
||||||
|
try:
|
||||||
|
provider = await self._get_provider(provider_name)
|
||||||
|
except ValueError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Unable to resolve metadata provider %s: %s", provider_name, exc
|
||||||
|
)
|
||||||
|
raise ValueError(
|
||||||
|
"CivitArchive is not available or not enabled. "
|
||||||
|
"Enable the CivitArchive API in settings to relink via CivArchive."
|
||||||
|
) from exc
|
||||||
|
else:
|
||||||
|
provider = await self._get_default_provider()
|
||||||
|
|
||||||
provider = await self._get_default_provider()
|
|
||||||
civitai_metadata = await provider.get_model_version(model_id, model_version_id)
|
civitai_metadata = await provider.get_model_version(model_id, model_version_id)
|
||||||
if not civitai_metadata:
|
if not civitai_metadata:
|
||||||
|
provider_label = (
|
||||||
|
"CivitArchive" if provider_name == "civarchive_api" else "CivitAI"
|
||||||
|
)
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Model version not found on CivitAI for ID: {model_id}"
|
f"Model version not found on {provider_label} for ID: {model_id}"
|
||||||
+ (f" with version: {model_version_id}" if model_version_id else "")
|
+ (f" with version: {model_version_id}" if model_version_id else "")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -33,6 +33,11 @@ class ModelCache:
|
|||||||
|
|
||||||
raw_data: List[Dict[str, Any]]
|
raw_data: List[Dict[str, Any]]
|
||||||
folders: List[str]
|
folders: List[str]
|
||||||
|
# Every directory under the model roots (including empty ones), as
|
||||||
|
# recorded by the last scan/hydration. ``None`` means "never recorded"
|
||||||
|
# (e.g. a persisted snapshot predating this field) and triggers a
|
||||||
|
# background filesystem backfill in the scanner.
|
||||||
|
all_folders: Optional[List[str]] = None
|
||||||
version_index: Dict[int, Dict[str, Any]] = field(default_factory=dict)
|
version_index: Dict[int, Dict[str, Any]] = field(default_factory=dict)
|
||||||
model_id_index: Dict[int, List[Dict[str, Any]]] = field(default_factory=dict)
|
model_id_index: Dict[int, List[Dict[str, Any]]] = field(default_factory=dict)
|
||||||
# Multi-valued companion to version_index: every local file entry of a
|
# Multi-valued companion to version_index: every local file entry of a
|
||||||
|
|||||||
@@ -2,13 +2,15 @@ import asyncio
|
|||||||
import fnmatch
|
import fnmatch
|
||||||
import os
|
import os
|
||||||
import logging
|
import logging
|
||||||
|
import shutil
|
||||||
from typing import Any, Dict, List, Optional, Sequence, Set
|
from typing import Any, Dict, List, Optional, Sequence, Set
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
|
||||||
from ..utils.utils import calculate_relative_path_for_model, remove_empty_dirs
|
from ..utils.utils import calculate_relative_path_for_model, remove_empty_dirs
|
||||||
from ..utils.constants import AUTO_ORGANIZE_BATCH_SIZE
|
from ..utils.constants import AUTO_ORGANIZE_BATCH_SIZE, MODEL_FILE_EXTENSIONS
|
||||||
from ..services.settings_manager import get_settings_manager
|
from ..services.settings_manager import get_settings_manager
|
||||||
from ..services.model_lifecycle_service import _require_path_in_library_roots
|
from ..services.model_lifecycle_service import _require_path_in_library_roots
|
||||||
|
from ..services.pending_delete_service import PENDING_DELETE_DIR_NAME
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -41,10 +43,22 @@ class AutoOrganizeResult:
|
|||||||
|
|
||||||
def to_dict(self) -> Dict[str, Any]:
|
def to_dict(self) -> Dict[str, Any]:
|
||||||
"""Convert result to dictionary"""
|
"""Convert result to dictionary"""
|
||||||
|
if self.operation_type == 'filename_template':
|
||||||
|
message = (
|
||||||
|
f'Filename template applied: {self.success_count} renamed, '
|
||||||
|
f'{self.skipped_count} skipped, {self.failure_count} failed '
|
||||||
|
f'out of {self.total} total'
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
message = (
|
||||||
|
f'Auto-organize {self.operation_type} completed: '
|
||||||
|
f'{self.success_count} moved, {self.skipped_count} skipped, '
|
||||||
|
f'{self.failure_count} failed out of {self.total} total'
|
||||||
|
)
|
||||||
result: Dict[str, Any] = {
|
result: Dict[str, Any] = {
|
||||||
'success': self.status != 'error',
|
'success': self.status != 'error',
|
||||||
'status': self.status,
|
'status': self.status,
|
||||||
'message': f'Auto-organize {self.operation_type} completed: {self.success_count} moved, {self.skipped_count} skipped, {self.failure_count} failed out of {self.total} total',
|
'message': message,
|
||||||
'summary': {
|
'summary': {
|
||||||
'total': self.total,
|
'total': self.total,
|
||||||
'success': self.success_count,
|
'success': self.success_count,
|
||||||
@@ -473,17 +487,368 @@ class ModelFileService:
|
|||||||
|
|
||||||
class ModelMoveService:
|
class ModelMoveService:
|
||||||
"""Service for handling individual model moves"""
|
"""Service for handling individual model moves"""
|
||||||
|
|
||||||
def __init__(self, scanner, model_type: str):
|
def __init__(self, scanner, model_type: str):
|
||||||
"""Initialize the service
|
"""Initialize the service
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
scanner: Model scanner instance
|
scanner: Model scanner instance
|
||||||
model_type: Type of model (e.g., 'lora', 'checkpoint')
|
model_type: Type of model (e.g., 'lora', 'checkpoint')
|
||||||
"""
|
"""
|
||||||
self.scanner = scanner
|
self.scanner = scanner
|
||||||
self.model_type = model_type
|
self.model_type = model_type
|
||||||
|
|
||||||
|
async def create_folder(self, folder_path: str) -> Dict[str, Any]:
|
||||||
|
"""Create a directory inside the model library roots.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
folder_path: Absolute path of the directory to create (business
|
||||||
|
path — symlinks are not resolved)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with success flag, the created path and the
|
||||||
|
library-relative folder name used by folder trees.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
if not folder_path or not str(folder_path).strip():
|
||||||
|
return {"success": False, "error": "Folder path is required"}
|
||||||
|
|
||||||
|
_require_path_in_library_roots(folder_path, self.scanner, label="Folder path")
|
||||||
|
|
||||||
|
absolute_path = os.path.abspath(folder_path)
|
||||||
|
already_exists = os.path.isdir(absolute_path)
|
||||||
|
os.makedirs(absolute_path, exist_ok=True)
|
||||||
|
|
||||||
|
relative_folder = self._calculate_relative_folder(absolute_path)
|
||||||
|
if relative_folder:
|
||||||
|
await self.scanner.add_known_folder(relative_folder)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"folder_path": absolute_path.replace(os.sep, "/"),
|
||||||
|
"folder": relative_folder,
|
||||||
|
"created": not already_exists,
|
||||||
|
}
|
||||||
|
except ValueError as exc:
|
||||||
|
return {"success": False, "error": str(exc)}
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"Error creating folder: {exc}", exc_info=True)
|
||||||
|
return {"success": False, "error": str(exc)}
|
||||||
|
|
||||||
|
def _calculate_relative_folder(self, absolute_path: str) -> str:
|
||||||
|
"""Return the library-relative folder for an absolute directory path."""
|
||||||
|
normalized = os.path.abspath(absolute_path)
|
||||||
|
for root in self.scanner.get_model_roots():
|
||||||
|
abs_root = os.path.abspath(root)
|
||||||
|
try:
|
||||||
|
rel = os.path.relpath(normalized, abs_root)
|
||||||
|
except ValueError:
|
||||||
|
continue
|
||||||
|
if rel == ".":
|
||||||
|
return ""
|
||||||
|
if not rel.startswith(".."):
|
||||||
|
return rel.replace(os.sep, "/")
|
||||||
|
return ""
|
||||||
|
|
||||||
|
async def delete_folder(self, folder_path: str, dry_run: bool = False) -> Dict[str, Any]:
|
||||||
|
"""Delete a model-free directory inside the model library roots.
|
||||||
|
|
||||||
|
Only directories whose subtree holds no model weight files can be
|
||||||
|
removed: a folder-level cascade would bypass the per-model lifecycle
|
||||||
|
bookkeeping (metadata sidecars, previews, cache entries, pending-delete
|
||||||
|
staging and recipe references), so it is deliberately refused. Leftover
|
||||||
|
non-model files (stray previews, sidecars, ``.bak`` files) are reported
|
||||||
|
in the manifest before they are removed.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
folder_path: Absolute path of the directory to remove (business
|
||||||
|
path — symlinks are not resolved)
|
||||||
|
dry_run: When true, only report what would be removed
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with the success flag plus a removal manifest
|
||||||
|
(``model_count``/``file_count``/``dir_count``/``symlink_count``/
|
||||||
|
``total_bytes``/``restorable``) on success.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
if not folder_path or not str(folder_path).strip():
|
||||||
|
return {"success": False, "error": "Folder path is required"}
|
||||||
|
|
||||||
|
_require_path_in_library_roots(folder_path, self.scanner, label="Folder path")
|
||||||
|
|
||||||
|
absolute_path = os.path.abspath(folder_path)
|
||||||
|
if os.path.islink(absolute_path):
|
||||||
|
# shutil.rmtree refuses symlinked roots, and silently deleting
|
||||||
|
# the link (leaving the real directory behind) is a separate
|
||||||
|
# decision we do not make here.
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": "Symlinked folders cannot be deleted",
|
||||||
|
}
|
||||||
|
if not os.path.isdir(absolute_path):
|
||||||
|
return {"success": False, "error": "Folder no longer exists"}
|
||||||
|
|
||||||
|
if self._is_model_root(absolute_path):
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": "The library root itself cannot be deleted",
|
||||||
|
}
|
||||||
|
|
||||||
|
manifest = self._collect_folder_manifest(absolute_path)
|
||||||
|
|
||||||
|
if manifest["pending_delete_job"]:
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"code": "busy",
|
||||||
|
"error": (
|
||||||
|
"A staged delete is still pending inside this folder; "
|
||||||
|
"wait for the undo window to expire"
|
||||||
|
),
|
||||||
|
"manifest": manifest,
|
||||||
|
}
|
||||||
|
|
||||||
|
if manifest["model_count"] > 0:
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"code": "not_empty",
|
||||||
|
"error": (
|
||||||
|
f"Folder still contains {manifest['model_count']} model "
|
||||||
|
"file(s); delete or move them first"
|
||||||
|
),
|
||||||
|
"manifest": manifest,
|
||||||
|
}
|
||||||
|
|
||||||
|
relative_folder = self._calculate_relative_folder(absolute_path)
|
||||||
|
|
||||||
|
if dry_run:
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"dry_run": True,
|
||||||
|
"folder_path": absolute_path.replace(os.sep, "/"),
|
||||||
|
"folder": relative_folder,
|
||||||
|
**manifest,
|
||||||
|
}
|
||||||
|
|
||||||
|
shutil.rmtree(absolute_path)
|
||||||
|
|
||||||
|
await self._forget_folder(relative_folder)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"dry_run": False,
|
||||||
|
"folder_path": absolute_path.replace(os.sep, "/"),
|
||||||
|
"folder": relative_folder,
|
||||||
|
**manifest,
|
||||||
|
}
|
||||||
|
except ValueError as exc:
|
||||||
|
return {"success": False, "error": str(exc)}
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"Error deleting folder: {exc}", exc_info=True)
|
||||||
|
return {"success": False, "error": str(exc)}
|
||||||
|
|
||||||
|
def _is_model_root(self, absolute_path: str) -> bool:
|
||||||
|
"""Return True when the path *is* one of the configured library roots."""
|
||||||
|
normalized = os.path.normpath(absolute_path)
|
||||||
|
for root in self.scanner.get_model_roots():
|
||||||
|
if os.path.normpath(os.path.abspath(root)) == normalized:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_model_file(file_name: str) -> bool:
|
||||||
|
"""Return True when the file name carries a model weight extension."""
|
||||||
|
return os.path.splitext(file_name)[1].lower() in MODEL_FILE_EXTENSIONS
|
||||||
|
|
||||||
|
def _collect_folder_manifest(self, absolute_path: str) -> Dict[str, Any]:
|
||||||
|
"""Describe everything a recursive delete of *absolute_path* removes.
|
||||||
|
|
||||||
|
Walking is intentional: the scanner cache can be stale, and a model file
|
||||||
|
that appeared on disk since the last scan must still block the delete.
|
||||||
|
Symbolic links are never followed (``os.walk`` default) and are counted
|
||||||
|
separately — ``shutil.rmtree`` unlinks them without touching their
|
||||||
|
targets.
|
||||||
|
"""
|
||||||
|
model_count = 0
|
||||||
|
file_count = 0
|
||||||
|
dir_count = 0
|
||||||
|
symlink_count = 0
|
||||||
|
total_bytes = 0
|
||||||
|
pending_delete_job = False
|
||||||
|
|
||||||
|
for dirpath, dirnames, filenames in os.walk(absolute_path):
|
||||||
|
if PENDING_DELETE_DIR_NAME in dirnames:
|
||||||
|
pending_delete_job = True
|
||||||
|
|
||||||
|
for name in dirnames:
|
||||||
|
if os.path.islink(os.path.join(dirpath, name)):
|
||||||
|
symlink_count += 1
|
||||||
|
else:
|
||||||
|
dir_count += 1
|
||||||
|
|
||||||
|
for name in filenames:
|
||||||
|
full_path = os.path.join(dirpath, name)
|
||||||
|
if os.path.islink(full_path):
|
||||||
|
symlink_count += 1
|
||||||
|
continue
|
||||||
|
if self._is_model_file(name):
|
||||||
|
model_count += 1
|
||||||
|
else:
|
||||||
|
file_count += 1
|
||||||
|
try:
|
||||||
|
total_bytes += os.path.getsize(full_path)
|
||||||
|
except OSError: # pragma: no cover - defensive
|
||||||
|
pass
|
||||||
|
|
||||||
|
return {
|
||||||
|
"model_count": model_count,
|
||||||
|
"file_count": file_count,
|
||||||
|
"dir_count": dir_count,
|
||||||
|
"symlink_count": symlink_count,
|
||||||
|
"total_bytes": total_bytes,
|
||||||
|
"pending_delete_job": pending_delete_job,
|
||||||
|
# A truly empty directory is the only case an "undo" can restore by
|
||||||
|
# simply recreating it; a folder holding stray files is gone for good.
|
||||||
|
"restorable": (
|
||||||
|
model_count == 0
|
||||||
|
and file_count == 0
|
||||||
|
and dir_count == 0
|
||||||
|
and symlink_count == 0
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
async def _forget_folder(self, relative_folder: str) -> None:
|
||||||
|
"""Drop a removed directory from the scanner's folder/cache records."""
|
||||||
|
if not relative_folder:
|
||||||
|
return
|
||||||
|
remove_known_folder = getattr(self.scanner, "remove_known_folder", None)
|
||||||
|
if callable(remove_known_folder):
|
||||||
|
await remove_known_folder(relative_folder)
|
||||||
|
|
||||||
|
async def rename_folder(self, folder_path: str, new_name: str) -> Dict[str, Any]:
|
||||||
|
"""Rename a directory inside the model library roots.
|
||||||
|
|
||||||
|
Unlike :meth:`delete_folder` this works on folders that hold models.
|
||||||
|
A rename keeps every file, so no per-model lifecycle step is bypassed:
|
||||||
|
the directory is renamed on disk and the affected folder, cache, hash
|
||||||
|
index and metadata-sidecar records are re-keyed onto the new prefix by
|
||||||
|
the scanner.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
folder_path: Absolute path of the directory to rename (business
|
||||||
|
path — symlinks are not resolved)
|
||||||
|
new_name: New leaf name; a single path segment, not a path
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with the success flag, the previous/next library-relative
|
||||||
|
folder names and whether the directory actually moved.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
if not folder_path or not str(folder_path).strip():
|
||||||
|
return {"success": False, "error": "Folder path is required"}
|
||||||
|
|
||||||
|
new_name = str(new_name or "").strip()
|
||||||
|
if not new_name:
|
||||||
|
return {"success": False, "error": "New folder name is required"}
|
||||||
|
if new_name in (".", "..") or any(
|
||||||
|
char in new_name for char in '/\\:*?"<>|'
|
||||||
|
):
|
||||||
|
return {"success": False, "error": "Invalid characters in folder name"}
|
||||||
|
|
||||||
|
_require_path_in_library_roots(folder_path, self.scanner, label="Folder path")
|
||||||
|
|
||||||
|
absolute_path = os.path.abspath(folder_path)
|
||||||
|
if os.path.islink(absolute_path):
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": "Symlinked folders cannot be renamed",
|
||||||
|
}
|
||||||
|
if not os.path.isdir(absolute_path):
|
||||||
|
return {"success": False, "error": "Folder no longer exists"}
|
||||||
|
|
||||||
|
if self._is_model_root(absolute_path):
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"error": "The library root itself cannot be renamed",
|
||||||
|
}
|
||||||
|
|
||||||
|
previous_relative = self._calculate_relative_folder(absolute_path)
|
||||||
|
target = os.path.join(os.path.dirname(absolute_path), new_name)
|
||||||
|
|
||||||
|
if os.path.normpath(target) == os.path.normpath(absolute_path):
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"renamed": False,
|
||||||
|
"folder": previous_relative,
|
||||||
|
"previous_folder": previous_relative,
|
||||||
|
"folder_path": absolute_path.replace(os.sep, "/"),
|
||||||
|
}
|
||||||
|
|
||||||
|
if os.path.exists(target):
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"code": "target_exists",
|
||||||
|
"error": f"A folder named \"{new_name}\" already exists here",
|
||||||
|
}
|
||||||
|
|
||||||
|
# A staging manifest records absolute original/staged paths, so
|
||||||
|
# moving a folder that holds one would break its undo and purge.
|
||||||
|
if self._has_pending_delete_job(absolute_path):
|
||||||
|
return {
|
||||||
|
"success": False,
|
||||||
|
"code": "busy",
|
||||||
|
"error": (
|
||||||
|
"A staged delete is still pending inside this folder; "
|
||||||
|
"wait for the undo window to expire"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
os.rename(absolute_path, target)
|
||||||
|
|
||||||
|
new_relative = self._calculate_relative_folder(target)
|
||||||
|
await self._rename_folder_records(
|
||||||
|
previous_relative, new_relative, absolute_path, target
|
||||||
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"success": True,
|
||||||
|
"renamed": True,
|
||||||
|
"folder": new_relative,
|
||||||
|
"previous_folder": previous_relative,
|
||||||
|
"folder_path": target.replace(os.sep, "/"),
|
||||||
|
}
|
||||||
|
except ValueError as exc:
|
||||||
|
return {"success": False, "error": str(exc)}
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(f"Error renaming folder: {exc}", exc_info=True)
|
||||||
|
return {"success": False, "error": str(exc)}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _has_pending_delete_job(absolute_path: str) -> bool:
|
||||||
|
"""Return True when a staged-delete batch lives inside the subtree."""
|
||||||
|
for _dirpath, dirnames, _filenames in os.walk(absolute_path):
|
||||||
|
if PENDING_DELETE_DIR_NAME in dirnames:
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
async def _rename_folder_records(
|
||||||
|
self,
|
||||||
|
previous_relative: str,
|
||||||
|
new_relative: str,
|
||||||
|
previous_path: str,
|
||||||
|
new_path: str,
|
||||||
|
) -> None:
|
||||||
|
"""Hand the rename to the scanner so folder/cache records follow it."""
|
||||||
|
if not previous_relative or not new_relative:
|
||||||
|
return
|
||||||
|
rename_known_folder = getattr(self.scanner, "rename_known_folder", None)
|
||||||
|
if callable(rename_known_folder):
|
||||||
|
await rename_known_folder(
|
||||||
|
previous_relative,
|
||||||
|
new_relative,
|
||||||
|
previous_path=previous_path,
|
||||||
|
new_path=new_path,
|
||||||
|
)
|
||||||
|
|
||||||
async def move_model(self, file_path: str, target_path: str, use_default_paths: bool = False) -> Dict[str, Any]:
|
async def move_model(self, file_path: str, target_path: str, use_default_paths: bool = False) -> Dict[str, Any]:
|
||||||
"""Move a single model file
|
"""Move a single model file
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
from typing import Dict, Optional, Set, List
|
from typing import Dict, Optional, Set, List
|
||||||
import os
|
import os
|
||||||
|
|
||||||
|
from ..utils.constants import is_empty_placeholder_hash
|
||||||
|
|
||||||
class ModelHashIndex:
|
class ModelHashIndex:
|
||||||
"""Index for looking up models by hash or filename"""
|
"""Index for looking up models by hash or filename"""
|
||||||
|
|
||||||
@@ -81,6 +83,8 @@ class ModelHashIndex:
|
|||||||
# mapping. First-time registrations stay O(1).
|
# mapping. First-time registrations stay O(1).
|
||||||
if autov3:
|
if autov3:
|
||||||
autov3 = autov3.lower()
|
autov3 = autov3.lower()
|
||||||
|
if is_empty_placeholder_hash(autov3):
|
||||||
|
autov3 = None
|
||||||
if is_re_registration and (existing_hash != sha256 or autov3):
|
if is_re_registration and (existing_hash != sha256 or autov3):
|
||||||
stale_autov3_keys = [
|
stale_autov3_keys = [
|
||||||
key for key, mapped_path in self._autov3_to_path.items()
|
key for key, mapped_path in self._autov3_to_path.items()
|
||||||
@@ -93,7 +97,7 @@ class ModelHashIndex:
|
|||||||
|
|
||||||
def add_autov3(self, autov3: str, file_path: str) -> None:
|
def add_autov3(self, autov3: str, file_path: str) -> None:
|
||||||
"""Add or update an AutoV3-only index entry (used when only AutoV3 is known)"""
|
"""Add or update an AutoV3-only index entry (used when only AutoV3 is known)"""
|
||||||
if not autov3:
|
if not autov3 or is_empty_placeholder_hash(autov3):
|
||||||
return
|
return
|
||||||
autov3 = autov3.lower()
|
autov3 = autov3.lower()
|
||||||
self._autov3_to_path[autov3] = file_path
|
self._autov3_to_path[autov3] = file_path
|
||||||
@@ -250,6 +254,8 @@ class ModelHashIndex:
|
|||||||
|
|
||||||
def has_hash(self, hash_value: str) -> bool:
|
def has_hash(self, hash_value: str) -> bool:
|
||||||
"""Check if hash exists in index (SHA256, AutoV2, or AutoV3)"""
|
"""Check if hash exists in index (SHA256, AutoV2, or AutoV3)"""
|
||||||
|
if is_empty_placeholder_hash(hash_value):
|
||||||
|
return False
|
||||||
normalized = hash_value.lower()
|
normalized = hash_value.lower()
|
||||||
if normalized in self._hash_to_path:
|
if normalized in self._hash_to_path:
|
||||||
return True
|
return True
|
||||||
@@ -261,6 +267,8 @@ class ModelHashIndex:
|
|||||||
|
|
||||||
def get_path(self, hash_value: str) -> Optional[str]:
|
def get_path(self, hash_value: str) -> Optional[str]:
|
||||||
"""Get file path for a hash (SHA256, AutoV2, or AutoV3)"""
|
"""Get file path for a hash (SHA256, AutoV2, or AutoV3)"""
|
||||||
|
if is_empty_placeholder_hash(hash_value):
|
||||||
|
return None
|
||||||
normalized = hash_value.lower()
|
normalized = hash_value.lower()
|
||||||
path = self._hash_to_path.get(normalized)
|
path = self._hash_to_path.get(normalized)
|
||||||
if path is not None:
|
if path is not None:
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from typing import Any, Awaitable, Callable, Dict, Iterable, List, Mapping, Optional, TYPE_CHECKING, cast
|
from typing import Any, Awaitable, Callable, Dict, Iterable, List, Mapping, Optional, TYPE_CHECKING, cast
|
||||||
@@ -17,6 +18,26 @@ if TYPE_CHECKING:
|
|||||||
from ..services.model_update_service import ModelUpdateService
|
from ..services.model_update_service import ModelUpdateService
|
||||||
|
|
||||||
|
|
||||||
|
async def load_local_metadata(metadata_path: str) -> Dict[str, Any]:
|
||||||
|
"""Load a metadata sidecar JSON, returning an empty dict when missing.
|
||||||
|
|
||||||
|
Thin equivalent of ``MetadataSyncService.load_local_metadata`` for callers
|
||||||
|
(download manager, use cases) that do not hold a sync-service instance.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not os.path.exists(metadata_path):
|
||||||
|
return {}
|
||||||
|
|
||||||
|
try:
|
||||||
|
with open(metadata_path, "r", encoding="utf-8") as handle:
|
||||||
|
payload = json.load(handle)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Failed to load metadata from %s: %s", metadata_path, exc)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
return payload if isinstance(payload, dict) else {}
|
||||||
|
|
||||||
|
|
||||||
async def delete_model_artifacts(
|
async def delete_model_artifacts(
|
||||||
target_dir: str, file_name: str, main_extension: str | None = None
|
target_dir: str, file_name: str, main_extension: str | None = None
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
@@ -404,6 +425,9 @@ class ModelLifecycleService:
|
|||||||
if metadata and new_metadata_path:
|
if metadata and new_metadata_path:
|
||||||
metadata["file_name"] = new_file_name
|
metadata["file_name"] = new_file_name
|
||||||
metadata["file_path"] = new_file_path
|
metadata["file_path"] = new_file_path
|
||||||
|
# Preserve the pre-rename stem so the original download filename
|
||||||
|
# stays recoverable after template-driven renames.
|
||||||
|
metadata.setdefault("original_file_name", old_file_name)
|
||||||
|
|
||||||
if metadata.get("preview_url"):
|
if metadata.get("preview_url"):
|
||||||
old_preview = str(metadata["preview_url"])
|
old_preview = str(metadata["preview_url"])
|
||||||
|
|||||||
@@ -66,6 +66,14 @@ class _RateLimitRetryHelper:
|
|||||||
except RateLimitError as exc:
|
except RateLimitError as exc:
|
||||||
attempt += 1
|
attempt += 1
|
||||||
|
|
||||||
|
# The downloader's rate-limit gate already applied the wait
|
||||||
|
# policy for this request (waited out the vendor window or
|
||||||
|
# deliberately refused because it exceeds the cap). Sleeping
|
||||||
|
# again here would double the wait — just propagate.
|
||||||
|
if getattr(exc, "gate_handled", False):
|
||||||
|
exc.provider = exc.provider or label
|
||||||
|
raise
|
||||||
|
|
||||||
# Determine effective retry limit based on rate-limit magnitude
|
# Determine effective retry limit based on rate-limit magnitude
|
||||||
effective_retry_limit = self._retry_limit # default: 3
|
effective_retry_limit = self._retry_limit # default: 3
|
||||||
if exc.retry_after is not None and exc.retry_after >= 120.0:
|
if exc.retry_after is not None and exc.retry_after >= 120.0:
|
||||||
@@ -101,6 +109,12 @@ class _RateLimitRetryHelper:
|
|||||||
|
|
||||||
return min(self._max_delay, max(0.0, base_delay))
|
return min(self._max_delay, max(0.0, base_delay))
|
||||||
|
|
||||||
|
|
||||||
|
# Labels of providers that are free to consult even while a network provider
|
||||||
|
# is rate-limited (local lookups, no vendor cost).
|
||||||
|
_LOCAL_PROVIDER_LABELS = frozenset({"sqlite"})
|
||||||
|
|
||||||
|
|
||||||
class ModelMetadataProvider(ABC):
|
class ModelMetadataProvider(ABC):
|
||||||
"""Base abstract class for all model metadata providers"""
|
"""Base abstract class for all model metadata providers"""
|
||||||
|
|
||||||
@@ -155,6 +169,17 @@ class ModelMetadataProvider(ABC):
|
|||||||
"""Published model count for the user; None when unsupported."""
|
"""Published model count for the user; None when unsupported."""
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
async def get_version_file_mini(
|
||||||
|
self, version_id: int, file_id: int
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Fetch raw stored file info via CivitAI's model-versions/mini endpoint.
|
||||||
|
|
||||||
|
Only the CivitAI provider implements this (#1100); other providers
|
||||||
|
already serve raw file names (CivArchive) or cannot resolve this
|
||||||
|
lookup (SQLite), so the default is None.
|
||||||
|
"""
|
||||||
|
return None
|
||||||
|
|
||||||
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
||||||
"""Provider that uses Civitai API for metadata"""
|
"""Provider that uses Civitai API for metadata"""
|
||||||
|
|
||||||
@@ -189,6 +214,11 @@ class CivitaiModelMetadataProvider(ModelMetadataProvider):
|
|||||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
return await self.client.get_creator_model_count(username)
|
return await self.client.get_creator_model_count(username)
|
||||||
|
|
||||||
|
async def get_version_file_mini(
|
||||||
|
self, version_id: int, file_id: int
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
return await self.client.get_version_file_mini(version_id, file_id)
|
||||||
|
|
||||||
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
class CivArchiveModelMetadataProvider(ModelMetadataProvider):
|
||||||
"""Provider that uses CivArchive API for metadata"""
|
"""Provider that uses CivArchive API for metadata"""
|
||||||
|
|
||||||
@@ -451,7 +481,14 @@ class SQLiteModelMetadataProvider(ModelMetadataProvider):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
class FallbackMetadataProvider(ModelMetadataProvider):
|
class FallbackMetadataProvider(ModelMetadataProvider):
|
||||||
"""Try providers in order, return first successful result."""
|
"""Try providers in order, return first successful result.
|
||||||
|
|
||||||
|
Rate-limit policy (#1085): once a *network* provider raises
|
||||||
|
``RateLimitError``, the chain stops consulting further network providers —
|
||||||
|
failing over would just spread the flood to the next vendor. Local-only
|
||||||
|
providers (see ``_LOCAL_PROVIDER_LABELS``) are still allowed as a last
|
||||||
|
resort because they cost the vendor nothing.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -486,7 +523,10 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
async def get_model_by_hash(self, model_hash: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||||
|
rate_limited = False
|
||||||
for provider, label in self._iter_providers():
|
for provider, label in self._iter_providers():
|
||||||
|
if rate_limited and label not in _LOCAL_PROVIDER_LABELS:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
result, error = await self._call_with_rate_limit(
|
result, error = await self._call_with_rate_limit(
|
||||||
label,
|
label,
|
||||||
@@ -496,8 +536,9 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
if result:
|
if result:
|
||||||
return result, error
|
return result, error
|
||||||
except RateLimitError as exc:
|
except RateLimitError as exc:
|
||||||
|
rate_limited = True
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Provider %s is rate-limited (retry_after=%.0fs); skipping to next provider",
|
"Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers",
|
||||||
label,
|
label,
|
||||||
exc.retry_after or 0,
|
exc.retry_after or 0,
|
||||||
)
|
)
|
||||||
@@ -505,11 +546,18 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug("Provider %s failed for get_model_by_hash: %s", label, e)
|
logger.debug("Provider %s failed for get_model_by_hash: %s", label, e)
|
||||||
continue
|
continue
|
||||||
|
if rate_limited:
|
||||||
|
# Distinct from "Model not found": callers must not mistake a
|
||||||
|
# rate-limited lookup for a confirmed deletion.
|
||||||
|
return None, "Rate limited"
|
||||||
return None, "Model not found"
|
return None, "Model not found"
|
||||||
|
|
||||||
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
|
async def get_model_versions(self, model_id: str) -> Optional[Dict[str, Any]]:
|
||||||
not_found_confirmed = False
|
not_found_confirmed = False
|
||||||
|
rate_limited = False
|
||||||
for provider, label in self._iter_providers():
|
for provider, label in self._iter_providers():
|
||||||
|
if rate_limited and label not in _LOCAL_PROVIDER_LABELS:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
result = await self._call_with_rate_limit(
|
result = await self._call_with_rate_limit(
|
||||||
label,
|
label,
|
||||||
@@ -519,8 +567,9 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
if result:
|
if result:
|
||||||
return result
|
return result
|
||||||
except RateLimitError as exc:
|
except RateLimitError as exc:
|
||||||
|
rate_limited = True
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Provider %s is rate-limited (retry_after=%.0fs); skipping to next provider",
|
"Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers",
|
||||||
label,
|
label,
|
||||||
exc.retry_after or 0,
|
exc.retry_after or 0,
|
||||||
)
|
)
|
||||||
@@ -539,7 +588,10 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
|
async def get_model_version(self, model_id: Optional[int] = None, version_id: Optional[int] = None) -> Optional[Dict[str, Any]]:
|
||||||
|
rate_limited = False
|
||||||
for provider, label in self._iter_providers():
|
for provider, label in self._iter_providers():
|
||||||
|
if rate_limited and label not in _LOCAL_PROVIDER_LABELS:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
result = await self._call_with_rate_limit(
|
result = await self._call_with_rate_limit(
|
||||||
label,
|
label,
|
||||||
@@ -550,8 +602,9 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
if result:
|
if result:
|
||||||
return result
|
return result
|
||||||
except RateLimitError as exc:
|
except RateLimitError as exc:
|
||||||
|
rate_limited = True
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Provider %s is rate-limited (retry_after=%.0fs); skipping to next provider",
|
"Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers",
|
||||||
label,
|
label,
|
||||||
exc.retry_after or 0,
|
exc.retry_after or 0,
|
||||||
)
|
)
|
||||||
@@ -562,7 +615,10 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
async def get_model_version_info(self, version_id: str) -> Tuple[Optional[Dict[str, Any]], Optional[str]]:
|
||||||
|
rate_limited = False
|
||||||
for provider, label in self._iter_providers():
|
for provider, label in self._iter_providers():
|
||||||
|
if rate_limited and label not in _LOCAL_PROVIDER_LABELS:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
result, error = await self._call_with_rate_limit(
|
result, error = await self._call_with_rate_limit(
|
||||||
label,
|
label,
|
||||||
@@ -572,8 +628,9 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
if result:
|
if result:
|
||||||
return result, error
|
return result, error
|
||||||
except RateLimitError as exc:
|
except RateLimitError as exc:
|
||||||
|
rate_limited = True
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Provider %s is rate-limited (retry_after=%.0fs); skipping to next provider",
|
"Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers",
|
||||||
label,
|
label,
|
||||||
exc.retry_after or 0,
|
exc.retry_after or 0,
|
||||||
)
|
)
|
||||||
@@ -581,12 +638,17 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug("Provider %s failed for get_model_version_info: %s", label, e)
|
logger.debug("Provider %s failed for get_model_version_info: %s", label, e)
|
||||||
continue
|
continue
|
||||||
|
if rate_limited:
|
||||||
|
return None, "Rate limited"
|
||||||
return None, "No provider could retrieve the data"
|
return None, "No provider could retrieve the data"
|
||||||
|
|
||||||
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[str, Any]]]:
|
) -> Optional[List[Dict[str, Any]]]:
|
||||||
|
rate_limited = False
|
||||||
for provider, label in self._iter_providers():
|
for provider, label in self._iter_providers():
|
||||||
|
if rate_limited and label not in _LOCAL_PROVIDER_LABELS:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
result = await self._call_with_rate_limit(
|
result = await self._call_with_rate_limit(
|
||||||
label,
|
label,
|
||||||
@@ -598,8 +660,9 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
except NotImplementedError:
|
except NotImplementedError:
|
||||||
continue
|
continue
|
||||||
except RateLimitError as exc:
|
except RateLimitError as exc:
|
||||||
|
rate_limited = True
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Provider %s is rate-limited (retry_after=%.0fs); skipping to next provider",
|
"Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers",
|
||||||
label,
|
label,
|
||||||
exc.retry_after or 0,
|
exc.retry_after or 0,
|
||||||
)
|
)
|
||||||
@@ -614,7 +677,10 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
|
async def get_user_models(self, username: str, cursor: Optional[str] = None) -> Optional[Dict[str, Any]]:
|
||||||
|
rate_limited = False
|
||||||
for provider, label in self._iter_providers():
|
for provider, label in self._iter_providers():
|
||||||
|
if rate_limited and label not in _LOCAL_PROVIDER_LABELS:
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
result = await self._call_with_rate_limit(
|
result = await self._call_with_rate_limit(
|
||||||
label,
|
label,
|
||||||
@@ -625,8 +691,9 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
if result is not None:
|
if result is not None:
|
||||||
return result
|
return result
|
||||||
except RateLimitError as exc:
|
except RateLimitError as exc:
|
||||||
|
rate_limited = True
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Provider %s is rate-limited (retry_after=%.0fs); skipping to next provider",
|
"Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers",
|
||||||
label,
|
label,
|
||||||
exc.retry_after or 0,
|
exc.retry_after or 0,
|
||||||
)
|
)
|
||||||
@@ -649,6 +716,37 @@ class FallbackMetadataProvider(ModelMetadataProvider):
|
|||||||
continue
|
continue
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
async def get_version_file_mini(
|
||||||
|
self, version_id: int, file_id: int
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
rate_limited = False
|
||||||
|
for provider, label in self._iter_providers():
|
||||||
|
if rate_limited and label not in _LOCAL_PROVIDER_LABELS:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
result = await self._call_with_rate_limit(
|
||||||
|
label,
|
||||||
|
provider.get_version_file_mini,
|
||||||
|
version_id,
|
||||||
|
file_id,
|
||||||
|
)
|
||||||
|
if result:
|
||||||
|
return result
|
||||||
|
except RateLimitError as exc:
|
||||||
|
rate_limited = True
|
||||||
|
logger.warning(
|
||||||
|
"Provider %s is rate-limited (retry_after=%.0fs); not failing over to other network providers",
|
||||||
|
label,
|
||||||
|
exc.retry_after or 0,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug(
|
||||||
|
"Provider %s failed for get_version_file_mini: %s", label, e
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
return None
|
||||||
|
|
||||||
def _iter_providers(self):
|
def _iter_providers(self):
|
||||||
return zip(self.providers, self._provider_labels)
|
return zip(self.providers, self._provider_labels)
|
||||||
|
|
||||||
@@ -740,6 +838,16 @@ class RateLimitRetryingProvider(ModelMetadataProvider):
|
|||||||
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
async def get_creator_model_count(self, username: str) -> Optional[int]:
|
||||||
return await self._provider.get_creator_model_count(username)
|
return await self._provider.get_creator_model_count(username)
|
||||||
|
|
||||||
|
async def get_version_file_mini(
|
||||||
|
self, version_id: int, file_id: int
|
||||||
|
) -> Optional[Dict[str, Any]]:
|
||||||
|
return await self._rate_limit_helper.run(
|
||||||
|
self._label,
|
||||||
|
self._provider.get_version_file_mini,
|
||||||
|
version_id,
|
||||||
|
file_id,
|
||||||
|
)
|
||||||
|
|
||||||
class ModelMetadataProviderManager:
|
class ModelMetadataProviderManager:
|
||||||
"""Manager for selecting and using model metadata providers"""
|
"""Manager for selecting and using model metadata providers"""
|
||||||
|
|
||||||
|
|||||||
+806
-150
File diff suppressed because it is too large
Load Diff
@@ -118,19 +118,24 @@ class ModelServiceFactory:
|
|||||||
|
|
||||||
|
|
||||||
def register_default_model_types():
|
def register_default_model_types():
|
||||||
"""Register the default model types (LoRA, Checkpoint, and Embedding)"""
|
"""Register the default model types (LoRA, Checkpoint, Embedding, and Other)"""
|
||||||
from ..services.lora_service import LoraService
|
from ..services.lora_service import LoraService
|
||||||
from ..services.checkpoint_service import CheckpointService
|
from ..services.checkpoint_service import CheckpointService
|
||||||
from ..services.embedding_service import EmbeddingService
|
from ..services.embedding_service import EmbeddingService
|
||||||
|
from ..services.other_model_service import OtherModelService
|
||||||
from ..routes.lora_routes import LoraRoutes
|
from ..routes.lora_routes import LoraRoutes
|
||||||
from ..routes.checkpoint_routes import CheckpointRoutes
|
from ..routes.checkpoint_routes import CheckpointRoutes
|
||||||
from ..routes.embedding_routes import EmbeddingRoutes
|
from ..routes.embedding_routes import EmbeddingRoutes
|
||||||
|
from ..routes.other_routes import OtherRoutes
|
||||||
|
|
||||||
# Register LoRA model type
|
# Register LoRA model type
|
||||||
ModelServiceFactory.register_model_type('lora', LoraService, LoraRoutes)
|
ModelServiceFactory.register_model_type('lora', LoraService, LoraRoutes)
|
||||||
|
|
||||||
# Register Checkpoint model type
|
# Register Checkpoint model type
|
||||||
ModelServiceFactory.register_model_type('checkpoint', CheckpointService, CheckpointRoutes)
|
ModelServiceFactory.register_model_type('checkpoint', CheckpointService, CheckpointRoutes)
|
||||||
|
|
||||||
# Register Embedding model type
|
# Register Embedding model type
|
||||||
ModelServiceFactory.register_model_type('embedding', EmbeddingService, EmbeddingRoutes)
|
ModelServiceFactory.register_model_type('embedding', EmbeddingService, EmbeddingRoutes)
|
||||||
|
|
||||||
|
# Register Other model type (VAE, upscaler, text encoder, ...)
|
||||||
|
ModelServiceFactory.register_model_type('other', OtherModelService, OtherRoutes)
|
||||||
@@ -0,0 +1,86 @@
|
|||||||
|
"""External model-source providers (Hugging Face, ModelScope, TensorArt).
|
||||||
|
|
||||||
|
This package is the single abstraction over "a site that hosts models and
|
||||||
|
a model card". See :mod:`py.services.model_sources.base` for the provider
|
||||||
|
protocol and :mod:`py.services.model_sources.registry` for the lookup and
|
||||||
|
metadata-normalisation helpers used across the codebase.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
GROUP_PREFIXES,
|
||||||
|
HTTP_TIMEOUT,
|
||||||
|
ModelCardContext,
|
||||||
|
ModelSource,
|
||||||
|
ModelSourceCache,
|
||||||
|
ModelSourceError,
|
||||||
|
SourceRef,
|
||||||
|
USER_AGENT,
|
||||||
|
clean_source_url,
|
||||||
|
fetch_json,
|
||||||
|
fetch_text,
|
||||||
|
filter_weight_files,
|
||||||
|
is_valid_source_id,
|
||||||
|
)
|
||||||
|
from .huggingface import HuggingFaceSource
|
||||||
|
from .hydration import (
|
||||||
|
hydrate_from_source,
|
||||||
|
load_model_card,
|
||||||
|
resolve_site_base_model,
|
||||||
|
)
|
||||||
|
from .modelscope import ModelScopeIntlSource, ModelScopeSource
|
||||||
|
from .registry import (
|
||||||
|
LEGACY_HF_URL_FIELD,
|
||||||
|
SOURCE_PLATFORM_FIELD,
|
||||||
|
SOURCE_URL_FIELD,
|
||||||
|
detect_source,
|
||||||
|
downloadable_sources,
|
||||||
|
get_download_source,
|
||||||
|
get_source,
|
||||||
|
get_source_platform,
|
||||||
|
has_external_source,
|
||||||
|
list_sources,
|
||||||
|
normalize_metadata_source,
|
||||||
|
resolve_source_ref,
|
||||||
|
source_group_key,
|
||||||
|
source_label,
|
||||||
|
)
|
||||||
|
from .tensorart import TensorArtSource
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"GROUP_PREFIXES",
|
||||||
|
"HTTP_TIMEOUT",
|
||||||
|
"LEGACY_HF_URL_FIELD",
|
||||||
|
"ModelCardContext",
|
||||||
|
"ModelSource",
|
||||||
|
"ModelSourceCache",
|
||||||
|
"ModelSourceError",
|
||||||
|
"HuggingFaceSource",
|
||||||
|
"ModelScopeIntlSource",
|
||||||
|
"ModelScopeSource",
|
||||||
|
"SOURCE_PLATFORM_FIELD",
|
||||||
|
"SOURCE_URL_FIELD",
|
||||||
|
"SourceRef",
|
||||||
|
"TensorArtSource",
|
||||||
|
"USER_AGENT",
|
||||||
|
"clean_source_url",
|
||||||
|
"detect_source",
|
||||||
|
"downloadable_sources",
|
||||||
|
"fetch_json",
|
||||||
|
"fetch_text",
|
||||||
|
"filter_weight_files",
|
||||||
|
"get_download_source",
|
||||||
|
"get_source",
|
||||||
|
"get_source_platform",
|
||||||
|
"has_external_source",
|
||||||
|
"hydrate_from_source",
|
||||||
|
"is_valid_source_id",
|
||||||
|
"list_sources",
|
||||||
|
"load_model_card",
|
||||||
|
"normalize_metadata_source",
|
||||||
|
"resolve_site_base_model",
|
||||||
|
"resolve_source_ref",
|
||||||
|
"source_group_key",
|
||||||
|
"source_label",
|
||||||
|
]
|
||||||
@@ -0,0 +1,446 @@
|
|||||||
|
"""Base types for the external model-source provider abstraction.
|
||||||
|
|
||||||
|
A *model source* is a third-party site that hosts model files and a model
|
||||||
|
card (README) describing them — Hugging Face, ModelScope, TensorArt, and
|
||||||
|
whatever gets added later. Everything the rest of the codebase needs to
|
||||||
|
know about such a site is expressed by :class:`ModelSource`:
|
||||||
|
|
||||||
|
* how to recognise one of its URLs (:meth:`ModelSource.parse`)
|
||||||
|
* the canonical page URL for a source id (:meth:`ModelSource.canonical_url`)
|
||||||
|
* how to fetch the model card (:meth:`ModelSource.fetch_model_card`)
|
||||||
|
* how to fetch the extras that live *outside* the README
|
||||||
|
(:meth:`ModelSource.fetch_model_card_context`)
|
||||||
|
* how to turn repository-relative asset paths into absolute URLs
|
||||||
|
(:meth:`ModelSource.asset_base_url`)
|
||||||
|
* which capabilities the site actually supports
|
||||||
|
(``supports_enrichment`` / ``supports_download``)
|
||||||
|
|
||||||
|
Keeping this in one place means the agent pipeline, the scanners, and the
|
||||||
|
HTTP handlers never need site-specific branching.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Dict, Iterable, Optional
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
|
||||||
|
from ...utils.constants import MODEL_FILE_EXTENSIONS
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
#: Shared HTTP timeout for model-card fetches.
|
||||||
|
HTTP_TIMEOUT = 30
|
||||||
|
|
||||||
|
#: User agent used for all model-source HTTP requests.
|
||||||
|
USER_AGENT = "ComfyUI-LoRA-Manager/1.0"
|
||||||
|
|
||||||
|
#: Platform → short prefix used when building version-group keys.
|
||||||
|
#: ``huggingface`` keeps the historical ``hf:`` prefix for backward
|
||||||
|
#: compatibility with already-cached group keys.
|
||||||
|
GROUP_PREFIXES: dict[str, str] = {
|
||||||
|
"huggingface": "hf",
|
||||||
|
"modelscope": "ms",
|
||||||
|
"modelscope-ai": "msai",
|
||||||
|
"tensorart": "ta",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SourceRef:
|
||||||
|
"""A parsed reference to a model hosted on an external site."""
|
||||||
|
|
||||||
|
platform: str
|
||||||
|
"""Canonical platform id, e.g. ``"huggingface"``."""
|
||||||
|
|
||||||
|
source_id: str
|
||||||
|
"""Site-specific identity, e.g. ``"user/repo"`` or ``"827823520299086029"``."""
|
||||||
|
|
||||||
|
url: str
|
||||||
|
"""Canonical URL of the model page."""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class ModelCardContext:
|
||||||
|
"""Site-specific extras that accompany a model's README model card.
|
||||||
|
|
||||||
|
A model card is not always just ``README.md``. ModelScope, for example,
|
||||||
|
keeps the author's summary, the site-curated tags, and the per-file
|
||||||
|
example images in its model-detail API rather than in the repository.
|
||||||
|
Sources with no such extras return an empty context (the default), so
|
||||||
|
every field here must be treated as optional by callers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
description: str = ""
|
||||||
|
"""Author-written summary shown on the model page, outside the README."""
|
||||||
|
|
||||||
|
model_name: str = ""
|
||||||
|
"""Site-published display name for the repository.
|
||||||
|
|
||||||
|
Sites publish this next to the repository id (ModelScope's ``Name``).
|
||||||
|
It is what a CivitAI download would store as the model's name, so the
|
||||||
|
card never has to fall back to the local filename.
|
||||||
|
"""
|
||||||
|
|
||||||
|
model_name_localized: str = ""
|
||||||
|
"""Site-published localized name (ModelScope's ``ChineseName``)."""
|
||||||
|
|
||||||
|
version_name: str = ""
|
||||||
|
"""Site-published label for the requested file's version.
|
||||||
|
|
||||||
|
Resolved per file, like :attr:`example_images`: a repository publishes
|
||||||
|
one label per checkpoint (ModelScope's ``modelVersion.showName``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
license: str = ""
|
||||||
|
"""License the site records for the repository."""
|
||||||
|
|
||||||
|
model_type: str = ""
|
||||||
|
"""Site-reported model type, e.g. ModelScope's ``AigcType`` (``LoRA``)."""
|
||||||
|
|
||||||
|
base_model: str = ""
|
||||||
|
"""Base model as reported by the site (possibly a site-local id)."""
|
||||||
|
|
||||||
|
base_model_aliases: list[str] = field(default_factory=list)
|
||||||
|
"""Other names the site uses for the same base model.
|
||||||
|
|
||||||
|
Sites often publish both a link-style id (``krea/Krea-2-Turbo``) and an
|
||||||
|
internal architecture enum (``KREA_2``). The enum usually normalises
|
||||||
|
cleanly onto this system's canonical vocabulary, so it is the better
|
||||||
|
resolution hint for :mod:`py.services.agent.base_model_resolver`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
official_tags: list[str] = field(default_factory=list)
|
||||||
|
"""Content tags curated by the site itself."""
|
||||||
|
|
||||||
|
example_images: list[str] = field(default_factory=list)
|
||||||
|
"""Absolute URLs of example images for the requested model file."""
|
||||||
|
|
||||||
|
trigger_words: list[str] = field(default_factory=list)
|
||||||
|
"""Trigger words the site records for the requested model file."""
|
||||||
|
|
||||||
|
def is_empty(self) -> bool:
|
||||||
|
"""Return ``True`` when the site contributed nothing extra."""
|
||||||
|
|
||||||
|
return not any(
|
||||||
|
(
|
||||||
|
self.description,
|
||||||
|
self.model_name,
|
||||||
|
self.model_name_localized,
|
||||||
|
self.version_name,
|
||||||
|
self.license,
|
||||||
|
self.model_type,
|
||||||
|
self.base_model,
|
||||||
|
self.base_model_aliases,
|
||||||
|
self.official_tags,
|
||||||
|
self.example_images,
|
||||||
|
self.trigger_words,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelSourceError(Exception):
|
||||||
|
"""Raised when a model source cannot satisfy a request.
|
||||||
|
|
||||||
|
Carries the HTTP status the API handler should answer with, so the
|
||||||
|
handlers stay free of per-site error mapping.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, message: str, status: int = 502) -> None:
|
||||||
|
super().__init__(message)
|
||||||
|
self.status = status
|
||||||
|
|
||||||
|
|
||||||
|
class ModelSourceCache:
|
||||||
|
"""Per-run memo shared between the agent pipeline and a model source.
|
||||||
|
|
||||||
|
A collection repository publishes many model files under a single source
|
||||||
|
id, so enriching each file re-fetches the same README and the same
|
||||||
|
repository metadata. One cache is created per enrichment run and thrown
|
||||||
|
away afterwards: nothing is retained across runs (a model card can change
|
||||||
|
at any time), and download URLs are never routed through it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
#: Provider-agnostic: ``"<platform>:<source_id>"`` → raw README text.
|
||||||
|
self.readmes: Dict[str, str] = {}
|
||||||
|
#: Provider-owned scratch space. Keys must be namespaced by the
|
||||||
|
#: provider (``(platform, kind, source_id)``) so two providers can
|
||||||
|
#: never collide. Only successful results should be stored, so a
|
||||||
|
#: transient failure is still retried for the next file.
|
||||||
|
self.provider: Dict[Any, Any] = {}
|
||||||
|
|
||||||
|
|
||||||
|
#: Repository ids are always exactly ``owner/name``. Components may contain
|
||||||
|
#: dots (``black-forest-labs/FLUX.1-dev``) but must not be empty, ``.`` / ``..``,
|
||||||
|
#: or start with a dot - the id is used as a path segment on disk.
|
||||||
|
_SOURCE_ID_COMPONENT = re.compile(r"^[A-Za-z0-9_][A-Za-z0-9_.\-]*$")
|
||||||
|
|
||||||
|
|
||||||
|
def is_valid_source_id(source_id: str) -> bool:
|
||||||
|
"""Return ``True`` when *source_id* is a safe ``owner/name`` repository id."""
|
||||||
|
|
||||||
|
if not source_id or not isinstance(source_id, str) or source_id.count("/") != 1:
|
||||||
|
return False
|
||||||
|
owner, name = source_id.split("/", 1)
|
||||||
|
return all(
|
||||||
|
part and part not in (".", "..") and _SOURCE_ID_COMPONENT.match(part)
|
||||||
|
for part in (owner, name)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def fetch_text(url: str, *, timeout: int = HTTP_TIMEOUT) -> str:
|
||||||
|
"""Fetch *url* and return its body as text, or ``""`` on any failure.
|
||||||
|
|
||||||
|
Network problems are expected (offline installs, rate limits, dead
|
||||||
|
repos) and must never bubble up into the pipeline, so every error is
|
||||||
|
logged at debug level and normalised to an empty string.
|
||||||
|
"""
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession(
|
||||||
|
headers={"User-Agent": USER_AGENT},
|
||||||
|
timeout=aiohttp.ClientTimeout(total=timeout),
|
||||||
|
) as session:
|
||||||
|
async with session.get(url) as resp:
|
||||||
|
if resp.status == 200:
|
||||||
|
return await resp.text()
|
||||||
|
logger.debug("Fetch %s returned HTTP %s", url, resp.status)
|
||||||
|
except Exception as exc: # pragma: no cover - network dependent
|
||||||
|
logger.debug("Failed to fetch %s: %s", url, exc)
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
async def fetch_json(
|
||||||
|
url: str, *, timeout: int = HTTP_TIMEOUT
|
||||||
|
) -> tuple[int, Any]:
|
||||||
|
"""Fetch *url* and return ``(status, parsed_body)``.
|
||||||
|
|
||||||
|
Unlike :func:`fetch_text` this reports the status, because callers such as
|
||||||
|
the file-listing endpoints need to distinguish "repo not found" (404) from
|
||||||
|
a transport failure. ``parsed_body`` is ``None`` when the response is not
|
||||||
|
JSON or the request failed outright (status ``0``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession(
|
||||||
|
headers={"User-Agent": USER_AGENT},
|
||||||
|
timeout=aiohttp.ClientTimeout(total=timeout),
|
||||||
|
) as session:
|
||||||
|
async with session.get(url) as resp:
|
||||||
|
if resp.status != 200:
|
||||||
|
return resp.status, None
|
||||||
|
try:
|
||||||
|
return resp.status, await resp.json(content_type=None)
|
||||||
|
except Exception:
|
||||||
|
return resp.status, None
|
||||||
|
except Exception as exc: # pragma: no cover - network dependent
|
||||||
|
logger.debug("Failed to fetch %s: %s", url, exc)
|
||||||
|
return 0, None
|
||||||
|
|
||||||
|
|
||||||
|
class ModelSource:
|
||||||
|
"""Description and I/O for one external model hosting site."""
|
||||||
|
|
||||||
|
#: Canonical platform id stored in metadata.
|
||||||
|
platform: str = ""
|
||||||
|
|
||||||
|
#: Human-readable name used in UI copy and prompts.
|
||||||
|
label: str = ""
|
||||||
|
|
||||||
|
#: Whether the agent skill can fetch a model card and run AI extraction.
|
||||||
|
supports_enrichment: bool = False
|
||||||
|
|
||||||
|
#: Whether models can be downloaded directly from this site.
|
||||||
|
supports_download: bool = False
|
||||||
|
|
||||||
|
#: Branch used when the caller does not pass an explicit revision.
|
||||||
|
default_revision: str = ""
|
||||||
|
|
||||||
|
#: Sub-directory the "use default paths" template places downloads in.
|
||||||
|
default_subdir: str = ""
|
||||||
|
|
||||||
|
#: Lenient pattern used to recognise URLs already stored in metadata.
|
||||||
|
#: Captures the site-specific source id in group ``id``.
|
||||||
|
url_pattern: re.Pattern[str] | None = None
|
||||||
|
|
||||||
|
#: Strict pattern used to validate user input. Must match the whole URL.
|
||||||
|
strict_url_pattern: re.Pattern[str] | None = None
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Parsing
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def parse(self, url: str, *, strict: bool = False) -> Optional[str]:
|
||||||
|
"""Return the source id contained in *url*, or ``None``.
|
||||||
|
|
||||||
|
With ``strict=True`` the URL must match this site's canonical shape
|
||||||
|
exactly (used when validating what a user pasted); with
|
||||||
|
``strict=False`` sub-paths such as ``/resolve/main/file.bin`` are
|
||||||
|
tolerated (used when normalising already-stored values).
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not url or not isinstance(url, str):
|
||||||
|
return None
|
||||||
|
candidate = url.strip()
|
||||||
|
if not candidate:
|
||||||
|
return None
|
||||||
|
pattern = self.strict_url_pattern if strict else self.url_pattern
|
||||||
|
if pattern is None:
|
||||||
|
return None
|
||||||
|
match = pattern.match(candidate)
|
||||||
|
return match.group("id") if match else None
|
||||||
|
|
||||||
|
def ref(self, url: str, *, strict: bool = False) -> Optional[SourceRef]:
|
||||||
|
"""Return a :class:`SourceRef` for *url*, or ``None`` if not ours."""
|
||||||
|
|
||||||
|
source_id = self.parse(url, strict=strict)
|
||||||
|
if not source_id:
|
||||||
|
return None
|
||||||
|
return SourceRef(
|
||||||
|
platform=self.platform,
|
||||||
|
source_id=source_id,
|
||||||
|
url=self.canonical_url(source_id),
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# URLs and content
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
def canonical_url(self, source_id: str) -> str:
|
||||||
|
"""Return the canonical model-page URL for *source_id*."""
|
||||||
|
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def asset_base_url(self, source_id: str, revision: str = "") -> str:
|
||||||
|
"""Base URL used to resolve repository-relative asset paths."""
|
||||||
|
|
||||||
|
return ""
|
||||||
|
|
||||||
|
def group_key(self, source_id: str) -> str:
|
||||||
|
"""Return the version-group key for *source_id*."""
|
||||||
|
|
||||||
|
prefix = GROUP_PREFIXES.get(self.platform, self.platform)
|
||||||
|
return f"{prefix}:{source_id}"
|
||||||
|
|
||||||
|
async def fetch_model_card(self, source_id: str) -> str:
|
||||||
|
"""Fetch the raw model card (README) markdown for *source_id*."""
|
||||||
|
|
||||||
|
return ""
|
||||||
|
|
||||||
|
async def fetch_model_card_context(
|
||||||
|
self,
|
||||||
|
source_id: str,
|
||||||
|
filename: str = "",
|
||||||
|
*,
|
||||||
|
sha256: str = "",
|
||||||
|
cache: Optional["ModelSourceCache"] = None,
|
||||||
|
) -> ModelCardContext:
|
||||||
|
"""Return the card extras the site keeps outside the README.
|
||||||
|
|
||||||
|
*filename* is the model file's basename (no directory) and *sha256*
|
||||||
|
its content hash; between them they select the right entry when a
|
||||||
|
repository holds several models. A site that records per-file hashes
|
||||||
|
should prefer *sha256*, because it is the only identifier that
|
||||||
|
survives the user renaming the weights.
|
||||||
|
|
||||||
|
*cache* is an optional per-run memo (see :class:`ModelSourceCache`)
|
||||||
|
that lets a provider avoid re-fetching repository-wide data for every
|
||||||
|
file in a collection repository.
|
||||||
|
|
||||||
|
Sites whose model card is fully described by :meth:`fetch_model_card`
|
||||||
|
need no override and inherit this empty context.
|
||||||
|
|
||||||
|
Implementations must never raise: enrichment treats a missing
|
||||||
|
context as "the site had nothing extra to say".
|
||||||
|
"""
|
||||||
|
|
||||||
|
return ModelCardContext()
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# Download support
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def list_files(
|
||||||
|
self, source_id: str, revision: str = ""
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""List downloadable weight files in *source_id*.
|
||||||
|
|
||||||
|
Returns ``[{"filename": <repo-relative path>, "size": <bytes>}]``,
|
||||||
|
largest first, filtered to :data:`MODEL_FILE_EXTENSIONS`. Sites
|
||||||
|
without download support return an empty list.
|
||||||
|
|
||||||
|
Raises :class:`ModelSourceError` when the repository cannot be read,
|
||||||
|
so the handler can surface "not found" separately from a transport
|
||||||
|
failure.
|
||||||
|
"""
|
||||||
|
|
||||||
|
return []
|
||||||
|
|
||||||
|
def file_download_url(
|
||||||
|
self, source_id: str, filename: str, revision: str = ""
|
||||||
|
) -> str:
|
||||||
|
"""Return the direct (redirecting) download URL for one file."""
|
||||||
|
|
||||||
|
raise ModelSourceError(
|
||||||
|
f"{self.label or self.platform} does not support downloads", status=400
|
||||||
|
)
|
||||||
|
|
||||||
|
def resolve_revision(self, revision: str = "") -> str:
|
||||||
|
"""Return *revision*, falling back to this site's default branch."""
|
||||||
|
|
||||||
|
return revision or self.default_revision
|
||||||
|
|
||||||
|
def page_url_for_file(self, source_id: str, filename: str) -> str:
|
||||||
|
"""Return the human-facing page for *filename* inside *source_id*."""
|
||||||
|
|
||||||
|
return self.canonical_url(source_id)
|
||||||
|
|
||||||
|
def __repr__(self) -> str: # pragma: no cover - debugging aid
|
||||||
|
return f"<ModelSource {self.platform}>"
|
||||||
|
|
||||||
|
|
||||||
|
def clean_source_url(url: Any) -> str:
|
||||||
|
"""Normalise a stored source URL value into a stripped string."""
|
||||||
|
|
||||||
|
if not isinstance(url, str):
|
||||||
|
return ""
|
||||||
|
return url.strip()
|
||||||
|
|
||||||
|
|
||||||
|
def filter_weight_files(entries: Iterable[tuple[str, int]]) -> list[dict[str, Any]]:
|
||||||
|
"""Keep model-weight files from ``(path, size)`` pairs, largest first.
|
||||||
|
|
||||||
|
Every site lists a lot more than weights (READMEs, configs, tokenizers,
|
||||||
|
…); the download picker only ever wants the files ComfyUI can load, which
|
||||||
|
is exactly :data:`MODEL_FILE_EXTENSIONS`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
files = [
|
||||||
|
{"filename": path, "size": int(size or 0)}
|
||||||
|
for path, size in entries
|
||||||
|
if path and os.path.splitext(path)[1].lower() in MODEL_FILE_EXTENSIONS
|
||||||
|
]
|
||||||
|
files.sort(key=lambda entry: entry["size"], reverse=True)
|
||||||
|
return files
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"GROUP_PREFIXES",
|
||||||
|
"HTTP_TIMEOUT",
|
||||||
|
"ModelCardContext",
|
||||||
|
"ModelSource",
|
||||||
|
"ModelSourceCache",
|
||||||
|
"ModelSourceError",
|
||||||
|
"SourceRef",
|
||||||
|
"USER_AGENT",
|
||||||
|
"clean_source_url",
|
||||||
|
"fetch_json",
|
||||||
|
"fetch_text",
|
||||||
|
"filter_weight_files",
|
||||||
|
"is_valid_source_id",
|
||||||
|
]
|
||||||
@@ -0,0 +1,106 @@
|
|||||||
|
"""Hugging Face model source."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
ModelSource,
|
||||||
|
ModelSourceError,
|
||||||
|
fetch_json,
|
||||||
|
fetch_text,
|
||||||
|
filter_weight_files,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
#: Lenient — used to normalise URLs already stored in metadata; tolerates
|
||||||
|
#: sub-paths such as ``/resolve/main/model.safetensors``.
|
||||||
|
_URL_PATTERN = re.compile(
|
||||||
|
r"https?://(?:www\.)?huggingface\.co/(?P<id>[^/?#\s]+/[^/?#\s]+)"
|
||||||
|
)
|
||||||
|
|
||||||
|
#: Strict — validates what the user pasted into the "link model" dialog.
|
||||||
|
_STRICT_URL_PATTERN = re.compile(
|
||||||
|
r"https?://(?:www\.)?huggingface\.co/(?P<id>[^/?#\s]+/[^/?#\s]+)/?$"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class HuggingFaceSource(ModelSource):
|
||||||
|
"""Hugging Face Hub (``huggingface.co``)."""
|
||||||
|
|
||||||
|
platform = "huggingface"
|
||||||
|
label = "Hugging Face"
|
||||||
|
supports_enrichment = True
|
||||||
|
supports_download = True
|
||||||
|
default_revision = "main"
|
||||||
|
default_subdir = "huggingface"
|
||||||
|
url_pattern = _URL_PATTERN
|
||||||
|
strict_url_pattern = _STRICT_URL_PATTERN
|
||||||
|
|
||||||
|
def canonical_url(self, source_id: str) -> str:
|
||||||
|
return f"https://huggingface.co/{source_id}"
|
||||||
|
|
||||||
|
def asset_base_url(self, source_id: str, revision: str = "") -> str:
|
||||||
|
return f"https://huggingface.co/{source_id}/resolve/{self.resolve_revision(revision)}"
|
||||||
|
|
||||||
|
async def fetch_model_card(self, source_id: str) -> str:
|
||||||
|
"""Fetch ``README.md`` from Hugging Face (tries ``main``, then ``master``)."""
|
||||||
|
|
||||||
|
for branch in ("main", "master"):
|
||||||
|
text = await fetch_text(
|
||||||
|
f"https://huggingface.co/{source_id}/raw/{branch}/README.md"
|
||||||
|
)
|
||||||
|
if text:
|
||||||
|
return text
|
||||||
|
return ""
|
||||||
|
|
||||||
|
async def list_files(
|
||||||
|
self, source_id: str, revision: str = ""
|
||||||
|
) -> list[dict]:
|
||||||
|
"""List weight files via the Hub tree API.
|
||||||
|
|
||||||
|
The tree endpoint (rather than the model-info endpoint) is used
|
||||||
|
because it reports accurate sizes for LFS-tracked files.
|
||||||
|
"""
|
||||||
|
|
||||||
|
revision = self.resolve_revision(revision)
|
||||||
|
status, payload = await fetch_json(
|
||||||
|
f"https://huggingface.co/api/models/{source_id}/tree/{revision}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if status == 404:
|
||||||
|
raise ModelSourceError(f"Repository '{source_id}' not found", status=404)
|
||||||
|
if status != 200 or not isinstance(payload, list):
|
||||||
|
raise ModelSourceError(
|
||||||
|
f"Hugging Face API error while listing '{source_id}' (HTTP {status})"
|
||||||
|
)
|
||||||
|
|
||||||
|
entries = []
|
||||||
|
for entry in payload:
|
||||||
|
if not isinstance(entry, dict):
|
||||||
|
continue
|
||||||
|
path = entry.get("path", "")
|
||||||
|
size = entry.get("size", 0) or 0
|
||||||
|
if not size and isinstance(entry.get("lfs"), dict):
|
||||||
|
size = entry["lfs"].get("size", 0) or 0
|
||||||
|
entries.append((path, size))
|
||||||
|
|
||||||
|
return filter_weight_files(entries)
|
||||||
|
|
||||||
|
def file_download_url(
|
||||||
|
self, source_id: str, filename: str, revision: str = ""
|
||||||
|
) -> str:
|
||||||
|
return (
|
||||||
|
f"https://huggingface.co/{source_id}/resolve/"
|
||||||
|
f"{self.resolve_revision(revision)}/{filename}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def page_url_for_file(self, source_id: str, filename: str) -> str:
|
||||||
|
return (
|
||||||
|
f"https://huggingface.co/{source_id}/blob/{self.default_revision}/{filename}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["HuggingFaceSource"]
|
||||||
@@ -0,0 +1,235 @@
|
|||||||
|
"""Deterministic metadata hydration for freshly downloaded source models.
|
||||||
|
|
||||||
|
A CivitAI download writes a fully-populated metadata sidecar as part of the
|
||||||
|
download itself: the name, the description, the tags, the trigger words and
|
||||||
|
the example images all arrive with the file. A download from an external
|
||||||
|
model source (ModelScope, Hugging Face) has the same information behind a
|
||||||
|
public API, but historically landed as a bare filename plus a source URL that
|
||||||
|
the user had to enrich by hand ("Enrich Metadata with AI").
|
||||||
|
|
||||||
|
This module closes that gap without involving an LLM. It fetches the linked
|
||||||
|
site's model card, hands it to the same :class:`~py.services.agent.post_processor.PostProcessor`
|
||||||
|
the AI skill uses, and writes the result. Everything it applies is data the
|
||||||
|
site published, so it is safe to run automatically on every download and to
|
||||||
|
treat as a fallback for the gaps the LLM would otherwise fill.
|
||||||
|
|
||||||
|
Nothing here may break a download: every failure is logged and normalised to
|
||||||
|
"the site had nothing to contribute".
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
from typing import TYPE_CHECKING, Optional
|
||||||
|
|
||||||
|
from .base import ModelCardContext, ModelSourceCache
|
||||||
|
from .registry import get_source, resolve_source_ref
|
||||||
|
|
||||||
|
if TYPE_CHECKING: # pragma: no cover - typing only
|
||||||
|
from .base import ModelSource, SourceRef
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
#: How long a fetched repository payload stays usable. A download batch walks
|
||||||
|
#: a repository's files one HTTP request at a time, and the README plus the
|
||||||
|
#: detail payload describe the *repository*, not the file, so re-fetching them
|
||||||
|
#: per file would be pure waste. They expire so an edited model card is still
|
||||||
|
#: picked up by the next batch.
|
||||||
|
SHARED_CACHE_TTL = 300.0
|
||||||
|
|
||||||
|
#: Upper bound on memoised repositories; a long-running server must not grow
|
||||||
|
#: without limit.
|
||||||
|
SHARED_CACHE_MAX_ENTRIES = 32
|
||||||
|
|
||||||
|
#: ``"<platform>:<source_id>"`` → ``(expiry, memo)``.
|
||||||
|
_shared_caches: dict[str, tuple[float, ModelSourceCache]] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def shared_source_cache(platform: str, source_id: str) -> ModelSourceCache:
|
||||||
|
"""Return a short-lived per-repository memo for download-time hydration."""
|
||||||
|
|
||||||
|
now = time.monotonic()
|
||||||
|
key = f"{platform}:{source_id}"
|
||||||
|
entry = _shared_caches.get(key)
|
||||||
|
if entry is not None and entry[0] > now:
|
||||||
|
return entry[1]
|
||||||
|
|
||||||
|
for expired in [k for k, (expiry, _) in _shared_caches.items() if expiry <= now]:
|
||||||
|
_shared_caches.pop(expired, None)
|
||||||
|
if len(_shared_caches) >= SHARED_CACHE_MAX_ENTRIES:
|
||||||
|
oldest = min(_shared_caches, key=lambda k: _shared_caches[k][0])
|
||||||
|
_shared_caches.pop(oldest, None)
|
||||||
|
|
||||||
|
cache = ModelSourceCache()
|
||||||
|
_shared_caches[key] = (now + SHARED_CACHE_TTL, cache)
|
||||||
|
return cache
|
||||||
|
|
||||||
|
|
||||||
|
def reset_shared_caches() -> None:
|
||||||
|
"""Drop every memoised repository — used by tests."""
|
||||||
|
|
||||||
|
_shared_caches.clear()
|
||||||
|
|
||||||
|
|
||||||
|
async def load_model_card(
|
||||||
|
source: "ModelSource",
|
||||||
|
source_id: str,
|
||||||
|
cache: Optional[ModelSourceCache] = None,
|
||||||
|
) -> str:
|
||||||
|
"""Return *source_id*'s README, reusing *cache* when one is supplied.
|
||||||
|
|
||||||
|
Only successful reads are memoised, leaving a transient failure to be
|
||||||
|
retried for the next file of the same repository.
|
||||||
|
"""
|
||||||
|
|
||||||
|
key = f"{source.platform}:{source_id}"
|
||||||
|
if cache is not None:
|
||||||
|
cached = cache.readmes.get(key)
|
||||||
|
if cached is not None:
|
||||||
|
return cached
|
||||||
|
|
||||||
|
readme = await source.fetch_model_card(source_id)
|
||||||
|
if cache is not None and readme:
|
||||||
|
cache.readmes[key] = readme
|
||||||
|
return readme or ""
|
||||||
|
|
||||||
|
|
||||||
|
async def resolve_site_base_model(context: ModelCardContext) -> str:
|
||||||
|
"""Resolve the site's base-model hints to a canonical name, or ``""``.
|
||||||
|
|
||||||
|
Sites name base models in their own vocabulary (ModelScope publishes both
|
||||||
|
``krea/Krea-2-Turbo`` and the ``KREA_2_TURBO`` enum). The resolver is
|
||||||
|
strict and only ever returns a name the canonical vocabulary already
|
||||||
|
contains, so an uncertain hint yields ``""`` rather than a plausible-looking
|
||||||
|
wrong value.
|
||||||
|
"""
|
||||||
|
|
||||||
|
hints = [*context.base_model_aliases, context.base_model]
|
||||||
|
if not any(hints):
|
||||||
|
return ""
|
||||||
|
|
||||||
|
# Imported lazily: pulling in the agent package at module scope would make
|
||||||
|
# the model-source package import itself while it is still initialising.
|
||||||
|
try:
|
||||||
|
from ...metadata_ops import list_base_models
|
||||||
|
from ..agent.base_model_resolver import resolve_base_model
|
||||||
|
|
||||||
|
known_names = await list_base_models()
|
||||||
|
except Exception as exc:
|
||||||
|
logger.warning("Could not resolve a site base model: %s", exc)
|
||||||
|
return ""
|
||||||
|
return resolve_base_model(hints, known_names)
|
||||||
|
|
||||||
|
|
||||||
|
async def hydrate_from_source(
|
||||||
|
file_path: str,
|
||||||
|
*,
|
||||||
|
ref: "SourceRef",
|
||||||
|
cache: Optional[ModelSourceCache] = None,
|
||||||
|
) -> list[str]:
|
||||||
|
"""Apply the linked site's published metadata to a downloaded model.
|
||||||
|
|
||||||
|
This is the deterministic counterpart of the ``enrich_hf_metadata`` skill:
|
||||||
|
it produces the same populated model card a CivitAI download produces,
|
||||||
|
without an LLM and without user action.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_path: The just-downloaded model file, whose sidecar already
|
||||||
|
carries the SHA256 used to match the right file in a collection
|
||||||
|
repository.
|
||||||
|
ref: The source the file came from.
|
||||||
|
cache: Optional per-call memo; defaults to a short-lived shared one so
|
||||||
|
a batch over one repository fetches its card only once.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The names of the metadata fields that changed. Never raises — a site
|
||||||
|
that is down, or an API that changed shape, must not fail a download.
|
||||||
|
"""
|
||||||
|
|
||||||
|
try:
|
||||||
|
source = get_source(ref.platform)
|
||||||
|
if source is None or not source.supports_enrichment:
|
||||||
|
return []
|
||||||
|
|
||||||
|
from ...metadata_ops import read_metadata
|
||||||
|
|
||||||
|
metadata = await read_metadata(file_path)
|
||||||
|
if not metadata:
|
||||||
|
logger.debug("No metadata to hydrate for %s", file_path)
|
||||||
|
return []
|
||||||
|
|
||||||
|
# Only a model that is actually linked to this repository may be
|
||||||
|
# updated. The download path writes those fields just before calling
|
||||||
|
# us; a file that merely shares a name with the requested one must not
|
||||||
|
# be given another model's card.
|
||||||
|
linked = resolve_source_ref(metadata)
|
||||||
|
if linked is None or (linked.platform, linked.source_id) != (
|
||||||
|
ref.platform,
|
||||||
|
ref.source_id,
|
||||||
|
):
|
||||||
|
logger.debug(
|
||||||
|
"Not hydrating %s: linked to %s, not %s",
|
||||||
|
file_path, linked.url if linked else "no model source", ref.url,
|
||||||
|
)
|
||||||
|
return []
|
||||||
|
|
||||||
|
memo = cache if cache is not None else shared_source_cache(
|
||||||
|
ref.platform, ref.source_id
|
||||||
|
)
|
||||||
|
readme = await load_model_card(source, ref.source_id, memo)
|
||||||
|
context = await source.fetch_model_card_context(
|
||||||
|
ref.source_id,
|
||||||
|
os.path.basename(file_path),
|
||||||
|
sha256=(metadata.get("sha256") or "").strip(),
|
||||||
|
cache=memo,
|
||||||
|
)
|
||||||
|
if context.is_empty() and not readme:
|
||||||
|
logger.debug(
|
||||||
|
"No published metadata for %s on %s", ref.source_id, ref.platform
|
||||||
|
)
|
||||||
|
return []
|
||||||
|
|
||||||
|
resolved_base_model = await resolve_site_base_model(context)
|
||||||
|
|
||||||
|
from ..agent.post_processor import PostProcessor
|
||||||
|
|
||||||
|
result = await PostProcessor().process(
|
||||||
|
skill_name="enrich_hf_metadata",
|
||||||
|
model_path=file_path,
|
||||||
|
llm_output={},
|
||||||
|
metadata=metadata,
|
||||||
|
readme_content=readme,
|
||||||
|
source_context=context,
|
||||||
|
resolved_base_model=resolved_base_model,
|
||||||
|
metadata_source=f"source:{ref.platform}",
|
||||||
|
)
|
||||||
|
if not result.get("success", True):
|
||||||
|
logger.debug(
|
||||||
|
"Hydration reported failure for %s: %s",
|
||||||
|
file_path, result.get("errors"),
|
||||||
|
)
|
||||||
|
return []
|
||||||
|
|
||||||
|
updated = list(result.get("updated_fields") or [])
|
||||||
|
logger.info(
|
||||||
|
"Hydrated %s from %s (%s): %s",
|
||||||
|
file_path, source.label or ref.platform, ref.source_id,
|
||||||
|
", ".join(updated) or "nothing to change",
|
||||||
|
)
|
||||||
|
return updated
|
||||||
|
except Exception as exc: # pragma: no cover - defensive by design
|
||||||
|
logger.warning("Source hydration failed for %s: %s", file_path, exc)
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"SHARED_CACHE_MAX_ENTRIES",
|
||||||
|
"SHARED_CACHE_TTL",
|
||||||
|
"hydrate_from_source",
|
||||||
|
"load_model_card",
|
||||||
|
"reset_shared_caches",
|
||||||
|
"resolve_site_base_model",
|
||||||
|
"shared_source_cache",
|
||||||
|
]
|
||||||
@@ -0,0 +1,644 @@
|
|||||||
|
"""ModelScope (魔搭社区) model sources.
|
||||||
|
|
||||||
|
ModelScope exposes the same "model card as README.md" convention as
|
||||||
|
Hugging Face, including a YAML frontmatter block that often carries
|
||||||
|
``base_model:`` and ``trigger_words:``. Four public endpoints are used,
|
||||||
|
none of which requires an API key for public models:
|
||||||
|
|
||||||
|
* ``/models/{owner}/{name}/resolve/{revision}/README.md`` — raw model card
|
||||||
|
* ``/api/v1/models/{owner}/{name}/repo?Revision=..&FilePath=README.md`` —
|
||||||
|
the same content through the API, used as a fallback when the resolve
|
||||||
|
URL is unavailable.
|
||||||
|
* ``/api/v1/models/{owner}/{name}`` — the model-detail payload behind the
|
||||||
|
model page. It carries the repository's display name (``Name`` /
|
||||||
|
``ChineseName``), the author's summary (``Description``), the license, the
|
||||||
|
AIGC type, the site tags (``OfficialTags``, falling back to ``Tags``), and,
|
||||||
|
per published version, the model filenames
|
||||||
|
(``MuseInfo.versions[].stats.fileList``) together with that version's label
|
||||||
|
(``modelVersion.showName``), example images (``coverImages``) and trigger
|
||||||
|
words. See :meth:`ModelScopeSource.fetch_model_card_context`.
|
||||||
|
* ``/api/v1/models/{owner}/{name}/repo/files?Revision=..`` — the file
|
||||||
|
listing backing the download picker. It reports real sizes for LFS
|
||||||
|
files (not the pointer size), so no extra HEAD request is needed.
|
||||||
|
|
||||||
|
Downloads go through ``/models/{owner}/{name}/resolve/{revision}/{path}``,
|
||||||
|
which redirects to a CDN URL carrying a time-limited ``auth_key``.
|
||||||
|
Requesting the resolve URL fresh on every attempt (which the shared
|
||||||
|
downloader does, including for resumable Range requests) keeps that key
|
||||||
|
valid; the CDN URL must never be cached.
|
||||||
|
|
||||||
|
The README and the detail payload both describe the whole repository rather
|
||||||
|
than one file, so a per-run ``ModelSourceCache`` keeps them from being read
|
||||||
|
again for every checkpoint of a collection repository.
|
||||||
|
|
||||||
|
Two deployments are served by this module. ``modelscope.cn`` (with
|
||||||
|
``modelscope.com`` as a redirect alias) and ``modelscope.ai`` are *separate
|
||||||
|
catalogues*, not mirrors, so they are registered as distinct sources:
|
||||||
|
:class:`ModelScopeSource` and :class:`ModelScopeIntlSource`. Every URL either
|
||||||
|
class builds is derived from its ``base_url``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import re
|
||||||
|
from typing import TYPE_CHECKING, Any, Iterable, Optional
|
||||||
|
|
||||||
|
from .base import (
|
||||||
|
ModelCardContext,
|
||||||
|
ModelSource,
|
||||||
|
ModelSourceError,
|
||||||
|
fetch_json,
|
||||||
|
fetch_text,
|
||||||
|
filter_weight_files,
|
||||||
|
)
|
||||||
|
|
||||||
|
if TYPE_CHECKING: # pragma: no cover - typing only
|
||||||
|
from .base import ModelSourceCache
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
#: ModelScope runs two independent catalogues. ``modelscope.com`` is a
|
||||||
|
#: redirect alias of the mainland site, but ``modelscope.ai`` is the
|
||||||
|
#: *international* deployment with its own repository catalogue — a repository
|
||||||
|
#: published on one is routinely absent from the other (``referall13/EM1``
|
||||||
|
#: exists only on ``.ai``, ``jj3550945163/Krea-2-LORA`` only on ``.cn``). The
|
||||||
|
#: host therefore decides which site, API and CDN a model belongs to, and the
|
||||||
|
#: two deployments are registered as separate sources rather than folded into
|
||||||
|
#: one id.
|
||||||
|
_MAINLAND_HOSTS = r"modelscope\.(?:cn|com)"
|
||||||
|
_INTERNATIONAL_HOSTS = r"modelscope\.ai"
|
||||||
|
|
||||||
|
#: Trailing view segments the site appends to a model URL; accepted verbatim
|
||||||
|
#: when the user pastes a browser tab URL.
|
||||||
|
_VIEW_SEGMENTS = r"(?:summary|files|model-file|readme|community|evaluation)?"
|
||||||
|
|
||||||
|
|
||||||
|
def _url_patterns(hosts: str) -> tuple[re.Pattern[str], re.Pattern[str]]:
|
||||||
|
"""Build the lenient and strict model-URL patterns for *hosts*."""
|
||||||
|
|
||||||
|
body = rf"https?://(?:www\.)?(?:{hosts})/models/(?P<id>[^/?#\s]+/[^/?#\s]+)"
|
||||||
|
return re.compile(body), re.compile(rf"{body}/?{_VIEW_SEGMENTS}/?$")
|
||||||
|
|
||||||
|
|
||||||
|
#: ``master`` is ModelScope's default branch; ``main`` is tried as a fallback
|
||||||
|
#: for repos imported from Hugging Face.
|
||||||
|
_REVISIONS = ("master", "main")
|
||||||
|
|
||||||
|
|
||||||
|
class ModelScopeSource(ModelSource):
|
||||||
|
"""ModelScope's mainland site (``modelscope.cn``).
|
||||||
|
|
||||||
|
``modelscope.com`` is accepted as an alias of it. The international
|
||||||
|
deployment is :class:`ModelScopeIntlSource`; everything below is written in
|
||||||
|
terms of ``base_url`` so both share one implementation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
platform = "modelscope"
|
||||||
|
label = "ModelScope"
|
||||||
|
supports_enrichment = True
|
||||||
|
supports_download = True
|
||||||
|
default_revision = "master"
|
||||||
|
default_subdir = "modelscope"
|
||||||
|
|
||||||
|
#: Origin every outgoing URL is built from.
|
||||||
|
base_url = "https://modelscope.cn"
|
||||||
|
|
||||||
|
url_pattern, strict_url_pattern = _url_patterns(_MAINLAND_HOSTS)
|
||||||
|
|
||||||
|
def canonical_url(self, source_id: str) -> str:
|
||||||
|
return f"{self.base_url}/models/{source_id}"
|
||||||
|
|
||||||
|
def asset_base_url(self, source_id: str, revision: str = "") -> str:
|
||||||
|
return (
|
||||||
|
f"{self.base_url}/models/{source_id}/resolve/"
|
||||||
|
f"{self.resolve_revision(revision)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
async def fetch_model_card(self, source_id: str) -> str:
|
||||||
|
"""Fetch the model card, preferring the raw resolve URL."""
|
||||||
|
|
||||||
|
for revision in _REVISIONS:
|
||||||
|
text = await fetch_text(
|
||||||
|
f"{self.base_url}/models/{source_id}/resolve/{revision}/README.md"
|
||||||
|
)
|
||||||
|
if text:
|
||||||
|
return text
|
||||||
|
|
||||||
|
# Fallback: the repo API proxies the same file and is reachable in
|
||||||
|
# environments where the CDN resolve host is blocked.
|
||||||
|
for revision in _REVISIONS:
|
||||||
|
text = await fetch_text(
|
||||||
|
f"{self.base_url}/api/v1/models/"
|
||||||
|
f"{source_id}/repo?Revision={revision}&FilePath=README.md"
|
||||||
|
)
|
||||||
|
if text:
|
||||||
|
return text
|
||||||
|
return ""
|
||||||
|
|
||||||
|
async def fetch_model_card_context(
|
||||||
|
self,
|
||||||
|
source_id: str,
|
||||||
|
filename: str = "",
|
||||||
|
*,
|
||||||
|
sha256: str = "",
|
||||||
|
cache: Optional["ModelSourceCache"] = None,
|
||||||
|
) -> ModelCardContext:
|
||||||
|
"""Read the model-detail API that backs the ModelScope model page.
|
||||||
|
|
||||||
|
ModelScope splits a model card in two: ``README.md`` holds the
|
||||||
|
long-form content, while the author's summary, the site-curated tags,
|
||||||
|
and the per-file example images live only here. AIGC repositories
|
||||||
|
frequently ship an auto-generated README ("the contributor provided
|
||||||
|
no further description") and put everything useful in ``Description``,
|
||||||
|
so enrichment that reads only the README comes back nearly empty.
|
||||||
|
|
||||||
|
The wanted file is identified by its sha256 when the caller knows it
|
||||||
|
and by *filename* otherwise; see :func:`_matching_versions`. The
|
||||||
|
images and trigger words returned belong to that exact
|
||||||
|
``.safetensors`` — essential for collection repositories, where every
|
||||||
|
checkpoint has its own sample image.
|
||||||
|
|
||||||
|
The detail payload describes the whole repository and is therefore
|
||||||
|
shared across every file in it, so it is read through *cache* when the
|
||||||
|
caller supplies one; only the per-file selection is redone.
|
||||||
|
"""
|
||||||
|
|
||||||
|
data = await self._fetch_detail(source_id, cache=cache)
|
||||||
|
if data is None:
|
||||||
|
return ModelCardContext()
|
||||||
|
return _build_card_context(data, filename, sha256)
|
||||||
|
|
||||||
|
async def _fetch_detail(
|
||||||
|
self,
|
||||||
|
source_id: str,
|
||||||
|
*,
|
||||||
|
cache: Optional["ModelSourceCache"] = None,
|
||||||
|
) -> Optional[dict[str, Any]]:
|
||||||
|
"""Fetch (or reuse) the model-detail payload for *source_id*."""
|
||||||
|
|
||||||
|
cache_key = (self.platform, "detail", source_id)
|
||||||
|
if cache is not None and cache_key in cache.provider:
|
||||||
|
return cache.provider[cache_key]
|
||||||
|
|
||||||
|
status, payload = await fetch_json(
|
||||||
|
f"{self.base_url}/api/v1/models/{source_id}"
|
||||||
|
)
|
||||||
|
if status != 200 or not isinstance(payload, dict):
|
||||||
|
logger.debug(
|
||||||
|
"ModelScope detail API returned HTTP %s for %s", status, source_id
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
data = payload.get("Data")
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
return None
|
||||||
|
|
||||||
|
if cache is not None:
|
||||||
|
cache.provider[cache_key] = data
|
||||||
|
return data
|
||||||
|
|
||||||
|
async def list_files(
|
||||||
|
self, source_id: str, revision: str = ""
|
||||||
|
) -> list[dict]:
|
||||||
|
"""List weight files via the repo files API.
|
||||||
|
|
||||||
|
``master`` is the only branch name the API accepts — even repos
|
||||||
|
imported from Hugging Face are addressed as ``master`` (``main``
|
||||||
|
returns 404) — so no fallback probing is done here.
|
||||||
|
"""
|
||||||
|
|
||||||
|
revision = self.resolve_revision(revision)
|
||||||
|
status, payload = await fetch_json(
|
||||||
|
f"{self.base_url}/api/v1/models/"
|
||||||
|
f"{source_id}/repo/files?Revision={revision}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if status == 404:
|
||||||
|
raise ModelSourceError(f"Repository '{source_id}' not found", status=404)
|
||||||
|
if status != 200 or not isinstance(payload, dict):
|
||||||
|
raise ModelSourceError(
|
||||||
|
f"ModelScope API error while listing '{source_id}' (HTTP {status})"
|
||||||
|
)
|
||||||
|
|
||||||
|
entries = []
|
||||||
|
for entry in (payload.get("Data") or {}).get("Files") or []:
|
||||||
|
if not isinstance(entry, dict) or entry.get("Type") != "blob":
|
||||||
|
continue
|
||||||
|
entries.append((entry.get("Path", ""), entry.get("Size", 0) or 0))
|
||||||
|
|
||||||
|
return filter_weight_files(entries)
|
||||||
|
|
||||||
|
def file_download_url(
|
||||||
|
self, source_id: str, filename: str, revision: str = ""
|
||||||
|
) -> str:
|
||||||
|
return (
|
||||||
|
f"{self.base_url}/models/{source_id}/resolve/"
|
||||||
|
f"{self.resolve_revision(revision)}/{filename}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def page_url_for_file(self, source_id: str, filename: str) -> str:
|
||||||
|
return (
|
||||||
|
f"{self.base_url}/models/{source_id}/file/view/"
|
||||||
|
f"{self.default_revision}/{filename}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ModelScopeIntlSource(ModelScopeSource):
|
||||||
|
"""ModelScope's international site (``modelscope.ai``).
|
||||||
|
|
||||||
|
A separate catalogue rather than a mirror, so it is registered under its
|
||||||
|
own platform id: the two deployments must not share a version group, a
|
||||||
|
"use default paths" directory, or a stored ``source_url``. The detail API,
|
||||||
|
the file listing, the resolve URLs and the CDN redirect all behave exactly
|
||||||
|
like the mainland site, which is why every URL here is derived from
|
||||||
|
:attr:`base_url` instead of being duplicated.
|
||||||
|
"""
|
||||||
|
|
||||||
|
platform = "modelscope-ai"
|
||||||
|
label = "ModelScope (International)"
|
||||||
|
default_subdir = "modelscope-ai"
|
||||||
|
base_url = "https://www.modelscope.ai"
|
||||||
|
|
||||||
|
url_pattern, strict_url_pattern = _url_patterns(_INTERNATIONAL_HOSTS)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["ModelScopeIntlSource", "ModelScopeSource"]
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Model-detail API parsing helpers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
#: Trigger-word values that mean "the author left this blank".
|
||||||
|
_EMPTY_TRIGGER_VALUES = frozenset({"none", "null", "n/a"})
|
||||||
|
|
||||||
|
#: Repository tags that only restate what the model *is* (its library, task or
|
||||||
|
#: framework) rather than what it depicts. ModelScope mixes both into the
|
||||||
|
#: plain ``Tags`` list, and a card tagged "lora" or "text-to-image" is noise.
|
||||||
|
_GENERIC_TAGS = frozenset(
|
||||||
|
{
|
||||||
|
"any-to-any",
|
||||||
|
"checkpoint",
|
||||||
|
"controlnet",
|
||||||
|
"diffusers",
|
||||||
|
"embedding",
|
||||||
|
"image-text-to-text",
|
||||||
|
"image-to-image",
|
||||||
|
"image-to-video",
|
||||||
|
"lora",
|
||||||
|
"lycoris",
|
||||||
|
"onnx",
|
||||||
|
"pytorch",
|
||||||
|
"safetensors",
|
||||||
|
"tensorflow",
|
||||||
|
"text-to-image",
|
||||||
|
"text-to-speech",
|
||||||
|
"text-to-video",
|
||||||
|
"textual-inversion",
|
||||||
|
"vae",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _clean_text(value: Any) -> str:
|
||||||
|
"""Return a stripped string for *value*, or ``""`` for anything else."""
|
||||||
|
|
||||||
|
return value.strip() if isinstance(value, str) else ""
|
||||||
|
|
||||||
|
|
||||||
|
def _first_string(value: Any) -> str:
|
||||||
|
"""Return the first non-empty string in a list, or ``""``."""
|
||||||
|
|
||||||
|
if isinstance(value, list):
|
||||||
|
for item in value:
|
||||||
|
text = _clean_text(item)
|
||||||
|
if text:
|
||||||
|
return text
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _build_card_context(
|
||||||
|
data: dict[str, Any], filename: str, sha256: str = ""
|
||||||
|
) -> ModelCardContext:
|
||||||
|
"""Turn a model-detail payload into a :class:`ModelCardContext`.
|
||||||
|
|
||||||
|
Separated from the HTTP fetch so the repository-wide payload can be cached
|
||||||
|
across the files of a collection repository while the per-file selection
|
||||||
|
is still redone for each one.
|
||||||
|
"""
|
||||||
|
|
||||||
|
context = ModelCardContext(
|
||||||
|
description=_clean_text(data.get("Description")),
|
||||||
|
model_name=_clean_text(data.get("Name")),
|
||||||
|
model_name_localized=_clean_text(data.get("ChineseName")),
|
||||||
|
license=_clean_text(data.get("License")),
|
||||||
|
model_type=_clean_text(data.get("AigcType")),
|
||||||
|
base_model=_first_string(data.get("BaseModel")),
|
||||||
|
base_model_aliases=_base_model_aliases(data),
|
||||||
|
official_tags=_official_tags(data),
|
||||||
|
)
|
||||||
|
|
||||||
|
versions = _matching_versions(
|
||||||
|
data.get("MuseInfo"),
|
||||||
|
filename,
|
||||||
|
digests=_file_digests(data),
|
||||||
|
sha256=sha256,
|
||||||
|
)
|
||||||
|
if versions:
|
||||||
|
context.version_name = _version_label(versions)
|
||||||
|
context.example_images = _cover_image_urls(versions)
|
||||||
|
context.trigger_words = _version_trigger_words(versions)
|
||||||
|
return context
|
||||||
|
|
||||||
|
|
||||||
|
def _base_model_aliases(data: dict[str, Any]) -> list[str]:
|
||||||
|
"""Return the site's own names for the base model.
|
||||||
|
|
||||||
|
ModelScope publishes a link-style id (``krea/Krea-2-Turbo``) plus its
|
||||||
|
internal architecture enums (``VisionFoundation: KREA_2``,
|
||||||
|
``SubVisionFoundation: KREA_2_TURBO``). The enums are the better
|
||||||
|
resolution hint because they normalise onto this system's canonical
|
||||||
|
vocabulary, so they come first; the owner prefix is also stripped from
|
||||||
|
the link-style ids.
|
||||||
|
"""
|
||||||
|
|
||||||
|
aliases: list[str] = []
|
||||||
|
for key in ("VisionFoundation", "SubVisionFoundation"):
|
||||||
|
value = _clean_text(data.get(key))
|
||||||
|
if value and value not in aliases:
|
||||||
|
aliases.append(value)
|
||||||
|
|
||||||
|
base_models = data.get("BaseModel")
|
||||||
|
if isinstance(base_models, list):
|
||||||
|
for item in base_models:
|
||||||
|
text = _clean_text(item)
|
||||||
|
leaf = text.rsplit("/", 1)[-1] if text else ""
|
||||||
|
if leaf and leaf not in aliases:
|
||||||
|
aliases.append(leaf)
|
||||||
|
return aliases
|
||||||
|
|
||||||
|
|
||||||
|
def _official_tags(data: dict[str, Any]) -> list[str]:
|
||||||
|
"""Return the content tags the site publishes for the repository.
|
||||||
|
|
||||||
|
``OfficialTags`` is ModelScope's curated content vocabulary and is
|
||||||
|
preferred whenever it is populated. Plenty of AIGC repositories leave it
|
||||||
|
empty and carry only the plain ``Tags`` list, which mixes content tags with
|
||||||
|
framework and task categories; those categories are dropped so a card is
|
||||||
|
not handed "lora" and "text-to-image" as if they described the model.
|
||||||
|
"""
|
||||||
|
|
||||||
|
curated = _dedupe(_tag_values(data.get("OfficialTags")))
|
||||||
|
if curated:
|
||||||
|
return curated
|
||||||
|
|
||||||
|
generic = set(_GENERIC_TAGS)
|
||||||
|
for value in (
|
||||||
|
data.get("AigcType"),
|
||||||
|
data.get("Libraries"),
|
||||||
|
data.get("Frameworks"),
|
||||||
|
):
|
||||||
|
for item in value if isinstance(value, list) else [value]:
|
||||||
|
text = _clean_text(item).lower()
|
||||||
|
if text:
|
||||||
|
generic.add(text)
|
||||||
|
|
||||||
|
return _dedupe(
|
||||||
|
tag for tag in _tag_values(data.get("Tags")) if tag.lower() not in generic
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _tag_values(value: Any) -> list[str]:
|
||||||
|
"""Return the tag strings from either shape ModelScope publishes.
|
||||||
|
|
||||||
|
``OfficialTags`` is a list of ``{"Tag": ..., "ChineseName": ...}`` dicts
|
||||||
|
carrying an English value; the plain ``Tags`` list is already strings.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not isinstance(value, list):
|
||||||
|
return []
|
||||||
|
tags: list[str] = []
|
||||||
|
for entry in value:
|
||||||
|
tag = _clean_text(entry.get("Tag") if isinstance(entry, dict) else entry)
|
||||||
|
if tag:
|
||||||
|
tags.append(tag)
|
||||||
|
return tags
|
||||||
|
|
||||||
|
|
||||||
|
def _dedupe(values: Iterable[str]) -> list[str]:
|
||||||
|
"""Drop empties and repeats, keeping the first spelling seen."""
|
||||||
|
|
||||||
|
unique: list[str] = []
|
||||||
|
for value in values:
|
||||||
|
if value and value not in unique:
|
||||||
|
unique.append(value)
|
||||||
|
return unique
|
||||||
|
|
||||||
|
|
||||||
|
def _version_files(version: dict[str, Any]) -> list[str]:
|
||||||
|
"""Return the model filenames covered by one ``MuseInfo.versions`` entry.
|
||||||
|
|
||||||
|
The listing normally sits in ``stats.fileList``; some payloads only
|
||||||
|
carry the same field as a JSON-encoded string under
|
||||||
|
``modelVersion.stats``, so both shapes are accepted.
|
||||||
|
"""
|
||||||
|
|
||||||
|
stats = version.get("stats")
|
||||||
|
files = stats.get("fileList") if isinstance(stats, dict) else None
|
||||||
|
|
||||||
|
if not isinstance(files, list):
|
||||||
|
model_version = version.get("modelVersion")
|
||||||
|
raw = model_version.get("stats") if isinstance(model_version, dict) else None
|
||||||
|
if isinstance(raw, str) and raw.strip():
|
||||||
|
try:
|
||||||
|
decoded = json.loads(raw)
|
||||||
|
except (json.JSONDecodeError, TypeError):
|
||||||
|
decoded = None
|
||||||
|
if isinstance(decoded, dict):
|
||||||
|
files = decoded.get("fileList")
|
||||||
|
|
||||||
|
if not isinstance(files, list):
|
||||||
|
return []
|
||||||
|
return [item for item in files if isinstance(item, str) and item]
|
||||||
|
|
||||||
|
|
||||||
|
def _version_show_name(version: dict[str, Any]) -> str:
|
||||||
|
"""Return the human-facing version label (e.g. ``c1-st1000``)."""
|
||||||
|
|
||||||
|
model_version = version.get("modelVersion")
|
||||||
|
if not isinstance(model_version, dict):
|
||||||
|
return ""
|
||||||
|
return _clean_text(model_version.get("showName")).lower()
|
||||||
|
|
||||||
|
|
||||||
|
def _version_label(versions: list[dict[str, Any]]) -> str:
|
||||||
|
"""Return the first published version label, preserving its spelling.
|
||||||
|
|
||||||
|
Unlike :func:`_version_show_name` this is for display, so the label is
|
||||||
|
not lowercased.
|
||||||
|
"""
|
||||||
|
|
||||||
|
for version in versions:
|
||||||
|
model_version = version.get("modelVersion")
|
||||||
|
if not isinstance(model_version, dict):
|
||||||
|
continue
|
||||||
|
label = _clean_text(model_version.get("showName"))
|
||||||
|
if label:
|
||||||
|
return label
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def _file_digests(data: dict[str, Any]) -> dict[str, str]:
|
||||||
|
"""Return ``basename -> sha256`` for every published weight file.
|
||||||
|
|
||||||
|
``ModelInfos`` groups the repository's files by kind (``safetensor``,
|
||||||
|
…) and records a real sha256 for each, which is what makes it possible to
|
||||||
|
recognise a file the user has renamed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
digests: dict[str, str] = {}
|
||||||
|
model_infos = data.get("ModelInfos")
|
||||||
|
if not isinstance(model_infos, dict):
|
||||||
|
return digests
|
||||||
|
for info in model_infos.values():
|
||||||
|
files = info.get("files") if isinstance(info, dict) else None
|
||||||
|
if not isinstance(files, list):
|
||||||
|
continue
|
||||||
|
for entry in files:
|
||||||
|
if not isinstance(entry, dict):
|
||||||
|
continue
|
||||||
|
name = _clean_text(entry.get("name"))
|
||||||
|
digest = _clean_text(entry.get("sha256"))
|
||||||
|
if name and digest:
|
||||||
|
digests.setdefault(os.path.basename(name).lower(), digest.lower())
|
||||||
|
return digests
|
||||||
|
|
||||||
|
|
||||||
|
def _matching_versions(
|
||||||
|
muse_info: Any,
|
||||||
|
filename: str,
|
||||||
|
*,
|
||||||
|
digests: dict[str, str] | None = None,
|
||||||
|
sha256: str = "",
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Return the ``versions`` entries that publish the wanted model file.
|
||||||
|
|
||||||
|
Strategies, in order:
|
||||||
|
|
||||||
|
1. **sha256** — the file's content hash, looked up through
|
||||||
|
:func:`_file_digests`. This is the only strategy that survives the
|
||||||
|
user renaming the weights, which is common once a model is filed away.
|
||||||
|
2. **Exact basename** against each version's ``stats.fileList``.
|
||||||
|
3. **``showName`` inside the file stem**, which absorbs the naming drift
|
||||||
|
ModelScope sometimes applies to uploaded weights.
|
||||||
|
|
||||||
|
A known-but-unmatched hash falls through to the filename strategies
|
||||||
|
rather than giving up, in case the local file was re-encoded. All matches
|
||||||
|
are returned so a file re-published across several versions contributes
|
||||||
|
all of its example images. With no *filename* and no *sha256*, only an
|
||||||
|
unambiguous single-version repository is used, because a per-file image
|
||||||
|
must never be attributed to the wrong file.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not isinstance(muse_info, dict):
|
||||||
|
return []
|
||||||
|
versions = muse_info.get("versions")
|
||||||
|
if not isinstance(versions, list):
|
||||||
|
return []
|
||||||
|
entries = [entry for entry in versions if isinstance(entry, dict)]
|
||||||
|
if not entries:
|
||||||
|
return []
|
||||||
|
|
||||||
|
target_hash = (sha256 or "").strip().lower()
|
||||||
|
if target_hash:
|
||||||
|
known = digests or {}
|
||||||
|
by_hash: list[dict[str, Any]] = []
|
||||||
|
for version in entries:
|
||||||
|
for path in _version_files(version):
|
||||||
|
if known.get(os.path.basename(path).lower()) == target_hash:
|
||||||
|
by_hash.append(version)
|
||||||
|
break
|
||||||
|
if by_hash:
|
||||||
|
return by_hash
|
||||||
|
|
||||||
|
if not filename:
|
||||||
|
return entries if len(entries) == 1 else []
|
||||||
|
|
||||||
|
target = os.path.basename(filename).strip().lower()
|
||||||
|
if not target:
|
||||||
|
return []
|
||||||
|
stem = os.path.splitext(target)[0]
|
||||||
|
|
||||||
|
exact: list[dict[str, Any]] = []
|
||||||
|
fuzzy: list[dict[str, Any]] = []
|
||||||
|
for version in entries:
|
||||||
|
files = {os.path.basename(path).lower() for path in _version_files(version)}
|
||||||
|
if target in files:
|
||||||
|
exact.append(version)
|
||||||
|
continue
|
||||||
|
show_name = _version_show_name(version)
|
||||||
|
if show_name and show_name in stem:
|
||||||
|
fuzzy.append(version)
|
||||||
|
|
||||||
|
return exact or fuzzy
|
||||||
|
|
||||||
|
|
||||||
|
def _cover_image_urls(versions: list[dict[str, Any]]) -> list[str]:
|
||||||
|
"""Collect the example-image URLs published by the given versions."""
|
||||||
|
|
||||||
|
urls: list[str] = []
|
||||||
|
for version in versions:
|
||||||
|
covers = version.get("coverImages")
|
||||||
|
if not isinstance(covers, list):
|
||||||
|
continue
|
||||||
|
for cover in covers:
|
||||||
|
if not isinstance(cover, dict):
|
||||||
|
continue
|
||||||
|
url = _clean_text(cover.get("url"))
|
||||||
|
if url and url not in urls:
|
||||||
|
urls.append(url)
|
||||||
|
return urls
|
||||||
|
|
||||||
|
|
||||||
|
def _version_trigger_words(versions: list[dict[str, Any]]) -> list[str]:
|
||||||
|
"""Return the first non-empty trigger-word list across *versions*."""
|
||||||
|
|
||||||
|
for version in versions:
|
||||||
|
model_version = version.get("modelVersion")
|
||||||
|
raw = (
|
||||||
|
model_version.get("triggerWords")
|
||||||
|
if isinstance(model_version, dict)
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
words = _parse_trigger_words(raw)
|
||||||
|
if words:
|
||||||
|
return words
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_trigger_words(raw: Any) -> list[str]:
|
||||||
|
"""Decode ModelScope's JSON-encoded trigger-word string list."""
|
||||||
|
|
||||||
|
if isinstance(raw, list):
|
||||||
|
candidates = raw
|
||||||
|
elif isinstance(raw, str) and raw.strip():
|
||||||
|
try:
|
||||||
|
decoded = json.loads(raw)
|
||||||
|
except (json.JSONDecodeError, TypeError):
|
||||||
|
return []
|
||||||
|
if not isinstance(decoded, list):
|
||||||
|
return []
|
||||||
|
candidates = decoded
|
||||||
|
else:
|
||||||
|
return []
|
||||||
|
|
||||||
|
words: list[str] = []
|
||||||
|
for item in candidates:
|
||||||
|
word = _clean_text(item)
|
||||||
|
if not word or word.lower() in _EMPTY_TRIGGER_VALUES:
|
||||||
|
continue
|
||||||
|
if word not in words:
|
||||||
|
words.append(word)
|
||||||
|
return words
|
||||||
@@ -0,0 +1,228 @@
|
|||||||
|
"""Registry and metadata helpers for external model sources.
|
||||||
|
|
||||||
|
The registry is the single place the rest of the codebase asks "which site
|
||||||
|
is this URL from?", "what is this model's source?", and "can we enrich it?".
|
||||||
|
Import from :mod:`py.services.model_sources` rather than this module
|
||||||
|
directly.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict, Mapping, Optional
|
||||||
|
|
||||||
|
from .base import GROUP_PREFIXES, ModelSource, SourceRef, clean_source_url
|
||||||
|
from .huggingface import HuggingFaceSource
|
||||||
|
from .modelscope import ModelScopeIntlSource, ModelScopeSource
|
||||||
|
from .tensorart import TensorArtSource
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
#: Order matters only for disambiguation; the URL patterns are disjoint.
|
||||||
|
#: ``modelscope.ai`` is a separate catalogue from ``modelscope.cn`` rather than
|
||||||
|
#: an alias, which is why it gets its own entry (see ``modelscope.py``).
|
||||||
|
_SOURCES: tuple[ModelSource, ...] = (
|
||||||
|
HuggingFaceSource(),
|
||||||
|
ModelScopeSource(),
|
||||||
|
ModelScopeIntlSource(),
|
||||||
|
TensorArtSource(),
|
||||||
|
)
|
||||||
|
|
||||||
|
_BY_PLATFORM: Dict[str, ModelSource] = {s.platform: s for s in _SOURCES}
|
||||||
|
|
||||||
|
#: Metadata keys that carry the canonical external-source identity.
|
||||||
|
SOURCE_PLATFORM_FIELD = "source_platform"
|
||||||
|
SOURCE_URL_FIELD = "source_url"
|
||||||
|
#: Legacy field kept as a read/write alias for Hugging Face models so that
|
||||||
|
#: older sidecars, cached rows, and third-party consumers keep working.
|
||||||
|
LEGACY_HF_URL_FIELD = "hf_url"
|
||||||
|
|
||||||
|
|
||||||
|
def list_sources() -> list[ModelSource]:
|
||||||
|
"""Return every known model source."""
|
||||||
|
|
||||||
|
return list(_SOURCES)
|
||||||
|
|
||||||
|
|
||||||
|
def get_source(platform: Optional[str]) -> Optional[ModelSource]:
|
||||||
|
"""Return the source registered for *platform*, or ``None``."""
|
||||||
|
|
||||||
|
if not platform or not isinstance(platform, str):
|
||||||
|
return None
|
||||||
|
return _BY_PLATFORM.get(platform.strip().lower())
|
||||||
|
|
||||||
|
|
||||||
|
def source_label(platform: Optional[str], default: str = "") -> str:
|
||||||
|
"""Return the human-readable label for *platform*."""
|
||||||
|
|
||||||
|
source = get_source(platform)
|
||||||
|
return source.label if source else default
|
||||||
|
|
||||||
|
|
||||||
|
def downloadable_sources() -> list[ModelSource]:
|
||||||
|
"""Return the sources whose repositories can be downloaded directly."""
|
||||||
|
|
||||||
|
return [source for source in _SOURCES if source.supports_download]
|
||||||
|
|
||||||
|
|
||||||
|
def get_download_source(platform: Optional[str]) -> Optional[ModelSource]:
|
||||||
|
"""Return the source for *platform*, but only when it supports downloads."""
|
||||||
|
|
||||||
|
source = get_source(platform)
|
||||||
|
if source is None or not source.supports_download:
|
||||||
|
return None
|
||||||
|
return source
|
||||||
|
|
||||||
|
|
||||||
|
def detect_source(url: Optional[str], *, strict: bool = False) -> Optional[SourceRef]:
|
||||||
|
"""Return the :class:`SourceRef` for *url*, or ``None`` if unsupported."""
|
||||||
|
|
||||||
|
if not url or not isinstance(url, str):
|
||||||
|
return None
|
||||||
|
for source in _SOURCES:
|
||||||
|
ref = source.ref(url, strict=strict)
|
||||||
|
if ref is not None:
|
||||||
|
return ref
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_source_ref(metadata: Mapping[str, Any]) -> Optional[SourceRef]:
|
||||||
|
"""Return the source reference described by a model's metadata.
|
||||||
|
|
||||||
|
Handles all three storage states found in the wild:
|
||||||
|
|
||||||
|
1. ``source_url`` + ``source_platform`` (current format)
|
||||||
|
2. ``hf_url`` only (legacy Hugging Face storage)
|
||||||
|
3. ``hf_url`` plus a newer ``source_url`` (both written by older builds)
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not isinstance(metadata, Mapping):
|
||||||
|
return None
|
||||||
|
|
||||||
|
platform = clean_source_url(metadata.get(SOURCE_PLATFORM_FIELD)).lower()
|
||||||
|
url = clean_source_url(metadata.get(SOURCE_URL_FIELD))
|
||||||
|
legacy = clean_source_url(metadata.get(LEGACY_HF_URL_FIELD))
|
||||||
|
|
||||||
|
source = get_source(platform)
|
||||||
|
if url:
|
||||||
|
if source is not None:
|
||||||
|
ref = source.ref(url)
|
||||||
|
if ref is not None:
|
||||||
|
return ref
|
||||||
|
ref = detect_source(url)
|
||||||
|
if ref is not None:
|
||||||
|
return ref
|
||||||
|
# Unknown platform but a URL is present: keep it addressable.
|
||||||
|
return SourceRef(platform=platform or "unknown", source_id="", url=url)
|
||||||
|
|
||||||
|
if legacy:
|
||||||
|
return detect_source(legacy)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_metadata_source(metadata: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Normalise the external-source fields on *metadata* in place.
|
||||||
|
|
||||||
|
Guarantees that ``source_url``/``source_platform`` are present and
|
||||||
|
consistent, and that ``hf_url`` mirrors ``source_url`` for Hugging Face
|
||||||
|
models (never for other platforms, so a stale alias can't make a
|
||||||
|
ModelScope model look like a Hugging Face one).
|
||||||
|
|
||||||
|
Returns the same dict for convenient chaining.
|
||||||
|
"""
|
||||||
|
|
||||||
|
if not isinstance(metadata, dict):
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
platform = clean_source_url(metadata.get(SOURCE_PLATFORM_FIELD)).lower()
|
||||||
|
url = clean_source_url(metadata.get(SOURCE_URL_FIELD))
|
||||||
|
legacy = clean_source_url(metadata.get(LEGACY_HF_URL_FIELD))
|
||||||
|
|
||||||
|
source = get_source(platform)
|
||||||
|
ref: Optional[SourceRef] = None
|
||||||
|
|
||||||
|
if url:
|
||||||
|
ref = source.ref(url) if source is not None else None
|
||||||
|
if ref is None:
|
||||||
|
ref = detect_source(url)
|
||||||
|
elif legacy:
|
||||||
|
ref = detect_source(legacy)
|
||||||
|
|
||||||
|
if ref is not None and ref.source_id:
|
||||||
|
platform = ref.platform
|
||||||
|
url = ref.url or url
|
||||||
|
|
||||||
|
if platform:
|
||||||
|
metadata[SOURCE_PLATFORM_FIELD] = platform
|
||||||
|
else:
|
||||||
|
metadata.setdefault(SOURCE_PLATFORM_FIELD, "")
|
||||||
|
|
||||||
|
metadata[SOURCE_URL_FIELD] = url
|
||||||
|
|
||||||
|
# Keep the legacy alias in sync, but only for Hugging Face.
|
||||||
|
if url and platform == "huggingface":
|
||||||
|
metadata[LEGACY_HF_URL_FIELD] = url
|
||||||
|
elif LEGACY_HF_URL_FIELD in metadata and platform and platform != "huggingface":
|
||||||
|
metadata[LEGACY_HF_URL_FIELD] = ""
|
||||||
|
elif legacy and not url:
|
||||||
|
metadata[LEGACY_HF_URL_FIELD] = legacy
|
||||||
|
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
|
||||||
|
def has_external_source(item: Mapping[str, Any]) -> bool:
|
||||||
|
"""Return ``True`` when *item* is linked to any external model site."""
|
||||||
|
|
||||||
|
if not isinstance(item, Mapping):
|
||||||
|
return False
|
||||||
|
return bool(
|
||||||
|
clean_source_url(item.get(SOURCE_URL_FIELD))
|
||||||
|
or clean_source_url(item.get(LEGACY_HF_URL_FIELD))
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_source_platform(item: Mapping[str, Any]) -> str:
|
||||||
|
"""Return the platform id stored on *item* (may be empty)."""
|
||||||
|
|
||||||
|
if not isinstance(item, Mapping):
|
||||||
|
return ""
|
||||||
|
platform = clean_source_url(item.get(SOURCE_PLATFORM_FIELD)).lower()
|
||||||
|
if platform:
|
||||||
|
return platform
|
||||||
|
ref = resolve_source_ref(item)
|
||||||
|
return ref.platform if ref else ""
|
||||||
|
|
||||||
|
|
||||||
|
def source_group_key(item: Mapping[str, Any]) -> Optional[str]:
|
||||||
|
"""Return the version-group key for *item*, or ``None``.
|
||||||
|
|
||||||
|
Hugging Face keeps the historical ``hf:{owner}/{repo}`` shape; other
|
||||||
|
platforms use their own short prefix (see :data:`GROUP_PREFIXES`).
|
||||||
|
"""
|
||||||
|
|
||||||
|
ref = resolve_source_ref(item)
|
||||||
|
if ref is None or not ref.source_id:
|
||||||
|
return None
|
||||||
|
source = get_source(ref.platform)
|
||||||
|
if source is None:
|
||||||
|
return None
|
||||||
|
return source.group_key(ref.source_id)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"GROUP_PREFIXES",
|
||||||
|
"LEGACY_HF_URL_FIELD",
|
||||||
|
"SOURCE_PLATFORM_FIELD",
|
||||||
|
"SOURCE_URL_FIELD",
|
||||||
|
"detect_source",
|
||||||
|
"downloadable_sources",
|
||||||
|
"get_download_source",
|
||||||
|
"get_source",
|
||||||
|
"get_source_platform",
|
||||||
|
"has_external_source",
|
||||||
|
"list_sources",
|
||||||
|
"normalize_metadata_source",
|
||||||
|
"resolve_source_ref",
|
||||||
|
"source_group_key",
|
||||||
|
"source_label",
|
||||||
|
]
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
"""TensorArt model source (link / provenance only).
|
||||||
|
|
||||||
|
TensorArt support is intentionally limited to *linking* a model to its
|
||||||
|
TensorArt page. Automatic metadata extraction is not possible without a
|
||||||
|
user session:
|
||||||
|
|
||||||
|
* ``tensor.art`` sits behind a Cloudflare managed challenge, so plain
|
||||||
|
HTTP clients (aiohttp, requests, curl) receive ``403 "Just a moment..."``.
|
||||||
|
* Its internal API (``ap-east-1.tensorart.cloud`` / ``cn.tensorart.net``)
|
||||||
|
answers every ``/v1/model/*`` route with
|
||||||
|
``{"code":100002,"message":"invalid authorization header"}``.
|
||||||
|
* The official TAMS API requires an AccessKey/SecretKey pair and request
|
||||||
|
signatures, which is a poor fit for a "paste a URL" workflow.
|
||||||
|
|
||||||
|
``supports_enrichment`` is therefore ``False``: the agent pipeline skips
|
||||||
|
these models with an explicit reason instead of failing silently, and the
|
||||||
|
UI keeps showing the "View on TensorArt" link. ``tusi.cn`` is TensorArt's
|
||||||
|
Chinese mirror and is accepted as the same platform.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import re
|
||||||
|
|
||||||
|
from .base import ModelSource
|
||||||
|
|
||||||
|
_DOMAINS = r"(?:tensor\.art|tusi\.cn)"
|
||||||
|
|
||||||
|
_URL_PATTERN = re.compile(
|
||||||
|
rf"https?://(?:www\.)?{_DOMAINS}/models/(?P<id>\d+)"
|
||||||
|
)
|
||||||
|
|
||||||
|
_STRICT_URL_PATTERN = re.compile(
|
||||||
|
rf"https?://(?:www\.)?{_DOMAINS}/models/(?P<id>\d+)(?:/[^/?#\s]+)?/?$"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TensorArtSource(ModelSource):
|
||||||
|
"""TensorArt (``tensor.art``)."""
|
||||||
|
|
||||||
|
platform = "tensorart"
|
||||||
|
label = "TensorArt"
|
||||||
|
supports_enrichment = False
|
||||||
|
supports_download = False
|
||||||
|
url_pattern = _URL_PATTERN
|
||||||
|
strict_url_pattern = _STRICT_URL_PATTERN
|
||||||
|
|
||||||
|
def canonical_url(self, source_id: str) -> str:
|
||||||
|
return f"https://tensor.art/models/{source_id}"
|
||||||
|
|
||||||
|
def asset_base_url(self, source_id: str, revision: str = "") -> str:
|
||||||
|
# Unreachable today: enrichment is disabled for this platform.
|
||||||
|
return f"https://tensor.art/models/{source_id}"
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["TensorArtSource"]
|
||||||
@@ -13,7 +13,7 @@ import sqlite3
|
|||||||
import time
|
import time
|
||||||
from dataclasses import dataclass, replace
|
from dataclasses import dataclass, replace
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Any, Dict, Iterable, List, Mapping, Optional, Sequence
|
from typing import Any, Dict, Iterable, Iterator, List, Mapping, Optional, Sequence
|
||||||
|
|
||||||
from .errors import RateLimitError, ResourceNotFoundError
|
from .errors import RateLimitError, ResourceNotFoundError
|
||||||
from .settings_manager import get_settings_manager
|
from .settings_manager import get_settings_manager
|
||||||
@@ -250,6 +250,51 @@ class ModelUpdateRecord:
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def has_update_for_local_bases(
|
||||||
|
self,
|
||||||
|
hide_early_access: bool = False,
|
||||||
|
hide_non_downloadable: bool = True,
|
||||||
|
hide_paid: bool = False,
|
||||||
|
) -> bool:
|
||||||
|
"""Return True when any locally-held base model scope has an update.
|
||||||
|
|
||||||
|
Aggregates :meth:`has_update_for_base` across every distinct base model
|
||||||
|
present among in-library versions. This mirrors the per-item evaluation
|
||||||
|
performed by ``BaseModelService._annotate_update_flags`` when the
|
||||||
|
``version_grouping`` setting is ``same_base``, so callers reporting
|
||||||
|
"how many models have updates" stay aligned with what the Updates
|
||||||
|
filter displays. Use this instead of :meth:`has_update` for such
|
||||||
|
summaries; see issue #1083.
|
||||||
|
|
||||||
|
When no local base model is known (nothing held locally, or versions
|
||||||
|
never seen in any remote listing), falls back to :meth:`has_update` so
|
||||||
|
a model the item-level filter may still flag is not silently dropped
|
||||||
|
from summaries.
|
||||||
|
"""
|
||||||
|
|
||||||
|
bases = {
|
||||||
|
_normalize_base_model(version.base_model)
|
||||||
|
for version in self.versions
|
||||||
|
if version.is_in_library
|
||||||
|
}
|
||||||
|
bases.discard(None)
|
||||||
|
if not bases:
|
||||||
|
return self.has_update(
|
||||||
|
hide_early_access=hide_early_access,
|
||||||
|
hide_non_downloadable=hide_non_downloadable,
|
||||||
|
hide_paid=hide_paid,
|
||||||
|
)
|
||||||
|
return any(
|
||||||
|
self.has_update_for_base(
|
||||||
|
None,
|
||||||
|
base,
|
||||||
|
hide_early_access=hide_early_access,
|
||||||
|
hide_non_downloadable=hide_non_downloadable,
|
||||||
|
hide_paid=hide_paid,
|
||||||
|
)
|
||||||
|
for base in bases
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ModelUpdateService:
|
class ModelUpdateService:
|
||||||
"""Persist and query remote model version metadata."""
|
"""Persist and query remote model version metadata."""
|
||||||
@@ -786,6 +831,11 @@ class ModelUpdateService:
|
|||||||
target_model_ids=target_filter,
|
target_model_ids=target_filter,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
local_base_models = await self._collect_local_version_bases(
|
||||||
|
scanner,
|
||||||
|
target_model_ids=target_filter,
|
||||||
|
)
|
||||||
|
|
||||||
results: Dict[int, ModelUpdateRecord] = {}
|
results: Dict[int, ModelUpdateRecord] = {}
|
||||||
prefetched: Dict[int, Mapping[Any, Any]] = {}
|
prefetched: Dict[int, Mapping[Any, Any]] = {}
|
||||||
|
|
||||||
@@ -838,6 +888,7 @@ class ModelUpdateService:
|
|||||||
force_refresh=force_refresh,
|
force_refresh=force_refresh,
|
||||||
prefetched_response=prefetched.get(model_id),
|
prefetched_response=prefetched.get(model_id),
|
||||||
all_local_version_ids=all_vids,
|
all_local_version_ids=all_vids,
|
||||||
|
local_base_models=local_base_models,
|
||||||
)
|
)
|
||||||
if scanner.is_cancelled():
|
if scanner.is_cancelled():
|
||||||
logger.info(f"{model_type.capitalize()} Update Service: Refresh cancelled by user")
|
logger.info(f"{model_type.capitalize()} Update Service: Refresh cancelled by user")
|
||||||
@@ -872,12 +923,14 @@ class ModelUpdateService:
|
|||||||
|
|
||||||
local_versions = await self._collect_local_versions(scanner)
|
local_versions = await self._collect_local_versions(scanner)
|
||||||
version_ids = local_versions.get(model_id, [])
|
version_ids = local_versions.get(model_id, [])
|
||||||
|
local_base_models = await self._collect_local_version_bases(scanner)
|
||||||
return await self._refresh_single_model(
|
return await self._refresh_single_model(
|
||||||
model_type,
|
model_type,
|
||||||
model_id,
|
model_id,
|
||||||
version_ids,
|
version_ids,
|
||||||
metadata_provider,
|
metadata_provider,
|
||||||
force_refresh=force_refresh,
|
force_refresh=force_refresh,
|
||||||
|
local_base_models=local_base_models,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def update_in_library_versions(
|
async def update_in_library_versions(
|
||||||
@@ -1053,6 +1106,7 @@ class ModelUpdateService:
|
|||||||
force_refresh: bool = False,
|
force_refresh: bool = False,
|
||||||
prefetched_response: Optional[Mapping[str, Any]] = None,
|
prefetched_response: Optional[Mapping[str, Any]] = None,
|
||||||
all_local_version_ids: Optional[Sequence[int]] = None,
|
all_local_version_ids: Optional[Sequence[int]] = None,
|
||||||
|
local_base_models: Optional[Mapping[int, str]] = None,
|
||||||
) -> Optional[ModelUpdateRecord]:
|
) -> Optional[ModelUpdateRecord]:
|
||||||
normalized_local = self._normalize_sequence(local_versions)
|
normalized_local = self._normalize_sequence(local_versions)
|
||||||
# When folder-filtering, this carries the cross-folder version set
|
# When folder-filtering, this carries the cross-folder version set
|
||||||
@@ -1177,6 +1231,7 @@ class ModelUpdateService:
|
|||||||
existing,
|
existing,
|
||||||
now,
|
now,
|
||||||
all_local_version_ids=normalized_all,
|
all_local_version_ids=normalized_all,
|
||||||
|
local_base_models=local_base_models,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
record = self._merge_with_local_versions(
|
record = self._merge_with_local_versions(
|
||||||
@@ -1383,27 +1438,17 @@ class ModelUpdateService:
|
|||||||
await self._enrich_version_entries(metadata_provider, aggregated)
|
await self._enrich_version_entries(metadata_provider, aggregated)
|
||||||
return aggregated
|
return aggregated
|
||||||
|
|
||||||
async def _collect_local_versions(
|
@staticmethod
|
||||||
self,
|
def _iter_local_civitai_items(
|
||||||
scanner,
|
cache,
|
||||||
*,
|
*,
|
||||||
target_model_ids: Optional[Sequence[int]] = None,
|
target_set: Optional[set[int]] = None,
|
||||||
folder_path: Optional[str] = None,
|
normalized_folder: Optional[str] = None,
|
||||||
) -> Dict[int, List[int]]:
|
) -> Iterator[tuple[int, int, Any]]:
|
||||||
cache = await scanner.get_cached_data()
|
"""Yield ``(modelId, versionId, base_model)`` for each scannable item."""
|
||||||
mapping: Dict[int, set[int]] = {}
|
|
||||||
if not cache or not getattr(cache, "raw_data", None):
|
if not cache or not getattr(cache, "raw_data", None):
|
||||||
return {}
|
return
|
||||||
|
|
||||||
target_set = None
|
|
||||||
if target_model_ids:
|
|
||||||
target_set = set(target_model_ids)
|
|
||||||
if not target_set:
|
|
||||||
return {}
|
|
||||||
|
|
||||||
normalized_folder = None
|
|
||||||
if folder_path is not None:
|
|
||||||
normalized_folder = folder_path.replace("\\", "/").strip("/")
|
|
||||||
|
|
||||||
for item in cache.raw_data:
|
for item in cache.raw_data:
|
||||||
# Apply folder filter first (cheapest check)
|
# Apply folder filter first (cheapest check)
|
||||||
@@ -1423,10 +1468,75 @@ class ModelUpdateService:
|
|||||||
continue
|
continue
|
||||||
if target_set is not None and model_id not in target_set:
|
if target_set is not None and model_id not in target_set:
|
||||||
continue
|
continue
|
||||||
|
yield model_id, version_id, item.get("base_model")
|
||||||
|
|
||||||
|
def _prepare_collection_filters(
|
||||||
|
self,
|
||||||
|
target_model_ids: Optional[Sequence[int]],
|
||||||
|
folder_path: Optional[str],
|
||||||
|
) -> tuple[Optional[set[int]], Optional[str]]:
|
||||||
|
target_set: Optional[set[int]] = None
|
||||||
|
if target_model_ids:
|
||||||
|
target_set = set(target_model_ids)
|
||||||
|
|
||||||
|
normalized_folder = None
|
||||||
|
if folder_path is not None:
|
||||||
|
normalized_folder = folder_path.replace("\\", "/").strip("/")
|
||||||
|
return target_set, normalized_folder
|
||||||
|
|
||||||
|
async def _collect_local_versions(
|
||||||
|
self,
|
||||||
|
scanner,
|
||||||
|
*,
|
||||||
|
target_model_ids: Optional[Sequence[int]] = None,
|
||||||
|
folder_path: Optional[str] = None,
|
||||||
|
) -> Dict[int, List[int]]:
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
mapping: Dict[int, set[int]] = {}
|
||||||
|
target_set, normalized_folder = self._prepare_collection_filters(
|
||||||
|
target_model_ids, folder_path
|
||||||
|
)
|
||||||
|
|
||||||
|
if target_model_ids and not target_set:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
for model_id, version_id, _base_model in self._iter_local_civitai_items(
|
||||||
|
cache, target_set=target_set, normalized_folder=normalized_folder
|
||||||
|
):
|
||||||
mapping.setdefault(model_id, set()).add(version_id)
|
mapping.setdefault(model_id, set()).add(version_id)
|
||||||
|
|
||||||
return {model_id: sorted(ids) for model_id, ids in mapping.items()}
|
return {model_id: sorted(ids) for model_id, ids in mapping.items()}
|
||||||
|
|
||||||
|
async def _collect_local_version_bases(
|
||||||
|
self,
|
||||||
|
scanner,
|
||||||
|
*,
|
||||||
|
target_model_ids: Optional[Sequence[int]] = None,
|
||||||
|
) -> Dict[int, str]:
|
||||||
|
"""Map version id -> base model from cache items.
|
||||||
|
|
||||||
|
Deliberately unfiltered by folder: synthesized in-library entries must
|
||||||
|
carry a base regardless of which folder triggered the refresh.
|
||||||
|
"""
|
||||||
|
|
||||||
|
cache = await scanner.get_cached_data()
|
||||||
|
bases: Dict[int, str] = {}
|
||||||
|
target_set, _normalized_folder = self._prepare_collection_filters(
|
||||||
|
target_model_ids, None
|
||||||
|
)
|
||||||
|
|
||||||
|
if target_model_ids and not target_set:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
for _model_id, version_id, base_model in self._iter_local_civitai_items(
|
||||||
|
cache, target_set=target_set
|
||||||
|
):
|
||||||
|
normalized_base = _normalize_string(base_model)
|
||||||
|
if normalized_base:
|
||||||
|
bases[version_id] = normalized_base
|
||||||
|
|
||||||
|
return bases
|
||||||
|
|
||||||
def _merge_with_local_versions(
|
def _merge_with_local_versions(
|
||||||
self,
|
self,
|
||||||
existing: Optional[ModelUpdateRecord],
|
existing: Optional[ModelUpdateRecord],
|
||||||
@@ -1506,6 +1616,7 @@ class ModelUpdateService:
|
|||||||
timestamp: float,
|
timestamp: float,
|
||||||
*,
|
*,
|
||||||
all_local_version_ids: Optional[Sequence[int]] = None,
|
all_local_version_ids: Optional[Sequence[int]] = None,
|
||||||
|
local_base_models: Optional[Mapping[int, str]] = None,
|
||||||
) -> ModelUpdateRecord:
|
) -> ModelUpdateRecord:
|
||||||
local_set = set(local_versions)
|
local_set = set(local_versions)
|
||||||
# When folder-filtering, also consider versions in other folders
|
# When folder-filtering, also consider versions in other folders
|
||||||
@@ -1552,6 +1663,7 @@ class ModelUpdateService:
|
|||||||
|
|
||||||
missing_local = local_set - seen_ids
|
missing_local = local_set - seen_ids
|
||||||
if missing_local:
|
if missing_local:
|
||||||
|
item_base_models = local_base_models or {}
|
||||||
for version_id in sorted(missing_local):
|
for version_id in sorted(missing_local):
|
||||||
existing_version = existing_map.get(version_id)
|
existing_version = existing_map.get(version_id)
|
||||||
if existing_version:
|
if existing_version:
|
||||||
@@ -1566,7 +1678,7 @@ class ModelUpdateService:
|
|||||||
ModelVersionRecord(
|
ModelVersionRecord(
|
||||||
version_id=version_id,
|
version_id=version_id,
|
||||||
name=None,
|
name=None,
|
||||||
base_model=None,
|
base_model=item_base_models.get(version_id),
|
||||||
released_at=None,
|
released_at=None,
|
||||||
size_bytes=None,
|
size_bytes=None,
|
||||||
preview_url=None,
|
preview_url=None,
|
||||||
|
|||||||
@@ -0,0 +1,81 @@
|
|||||||
|
import os
|
||||||
|
import logging
|
||||||
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
from .base_model_service import BaseModelService
|
||||||
|
from .auto_tag_service import extract_auto_tags
|
||||||
|
from ..utils.models import OtherModelMetadata
|
||||||
|
from ..config import config
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class OtherModelService(BaseModelService):
|
||||||
|
"""Other-model-specific service implementation (VAE, upscaler, text encoder, ...)"""
|
||||||
|
|
||||||
|
def __init__(self, scanner, update_service=None):
|
||||||
|
"""Initialize Other-model service
|
||||||
|
|
||||||
|
Args:
|
||||||
|
scanner: Other-model scanner instance
|
||||||
|
update_service: Optional service for remote update tracking.
|
||||||
|
"""
|
||||||
|
super().__init__("other", scanner, OtherModelMetadata, update_service=update_service)
|
||||||
|
|
||||||
|
async def format_response(self, model_data: Dict[str, Any]) -> Optional[Dict[str, Any]]:
|
||||||
|
"""Format other-model data for API response.
|
||||||
|
|
||||||
|
Returns None when the entry is missing critical fields (corrupted cache
|
||||||
|
row), so the handler layer can filter it out. See issue #730.
|
||||||
|
"""
|
||||||
|
# Guard against corrupted cache entries missing critical fields
|
||||||
|
file_path = model_data.get("file_path")
|
||||||
|
if not file_path or not isinstance(file_path, str):
|
||||||
|
logger.warning(
|
||||||
|
"Skipping corrupted other-model entry (missing file_path): %s",
|
||||||
|
model_data.get("file_name", "<unknown>"),
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Get sub_type from cache entry (new canonical field)
|
||||||
|
sub_type = model_data.get("sub_type", "vae")
|
||||||
|
|
||||||
|
file_name = model_data.get("file_name") or ""
|
||||||
|
model_name = model_data.get("model_name") or file_name
|
||||||
|
folder = model_data.get("folder") or ""
|
||||||
|
|
||||||
|
return {
|
||||||
|
"model_name": model_name,
|
||||||
|
"file_name": file_name,
|
||||||
|
"preview_url": config.get_preview_static_url(model_data.get("preview_url", "")),
|
||||||
|
"preview_nsfw_level": model_data.get("preview_nsfw_level", 0),
|
||||||
|
"base_model": model_data.get("base_model", ""),
|
||||||
|
"folder": folder,
|
||||||
|
"sha256": model_data.get("sha256", ""),
|
||||||
|
"autov3": model_data.get("autov3"),
|
||||||
|
"file_path": file_path.replace(os.sep, "/"),
|
||||||
|
"file_size": model_data.get("size", 0),
|
||||||
|
"modified": model_data.get("modified", ""),
|
||||||
|
"tags": model_data.get("tags", []),
|
||||||
|
"from_civitai": model_data.get("from_civitai", True),
|
||||||
|
"notes": model_data.get("notes", ""),
|
||||||
|
"sub_type": sub_type,
|
||||||
|
"favorite": model_data.get("favorite", False),
|
||||||
|
"exclude": bool(model_data.get("exclude", False)),
|
||||||
|
"update_available": bool(model_data.get("update_available", False)),
|
||||||
|
"skip_metadata_refresh": bool(model_data.get("skip_metadata_refresh", False)),
|
||||||
|
"civitai": self.filter_civitai_data(model_data.get("civitai", {}), minimal=True),
|
||||||
|
"auto_tags": model_data.get("auto_tags") or extract_auto_tags(model_data),
|
||||||
|
"version_count": model_data.get("version_count"),
|
||||||
|
"source_platform": model_data.get("source_platform", ""),
|
||||||
|
"source_url": model_data.get("source_url", ""),
|
||||||
|
"hf_url": model_data.get("hf_url", ""),
|
||||||
|
}
|
||||||
|
|
||||||
|
def find_duplicate_hashes(self) -> Dict[str, Any]:
|
||||||
|
"""Find other models with duplicate SHA256 hashes"""
|
||||||
|
return self.scanner._hash_index.get_duplicate_hashes()
|
||||||
|
|
||||||
|
def find_duplicate_filenames(self) -> Dict[str, Any]:
|
||||||
|
"""Find other models with conflicting filenames"""
|
||||||
|
return self.scanner._hash_index.get_duplicate_filenames()
|
||||||
@@ -0,0 +1,478 @@
|
|||||||
|
# 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 json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
from ..utils.models import OtherModelMetadata
|
||||||
|
from ..utils.file_utils import find_preview_file, normalize_path, calculate_autov3
|
||||||
|
from ..utils.metadata_manager import MetadataManager
|
||||||
|
from ..config import config
|
||||||
|
from .model_scanner import ModelScanner, _is_excluded_dir
|
||||||
|
from .model_hash_index import ModelHashIndex
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class OtherScanner(ModelScanner):
|
||||||
|
"""Service for scanning and managing "other" model files.
|
||||||
|
|
||||||
|
Aggregates every enabled folder_paths category from
|
||||||
|
OTHER_MODEL_FOLDER_SUBTYPES (VAE, upscalers, text encoders, CLIP vision,
|
||||||
|
opt-in ControlNet) into one scanner; sub_type is derived from the root
|
||||||
|
containing the file (mirrors CheckpointScanner's checkpoints/unet split).
|
||||||
|
|
||||||
|
Hashing is lazy (checkpoint-style): text encoders can be ~10 GB, so the
|
||||||
|
initial scan records hash_status="pending" and the SHA256 is computed
|
||||||
|
on-demand via calculate_hash_for_model (e.g. when fetching CivitAI
|
||||||
|
metadata).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
# Same extension set as CheckpointScanner (ComfyUI's
|
||||||
|
# supported_pt_extensions plus ".gguf").
|
||||||
|
file_extensions = {
|
||||||
|
".ckpt",
|
||||||
|
".pt",
|
||||||
|
".pt2",
|
||||||
|
".bin",
|
||||||
|
".pth",
|
||||||
|
".safetensors",
|
||||||
|
".pkl",
|
||||||
|
".sft",
|
||||||
|
".gguf",
|
||||||
|
}
|
||||||
|
super().__init__(
|
||||||
|
model_type="other",
|
||||||
|
model_class=OtherModelMetadata,
|
||||||
|
file_extensions=file_extensions,
|
||||||
|
hash_index=ModelHashIndex(),
|
||||||
|
)
|
||||||
|
if not hasattr(self, "_hash_calculation_lock"):
|
||||||
|
self._hash_calculation_lock = asyncio.Lock()
|
||||||
|
self._hash_calculation_tasks: dict[str, asyncio.Task[Optional[str]]] = {}
|
||||||
|
|
||||||
|
async def _create_default_metadata(
|
||||||
|
self, file_path: str
|
||||||
|
) -> Optional[OtherModelMetadata]:
|
||||||
|
"""Create default metadata without calculating hash (lazy hash).
|
||||||
|
|
||||||
|
Other models include multi-GB text encoders, so hash calculation is
|
||||||
|
deferred until on-demand (e.g. CivitAI metadata fetch).
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
real_path = os.path.realpath(file_path)
|
||||||
|
if not os.path.exists(real_path):
|
||||||
|
logger.error(f"File not found: {file_path}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
base_name = os.path.splitext(os.path.basename(file_path))[0]
|
||||||
|
dir_path = os.path.dirname(file_path)
|
||||||
|
|
||||||
|
# Find preview image
|
||||||
|
preview_url = find_preview_file(base_name, dir_path)
|
||||||
|
|
||||||
|
# AutoV3 reads only the safetensors header, so it is cheap even for
|
||||||
|
# large files; record the checked state at creation time ("" =
|
||||||
|
# checked but unavailable).
|
||||||
|
autov3 = calculate_autov3(real_path)
|
||||||
|
|
||||||
|
# Create metadata WITHOUT calculating hash
|
||||||
|
metadata = OtherModelMetadata(
|
||||||
|
file_name=base_name,
|
||||||
|
model_name=base_name,
|
||||||
|
file_path=normalize_path(file_path),
|
||||||
|
size=os.path.getsize(real_path),
|
||||||
|
modified=datetime.now().timestamp(),
|
||||||
|
sha256="", # Empty hash - will be calculated on-demand
|
||||||
|
base_model="Unknown",
|
||||||
|
preview_url=normalize_path(preview_url),
|
||||||
|
tags=[],
|
||||||
|
modelDescription="",
|
||||||
|
sub_type=self.resolve_sub_type_for_path(file_path) or "vae",
|
||||||
|
from_civitai=False, # Mark as local model since no hash yet
|
||||||
|
hash_status="pending", # Mark hash as pending
|
||||||
|
autov3=autov3 or "",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Save the created metadata
|
||||||
|
logger.info(f"Creating other-model metadata (hash pending) for {file_path}")
|
||||||
|
await MetadataManager.save_metadata(file_path, metadata)
|
||||||
|
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
f"Error creating default other-model metadata for {file_path}: {e}"
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def calculate_hash_for_model(self, file_path: str) -> Optional[str]:
|
||||||
|
"""Calculate hash for a model on-demand with per-file singleflight.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_path: Path to the model file
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
SHA256 hash string, or None if calculation failed
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
real_path = os.path.realpath(file_path)
|
||||||
|
if not os.path.exists(real_path):
|
||||||
|
logger.error(f"File not found for hash calculation: {file_path}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
metadata, _ = await MetadataManager.load_metadata(
|
||||||
|
file_path, self.model_class
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
metadata is not None
|
||||||
|
and metadata.hash_status == "completed"
|
||||||
|
and metadata.sha256
|
||||||
|
):
|
||||||
|
# Ensure the in-memory hash index is populated even when
|
||||||
|
# the hash was already computed and persisted to the metadata
|
||||||
|
# file. Without this, usage tracking (and any other caller
|
||||||
|
# that queries get_hash_by_filename first) will miss on every
|
||||||
|
# lookup and keep calling back into this method, creating a
|
||||||
|
# tight loop that never populates the index.
|
||||||
|
self._hash_index.add_entry(
|
||||||
|
metadata.sha256.lower(),
|
||||||
|
file_path,
|
||||||
|
getattr(metadata, "autov3", None) or None,
|
||||||
|
)
|
||||||
|
return metadata.sha256
|
||||||
|
|
||||||
|
async with self._hash_calculation_lock:
|
||||||
|
metadata, _ = await MetadataManager.load_metadata(
|
||||||
|
file_path, self.model_class
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
metadata is not None
|
||||||
|
and metadata.hash_status == "completed"
|
||||||
|
and metadata.sha256
|
||||||
|
):
|
||||||
|
self._hash_index.add_entry(
|
||||||
|
metadata.sha256.lower(),
|
||||||
|
file_path,
|
||||||
|
getattr(metadata, "autov3", None) or None,
|
||||||
|
)
|
||||||
|
return metadata.sha256
|
||||||
|
|
||||||
|
task = self._hash_calculation_tasks.get(real_path)
|
||||||
|
if task is None:
|
||||||
|
task = asyncio.create_task(
|
||||||
|
self._run_hash_calculation_task(file_path, real_path)
|
||||||
|
)
|
||||||
|
self._hash_calculation_tasks[real_path] = task
|
||||||
|
|
||||||
|
return await asyncio.shield(task)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error calculating hash for {file_path}: {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def _run_hash_calculation_task(
|
||||||
|
self, file_path: str, real_path: str
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Run a hash calculation task and remove it from the in-flight map."""
|
||||||
|
try:
|
||||||
|
return await self._calculate_hash_for_model_uncached(file_path, real_path)
|
||||||
|
finally:
|
||||||
|
task = asyncio.current_task()
|
||||||
|
async with self._hash_calculation_lock:
|
||||||
|
if self._hash_calculation_tasks.get(real_path) is task:
|
||||||
|
del self._hash_calculation_tasks[real_path]
|
||||||
|
|
||||||
|
async def _calculate_hash_for_model_uncached(
|
||||||
|
self, file_path: str, real_path: str
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Calculate hash for a model without checking in-flight tasks."""
|
||||||
|
from ..utils.file_utils import calculate_sha256
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Load current metadata
|
||||||
|
metadata, should_skip = await MetadataManager.load_metadata(
|
||||||
|
file_path, self.model_class
|
||||||
|
)
|
||||||
|
if metadata is None:
|
||||||
|
if should_skip:
|
||||||
|
logger.error(f"Invalid metadata found for {file_path}")
|
||||||
|
return None
|
||||||
|
created_metadata = await self._create_default_metadata(file_path)
|
||||||
|
if created_metadata is None:
|
||||||
|
logger.error(f"No metadata found for {file_path}")
|
||||||
|
return None
|
||||||
|
metadata = created_metadata
|
||||||
|
|
||||||
|
# Check if hash is already calculated
|
||||||
|
if metadata.hash_status == "completed" and metadata.sha256:
|
||||||
|
# Populate the in-memory hash index even for pre-computed
|
||||||
|
# hashes, mirroring the fix in calculate_hash_for_model.
|
||||||
|
self._hash_index.add_entry(
|
||||||
|
metadata.sha256.lower(),
|
||||||
|
file_path,
|
||||||
|
getattr(metadata, "autov3", None) or None,
|
||||||
|
)
|
||||||
|
return metadata.sha256
|
||||||
|
|
||||||
|
# Update status to calculating
|
||||||
|
metadata.hash_status = "calculating"
|
||||||
|
await MetadataManager.save_metadata(file_path, metadata)
|
||||||
|
|
||||||
|
# Calculate hash
|
||||||
|
logger.info(f"Calculating hash for other model: {file_path}")
|
||||||
|
sha256 = await calculate_sha256(real_path)
|
||||||
|
|
||||||
|
# Update metadata with hash
|
||||||
|
metadata.sha256 = sha256
|
||||||
|
metadata.hash_status = "completed"
|
||||||
|
await MetadataManager.save_metadata(file_path, metadata)
|
||||||
|
|
||||||
|
# Update hash index
|
||||||
|
self._hash_index.add_entry(
|
||||||
|
sha256.lower(),
|
||||||
|
file_path,
|
||||||
|
getattr(metadata, "autov3", None) or None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Update the in-memory cache entry so that subsequent
|
||||||
|
# _persist_current_cache / _save_persistent_cache calls
|
||||||
|
# write the hash back to the SQLite models table. Without
|
||||||
|
# this the hash only lives in the metadata file and the
|
||||||
|
# in-memory hash index, both of which are lost across
|
||||||
|
# restarts, causing the same re-computation loop on the
|
||||||
|
# next session.
|
||||||
|
if self._cache is not None and self._cache.raw_data:
|
||||||
|
for entry in self._cache.raw_data:
|
||||||
|
if entry.get("file_path") == file_path:
|
||||||
|
entry["sha256"] = sha256.lower()
|
||||||
|
entry["hash_status"] = "completed"
|
||||||
|
self.bump_cache_version()
|
||||||
|
break
|
||||||
|
|
||||||
|
logger.info(f"Hash calculated for other model: {file_path}")
|
||||||
|
return sha256
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error calculating hash for {file_path}: {e}")
|
||||||
|
# Update status to failed
|
||||||
|
try:
|
||||||
|
metadata, _ = await MetadataManager.load_metadata(
|
||||||
|
file_path, self.model_class
|
||||||
|
)
|
||||||
|
if metadata:
|
||||||
|
metadata.hash_status = "failed"
|
||||||
|
await MetadataManager.save_metadata(file_path, metadata)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def calculate_all_pending_hashes(
|
||||||
|
self, progress_callback=None
|
||||||
|
) -> Dict[str, int]:
|
||||||
|
"""Calculate hashes for all other models with pending hash status.
|
||||||
|
|
||||||
|
If cache is not initialized, scans filesystem directly for metadata files
|
||||||
|
with hash_status != 'completed'.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
progress_callback: Optional callback(progress, total, current_file)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with 'completed', 'failed', 'total' counts
|
||||||
|
"""
|
||||||
|
# Try to get from cache first
|
||||||
|
cache = await self.get_cached_data()
|
||||||
|
|
||||||
|
if cache and cache.raw_data:
|
||||||
|
# Use cache if available
|
||||||
|
pending_models = [
|
||||||
|
item
|
||||||
|
for item in cache.raw_data
|
||||||
|
if item.get("hash_status") != "completed" or not item.get("sha256")
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
# Cache not initialized, scan filesystem directly
|
||||||
|
pending_models = await self._find_pending_models_from_filesystem()
|
||||||
|
|
||||||
|
if not pending_models:
|
||||||
|
return {"completed": 0, "failed": 0, "total": 0}
|
||||||
|
|
||||||
|
total = len(pending_models)
|
||||||
|
completed = 0
|
||||||
|
failed = 0
|
||||||
|
|
||||||
|
for i, model_data in enumerate(pending_models):
|
||||||
|
file_path = model_data.get("file_path")
|
||||||
|
if not file_path:
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
sha256 = await self.calculate_hash_for_model(file_path)
|
||||||
|
if sha256:
|
||||||
|
completed += 1
|
||||||
|
else:
|
||||||
|
failed += 1
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"Error calculating hash for {file_path}: {e}")
|
||||||
|
failed += 1
|
||||||
|
|
||||||
|
if progress_callback:
|
||||||
|
try:
|
||||||
|
await progress_callback(i + 1, total, file_path)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
return {"completed": completed, "failed": failed, "total": total}
|
||||||
|
|
||||||
|
async def _find_pending_models_from_filesystem(self) -> List[Dict[str, Any]]:
|
||||||
|
"""Scan filesystem for other-model metadata files with pending hash status."""
|
||||||
|
pending_models = []
|
||||||
|
|
||||||
|
for root_path in self.get_model_roots():
|
||||||
|
if not os.path.exists(root_path):
|
||||||
|
continue
|
||||||
|
|
||||||
|
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:
|
||||||
|
if not filename.endswith(".metadata.json"):
|
||||||
|
continue
|
||||||
|
|
||||||
|
metadata_path = os.path.join(dirpath, filename)
|
||||||
|
try:
|
||||||
|
with open(metadata_path, "r", encoding="utf-8") as f:
|
||||||
|
data = json.load(f)
|
||||||
|
|
||||||
|
# Check if hash is pending
|
||||||
|
hash_status = data.get("hash_status", "completed")
|
||||||
|
sha256 = data.get("sha256", "")
|
||||||
|
|
||||||
|
if hash_status != "completed" or not sha256:
|
||||||
|
# Find corresponding model file
|
||||||
|
model_name = filename.replace(".metadata.json", "")
|
||||||
|
model_path = None
|
||||||
|
|
||||||
|
# Look for model file with matching name
|
||||||
|
for ext in self.file_extensions:
|
||||||
|
potential_path = os.path.join(dirpath, model_name + ext)
|
||||||
|
if os.path.exists(potential_path):
|
||||||
|
model_path = potential_path
|
||||||
|
break
|
||||||
|
|
||||||
|
if model_path:
|
||||||
|
pending_models.append(
|
||||||
|
{
|
||||||
|
"file_path": model_path.replace(os.sep, "/"),
|
||||||
|
"hash_status": hash_status,
|
||||||
|
"sha256": sha256,
|
||||||
|
**{
|
||||||
|
k: v
|
||||||
|
for k, v in data.items()
|
||||||
|
if k
|
||||||
|
not in [
|
||||||
|
"file_path",
|
||||||
|
"hash_status",
|
||||||
|
"sha256",
|
||||||
|
]
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
except (json.JSONDecodeError, Exception) as e:
|
||||||
|
logger.debug(
|
||||||
|
f"Error reading metadata file {metadata_path}: {e}"
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
return pending_models
|
||||||
|
|
||||||
|
def _root_sub_type_map(self) -> Dict[str, str]:
|
||||||
|
"""Return the configured business root -> sub_type map."""
|
||||||
|
root_map = getattr(config, "other_root_subtypes", None)
|
||||||
|
return root_map if isinstance(root_map, dict) else {}
|
||||||
|
|
||||||
|
def _resolve_sub_type(self, root_path: Optional[str]) -> Optional[str]:
|
||||||
|
"""Resolve the sub_type for a configured root path."""
|
||||||
|
if not root_path:
|
||||||
|
return None
|
||||||
|
|
||||||
|
normalized_root = self._normalize_path_value(root_path)
|
||||||
|
for root, sub_type in self._root_sub_type_map().items():
|
||||||
|
if self._normalize_path_value(root) == normalized_root:
|
||||||
|
return sub_type
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
def resolve_sub_type_for_path(self, file_path: Optional[str]) -> Optional[str]:
|
||||||
|
"""Resolve sub_type from the configured root that contains the file.
|
||||||
|
|
||||||
|
Uses the longest-prefix match so nested roots (e.g. a controlnet root
|
||||||
|
inside a vae root) resolve to the most specific category.
|
||||||
|
"""
|
||||||
|
normalized_path = self._normalize_path_value(file_path)
|
||||||
|
if not normalized_path:
|
||||||
|
return None
|
||||||
|
|
||||||
|
best_length = 0
|
||||||
|
best_sub_type: Optional[str] = None
|
||||||
|
for root, sub_type in self._root_sub_type_map().items():
|
||||||
|
normalized_root = self._normalize_path_value(root)
|
||||||
|
if not normalized_root:
|
||||||
|
continue
|
||||||
|
if (
|
||||||
|
normalized_path == normalized_root
|
||||||
|
or normalized_path.startswith(f"{normalized_root}/")
|
||||||
|
) and len(normalized_root) > best_length:
|
||||||
|
best_length = len(normalized_root)
|
||||||
|
best_sub_type = sub_type
|
||||||
|
|
||||||
|
return best_sub_type
|
||||||
|
|
||||||
|
def adjust_metadata(self, metadata, file_path, root_path):
|
||||||
|
"""Adjust metadata during scanning to set sub_type."""
|
||||||
|
sub_type = self._resolve_sub_type(root_path) or self.resolve_sub_type_for_path(
|
||||||
|
file_path
|
||||||
|
)
|
||||||
|
if sub_type:
|
||||||
|
metadata.sub_type = sub_type
|
||||||
|
return metadata
|
||||||
|
|
||||||
|
def adjust_cached_entry(self, entry: Dict[str, Any]) -> Dict[str, Any]:
|
||||||
|
"""Adjust entries loaded from the persisted cache to ensure sub_type is set.
|
||||||
|
|
||||||
|
sub_type is location-derived: it is re-derived on cache load, never
|
||||||
|
trusted from the persisted snapshot.
|
||||||
|
"""
|
||||||
|
sub_type = self.resolve_sub_type_for_path(entry.get("file_path"))
|
||||||
|
if sub_type:
|
||||||
|
entry["sub_type"] = sub_type
|
||||||
|
return entry
|
||||||
|
|
||||||
|
def _should_keep_cached_entry(self, entry: Dict[str, Any]) -> bool:
|
||||||
|
"""Drop persisted entries whose folder is no longer a managed root.
|
||||||
|
|
||||||
|
sub_type is location-derived and config only maps enabled roots, so a
|
||||||
|
file under a disabled sub_type - or under any other root while the
|
||||||
|
feature is off - resolves to None here and is filtered out while the
|
||||||
|
persisted cache is hydrated.
|
||||||
|
"""
|
||||||
|
return self.resolve_sub_type_for_path(entry.get("file_path")) is not None
|
||||||
|
|
||||||
|
def get_model_roots(self) -> List[str]:
|
||||||
|
"""Get other-model root directories"""
|
||||||
|
roots: List[str] = []
|
||||||
|
roots.extend(config.other_roots or [])
|
||||||
|
# Remove duplicates while preserving order
|
||||||
|
seen: set[str] = set()
|
||||||
|
unique_roots: List[str] = []
|
||||||
|
for root in roots:
|
||||||
|
if root and root not in seen:
|
||||||
|
seen.add(root)
|
||||||
|
unique_roots.append(root)
|
||||||
|
return unique_roots
|
||||||
@@ -59,6 +59,7 @@ _MODEL_TYPE_PAGE_MAP = {
|
|||||||
"lora": "loras",
|
"lora": "loras",
|
||||||
"checkpoint": "checkpoints",
|
"checkpoint": "checkpoints",
|
||||||
"embedding": "embeddings",
|
"embedding": "embeddings",
|
||||||
|
"other": "other",
|
||||||
}
|
}
|
||||||
|
|
||||||
# Module-level alias so tests can spy on timer task creation without patching
|
# Module-level alias so tests can spy on timer task creation without patching
|
||||||
@@ -274,17 +275,21 @@ class PendingDeleteService:
|
|||||||
async def merge_batches(self, batch_ids: Sequence[str]) -> Optional[str]:
|
async def merge_batches(self, batch_ids: Sequence[str]) -> Optional[str]:
|
||||||
"""Merge several batches into the first batch's manifest.
|
"""Merge several batches into the first batch's manifest.
|
||||||
|
|
||||||
Winner is ``batch_ids[0]``. The staged files of losing batches are
|
Winner is ``batch_ids[0]``. Merging is MANIFEST-ONLY: staged files
|
||||||
MOVED (os.rename) into the winner's batch dir and their ``staged``
|
are NEVER moved, so the merge is a pure metadata operation with zero
|
||||||
paths rewritten in the merged manifest BEFORE any loser dir is
|
data IO and is inherently cross-volume safe (no EXDEV, no rollback).
|
||||||
removed. ``expires_at`` is re-anchored to ``now + TTL`` at merge time
|
Every loser's entries are appended to the winner's manifest with
|
||||||
and a FRESH purge timer is armed for the winner.
|
their ``staged`` paths unchanged (files keep living in the loser's
|
||||||
|
own batch dir - the sibling-of-model staging location), each loser
|
||||||
|
dir is recorded in the winner manifest's ``merged_sources``, and each
|
||||||
|
loser manifest is stamped ``merged_into`` so its own purge timer, a
|
||||||
|
post-restart sweep or a direct undo call no-op. ``expires_at`` is
|
||||||
|
re-anchored to ``now + TTL`` at merge time and a FRESH purge timer is
|
||||||
|
armed for the winner.
|
||||||
|
|
||||||
On any move failure every already-moved file is moved BACK and the
|
Returns the winner id, or ``None`` when the winner batch cannot be
|
||||||
original batch dirs/manifests are left intact; ``None`` is returned so
|
resolved (callers then fall back to the ``batch_ids`` array
|
||||||
callers fall back to the ``batch_ids`` array contract. Cross-volume
|
contract).
|
||||||
merges hit EXDEV here - expected and fine (the fallback is the normal
|
|
||||||
path for those bulks).
|
|
||||||
"""
|
"""
|
||||||
if not batch_ids:
|
if not batch_ids:
|
||||||
return None
|
return None
|
||||||
@@ -298,77 +303,68 @@ class PendingDeleteService:
|
|||||||
if winner_manifest is None:
|
if winner_manifest is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Track (entry, original_staged_path, loser_dir) for rollback.
|
# Build the merged manifest in memory: loser entries are appended
|
||||||
moved: List[Tuple[Dict[str, Any], str, str]] = []
|
# with their staged paths UNCHANGED - no file moves, no IO, no
|
||||||
processed_losers: List[Tuple[str, str]] = [] # (loser_id, loser_dir)
|
# EXDEV. Loser dirs remain as physical storage until the merged
|
||||||
|
# batch is undone or purged.
|
||||||
try:
|
merged_sources: List[str] = []
|
||||||
for loser_id in batch_ids[1:]:
|
seen_loser_dirs: Set[str] = set()
|
||||||
loser_dir = await self._find_batch_dir(loser_id)
|
for loser_id in batch_ids[1:]:
|
||||||
if not loser_dir or os.path.normpath(loser_dir) == os.path.normpath(
|
loser_dir = await self._find_batch_dir(loser_id)
|
||||||
winner_dir
|
if not loser_dir or os.path.normpath(loser_dir) == os.path.normpath(
|
||||||
):
|
winner_dir
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
loser_abs = os.path.abspath(loser_dir)
|
||||||
|
if loser_abs in seen_loser_dirs:
|
||||||
|
continue
|
||||||
|
seen_loser_dirs.add(loser_abs)
|
||||||
|
loser_manifest = self._read_manifest(loser_dir)
|
||||||
|
if loser_manifest is None:
|
||||||
|
# Corrupted loser: leave it for the sweep to quarantine.
|
||||||
|
continue
|
||||||
|
for entry in loser_manifest.get("entries") or []:
|
||||||
|
if entry.get("restored"):
|
||||||
continue
|
continue
|
||||||
loser_manifest = self._read_manifest(loser_dir)
|
staged_path = entry.get("staged")
|
||||||
if loser_manifest is None:
|
if not staged_path or not os.path.exists(staged_path):
|
||||||
# Corrupted loser: leave it for the sweep to quarantine.
|
|
||||||
continue
|
continue
|
||||||
for entry in loser_manifest.get("entries") or []:
|
winner_manifest["entries"].append(entry)
|
||||||
if entry.get("restored"):
|
merged_sources.append(loser_abs)
|
||||||
continue
|
|
||||||
staged_path = entry.get("staged")
|
|
||||||
if not staged_path or not os.path.exists(staged_path):
|
|
||||||
continue
|
|
||||||
new_staged = os.path.join(
|
|
||||||
winner_dir, os.path.basename(staged_path)
|
|
||||||
)
|
|
||||||
if os.path.exists(new_staged):
|
|
||||||
# os.rename would silently overwrite the existing
|
|
||||||
# staged file on POSIX - never drop a staged file.
|
|
||||||
# Abort the merge so callers fall back to the
|
|
||||||
# batch_ids array contract.
|
|
||||||
raise OSError(
|
|
||||||
f"Merge collision: {os.path.basename(staged_path)} "
|
|
||||||
"already staged in winner batch"
|
|
||||||
)
|
|
||||||
os.rename(staged_path, new_staged)
|
|
||||||
original_staged = entry["staged"]
|
|
||||||
entry["staged"] = os.path.abspath(new_staged)
|
|
||||||
winner_manifest["entries"].append(entry)
|
|
||||||
moved.append((entry, original_staged, loser_dir))
|
|
||||||
processed_losers.append((loser_id, loser_dir))
|
|
||||||
except OSError as exc:
|
|
||||||
logger.warning(
|
|
||||||
"Merge of %s failed after moving files: %s; rolling back",
|
|
||||||
list(batch_ids),
|
|
||||||
exc,
|
|
||||||
)
|
|
||||||
self._rollback_merge_moves(moved)
|
|
||||||
return None
|
|
||||||
|
|
||||||
# Re-anchor expiry and persist the merged manifest atomically.
|
# Re-anchor expiry and persist the merged manifest atomically - it
|
||||||
|
# becomes the ONLY source of truth for every merged file, wherever
|
||||||
|
# it physically lives.
|
||||||
winner_manifest["expires_at"] = (
|
winner_manifest["expires_at"] = (
|
||||||
int(time.time()) + PENDING_DELETE_TTL_SECONDS
|
int(time.time()) + PENDING_DELETE_TTL_SECONDS
|
||||||
)
|
)
|
||||||
|
if merged_sources:
|
||||||
|
winner_manifest["merged_sources"] = merged_sources
|
||||||
try:
|
try:
|
||||||
self._write_manifest_atomic(winner_dir, winner_manifest)
|
self._write_manifest_atomic(winner_dir, winner_manifest)
|
||||||
except OSError as exc:
|
except OSError as exc:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Failed to write merged manifest for %s: %s; rolling back",
|
"Failed to write merged manifest for %s: %s",
|
||||||
winner_id,
|
winner_id,
|
||||||
exc,
|
exc,
|
||||||
)
|
)
|
||||||
self._rollback_merge_moves(moved)
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# All moves committed: remove loser dirs (must be empty by now)
|
# Stamp each loser manifest so its own purge timer / a later sweep
|
||||||
# and drop them from the registry. Skipped losers (missing /
|
# / a direct undo call no-op: the winner owns those files from
|
||||||
# corrupted / same-dir) stay registered so the sweep still
|
# here on. Best-effort coordination; a failed stamp only risks the
|
||||||
# quarantines them, exactly as before the registry existed.
|
# loser being swept at its own (earlier) expiry after a restart.
|
||||||
for loser_id, loser_dir in processed_losers:
|
for loser_dir in merged_sources:
|
||||||
self._remove_manifest(loser_dir)
|
try:
|
||||||
self._remove_empty_dir(loser_dir)
|
self._mark_merged(loser_dir, winner_id)
|
||||||
await self._forget_batch(loser_id)
|
except OSError as exc: # pragma: no cover - best-effort
|
||||||
|
logger.warning(
|
||||||
|
"Failed to mark merged loser %s: %s", loser_dir, exc
|
||||||
|
)
|
||||||
|
|
||||||
|
# Losers are no longer independently managed.
|
||||||
|
for loser_dir in merged_sources:
|
||||||
|
await self._forget_batch(os.path.basename(loser_dir))
|
||||||
await self._remember_batch(winner_id, winner_dir)
|
await self._remember_batch(winner_id, winner_dir)
|
||||||
|
|
||||||
# Arm a fresh purge timer for the winner with the re-anchored
|
# Arm a fresh purge timer for the winner with the re-anchored
|
||||||
@@ -397,6 +393,16 @@ class PendingDeleteService:
|
|||||||
if manifest is None:
|
if manifest is None:
|
||||||
raise ValueError(f"Manifest missing for batch {batch_id}")
|
raise ValueError(f"Manifest missing for batch {batch_id}")
|
||||||
|
|
||||||
|
merged_into = manifest.get("merged_into")
|
||||||
|
if merged_into:
|
||||||
|
# The batch was merged into another batch: its staged files
|
||||||
|
# are owned by the winner's manifest. Undo via the winner so
|
||||||
|
# the whole merged batch stays consistent.
|
||||||
|
raise ValueError(
|
||||||
|
f"Batch {batch_id} was merged into batch {merged_into}; "
|
||||||
|
"undo that batch instead"
|
||||||
|
)
|
||||||
|
|
||||||
if manifest.get("state") == "restored":
|
if manifest.get("state") == "restored":
|
||||||
return self._undo_result(manifest)
|
return self._undo_result(manifest)
|
||||||
|
|
||||||
@@ -448,6 +454,10 @@ class PendingDeleteService:
|
|||||||
self._remove_manifest(batch_dir)
|
self._remove_manifest(batch_dir)
|
||||||
self._remove_empty_dir(batch_dir)
|
self._remove_empty_dir(batch_dir)
|
||||||
await self._forget_batch(batch_id)
|
await self._forget_batch(batch_id)
|
||||||
|
# Clean up merged loser dirs (their staged files were restored
|
||||||
|
# above) and drop them from the registry too.
|
||||||
|
for loser_id in self._remove_merged_batch_dirs(manifest):
|
||||||
|
await self._forget_batch(loser_id)
|
||||||
|
|
||||||
logger.info("Restored pending-delete batch %s", batch_id)
|
logger.info("Restored pending-delete batch %s", batch_id)
|
||||||
return self._undo_result(manifest)
|
return self._undo_result(manifest)
|
||||||
@@ -500,10 +510,36 @@ class PendingDeleteService:
|
|||||||
QUARANTINE them (preserving the pre-registry sweep semantics). The
|
QUARANTINE them (preserving the pre-registry sweep semantics). The
|
||||||
walk only descends into dirs literally named ``.lm-pending-delete``,
|
walk only descends into dirs literally named ``.lm-pending-delete``,
|
||||||
so false positives are structurally limited.
|
so false positives are structurally limited.
|
||||||
|
|
||||||
|
The filesystem walk itself runs in a worker thread so a large or slow
|
||||||
|
library cannot block the event loop at startup; only the (rare) batch
|
||||||
|
registration awaits run on the loop.
|
||||||
|
"""
|
||||||
|
roots = await self._get_all_model_roots()
|
||||||
|
loop = asyncio.get_event_loop()
|
||||||
|
staging_parents = await loop.run_in_executor(
|
||||||
|
None, # Use default thread pool
|
||||||
|
self._collect_staging_parents, # Run the tree walk off the loop
|
||||||
|
roots,
|
||||||
|
)
|
||||||
|
for staging_parent in staging_parents:
|
||||||
|
await self._register_batch_candidates(staging_parent)
|
||||||
|
|
||||||
|
def _collect_staging_parents(self, roots: Sequence[str]) -> List[str]:
|
||||||
|
"""Walk every model root and return its staging-parent dirs.
|
||||||
|
|
||||||
|
Pure synchronous filesystem discovery with no awaits: walks with
|
||||||
|
``followlinks=True, topdown=True``, prunes symlink cycles via a
|
||||||
|
per-root ``visited`` realpath set (realpath is used ONLY for this
|
||||||
|
dedup set - the returned paths are the unresolved business paths),
|
||||||
|
filters out :func:`_is_excluded_dir` dirs, and collects every dir
|
||||||
|
named ``.lm-pending-delete`` (including the case where a model root
|
||||||
|
itself is one). Results are returned in walk order.
|
||||||
"""
|
"""
|
||||||
from .model_scanner import _is_excluded_dir
|
from .model_scanner import _is_excluded_dir
|
||||||
|
|
||||||
for root in await self._get_all_model_roots():
|
staging_parents: List[str] = []
|
||||||
|
for root in roots:
|
||||||
if not os.path.isdir(root):
|
if not os.path.isdir(root):
|
||||||
continue
|
continue
|
||||||
visited: Set[str] = set()
|
visited: Set[str] = set()
|
||||||
@@ -518,21 +554,20 @@ class PendingDeleteService:
|
|||||||
visited.add(real_dir)
|
visited.add(real_dir)
|
||||||
if os.path.basename(dirpath) == PENDING_DELETE_DIR_NAME:
|
if os.path.basename(dirpath) == PENDING_DELETE_DIR_NAME:
|
||||||
# The current dir IS a staging parent (reachable only when
|
# The current dir IS a staging parent (reachable only when
|
||||||
# a model root itself is one): register its batches.
|
# a model root itself is one): collect its batches.
|
||||||
await self._register_batch_candidates(dirpath)
|
staging_parents.append(dirpath)
|
||||||
dirnames[:] = []
|
dirnames[:] = []
|
||||||
continue
|
continue
|
||||||
next_dirs: List[str] = []
|
next_dirs: List[str] = []
|
||||||
for name in dirnames:
|
for name in dirnames:
|
||||||
if name == PENDING_DELETE_DIR_NAME:
|
if name == PENDING_DELETE_DIR_NAME:
|
||||||
await self._register_batch_candidates(
|
staging_parents.append(os.path.join(dirpath, name))
|
||||||
os.path.join(dirpath, name)
|
|
||||||
)
|
|
||||||
elif _is_excluded_dir(name):
|
elif _is_excluded_dir(name):
|
||||||
continue
|
continue
|
||||||
else:
|
else:
|
||||||
next_dirs.append(name)
|
next_dirs.append(name)
|
||||||
dirnames[:] = next_dirs
|
dirnames[:] = next_dirs
|
||||||
|
return staging_parents
|
||||||
|
|
||||||
async def _register_batch_candidates(self, staging_parent: str) -> None:
|
async def _register_batch_candidates(self, staging_parent: str) -> None:
|
||||||
"""Register every non-orphaned batch subdir of a staging parent."""
|
"""Register every non-orphaned batch subdir of a staging parent."""
|
||||||
@@ -754,25 +789,35 @@ class PendingDeleteService:
|
|||||||
"Failed to remove staged copy %s: %s", staged_path, exc
|
"Failed to remove staged copy %s: %s", staged_path, exc
|
||||||
)
|
)
|
||||||
|
|
||||||
def _rollback_merge_moves(
|
def _mark_merged(self, loser_dir: str, winner_id: str) -> None:
|
||||||
self, moved: Sequence[Tuple[Dict[str, Any], str, str]]
|
"""Stamp ``merged_into`` on a loser manifest (best-effort).
|
||||||
) -> None:
|
|
||||||
"""Move already-merged files back to their original loser batch dirs."""
|
The stamp makes the loser's own purge timer, post-restart sweeps and
|
||||||
for _entry, original_staged, _loser_dir in reversed(list(moved)):
|
direct undo calls no-op, so the winner's merged batch stays the only
|
||||||
current = _entry.get("staged")
|
owner of the loser's staged files until it is undone or purged.
|
||||||
if not current or not original_staged:
|
"""
|
||||||
|
loser_manifest = self._read_manifest(loser_dir)
|
||||||
|
if loser_manifest is None:
|
||||||
|
return
|
||||||
|
loser_manifest["merged_into"] = winner_id
|
||||||
|
self._write_manifest_atomic(loser_dir, loser_manifest)
|
||||||
|
|
||||||
|
def _remove_merged_batch_dirs(self, manifest: Dict[str, Any]) -> List[str]:
|
||||||
|
"""Remove merged loser batch dirs once their files were handled.
|
||||||
|
|
||||||
|
Called after a merged batch has been fully undone or purged: each
|
||||||
|
loser manifest (stamped ``merged_into``) and its now-empty dir are
|
||||||
|
removed so the sweep never quarantines an orphaned staging dir.
|
||||||
|
Best-effort - returns the removed batch ids for registry cleanup.
|
||||||
|
"""
|
||||||
|
removed: List[str] = []
|
||||||
|
for src in manifest.get("merged_sources") or []:
|
||||||
|
if not isinstance(src, str) or not src:
|
||||||
continue
|
continue
|
||||||
if not os.path.exists(current):
|
self._remove_manifest(src)
|
||||||
continue
|
self._remove_empty_dir(src)
|
||||||
try:
|
removed.append(os.path.basename(src))
|
||||||
os.rename(current, original_staged)
|
return removed
|
||||||
except OSError as exc: # pragma: no cover - best-effort rollback
|
|
||||||
logger.warning(
|
|
||||||
"Failed to roll back merge move %s -> %s: %s",
|
|
||||||
current,
|
|
||||||
original_staged,
|
|
||||||
exc,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _purge_batch_dir(self, batch_dir: str) -> bool:
|
def _purge_batch_dir(self, batch_dir: str) -> bool:
|
||||||
"""Purge one batch dir. Returns True when the batch was purged/removed."""
|
"""Purge one batch dir. Returns True when the batch was purged/removed."""
|
||||||
@@ -786,6 +831,13 @@ class PendingDeleteService:
|
|||||||
self._quarantine_batch_dir(batch_dir)
|
self._quarantine_batch_dir(batch_dir)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
if manifest.get("merged_into"):
|
||||||
|
# Merged into another batch: the winner owns these staged files.
|
||||||
|
# The loser's own purge timer / post-restart sweep must not remove
|
||||||
|
# them early (the winner re-anchored the merged expiry to give the
|
||||||
|
# whole bulk one undo window).
|
||||||
|
return False
|
||||||
|
|
||||||
if manifest.get("state") == "restored":
|
if manifest.get("state") == "restored":
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@@ -817,6 +869,7 @@ class PendingDeleteService:
|
|||||||
|
|
||||||
self._remove_manifest(batch_dir)
|
self._remove_manifest(batch_dir)
|
||||||
self._remove_empty_dir(batch_dir)
|
self._remove_empty_dir(batch_dir)
|
||||||
|
self._remove_merged_batch_dirs(manifest)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def _quarantine_batch_dir(self, batch_dir: str) -> str:
|
def _quarantine_batch_dir(self, batch_dir: str) -> str:
|
||||||
@@ -931,6 +984,7 @@ class PendingDeleteService:
|
|||||||
"get_lora_scanner",
|
"get_lora_scanner",
|
||||||
"get_checkpoint_scanner",
|
"get_checkpoint_scanner",
|
||||||
"get_embedding_scanner",
|
"get_embedding_scanner",
|
||||||
|
"get_other_scanner",
|
||||||
):
|
):
|
||||||
getter = getattr(ServiceRegistry, getter_name, None)
|
getter = getattr(ServiceRegistry, getter_name, None)
|
||||||
if not callable(getter):
|
if not callable(getter):
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user